UNIQ-DEV/Image-Classification-Benchmark
0
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 