"""BEIR-style retrieval eval, identical protocol for POG and zembed-1. Usage: eval_beir.py --model pog --tasks nfcorpus scifact fiqa --dims 256 2560 eval_beir.py --model zembed --tasks nfcorpus scifact fiqa --dims 320 2560 """ import argparse import json import math import os import sys import numpy as np import torch from datasets import load_dataset sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "adapter")) GGUF = "/home/anon/models/gemma-4-E4B-it-qat-GGUF/gemma-4-E4B-qat-L38-trim.gguf" EXTRACTOR = "/home/anon/pog/extractor/pog-extract" def load_beir(task): corpus = load_dataset(f"BeIR/{task}", "corpus", split="corpus") queries = load_dataset(f"BeIR/{task}", "queries", split="queries") qrels = load_dataset(f"BeIR/{task}-qrels", split="test") docs = {str(r["_id"]): ((r.get("title") or "") + " " + (r.get("text") or "")).strip() for r in corpus} qmap = {str(r["_id"]): r["text"] for r in queries} rel = {} for r in qrels: rel.setdefault(str(r["query-id"]), {})[str(r["corpus-id"])] = int(r["score"]) rel = {q: d for q, d in rel.items() if q in qmap and any(s > 0 for s in d.values())} return docs, qmap, rel def ndcg_at_k(ranked_ids, rels, k=10): dcg = sum(rels.get(d, 0) / math.log2(i + 2) for i, d in enumerate(ranked_ids[:k])) ideal = sorted(rels.values(), reverse=True)[:k] idcg = sum(s / math.log2(i + 2) for i, s in enumerate(ideal)) return dcg / idcg if idcg > 0 else 0.0 def recall_at_k(ranked_ids, rels, k=10): pos = {d for d, s in rels.items() if s > 0} return len(pos & set(ranked_ids[:k])) / len(pos) def evaluate(q_emb, d_emb, qids, dids, rel, block=512): """Cosine retrieval (embeddings already L2-normalized).""" d_t = torch.from_numpy(d_emb).cuda() res_ndcg, res_rec = [], [] for i in range(0, len(q_emb), block): sims = torch.from_numpy(q_emb[i:i + block]).cuda() @ d_t.T top = sims.topk(min(100, len(dids)), dim=-1).indices.cpu().numpy() for j, row in enumerate(top): qid = qids[i + j] ranked = [dids[t] for t in row] res_ndcg.append(ndcg_at_k(ranked, rel[qid])) res_rec.append(recall_at_k(ranked, rel[qid])) return float(np.mean(res_ndcg)), float(np.mean(res_rec)) def get_encoder(name, adapter_path): """Returns (encode_queries, encode_docs) producing FULL-dim L2-normalized embeddings.""" if name == "pog": from pog_embed import POGEmbedder emb = POGEmbedder(GGUF, adapter_path, EXTRACTOR) return (lambda t: emb.encode_query(t, max_tokens=512), lambda t: emb.encode_document(t, max_tokens=512)) if name == "zembed": from sentence_transformers import SentenceTransformer m = SentenceTransformer("/home/anon/models/zembed-1", trust_remote_code=True, model_kwargs={"torch_dtype": torch.bfloat16}) m.max_seq_length = 2048 # covers p99 of these corpora; avoids 32k-ctx OOM return (lambda t: m.encode_query(t, batch_size=4, show_progress_bar=False, normalize_embeddings=True), lambda t: m.encode_document(t, batch_size=4, show_progress_bar=False, normalize_embeddings=True)) raise ValueError(name) def truncate_renorm(emb, dim): if dim is None or dim >= emb.shape[1]: return emb e = emb[:, :dim].copy() return e / np.linalg.norm(e, axis=1, keepdims=True) def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", required=True, choices=["pog", "zembed"]) ap.add_argument("--adapter", default="/home/anon/pog/checkpoints/pog-v1") ap.add_argument("--label", default=None, help="results key prefix (default: --model)") ap.add_argument("--tasks", nargs="+", default=["nfcorpus", "scifact"]) ap.add_argument("--dims", nargs="+", type=int, default=[2560]) ap.add_argument("--save-emb", default=None, help="dir to cache full-dim embeddings per task") ap.add_argument("--out", default="/home/anon/pog/eval/results.json") args = ap.parse_args() label = args.label or args.model results = {} if os.path.exists(args.out): results = json.load(open(args.out)) enc_q, enc_d = get_encoder(args.model, args.adapter) for task in args.tasks: print(f"=== {task} ===", flush=True) docs, qmap, rel = load_beir(task) dids = list(docs.keys()) qids = list(rel.keys()) print(f"docs={len(dids)} queries={len(qids)}", flush=True) d_full = enc_d([docs[d] for d in dids]).astype(np.float32) q_full = enc_q([qmap[q] for q in qids]).astype(np.float32) if args.save_emb: os.makedirs(args.save_emb, exist_ok=True) np.savez(os.path.join(args.save_emb, f"{task}_{label}.npz"), q=q_full.astype(np.float16), d=d_full.astype(np.float16), qids=np.array(qids), dids=np.array(dids)) for dim in args.dims: ndcg, rec = evaluate(truncate_renorm(q_full, dim), truncate_renorm(d_full, dim), qids, dids, rel) key = f"{label}@{dim}" results.setdefault(task, {})[key] = {"ndcg@10": round(ndcg, 4), "recall@10": round(rec, 4)} print(f"{task} {key}: NDCG@10={ndcg:.4f} R@10={rec:.4f}", flush=True) json.dump(results, open(args.out, "w"), indent=2) print(json.dumps(results, indent=2)) if __name__ == "__main__": main()