"""Semantic / concept matching features using CLIP or SigLIP.""" import numpy as np from PIL import Image import torch import torch.nn.functional as F from typing import Dict, Optional, Tuple class SemanticFeatureExtractor: """Extracts concept-match features from image + user text using CLIP.""" def __init__(self, model, processor, device="cpu"): self.model = model self.processor = processor self.device = device @classmethod def from_clip(cls, model_id="openai/clip-vit-base-patch32", device="cpu"): """Load a standard CLIP model from transformers.""" try: from transformers import CLIPModel, CLIPProcessor model = CLIPModel.from_pretrained(model_id) processor = CLIPProcessor.from_pretrained(model_id) model = model.to(device) if device == "cpu": model = model.float() # avoid half precision issues on CPU return cls(model, processor, device) except Exception as e: print(f"[SemanticFeatureExtractor] CLIP load failed: {e}") return cls(None, None, device) def _encode_image(self, img: Image.Image) -> Optional[torch.Tensor]: """Encode image with CLIP via vision_model + visual_projection.""" if self.model is None or self.processor is None: return None try: inputs = self.processor(images=img, return_tensors="pt").to(self.device) with torch.no_grad(): vision_outputs = self.model.vision_model(pixel_values=inputs["pixel_values"]) image_embeds = vision_outputs.pooler_output image_features = self.model.visual_projection(image_embeds) return F.normalize(image_features, dim=-1) except Exception as e: print(f"[SemanticFeatureExtractor] Image encode failed: {e}") return None def _encode_text(self, texts: list) -> Optional[torch.Tensor]: """Encode text prompts with CLIP via text_model + text_projection.""" if self.model is None or self.processor is None: return None try: inputs = self.processor(text=texts, return_tensors="pt", padding=True).to(self.device) with torch.no_grad(): text_outputs = self.model.text_model(**inputs) text_embeds = text_outputs.pooler_output text_features = self.model.text_projection(text_embeds) return F.normalize(text_features, dim=-1) except Exception as e: print(f"[SemanticFeatureExtractor] Text encode failed: {e}") return None def compute(self, img: Image.Image, concept: str, audience: str, use_case: str) -> Dict[str, float]: """ Compute concept-match features. Returns dict with keys: concept_cosine, audience_cosine, usecase_cosine, composite_cosine, clip_image_norm, clip_image_std """ img_features = self._encode_image(img) if img_features is None: # Fallback: heuristic based on image size and color variety arr = np.array(img.convert("RGB")) return { "concept_cosine": 0.0, "audience_cosine": 0.0, "usecase_cosine": 0.0, "composite_cosine": 0.0, "clip_image_norm": 0.0, "clip_image_std": 0.0, } # Build prompts prompts = [ concept or "an image", f"{concept or 'an image'} for {audience}" if audience else concept or "an image", f"{use_case}: {concept or 'an image'}" if use_case else concept or "an image", ] text_features = self._encode_text(prompts) if text_features is None: return { "concept_cosine": 0.0, "audience_cosine": 0.0, "usecase_cosine": 0.0, "composite_cosine": 0.0, "clip_image_norm": float(img_features.norm(dim=-1).item()), "clip_image_std": float(img_features.std(dim=-1).item()), } # Cosine similarities (already normalized) sims = (img_features @ text_features.T).squeeze(0) return { "concept_cosine": float(sims[0].item()), "audience_cosine": float(sims[1].item()), "usecase_cosine": float(sims[2].item()), "composite_cosine": float(sims.mean().item()), "clip_image_norm": float(img_features.norm(dim=-1).item()), "clip_image_std": float(img_features.std(dim=-1).item()), } def compute_semantic_features(img: Image.Image, concept: str, audience: str, use_case: str, extractor: Optional[SemanticFeatureExtractor] = None) -> Dict[str, float]: """ Compute semantic/concept-match features. If extractor not provided, creates one (slow on first call). """ if extractor is None: extractor = SemanticFeatureExtractor.from_clip(device="cpu") return extractor.compute(img, concept, audience, use_case)