""" CLIP zero-shot content-type classifier. Uses the native OpenCLIP PyTorch text encoder for high-quality text embeddings (the ONNX text encoder has degraded quality due to the eot_indices workaround). Image embeddings use the ONNX visual encoder which works well. """ import logging import numpy as np import torch from app.config import VisionSettings from app.services.vision.base import ContentClassifier, ClassificationResult logger = logging.getLogger(__name__) CATEGORY_PROMPTS = { "screenshot": [ "a screenshot of a computer screen", "a screenshot of a phone screen", "a screen capture of a user interface", ], "document": [ "a scanned document", "a photo of a document with printed text", "a photo of a page of text on paper", ], "receipt": [ "a photo of a receipt", "a photo of a bill or invoice", ], "meme": [ "an internet meme with text overlay", "a funny image with caption text", ], "artwork": [ "a painting or drawing", "a sketch or illustration", "digital art or graphic design", ], "photograph": [ "a photograph taken with a camera", "a real photo of a real scene or person", "a candid photograph", ], } class CLIPContentClassifier(ContentClassifier): """Zero-shot content classifier using CLIP text-image similarity. Uses native PyTorch for text encoding, ONNX for image encoding.""" def __init__(self, settings: VisionSettings): import open_clip self._min_confidence = settings.classifier.min_confidence # Load native model for text encoding only. # Use whichever model family the embedder is configured for so # the classification text vectors live in the same space as the # image embeddings. embedder_name = settings.embedder.name if embedder_name.startswith("siglip2"): model_arch = "ViT-B-16-SigLIP2" pretrained = "webli" else: model_arch = "ViT-B-32" pretrained = "laion2b_s34b_b79k" logger.info("Loading %s text encoder for content classification", model_arch) model, _, _ = open_clip.create_model_and_transforms( model_arch, pretrained=pretrained ) model.eval() self._model = model self._tokenizer = open_clip.get_tokenizer(model_arch) # Get the ONNX image embedder from the registry from app.services.vision.registry import registry self._embedder = registry.get_embedder() # Pre-compute text embeddings for each category self._category_embeddings: dict[str, np.ndarray] = {} for category, prompts in CATEGORY_PROMPTS.items(): tokens = self._tokenizer(prompts) with torch.no_grad(): text_features = model.encode_text(tokens) text_features /= text_features.norm(dim=-1, keepdim=True) avg = text_features.mean(dim=0) avg /= avg.norm() self._category_embeddings[category] = avg.numpy().astype(np.float32) logger.info("Content classifier ready with %d categories", len(self._category_embeddings)) def classify(self, image: np.ndarray) -> list[ClassificationResult]: img_vec = self._embedder.embed_image(image) # Cosine similarity against each category scores = {} for category, cat_vec in self._category_embeddings.items(): scores[category] = float(np.dot(img_vec, cat_vec)) # Sort by score descending ranked = sorted(scores.items(), key=lambda x: -x[1]) best_cat, best_score = ranked[0] second_score = ranked[1][1] margin = best_score - second_score # Normalize: 0.01 margin → ~0.5 confidence, 0.03+ → ~1.0 confidence = min(1.0, margin * 30) if confidence >= self._min_confidence: return [ClassificationResult(label=best_cat, confidence=confidence)] return []