Team Ai
Apppublic

UNIQ-DEV/Image-Classification-Benchmark

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
google_vit.py39 linesDownload Raw Back to models
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