Views
No views yet
| Epoch | Training Loss | Validation Accuracy |
|---|---|---|
| 1 | 3.36 | 41.75% |
| 2 | 2.78 | 47.14% |
| 3 | 2.64 | 47.40% |
1import torch
2import torchvision.transforms as transforms
3import timm
4import requests
5import json
6from PIL import Image
7
8config = json.loads(
9 requests.get(
10 "https://huggingface.co/WhirlwindAI/GVM/resolve/main/config.json"
11 ).text
12)
13
14model = timm.create_model(
15 "mobilenetv2_100",
16 pretrained=False,
17 num_classes=config["num_classes"]
18)
19
20state = torch.hub.load_state_dict_from_url(
21 "https://huggingface.co/WhirlwindAI/GVM/resolve/main/model.pth",
22 map_location="cpu"
23)
24
25model.load_state_dict(state)
26model.eval()
27
28transform = transforms.Compose([
29 transforms.Resize(256),
30 transforms.CenterCrop(224),
31 transforms.ToTensor(),
32 transforms.Normalize(
33 mean=[0.485,0.456,0.406],
34 std=[0.229,0.224,0.225]
35 )
36])
37
38image = Image.open("image.jpg").convert("RGB")
39tensor = transform(image).unsqueeze(0)
40
41prediction = model(tensor).argmax(1).item()
42
43print(config["class_names"][prediction])| Architecture | MobileNetV2 |
| Dataset | CIFAR-100 |
| Classes | 100 |
| Model Size | 14 MB |
| Framework | PyTorch |
| Inference | CPU & GPU Friendly |
model.pth
config.json
README.md