Downloads · 30 days
1
1% of all-time downloads
rexologue/vit_large_384_for_trees
vit_large_384_for_trees is a image classification model from rexologue. Use it when you need a label for an image. It is set up for timm. The card lists the license as mit.
This repository hosts a fine-tuned vitlargepatch16384 classifier
Downloads · 30 days
1
1% of all-time downloads
All-time downloads
122
Public
Repo size
1.2 GB
Likes
0
Public
Click a slice to open those files.
.bin1.2 GB · 100%
From the Hugging Face model README
This repository hosts a fine-tuned vit_large_patch16_384 classifier
import json, torch, timm
from huggingface_hub import hf_hub_download
from timm.data.transforms_factory import create_transform
from timm.data.constants import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD
from PIL import Image
REPO = "rexologue/vit_large_384_for_trees"
MODEL_NAME = "vit_large_patch16_384"
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
# 1) labels
labels_path = hf_hub_download(REPO, filename="labels.json")
with open(labels_path, "r", encoding="utf-8") as f:
raw = json.load(f)
labels = [raw[str(i)] for i in range(len(raw))] if isinstance(raw, dict) else list(raw)
# 2) weights
ckpt_path = hf_hub_download(REPO, filename="pytorch_model.bin")
state = torch.load(ckpt_path, map_location="cpu")
if any(k.startswith("module.") for k in state): # DDP fix
state = {k.replace("module.", "", 1): v for k, v in state.items()}
# 3) model
model = timm.create_model(MODEL_NAME, num_classes=len(labels), pretrained=False)
model.load_state_dict(state, strict=True)
model.to(DEVICE).eval()
# 4) preprocessing (ViT-L/16 @ 384 w/ ImageNet mean/std + bicubic)
transform = create_transform(
input_size=(3, 384, 384),
interpolation="bicubic",
mean=IMAGENET_DEFAULT_MEAN,
std=IMAGENET_DEFAULT_STD,
)
# 5) run
img = Image.open("your_image.jpg").convert("RGB")
x = transform(img).unsqueeze(0).to(DEVICE)
with torch.no_grad():
logits = model(x)
probs = torch.softmax(logits, dim=1)[0].cpu()
topk = probs.topk(k=min(5, len(labels)))
print([(labels[i], float(probs[i])) for i in topk.indices])