Download scripts/verify.py from vespa-engine/embeddinggemma-2-ONNX: direct link, hf CLI and curl.
- Browser
- Download file 3.61 kB
-
https://huggingface.co/vespa-engine/embeddinggemma-2-ONNX/resolve/main/scripts/verify.py
- Command line
-
hf download hf://vespa-engine/embeddinggemma-2-ONNX/scripts/verify.py
-
curl -L -o verify.py https://huggingface.co/vespa-engine/embeddinggemma-2-ONNX/resolve/main/scripts/verify.py
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}") | |