Views
No views yet
google/vit-large-patch16-224-in21k on ImageNet-1K.
本仓库提供了一个基于 google/vit-large-patch16-224-in21k 模型,使用 LoRA 技术在 ImageNet-1K 数据集上微调后的版本。| Metric 指标 | Score 分数 |
|---|---|
| Top-1 Accuracy | 82.66% |
| Top-5 Accuracy | 93.10% |
| Item 项目 | Value 配置 |
|---|---|
| Base Model | google/vit-large-patch16-224-in21k |
| Method | LoRA |
| LoRA Rank (r) | 64 |
| Trainable Params | 28,311,552 (trainable%: 8.5112) |
| Dataset | ImageNet-1K |
| Epochs | 20 (Early Stop at 16) |
| Optimizer | AdamW |
| Scheduler | CosineAnnealingWarmRestarts |
| Precision | AMP (FP16/BF16) |
| Input Resolution | 224×224 |
1from transformers import ViTForImageClassification
2from peft import PeftModel
3import torch
4
5device = "cuda" if torch.cuda.is_available() else "cpu"
6
7base = ViTForImageClassification.from_pretrained(
8 "google/vit-large-patch16-224-in21k",
9 num_labels=1000
10)
11
12model = PeftModel.from_pretrained(base, "xn6o/lora-vit-large-patch16-224-in21k-r64-imagenet1k")
13model.to(device).eval()1state_dict = torch.load("vit_classifier_best.pt", map_location=device)
2model.base_model.classifier.load_state_dict(state_dict)1merged = model.merge_and_unload()
2merged.save_pretrained("vit_lora_merged_best")1from PIL import Image
2from torchvision import transforms
3import torch.nn.functional as F
4
5img = Image.open("your_image.jpg")
6
7tfm = transforms.Compose([
8 transforms.Resize(256),
9 transforms.CenterCrop(224),
10 transforms.ToTensor(),
11 transforms.Normalize([0.485, 0.456, 0.406],
12 [0.229, 0.224, 0.225]),
13])
14
15pixel_values = tfm(img).unsqueeze(0).to(device)
16
17with torch.no_grad():
18 logits = model(pixel_values=pixel_values).logits
19 probs = F.softmax(logits, dim=1)
20
21top5 = probs.topk(5)
22print("Top-5 Classes:", top5.indices.squeeze().tolist())
23print("Top-5 Probabilities:", top5.values.squeeze().tolist())|-- adapter/ # LoRA adapter weights
|-- vit_classifier.pt # Fine-tuned classifier head
|-- README.md # This model card