UNIQ-DEV/Image-Classification-Benchmark
0
1from typing import List2from src.interface import ModelInterface3from src.data.classification_result import ClassificationResult4from transformers import ViTFeatureExtractor, ViTForImageClassification, ViTImageProcessor5import torch6 7class GoogleVit(ModelInterface):8 def __init__(self):9 print('init... google vit model')10 # Load ViT model and feature extractor11 self.feature_extractor = ViTFeatureExtractor.from_pretrained('google/vit-base-patch16-224')12 self.model = ViTForImageClassification.from_pretrained('google/vit-base-patch16-224')13 self.processor = ViTImageProcessor.from_pretrained('google/vit-base-patch16-224')14 15 def classify_image(self, image) -> List[ClassificationResult]:16 # Preprocess the image17 inputs = self.processor(images=image, return_tensors="pt")18 19 # Perform inference20 outputs = self.model(**inputs)21 logits = outputs.logits.detach().numpy()22 23 # Convert logits to probabilities using softmax (using PyTorch)24 probabilities = torch.nn.functional.softmax(torch.from_numpy(logits), dim=-1).numpy()25 26 # Get the top 5 predictions27 top_5 = torch.argsort(torch.from_numpy(probabilities), axis=-1, descending=True)[0][:5].numpy()28 29 # Create ClassificationResult objects with confidence information30 results = [31 ClassificationResult(32 class_name=self.model.config.id2label[top_5[i]],33 confidence=float(probabilities[0][top_5[i]])34 )35 for i in range(5)36 ]37 38 39 return results