import torch import torchvision.transforms as T from PIL import Image from transformers import AutoModelForImageClassification import requests from io import BytesIO MEAN = [0.5, 0.5, 0.5] STD = [0.5, 0.5, 0.5] transform = T.Compose([ T.Resize((32, 32), interpolation=T.InterpolationMode.BICUBIC), T.ToTensor(), T.Normalize(mean=MEAN, std=STD) ]) def load_image(source: str) -> Image.Image: if source.startswith("http://") or source.startswith("https://"): response = requests.get(source) return Image.open(BytesIO(response.content)).convert("RGB") return Image.open(source).convert("RGB") def predict(model, image_source: str, device: str = "cpu") -> dict: image = load_image(image_source) x = transform(image).unsqueeze(0).to(device) with torch.no_grad(): logits = model(pixel_values=x).logits probs = torch.softmax(logits, dim=-1)[0] top5 = probs.topk(5) return { model.config.id2label[i.item()]: f"{p.item()*100:.2f}%" for i, p in zip(top5.indices, top5.values) } if __name__ == "__main__": import sys source = sys.argv[1] if len(sys.argv) > 1 else "https://upload.wikimedia.org/wikipedia/commons/a/a7/Camponotus_flavomarginatus_ant.jpg" DEVICE = "cuda" if torch.cuda.is_available() else "cpu" model = AutoModelForImageClassification.from_pretrained( "MamaPearl/nula-cifar10-robust-v0", trust_remote_code=True ).to(DEVICE) model.eval() results = predict(model, source, device=DEVICE) for label, prob in results.items(): print(f"{label:15} {prob}")