1from timm import create_model
2import torch
3from torchvision import transforms
4from PIL import Image
5
6# Load model
7model = create_model('vit_base_patch16_224', pretrained=False, num_classes=3)
8model.load_state_dict(torch.load("pytorch_model.bin"))
9model.eval()
10
11# Transform
12transform = transforms.Compose([
13 transforms.Resize((224, 224)),
14 transforms.ToTensor(),
15 transforms.Normalize(mean=[0.5]*3, std=[0.5]*3),
16])
17
18# Inference
19image = Image.open("example_mri.jpg").convert("RGB")
20tensor = transform(image).unsqueeze(0)
21output = model(tensor)
22pred = torch.argmax(output, dim=1)
23print("Predicted class:", pred.item())