textmachine/eval/retrieval_bench.py

326 lines
14 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/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()