import torch from PIL import Image from transformers import CLIPProcessor, CLIPModel from app.config.settings import settings from app.ml.image.clip_prediction_result import ClipPredictionResult class ClipClassifier: LABELS = [ "a normal photo", "a photo containing marijuana", "a photo containing drugs", "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 = [] for index, probability in enumerate(probs): label = self.LABELS[index] score = float(probability) scores[label] = score if label != "a normal photo" and score >= settings.CLIP_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 )