"""Verify text-only ONNX variants against the sentence-transformers fp32 reference. ONNX side mimics Vespa's hugging-face-embedder: tokenizer.json via `tokenizers` (add_special_tokens=True), batch size 1, mean pooling over last_hidden_state with the attention mask, then L2 normalize (and truncate+renormalize for MRL dims). """ import sys import numpy as np import onnxruntime as ort import torch from sentence_transformers import SentenceTransformer from tokenizers import Tokenizer OUT = sys.argv[1] TOKENIZER = sys.argv[2] Q = "task: search result | query: " D = "title: none | text: " long_doc = " ".join( f"Paragraph {i}: The aurora borealis appears when charged solar particles collide with " f"gases in Earth's upper atmosphere, producing light in green, red and violet hues." for i in range(80) ) texts = [ Q + "What causes the northern lights?", D + "The northern lights are caused by charged particles from the sun.", Q + "Which planet is known as the Red Planet?", D + "Mars, known for its reddish appearance, is often referred to as the Red Planet.", D + "Venus is often called Earth's twin because of its similar size and proximity.", Q + "Hvordan lager man brunost?", D + "Brunost lages ved å koke myse til sukkeret karamelliseres. 日本語のテキストも含む。", "task: code retrieval | query: reverse a linked list in python", D + "def reverse(head):\n prev = None\n while head:\n head.next, prev, head = prev, head, head.next\n return prev", "a", D + long_doc, ] st = SentenceTransformer( "google/embeddinggemma-2", model_kwargs={"torch_dtype": torch.float32}, config_kwargs={"vision_config": None, "audio_config": None}, ) with torch.no_grad(): ref = st.encode(texts, normalize_embeddings=True, batch_size=1, convert_to_numpy=True) tok = Tokenizer.from_file(TOKENIZER) st_ids = [st.tokenize([t])["input_ids"][0].tolist() for t in texts] vespa_ids = [tok.encode(t, add_special_tokens=True).ids for t in texts] print("token ids identical:", all(a == b for a, b in zip(st_ids, vespa_ids))) print("long doc tokens:", len(vespa_ids[-1]), "| sample ids:", vespa_ids[0][:4], "...", vespa_ids[0][-2:]) def norm(x): return x / np.linalg.norm(x, axis=-1, keepdims=True) for variant in ["fp32", "int8", "q4"]: sess = ort.InferenceSession(f"{OUT}/{variant}/model.onnx", providers=["CPUExecutionProvider"]) embs, sent = [], [] for ids in vespa_ids: ii = np.array([ids], dtype=np.int64) am = np.ones_like(ii) lhs, se = sess.run(["last_hidden_state", "sentence_embedding"], {"input_ids": ii, "attention_mask": am}) embs.append((lhs[0] * am[0, :, None]).sum(0) / am.sum()) sent.append(se[0]) embs = norm(np.stack(embs)) sent = norm(np.stack(sent)) print(f"\n== {variant}") print(" any NaN:", bool(np.isnan(embs).any())) print(" meanpool vs sentence_embedding output, min cos:", float((embs * sent).sum(-1).min())) for dim in [768, 512, 256, 128]: a, b = norm(embs[:, :dim]), norm(ref[:, :dim]) cos = (a * b).sum(-1) print(f" dim {dim}: cos vs ST min={cos.min():.6f} mean={cos.mean():.6f} | max abs diff={np.abs(a - b).max():.2e}") # Retrieval sanity: q0->d1, q2->d3 should rank first among docs; compare score matrices qi, di = [0, 2, 5, 7], [1, 3, 4, 6, 8] s_onnx, s_ref = embs[qi] @ embs[di].T, ref[qi] @ ref[di].T print(" ranking identical to ST:", bool((s_onnx.argsort(1) == s_ref.argsort(1)).all()), "| max score diff:", f"{np.abs(s_onnx - s_ref).max():.2e}")