Views
No views yet
1import torch
2from torchvision import transforms
3from PIL import Image
4import json
5
6# Load class names
7with open("class_names.json") as f:
8 class_names = json.load(f)
9
10# Load model
11checkpoint = torch.load("best_model.pth", map_location="cpu")
12
13# Preprocess image
14transform = transforms.Compose([
15 transforms.Resize((224, 224)),
16 transforms.ToTensor(),
17 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
18])
19
20image = Image.open("your_leaf.jpg").convert("RGB")
21tensor = transform(image).unsqueeze(0) # add batch dimension
22
23# Predict
24with torch.no_grad():
25 outputs = model(tensor)
26 _, predicted = outputs.max(1)
27 print(f"Prediction: {class_names[predicted.item()]}")