Views
No views yet
1from torchvision import transforms
2from PIL import Image
3import torch
4from torchvision.models import resnet18
5
6model = resnet18(weights=None)
7model.load_state_dict(torch.load('path_to_model/pytorch_model.bin'))
8model.eval()
9
10transform = transforms.Compose([
11 transforms.Resize((224, 224)),
12 transforms.ToTensor(),
13 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
14])
15
16image = Image.open('path_to_image.jpg')
17image = transform(image).unsqueeze(0)
18
19with torch.no_grad():
20 output = model(image)
21 _, predicted = torch.max(output.data, 1)
22 print(predicted.item())