Spaces:
Running
Running
Download siglip.py from primerz/pixagram-search-v4: direct link, hf CLI and curl.
- Browser
- Download file 12.7 kB
-
https://huggingface.co/spaces/primerz/pixagram-search-v4/resolve/main/siglip.py
- Command line
-
hf download hf://spaces/primerz/pixagram-search-v4/siglip.py
-
curl -L -o siglip.py https://huggingface.co/spaces/primerz/pixagram-search-v4/resolve/main/siglip.py
12.7 kB
| """ | |
| SigLIP embedder shared by app.py (Space) and handler.py (Inference Endpoint). | |
| One model, two towers: images for the indexing pipeline, texts for search queries — both land | |
| in the same vector space, which is what makes text → image search work. Vectors are | |
| L2-normalised so cosine similarity is a dot product (Vectorize metric: cosine). | |
| Text is lowercased here, before tokenisation. SigLIP 2 was trained on lowercased text and its | |
| Gemma tokenizer is case-sensitive: on Pixagram artworks, Title Case queries drop from 97 % to | |
| 3 % recall@1. The original multilingual SigLIP tokenizer lowercases on its own, so this is a | |
| no-op for it. Doing it here covers every caller (search queries, blog résumés, the demo page). | |
| NaFlex models (google/siglip2-*-naflex) keep each image's aspect ratio: the processor resizes it to | |
| the largest size that fits MAX_NUM_PATCHES patches of 16x16 px, instead of squashing it to a fixed | |
| square. A 274x183 artwork is seen as about 19x13 patches rather than distorted to 16x16. | |
| Environment: | |
| MODEL_ID Hub id of a SigLIP/CLIP-family model (default: multilingual SigLIP base, 768-d; | |
| the v2 Space primerz/pixagram-siglip2 sets google/siglip2-base-patch16-256, the | |
| v3 Space google/siglip2-base-patch16-naflex) | |
| MAX_NUM_PATCHES NaFlex only: patch budget per image (default 256, the model's training default). | |
| Reported in every reply so the Worker can check it matches EMBED_PATCHES. | |
| MAX_BATCH images/texts per forward pass (default 32) | |
| MAX_CONCURRENCY embeddings computed at once (default: one per CPU, effective_cpus()); each | |
| gets CPUs / MAX_CONCURRENCY PyTorch threads. One per CPU gives the most | |
| throughput (a small forward pass does not spread well over many threads); | |
| fewer, with more threads each, answer a single query sooner. | |
| TORCH_THREADS PyTorch threads per embedding, overriding CPUs / MAX_CONCURRENCY | |
| EMBED_STUB "1" returns deterministic pseudo-vectors without loading any model — wiring tests only | |
| """ | |
| from __future__ import annotations | |
| import base64 | |
| import hashlib | |
| import io | |
| import math | |
| import os | |
| import threading | |
| from typing import Any, Dict, List, Optional | |
| from PIL import Image | |
| # Default = what the production Space runs. Change models per Space with the MODEL_ID variable, | |
| # so pushing hf/ never switches production's model by accident. | |
| MODEL_ID = os.environ.get("MODEL_ID", "google/siglip-base-patch16-256-multilingual") | |
| MAX_BATCH = int(os.environ.get("MAX_BATCH", "32")) | |
| MAX_NUM_PATCHES = int(os.environ.get("MAX_NUM_PATCHES", "256")) | |
| MAX_TEXT_TOKENS = 64 # SigLIP text-tower context | |
| STUB = os.environ.get("EMBED_STUB", "") == "1" | |
| STUB_DIM = int(os.environ.get("EMBED_STUB_DIM", "768")) | |
| def decode_image(b64: str) -> Image.Image: | |
| """base64 (optionally a data URI) → RGB PIL image, transparency composited over white.""" | |
| if b64.startswith("data:"): | |
| b64 = b64.split(",", 1)[1] | |
| raw = base64.b64decode(b64) | |
| img = Image.open(io.BytesIO(raw)) | |
| img.load() | |
| return to_rgb(img) | |
| def to_rgb(img: Image.Image) -> Image.Image: | |
| if img.mode in ("RGBA", "LA", "P"): | |
| rgba = img.convert("RGBA") | |
| bg = Image.new("RGBA", rgba.size, (255, 255, 255, 255)) | |
| img = Image.alpha_composite(bg, rgba) | |
| return img.convert("RGB") | |
| class Embedder: | |
| """Lazy-loading, thread-safe wrapper around a SigLIP model.""" | |
| def __init__(self, model_id: str = MODEL_ID, max_num_patches: int = MAX_NUM_PATCHES) -> None: | |
| self.model_id = model_id | |
| self._lock = threading.Lock() | |
| self._model = None | |
| self._processor = None | |
| self.device = "cpu" | |
| self.dim = STUB_DIM if STUB else 0 | |
| self.stub = STUB | |
| # NaFlex: set on load when the image processor takes a patch budget | |
| self.naflex = False | |
| self.max_num_patches = max_num_patches | |
| # sigmoid(logit_scale * cos + logit_bias) is the model's own text/image match probability | |
| self.calibration: Optional[Dict[str, float]] = None | |
| # ---- loading -------------------------------------------------------------------------- | |
| def load(self) -> "Embedder": | |
| if self.stub or self._model is not None: | |
| return self | |
| with self._lock: | |
| if self._model is not None: | |
| return self | |
| import torch | |
| from transformers import AutoModel, AutoProcessor | |
| self.device = "cuda" if torch.cuda.is_available() else "cpu" | |
| # os.cpu_count() reports the HOST's cores inside a container (often 64-96 on Spaces) | |
| # while the cgroup allows 2-8; oversubscribing PyTorch's thread pool that badly turns | |
| # an 80 ms text embedding into 8-12 s of spin-waiting. Size the pool to what the | |
| # container can actually run, shared by the embeddings computed at once. | |
| torch.set_num_threads(threads_per_job()) | |
| model = AutoModel.from_pretrained(self.model_id).to(self.device).eval() | |
| self._processor = AutoProcessor.from_pretrained(self.model_id) | |
| self.naflex = hasattr(getattr(self._processor, "image_processor", None), "max_num_patches") | |
| cfg = model.config | |
| self.dim = int(getattr(cfg, "projection_dim", 0) or getattr(cfg.text_config, "projection_size", 0) or cfg.text_config.hidden_size) | |
| # SigLIP's learned temperature and bias (exp() of the stored log-scale). CLIP-family models | |
| # without a bias report none. | |
| try: | |
| scale = float(model.logit_scale.exp().item()) | |
| bias = float(model.logit_bias.item()) if getattr(model, "logit_bias", None) is not None else None | |
| self.calibration = {"logit_scale": round(scale, 4), "logit_bias": round(bias, 4)} if bias is not None else None | |
| except Exception: | |
| self.calibration = None | |
| self._model = model | |
| # First inference pays one-off kernel/allocator warm-up (~2 s); do it here, not on a user query. | |
| try: | |
| self.embed_texts(["warm up"]) | |
| self.embed_images([Image.new("RGB", (32, 32), (128, 128, 128))]) | |
| except Exception: | |
| pass | |
| return self | |
| def ready(self) -> bool: | |
| return self.stub or self._model is not None | |
| # ---- embedding ------------------------------------------------------------------------ | |
| def embed_images(self, images: List[Image.Image]) -> List[List[float]]: | |
| if self.stub: | |
| return [_stub_vector(hashlib.sha256(im.tobytes()).hexdigest()) for im in images] | |
| self.load() | |
| import torch | |
| out: List[List[float]] = [] | |
| # NaFlex: images of any aspect ratio share a batch; the processor pads them to the patch | |
| # budget and passes the attention mask and each image's patch grid to the model. | |
| kw = {"max_num_patches": self.max_num_patches} if self.naflex else {} | |
| with torch.inference_mode(): | |
| for i in range(0, len(images), MAX_BATCH): | |
| batch = [to_rgb(im) for im in images[i : i + MAX_BATCH]] | |
| inputs = self._processor(images=batch, return_tensors="pt", **kw).to(self.device) | |
| feats = self._model.get_image_features(**inputs) | |
| feats = _pooled(feats) | |
| feats = torch.nn.functional.normalize(feats, dim=-1) | |
| out.extend(feats.cpu().float().tolist()) | |
| return out | |
| def embed_texts(self, texts: List[str]) -> List[List[float]]: | |
| if self.stub: | |
| return [_stub_vector(hashlib.sha256(t.strip().lower().encode()).hexdigest()) for t in texts] | |
| self.load() | |
| import torch | |
| out: List[List[float]] = [] | |
| with torch.inference_mode(): | |
| for i in range(0, len(texts), MAX_BATCH): | |
| # Lowercase: SigLIP 2 was trained on lowercased text (see module docstring). | |
| batch = [t.lower() if t.strip() else " " for t in texts[i : i + MAX_BATCH]] | |
| # SigLIP was trained with padding="max_length"; keep it so query vectors match. | |
| inputs = self._processor( | |
| text=batch, padding="max_length", truncation=True, max_length=MAX_TEXT_TOKENS, return_tensors="pt" | |
| ).to(self.device) | |
| feats = self._model.get_text_features(**inputs) | |
| feats = _pooled(feats) | |
| feats = torch.nn.functional.normalize(feats, dim=-1) | |
| out.extend(feats.cpu().float().tolist()) | |
| return out | |
| # ---- request handling shared by the Space route and the Endpoint handler ----------------- | |
| def handle(self, payload: Dict[str, Any]) -> Dict[str, Any]: | |
| """ | |
| {"inputs": {"images": [b64, ...], "texts": [str, ...]}} (either key optional) | |
| → {"model", "dim", "embeddings": [image vectors..., text vectors...], "calibration": {"logit_scale", "logit_bias"}} | |
| """ | |
| inputs = payload.get("inputs", payload) if isinstance(payload, dict) else payload | |
| if isinstance(inputs, str): | |
| inputs = {"texts": [inputs]} | |
| if not isinstance(inputs, dict): | |
| raise ValueError("body must be {\"inputs\": {\"images\": [...], \"texts\": [...]}}") | |
| images_b64 = inputs.get("images") or [] | |
| texts = inputs.get("texts") or [] | |
| if isinstance(images_b64, str): | |
| images_b64 = [images_b64] | |
| if isinstance(texts, str): | |
| texts = [texts] | |
| if not images_b64 and not texts: | |
| raise ValueError("provide inputs.images (base64) and/or inputs.texts") | |
| if len(images_b64) + len(texts) > 256: | |
| raise ValueError("at most 256 items per request") | |
| embeddings: List[List[float]] = [] | |
| if images_b64: | |
| embeddings.extend(self.embed_images([decode_image(b) for b in images_b64])) | |
| if texts: | |
| embeddings.extend(self.embed_texts([str(t) for t in texts])) | |
| out: Dict[str, Any] = {"model": self.model_id if not self.stub else "stub", "dim": len(embeddings[0]) if embeddings else self.dim, "embeddings": embeddings} | |
| if self.calibration: | |
| out["calibration"] = self.calibration | |
| if self.naflex: | |
| out["max_num_patches"] = self.max_num_patches | |
| return out | |
| def effective_cpus(cap: int = 16) -> int: | |
| """CPUs this process may really use: affinity mask ∩ cgroup quota, capped.""" | |
| n = os.cpu_count() or 1 | |
| try: | |
| n = min(n, len(os.sched_getaffinity(0))) | |
| except Exception: | |
| pass | |
| try: # cgroup v2 | |
| quota, period = open("/sys/fs/cgroup/cpu.max").read().split()[:2] | |
| if quota != "max": | |
| n = min(n, max(1, int(int(quota) / int(period)))) | |
| except Exception: | |
| try: # cgroup v1 | |
| quota = int(open("/sys/fs/cgroup/cpu/cpu.cfs_quota_us").read()) | |
| period = int(open("/sys/fs/cgroup/cpu/cpu.cfs_period_us").read()) | |
| if quota > 0: | |
| n = min(n, max(1, quota // period)) | |
| except Exception: | |
| pass | |
| return max(1, min(n, cap)) | |
| def concurrency() -> int: | |
| """Embeddings computed at once: MAX_CONCURRENCY, else one per CPU.""" | |
| env = os.environ.get("MAX_CONCURRENCY", "") | |
| return max(1, int(env)) if env.isdigit() else effective_cpus() | |
| def threads_per_job() -> int: | |
| """PyTorch threads per embedding: TORCH_THREADS, else the CPUs shared by the embeddings at once.""" | |
| env = os.environ.get("TORCH_THREADS", "") | |
| if env.isdigit() and int(env) > 0: | |
| return int(env) | |
| return max(1, effective_cpus() // concurrency()) | |
| def _pooled(feats): | |
| """transformers ≥5 may return a ModelOutput; older versions return the tensor directly.""" | |
| if hasattr(feats, "pooler_output") and feats.pooler_output is not None: | |
| return feats.pooler_output | |
| if hasattr(feats, "last_hidden_state") and not hasattr(feats, "shape"): | |
| return feats.last_hidden_state | |
| return feats | |
| def _stub_vector(seed_hex: str) -> List[float]: | |
| """Deterministic unit vector from a hash — lets the Worker ↔ Space wiring be tested without weights.""" | |
| vals: List[float] = [] | |
| h = seed_hex | |
| while len(vals) < STUB_DIM: | |
| h = hashlib.sha256(h.encode()).hexdigest() | |
| vals.extend(int(h[i : i + 2], 16) / 127.5 - 1.0 for i in range(0, 64, 2)) | |
| vals = vals[:STUB_DIM] | |
| norm = math.sqrt(sum(v * v for v in vals)) or 1.0 | |
| return [v / norm for v in vals] | |
| _shared: Optional[Embedder] = None | |
| def shared() -> Embedder: | |
| global _shared | |
| if _shared is None: | |
| _shared = Embedder() | |
| return _shared | |