diff --git a/verify.py b/verify.py index 9aee3a7..38224de 100644 --- a/verify.py +++ b/verify.py @@ -10,6 +10,14 @@ Usage: HF_TOKEN=hf_... python3 verify.py events.jsonl # from client -out Requires HF_TOKEN env var (https://huggingface.co/google/embeddinggemma-300m). + +Precision strategy: Divepool runs EmbeddingGemma on M1 Mac Studios via MLX. +MLX stores weights in bf16 but computes all GEMM in float32 — Metal's +simdgroup_matrix uses the FP32 ALU pipeline and MLX kernels hardcode +AccumType=float. This script matches that: loads weights in bf16 for identical +rounding, then upcasts to float32 for compute. Expected cosine similarity +~0.995+. The remaining gap is from different attention implementations +(HuggingFace vs MLX Metal) and server-side batch padding context. """ import json @@ -144,16 +152,18 @@ class EmbeddingModel: def __init__(self, token: str): self.tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=token) - # bfloat16 to match Divepool's MLX server (mlx-community/embeddinggemma-300m-bf16) + # Load in bf16 for correct weight rounding, then upcast to float32 for compute. + # MLX on M1 does the same: bf16 storage, float32 GEMM (AccumType=float in Metal kernels). self.model = AutoModel.from_pretrained( MODEL_ID, token=token, trust_remote_code=True, dtype=torch.bfloat16, - ) + ).float() self.model.eval() d0_path = _download_dense_layer(token, "2_Dense/model.safetensors") d1_path = _download_dense_layer(token, "3_Dense/model.safetensors") - self.dense0 = load_file(d0_path)["linear.weight"].to(torch.bfloat16) - self.dense1 = load_file(d1_path)["linear.weight"].to(torch.bfloat16) + # bf16→float32 round-trip matches server weight precision + self.dense0 = load_file(d0_path)["linear.weight"].to(torch.bfloat16).float() + self.dense1 = load_file(d1_path)["linear.weight"].to(torch.bfloat16).float() def encode(self, text: str, prefix: str) -> np.ndarray: inputs = self.tokenizer( @@ -166,15 +176,15 @@ class EmbeddingModel: with torch.no_grad(): out = self.model(**inputs) - # Mean pooling in float32 (matches mlx_embeddings mean_pooling which casts mask to float32) + # Mean pooling (all float32 — model and weights already upcast in __init__) h = out.last_hidden_state mask = inputs["attention_mask"].unsqueeze(-1).float() - emb = (h.float() * mask).sum(dim=1) / mask.sum(dim=1) + emb = (h * mask).sum(dim=1) / mask.sum(dim=1) emb = emb.squeeze() # Dense projection (identity activation, no bias) - emb = emb @ self.dense0.float().T - emb = emb @ self.dense1.float().T + emb = emb @ self.dense0.T + emb = emb @ self.dense1.T # L2 normalize emb = emb / emb.norm()