Something went wrong. Try again.
sentence embeddings in pure zig: bge-small with an HF-exact tokenizer and an SME matmul
Something went wrong. Try again.
12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879# /// script# requires-python = ">=3.12"# dependencies = ["sentence-transformers>=5", "torch>=2.6", "numpy"]# ///"""Dump the reference implementation's view of a corpus, and time it.
usage: uv run scripts/reference.py CORPUS.jsonl OUT_DIR [--model HF_ID] [--bench]
writes, in OUT_DIR: model_path.txt local snapshot directory holding model.safetensors + tokenizer.json ids.jsonl one JSON array of token ids per post (what the model actually sees) embeddings.f32 N x DIM little-endian float32, row per post, as the model outputs them (normalized only if its modules include Normalize) meta.json model id, dims, count, torch version, thread count, timings"""
import jsonimport sysimport timefrom pathlib import Path
import numpy as npimport torchfrom huggingface_hub import snapshot_downloadfrom sentence_transformers import SentenceTransformer
DEFAULT_MODEL = "BAAI/bge-small-en-v1.5"BATCH = 64
def main() -> None: corpus, out_dir = Path(sys.argv[1]), Path(sys.argv[2]) bench = "--bench" in sys.argv model_id = sys.argv[sys.argv.index("--model") + 1] if "--model" in sys.argv else DEFAULT_MODEL out_dir.mkdir(parents=True, exist_ok=True) texts = [json.loads(line)["text"] for line in corpus.read_text().splitlines()]
path = snapshot_download(model_id, allow_patterns=["*.json", "*.safetensors", "*.txt", "1_Pooling/*"]) (out_dir / "model_path.txt").write_text(path)
model = SentenceTransformer(model_id, device="cpu") tok = model.tokenizer max_len = model.max_seq_length
with (out_dir / "ids.jsonl").open("w") as f: for t in texts: ids = tok(t, truncation=True, max_length=max_len)["input_ids"] f.write(json.dumps(ids) + "\n")
emb = model.encode(texts, batch_size=BATCH, convert_to_numpy=True) emb.astype("<f4").tofile(out_dir / "embeddings.f32")
meta = { "model": model_id, "count": len(texts), "dim": int(emb.shape[1]), "max_seq_length": max_len, "modules": [type(m).__name__ for m in model], "torch": torch.__version__, }
if bench: timings = {} for threads in (1, torch.get_num_threads()): torch.set_num_threads(threads) model.encode(texts[:256], batch_size=BATCH) t0 = time.perf_counter() model.encode(texts, batch_size=BATCH) dt = time.perf_counter() - t0 timings[f"threads={threads}"] = {"seconds": dt, "posts_per_sec": len(texts) / dt} meta["bench"] = timings
(out_dir / "meta.json").write_text(json.dumps(meta, indent=2)) print(json.dumps(meta, indent=2))
if __name__ == "__main__": main()