#!/usr/bin/env python3 """Measure held-out triplet accuracy and vector cost at every Matryoshka dimension.""" from __future__ import annotations import argparse import json from pathlib import Path import numpy as np import torch from datasets import load_dataset from sentence_transformers import SentenceTransformer def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--models", nargs="+", required=True) parser.add_argument("--labels", nargs="+", required=True) parser.add_argument("--data", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--dims", nargs="+", type=int, default=[768, 512, 384, 256, 128, 64]) parser.add_argument("--batch-size", type=int, default=32) return parser.parse_args() def normalized_prefix(embeddings: np.ndarray, dim: int) -> np.ndarray: truncated = embeddings[:, :dim].astype("float32", copy=False) norms = np.linalg.norm(truncated, axis=1, keepdims=True) return truncated / np.maximum(norms, 1e-12) def evaluate_model( model_id: str, label: str, rows, dims: list[int], batch_size: int ) -> dict[str, object]: model = SentenceTransformer(model_id, model_kwargs={"dtype": torch.bfloat16}) model.max_seq_length = 256 queries = [f"query: {text}" for text in rows["query"]] positives = [f"passage: {text}" for text in rows["positive"]] negatives = [f"passage: {text}" for text in rows["negative"]] encoded = [ model.encode( texts, batch_size=batch_size, convert_to_numpy=True, normalize_embeddings=False, show_progress_bar=True, ) for texts in (queries, positives, negatives) ] full_dim = encoded[0].shape[1] results = [] for dim in dims: if dim > full_dim: raise ValueError(f"Dimension {dim} exceeds model output dimension {full_dim}") query, positive, negative = [normalized_prefix(values, dim) for values in encoded] positive_scores = np.sum(query * positive, axis=1) negative_scores = np.sum(query * negative, axis=1) results.append( { "dimension": dim, "triplet_accuracy": float(np.mean(positive_scores > negative_scores)), "mean_positive_cosine": float(np.mean(positive_scores)), "mean_negative_cosine": float(np.mean(negative_scores)), "float32_bytes_per_vector": dim * 4, "relative_index_size_vs_768": dim / 768, } ) del model, encoded torch.cuda.empty_cache() return {"label": label, "model": model_id, "dimensions": results} def main() -> None: args = parse_args() if len(args.models) != len(args.labels): raise ValueError("--models and --labels must contain the same number of values") rows = load_dataset("parquet", data_files=str(args.data), split="train") payload = { "evaluation": "held_out_hard_negative_triplet_accuracy", "rows": len(rows), "models": [ evaluate_model(model, label, rows, args.dims, args.batch_size) for model, label in zip(args.models, args.labels, strict=True) ], } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8") print(json.dumps(payload, indent=2)) if __name__ == "__main__": main()