1import torch
2import timm
3from PIL import Image
4from torchvision import transforms
5
6# モデル読み込み
7model = timm.create_model("vit_tiny_patch16_224", pretrained=False, num_classes=2)
8state_dict = torch.hub.load_state_dict_from_url(
9 "https://huggingface.co/shunya510hi/vit-cifar10-classifier/resolve/main/pytorch_model.bin",
10 map_location="cpu"
11)
12model.load_state_dict(state_dict)
13model.eval()
14
15# 画像前処理
16transform = transforms.Compose([
17 transforms.Resize((224, 224)),
18 transforms.ToTensor(),
19 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
20])
21
22# 推論
23image = Image.open("test_image.jpg").convert("RGB")
24input_tensor = transform(image).unsqueeze(0)
25
26with torch.no_grad():
27 output = model(input_tensor)
28 probabilities = torch.softmax(output, dim=1)
29 prediction = output.argmax(1).item()
30
31class_names = ["fake", "real"]
32print(f"Prediction: {class_names[prediction]}")
33print(f"Confidence: {probabilities[0][prediction].item():.2%}")