#!/usr/bin/env python3 """Retrieval benchmark for the TextMachine book-memory-bank (session «Валидация банка памяти», docs/research/13-memory-bank-validation.md). Task: given a source chunk (query) → retrieve the relevant glossary/summary entries (docs) from the memory bank. Compares three embedding models mandated by the session (bge-m3 vs Qwen3-Embedding-0.6B vs LaBSE) against two non-neural baselines that model the real design (deterministic substring/alias key-match = the "hot path" of Р3; BM25 char-level = the lexical leg of a hybrid). Answers, empirically, on real zh/ja PD literary text + constructed wrong-sense traps: Q1 Can a fixed cosine threshold separate "relevant" from "garbage"? (ROC-AUC, precision-0.9 threshold, recall there, behaviour on a fully-unrelated query.) Q3 Dense vs deterministic-substring vs BM25 — who retrieves what, and the exact-term vs semantic/crosslingual split. Q4 Wrong-sense / homonym false-fire: 送灶 (festival vs literal), 修 (cultivate vs repair), 炎 (name vs flame), 气 (qi vs air) — does dense inject the wrong entry that exact-substring correctly rejects? Q2 Brute-force KNN latency vs N (model-agnostic, per embedding dim) → the migration trigger to a real vector layer. Runs on CPU by default (GPU on the stand is held by the polygon's resident model). Usage: eval/.venv/bin/python eval/retrieval_bench.py [--models bge-m3,qwen3-0.6b,labse] [--skip-scale] Output: stdout tables + eval/data/retrieval_bench/results.json """ from __future__ import annotations import argparse import json import re import time from pathlib import Path import numpy as np ROOT = Path(__file__).resolve().parent CORPUS = ROOT / "data" / "retrieval_bench" / "corpus.json" OUT = ROOT / "data" / "retrieval_bench" / "results.json" # HF ids + how queries must be prompted. Qwen3-Embedding uses a query instruction # (asymmetric); bge-m3 and LaBSE are symmetric (plain encode both sides). MODELS = { "bge-m3": {"hf": "BAAI/bge-m3", "query_prompt": None}, "qwen3-0.6b": {"hf": "Qwen/Qwen3-Embedding-0.6B", "query_prompt": "Instruct: Given a passage of a novel, retrieve the glossary term or summary it refers to\nQuery: "}, "labse": {"hf": "sentence-transformers/LaBSE", "query_prompt": None}, } GENUINE_KINDS = {"real", "semantic", "crosslingual"} # queries that have gold # ---------- corpus ---------- def load_corpus(): data = json.loads(CORPUS.read_text(encoding="utf-8")) entries = data["entries"] for e in entries: note = e.get("note", "") e["embed_text"] = f"{e['src']} — {e['dst']}." + (f" {note}" if note else "") return data["queries"], entries # ---------- metrics ---------- def rank_metrics(scores_row, entry_ids, gold, ks=(1, 3, 5)): order = np.argsort(-scores_row) ranked = [entry_ids[i] for i in order] goldset = set(gold) out = {} for k in ks: topk = set(ranked[:k]) out[f"R@{k}"] = len(topk & goldset) / len(goldset) # MRR (first relevant) mrr = 0.0 for r, eid in enumerate(ranked, 1): if eid in goldset: mrr = 1.0 / r break out["MRR"] = mrr # nDCG@5 binary dcg = 0.0 for r, eid in enumerate(ranked[:5], 1): if eid in goldset: dcg += 1.0 / np.log2(r + 1) idcg = sum(1.0 / np.log2(r + 1) for r in range(1, min(len(goldset), 5) + 1)) out["nDCG@5"] = dcg / idcg if idcg else 0.0 return out, ranked def threshold_analysis(sim, queries, entries): """Pool (score,label) over genuine queries × all entries; find the cosine threshold giving precision≥0.9 and the recall there; ROC-AUC.""" from sklearn.metrics import roc_auc_score eid = [e["id"] for e in entries] scores, labels = [], [] rel_scores, irr_scores = [], [] for qi, q in enumerate(queries): if q["kind"] not in GENUINE_KINDS: continue gold = set(q["gold"]) for ei, e in enumerate(entries): s = float(sim[qi, ei]) lab = 1 if e["id"] in gold else 0 scores.append(s); labels.append(lab) (rel_scores if lab else irr_scores).append(s) scores = np.array(scores); labels = np.array(labels) auc = float(roc_auc_score(labels, scores)) if labels.min() != labels.max() else float("nan") # sweep thresholds for precision≥0.9 with max recall best = {"t": None, "precision": 0.0, "recall": 0.0} for t in np.unique(scores): pred = scores >= t tp = int((pred & (labels == 1)).sum()) fp = int((pred & (labels == 0)).sum()) fn = int((~pred & (labels == 1)).sum()) prec = tp / (tp + fp) if (tp + fp) else 1.0 rec = tp / (tp + fn) if (tp + fn) else 0.0 if prec >= 0.9 and rec > best["recall"]: best = {"t": float(t), "precision": prec, "recall": rec} return { "roc_auc": auc, "rel_mean": float(np.mean(rel_scores)), "rel_min": float(np.min(rel_scores)), "irr_mean": float(np.mean(irr_scores)), "irr_p95": float(np.percentile(irr_scores, 95)), "irr_max": float(np.max(irr_scores)), "prec90_threshold": best["t"], "recall_at_prec90": best["recall"], } def trap_analysis(sim, queries, entries, t_star): """At the precision-0.9 threshold t*, do wrong-sense trap targets get injected? Also the fully-unrelated negative query's top-1 score (abstention).""" eidx = {e["id"]: i for i, e in enumerate(entries)} traps = [] for qi, q in enumerate(queries): if q["kind"] != "trap": continue for tgt in q["traps"]: s = float(sim[qi, eidx[tgt]]) traps.append({"q": q["id"], "target": tgt, "cos": round(s, 3), "injected_at_t*": (t_star is not None and s >= t_star)}) neg = None for qi, q in enumerate(queries): if q["kind"] == "negative": top = float(np.max(sim[qi])) neg = {"q": q["id"], "top1_cos": round(top, 3), "abstains_at_t*": (t_star is None or top < t_star)} inj = sum(1 for t in traps if t["injected_at_t*"]) return {"traps": traps, "n_injected_at_t*": inj, "negative": neg} # ---------- neural ---------- def run_embed_model(name, spec, queries, entries): from sentence_transformers import SentenceTransformer t0 = time.time() model = SentenceTransformer(spec["hf"], device="cpu") load_s = time.time() - t0 doc_texts = [e["embed_text"] for e in entries] q_texts = [q["text"] for q in queries] enc = dict(normalize_embeddings=True, batch_size=16, show_progress_bar=False, convert_to_numpy=True) doc_emb = model.encode(doc_texts, **enc) t1 = time.time() if spec["query_prompt"]: q_emb = model.encode(q_texts, prompt=spec["query_prompt"], **enc) else: q_emb = model.encode(q_texts, **enc) enc_s = time.time() - t1 dim = int(doc_emb.shape[1]) sim = q_emb @ doc_emb.T # cosine (normalized) entry_ids = [e["id"] for e in entries] per_q, agg = [], {} for qi, q in enumerate(queries): if q["kind"] not in GENUINE_KINDS: continue m, ranked = rank_metrics(sim[qi], entry_ids, q["gold"]) per_q.append({"q": q["id"], "kind": q["kind"], **{k: round(v, 3) for k, v in m.items()}, "top5": ranked[:5]}) for k, v in m.items(): agg.setdefault(k, []).append(v) agg = {k: round(float(np.mean(v)), 3) for k, v in agg.items()} thr = threshold_analysis(sim, queries, entries) trp = trap_analysis(sim, queries, entries, thr["prec90_threshold"]) return {"model": name, "hf": spec["hf"], "dim": dim, "load_s": round(load_s, 1), "encode_s": round(enc_s, 2), "aggregate": agg, "threshold": thr, "trap": trp, "per_query": per_q} # ---------- baselines ---------- def tok(s): """char-level for CJK, lowercase word-level for latin/cyrillic.""" out = [] for m in re.finditer(r"[぀-ヿ㐀-鿿]|[A-Za-zА-Яа-яЁё0-9]+", s): out.append(m.group(0).lower()) return out def run_substring(queries, entries): """Deterministic exact substring/alias key-match = the Р3 hot path.""" def keys(e): return [e["src"]] + e.get("aliases", []) if e["type"] != "summary" else [] tp = fp = fn = 0 per_q, trap_fires = [], 0 for q in queries: hit = [e["id"] for e in entries if any(k and k in q["text"] for k in keys(e))] if q["kind"] in GENUINE_KINDS: gold = set(q["gold"]); h = set(hit) tp += len(h & gold); fp += len(h - gold); fn += len(gold - h) per_q.append({"q": q["id"], "kind": q["kind"], "hits": hit, "recall": round(len(h & gold) / len(gold), 3), "missed": sorted(gold - h)}) if q["kind"] == "trap": fired = [e for e in hit if e in q["traps"]] trap_fires += len(fired) per_q.append({"q": q["id"], "kind": "trap", "hits": hit, "false_fire": fired}) prec = tp / (tp + fp) if (tp + fp) else 1.0 rec = tp / (tp + fn) if (tp + fn) else 0.0 return {"method": "substring-exact", "precision": round(prec, 3), "recall": round(rec, 3), "trap_false_fires": trap_fires, "per_query": per_q} def run_bm25(queries, entries): from rank_bm25 import BM25Okapi corpus = [tok(e["embed_text"]) for e in entries] bm = BM25Okapi(corpus) entry_ids = [e["id"] for e in entries] agg, per_q, trap_top5 = {}, [], 0 for q in queries: sc = np.array(bm.get_scores(tok(q["text"]))) if q["kind"] in GENUINE_KINDS: m, ranked = rank_metrics(sc, entry_ids, q["gold"]) per_q.append({"q": q["id"], "kind": q["kind"], **{k: round(v, 3) for k, v in m.items()}}) for k, v in m.items(): agg.setdefault(k, []).append(v) if q["kind"] == "trap": order = np.argsort(-sc) top5 = {entry_ids[i] for i in order[:5]} if set(q["traps"]) & top5: trap_top5 += 1 agg = {k: round(float(np.mean(v)), 3) for k, v in agg.items()} return {"method": "bm25-char", "aggregate": agg, "trap_in_top5": trap_top5, "per_query": per_q} # ---------- scale ---------- def scale_bench(dims=(768, 1024), Ns=(1000, 10000, 50000, 100000, 200000), reps=20): """Model-agnostic brute-force cosine KNN latency vs N (single query, top-10).""" rng = np.random.default_rng(42) res = {} for d in dims: row = {} for N in Ns: mat = rng.standard_normal((N, d)).astype(np.float32) mat /= np.linalg.norm(mat, axis=1, keepdims=True) q = rng.standard_normal(d).astype(np.float32); q /= np.linalg.norm(q) ts = [] for _ in range(reps): t = time.perf_counter() s = mat @ q np.argpartition(-s, 10)[:10] ts.append((time.perf_counter() - t) * 1000) row[N] = round(float(np.median(ts)), 2) res[d] = row return res # ---------- main ---------- def main(): ap = argparse.ArgumentParser() ap.add_argument("--models", default="bge-m3,qwen3-0.6b,labse") ap.add_argument("--skip-scale", action="store_true") args = ap.parse_args() queries, entries = load_corpus() n_gold = sum(1 for q in queries if q["kind"] in GENUINE_KINDS) print(f"corpus: {len(entries)} entries, {len(queries)} queries " f"({n_gold} with gold, {sum(1 for q in queries if q['kind']=='trap')} traps)\n") results = {"embedding": [], "baselines": {}, "scale": None} print("=== baselines ===") sub = run_substring(queries, entries) results["baselines"]["substring"] = sub print(f"substring-exact: precision={sub['precision']} recall={sub['recall']} " f"trap_false_fires={sub['trap_false_fires']}") bm = run_bm25(queries, entries) results["baselines"]["bm25"] = bm print(f"bm25-char: {bm['aggregate']} trap_in_top5={bm['trap_in_top5']}\n") print("=== embedding models (CPU) ===") for name in [m.strip() for m in args.models.split(",") if m.strip()]: spec = MODELS[name] try: r = run_embed_model(name, spec, queries, entries) except Exception as e: print(f"{name}: FAILED — {type(e).__name__}: {e}") results["embedding"].append({"model": name, "error": f"{type(e).__name__}: {e}"}) continue results["embedding"].append(r) t = r["threshold"] print(f"\n[{name}] dim={r['dim']} load={r['load_s']}s encode={r['encode_s']}s") print(f" retrieval: {r['aggregate']}") print(f" threshold: ROC-AUC={t['roc_auc']:.3f} rel_mean={t['rel_mean']:.3f} " f"rel_min={t['rel_min']:.3f} irr_mean={t['irr_mean']:.3f} irr_p95={t['irr_p95']:.3f}") print(f" prec≥0.9 @ cos≥{t['prec90_threshold']} → recall={t['recall_at_prec90']:.3f}") tr = r["trap"] print(f" wrong-sense traps injected at t*: {tr['n_injected_at_t*']}/{len(tr['traps'])} " f"| {[(x['q'], x['target'], x['cos']) for x in tr['traps']]}") if tr["negative"]: print(f" unrelated-query top1 cos: {tr['negative']['top1_cos']} " f"(abstains at t*: {tr['negative']['abstains_at_t*']})") if not args.skip_scale: print("\n=== brute-force KNN latency (ms, single query, top-10, median) ===") sc = scale_bench() results["scale"] = sc for d, row in sc.items(): print(f" dim={d}: " + " ".join(f"N={N}:{ms}ms" for N, ms in row.items())) OUT.write_text(json.dumps(results, ensure_ascii=False, indent=2), encoding="utf-8") print(f"\nsaved → {OUT}") if __name__ == "__main__": main()