thomasht86's picture
Text-only ONNX (fp32 + int8) of google/embeddinggemma-2 for Vespa
7dd7a5b verified
Raw History Blame Contribute Delete
3.61 kB
"""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}")