Views
No views yet
1from datasets import load_dataset
2from transformers import ViTImageProcessor, ViTForImageClassification
3import torch
4
5# Convert the image to RGB
6example = example["image"].convert('RGB')
7
8model_name = "orkungedik/hr-onboaring-doc-classifier"
9processor = ViTImageProcessor.from_pretrained(model_name)
10model = ViTForImageClassification.from_pretrained(model_name)
11
12inputs = processor(images=example, return_tensors="pt")
13
14with torch.no_grad():
15 outputs = model(**inputs)
16 logits = outputs.logits
17
18predicted_class_idx = logits.argmax(-1).item()
19label = model.config.id2label[predicted_class_idx]
20
21print(f"Predicted class: {label}")
22
23probs = torch.nn.functional.softmax(logits, dim=-1)
24top5 = torch.topk(probs, 5)
25
26for i in range(5):
27 idx = top5.indices[0][i].item()
28 prob = top5.values[0][i].item()
29 print(f"{model.config.id2label[idx]}: {prob:.4f}")