Views
No views yet
1from transformers.modeling_outputs import ImageClassifierOutput
2from transformers import ViTImageProcessor, ViTForImageClassification
3import torch
4from PIL import Image
5
6model_name_or_path = "Ojimi/vit-anime-caption"
7processor = ViTImageProcessor.from_pretrained(model_name_or_path)
8model = ViTForImageClassification.from_pretrained(model_name_or_path)
9threshold = 0.3
10
11device = torch.device('cuda')
12
13image = Image.open(YOUR_IMAGE_PATH)
14
15inputs = processor(image, return_tensors='pt')
16
17model.to(device=device)
18model.eval()
19
20
21with torch.no_grad():
22 pixel_values = inputs['pixel_values'].to(device=device)
23
24 outputs : ImageClassifierOutput = model(pixel_values=pixel_values)
25
26 logits = outputs.logits # The raw scores before applying any activation
27 sigmoid = torch.nn.Sigmoid() # Sigmoid function to convert logits to probabilities
28 logits : torch.FloatTensor = sigmoid(logits) # Applying sigmoid activation
29
30 predictions = [] # List to store predictions
31
32 for idx, p in enumerate(logits[0]):
33 if p > threshold: # Applying a threshold of 0.3 to consider a class prediction
34 predictions.append((model.config.id2label[idx], p.item())) # Storing class label and probability
35
36for tag in predictions:
37 print(tag)
38
39Sigmoid?