File size: 3,737 Bytes
56749a4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
50ba853
 
56749a4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
"""
SigLIP text-tower query encoder via ONNX Runtime (CPU, no torch).

Mirrors ingest/embedder.py's image tower — same checkpoint, same `pooler_output`
head, same L2-normalization — so a query vector and a stored image vector share
one cosine space. The only differences: the text tower instead of the vision
tower, and a tokenizer instead of pixel preprocessing.
"""
from __future__ import annotations

import functools
import os

import numpy as np
import onnxruntime as ort
from huggingface_hub import hf_hub_download
from transformers import AutoTokenizer

# Same checkpoint the image pipeline embeds with (ingest/embedder.py) and the web
# app references (env.server.ts SIGLIP_MODEL_ID). Must stay in lockstep or text
# and image vectors stop sharing a cosine space.
MODEL_REPO = os.environ.get("SIGLIP_MODEL_ID", "onnx-community/siglip2-base-patch16-256-ONNX")
# int8 text tower (283 MB) — matches the web app's q8 dtype.
ONNX_FILE = "onnx/text_model_quantized.onnx"
EMBED_DIM = 768
# SigLIP2's GemmaTokenizer has a sentinel model_max_length; the real training pad
# length lives in tokenizer_config.json (64). Must match the web tokenizer call.
MAX_LENGTH = 64
CACHE_SIZE = 512

_TEXT_EMBED_OUTPUT = "pooler_output"
_INPUT_IDS = "input_ids"

_SESSION: ort.InferenceSession | None = None
_TOKENIZER = None


def load_session() -> ort.InferenceSession:
    """Resolve the text ONNX (baked into the image in production, HF-hub-cached
    elsewhere) and build a CPU session once."""
    global _SESSION
    if _SESSION is None:
        path = hf_hub_download(MODEL_REPO, ONNX_FILE)
        _SESSION = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
    return _SESSION


def load_tokenizer():
    """Load the GemmaTokenizer for the checkpoint once."""
    global _TOKENIZER
    if _TOKENIZER is None:
        _TOKENIZER = AutoTokenizer.from_pretrained(MODEL_REPO)
    return _TOKENIZER


def warmup() -> None:
    """Load both heavy objects so the first real request doesn't pay for it."""
    load_session()
    load_tokenizer()


@functools.lru_cache(maxsize=CACHE_SIZE)
def _encode_cached(query: str) -> tuple[float, ...]:
    session = load_session()
    tokenizer = load_tokenizer()
    enc = tokenizer(
        [query],
        padding="max_length",
        truncation=True,
        max_length=MAX_LENGTH,
        return_tensors="np",
    )
    output_names = [o.name for o in session.get_outputs()]
    # Fail loudly rather than embedding from the wrong head if a future export
    # drops/renames pooler_output (matches ingest/embedder.py's guard).
    if _TEXT_EMBED_OUTPUT not in output_names:
        raise ValueError(
            f"ONNX model does not expose '{_TEXT_EMBED_OUTPUT}'; available: {output_names}"
        )
    # Feed only the inputs the model declares (SigLIP text uses input_ids; some
    # exports also take attention_mask). int64 is what ONNX expects for ids.
    input_names = {i.name for i in session.get_inputs()}
    feed = {name: enc[name].astype(np.int64) for name in input_names if name in enc}
    if _INPUT_IDS not in feed:
        raise ValueError(
            f"ONNX model does not accept '{_INPUT_IDS}'; inputs: {sorted(input_names)}"
        )
    out = session.run([_TEXT_EMBED_OUTPUT], feed)[0]
    vec = np.asarray(out, dtype=np.float32).reshape(EMBED_DIM)
    norm = float(np.linalg.norm(vec))
    vec = vec / max(norm, 1e-12)
    return tuple(float(x) for x in vec)


def encode(query: str) -> list[float]:
    """Encode a query into a 768-d L2-normalized vector, cosine-comparable to
    photos.embedding. Cached by exact query string."""
    if not query or not query.strip():
        raise ValueError("encode: empty query")
    return list(_encode_cached(query))