326 lines
14 KiB
Python
326 lines
14 KiB
Python
#!/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()
|