POG-E4B-v1 / eval /eval_beir.py
LordAce9's picture
v2 program: v2.1 multi-dataset variant, GGUF layer-trim early-exit (bit-exact), server-mode extractor, 5-task BEIR sweep, binary/int8 compression study, full ablation writeup
d27cbf8 verified
Raw
History Blame Contribute Delete
5.64 kB
"""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()