1import torch
2from torchvision import transforms
3from PIL import Image
4
5# 1. 加载模型
6model = torch.load("ViT-full120.pt")
7model.eval()
8
9# 2. 图像预处理
10transform = transforms.Compose([
11 transforms.Resize((224, 224)),
12 transforms.ToTensor(),
13 transforms.Normalize(
14 mean=[0.485, 0.456, 0.406],
15 std=[0.229, 0.224, 0.225]
16 ),
17])
18
19# 3. 推理
20image = Image.open("dog.jpg").convert("RGB")
21input_tensor = transform(image).unsqueeze(0)
22
23with torch.no_grad():
24 output = model(input_tensor)
25 predicted_class = output.argmax(dim=1).item()
26 print(f"预测类别: {predicted_class}")