| """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 |
| 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() |
|
|