import torch from PIL import Image from transformers import ( CLIPProcessor, CLIPModel ) from app.ml.image.clip_prediction_result import ClipPredictionResult from app.config.settings import settings class ClipClassifier: LABELS = [ "a normal photo", "a photo containing marijuana", "a photo containing drugs", "a photo containing weapons", "a pornographic photo", "a photo containing violence" ] def __init__(self): self.device = ( "cuda" if torch.cuda.is_available() else "cpu" ) print( f"Loading CLIP on {self.device}" ) self.model = CLIPModel.from_pretrained( settings.IMAGE_CLIP_MODEL ) self.processor = CLIPProcessor.from_pretrained( settings.IMAGE_CLIP_MODEL ) self.model.to(self.device) self.model.eval() def predict( self, image: Image.Image ) -> ClipPredictionResult: inputs = self.processor( text=self.LABELS, images=image, return_tensors="pt", padding=True ) inputs = { k: v.to(self.device) for k, v in inputs.items() } with torch.no_grad(): outputs = self.model(**inputs) logits = ( outputs .logits_per_image ) probs = ( logits.softmax(dim=1)[0] ) scores = {} detected_labels = [] THRESHOLD = 0.75 for index, probability in enumerate(probs): label = self.LABELS[index] score = float(probability) scores[label] = score if ( label != "a normal photo" and score >= THRESHOLD ): detected_labels.append(label) max_index = torch.argmax(probs).item() return ClipPredictionResult( label=self.LABELS[max_index], score=float( probs[max_index] ), detected_labels=detected_labels, scores=scores )