Team Ai
Apppublic

UNIQ-DEV/Image-Classification-Benchmark

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
classification_model.py52 linesDownload Raw Back to src
1from typing import List2from urllib.request import urlopen3from PIL import Image4from .data.model_data import ModelData5from .models.mobilenet_v3 import MobilenetV36from .models.clip_vit import ClipVit7from .models.google_vit import GoogleVit8from .models.resnet_50 import Resnet509 10from .data.classification_result import ClassificationResult11 12class ClassificationModel:13    """14    Base class for all classification models.15    """16 17    def __init__(self):18        self.load_model()19 20    def get_model_names(self):21        return [model.name for model in self.models]22 23    def get_model_data(self, model_name):24        for model in self.models:25            if model.name == model_name:26                return model27        raise Exception(f'Model {model_name} not found')28 29    def load_model(self):30        self.models = [31            ModelData('clip-vit-base-patch32', model_class=ClipVit()), 32            ModelData('mobilenet_v3', model_class=MobilenetV3()),33            ModelData('google-vit-base-patch16-224', model_class=GoogleVit()),34            ModelData('microsoft/resnet-50', model_class=Resnet50())35            ]36 37    def classify(self, model_name, image) -> List[ClassificationResult]:38        #print type of image 39        print('>> image type -->',type(image))40 41        #convert image to pil42        img = self.image_to_pil(image)43 44        model = self.get_model_data(model_name)45        return model.model_class.classify_image(img)      46 47    def image_to_pil(self, image):48        #if image is starts with https (means url), then download it 49        if image.startswith('https'):50            return Image.open(urlopen(image))51        return Image.open(image)52