Views
No views yet
1import torch
2from PIL import Image
3import torchvision.transforms as transforms
4
5# Load model
6model = torch.load('cifar10_cnn.pth')
7model.eval()
8
9# Prepare image
10transform = transforms.Compose([
11 transforms.Resize((32, 32)),
12 transforms.ToTensor(),
13 transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
14])
15
16img = Image.open('your_image.jpg')
17img_tensor = transform(img).unsqueeze(0)
18
19# Predict
20with torch.no_grad():
21 output = model(img_tensor)
22 predicted_class = output.argmax(dim=1).item()