Views
No views yet
| Property | Value |
|---|---|
| Model Type | FIRE Vision Transformer |
| Dataset | imagenet100 |
| Best Accuracy | 74.22% |
| Image Size | 224 |
| Patch Size | 16 |
| Hidden Dim | 192 |
| Depth | 12 |
| Num Heads | 3 |
| MLP Dim | 768 |
| Num Classes | 100 |
1import torch
2from models import FIRESimpleVisionTransformer
3
4# Initialize model
5model = FIRESimpleVisionTransformer(
6 image_size=224,
7 patch_size=16,
8 num_layers=12,
9 num_heads=3,
10 hidden_dim=192,
11 mlp_dim=768,
12 num_classes=100,
13)
14
15# Load checkpoint
16checkpoint = torch.load('fire_vit_imagenet100_best.pth', map_location='cpu')
17state_dict = checkpoint['state_dict']
18
19# Remove 'module.' prefix if present (from DDP training)
20state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}
21model.load_state_dict(state_dict)
22model.eval()
23
24# Inference
25from torchvision import transforms
26from PIL import Image
27
28transform = transforms.Compose([
29 transforms.Resize(256),
30 transforms.CenterCrop(224),
31 transforms.ToTensor(),
32 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
33])
34
35image = Image.open('your_image.jpg').convert('RGB')
36input_tensor = transform(image).unsqueeze(0)
37
38with torch.no_grad():
39 output = model(input_tensor)
40 prediction = output.argmax(dim=1)1@misc{vit-analysis,
2 title={Vision Transformer Position Encoding Analysis},
3 year={2024},
4 url={https://github.com/your-repo/vit-analysis}
5}