Natural Memory NM2.1: 记忆路由器分叉、数据集缺陷修复与全轴评测证据

- 引入 MemoryRouterXL 与 v5/v6 流式多线程训练/编码管线
- 修复 prepare_memory_router_dataset 候选池重建缺陷(mega 家族 3568x 加速,输出逐字节相同)
- 修复 v5 被破坏的拒答与多跳标签(train 未知样本 319 -> 16319,multi_hop 平均正例 1.00 -> 2.00)
- 同存储预算下 V2-128 v6 逐轴 22/22 通过:Top-1 41.12% -> 94.62%,未知拒答 0.00% -> 100.00%
- 记录三条被实测推翻的显然优化(logits_to_keep=1 反而慢 55%、XL 容量未带来收益)
- 记忆手术跨架构可移植性 14/14,读写关闭时与原生模型逐位相同
This commit is contained in:
WpyQwq
2026-09-19 11:11:31 +08:00
commit 643e22ecb9
484 changed files with 306821 additions and 0 deletions
+245
View File
@@ -0,0 +1,245 @@
"""Convert the LoCoMo benchmark into this harness' corpus format.
LoCoMo (snap-research/locomo) is a real long-term conversational-memory
benchmark: 10 multi-session human dialogues, 5,882 turns, 1,986 question/answer
pairs over five categories. Nothing here is generated by this project -- the
questions, answers and dialogue are the published dataset.
Mapping onto the harness:
* **query** <- the benchmark question, verbatim
* **candidates** <- real dialogue turns, formatted ``<speaker>: <text>``
* **positives** <- the turns the benchmark cites in ``evidence`` (dia_id refs)
* **acceptable** <- the benchmark's own answer string
* **category** <- 1 multi_hop, 2 temporal, 3 open_domain, 4 single_hop,
5 adversarial (no answer exists -> the system must refuse)
Two deliberate choices that must be stated with any number produced from this
corpus, because they bound what the number means:
1. **Bounded candidate pool.** Each question gets a pool of at most
``--pool`` turns: the cited evidence, other turns from the same session, and
turns sampled from other sessions. Running the full 5,882-turn history per
question is a long-context retrieval benchmark of a different shape; this
pool tests whether the memory layer reads and synthesises the right evidence
while real distractors compete with it.
2. **Category 5 is scored as "must refuse".** Those questions are written so
that no answer exists in the dialogue; the dataset supplies a plausible
``adversarial_answer`` which a system must *not* assert.
Answers that a model phrases differently from the benchmark string are a known
limitation of anchor matching, so ``make_locomo_corpus`` also records answer
content tokens for a secondary, paraphrase-tolerant metric.
"""
from __future__ import annotations
import argparse
import json
import random
import re
import sys
from collections import Counter, defaultdict
from pathlib import Path
CATEGORY_NAMES = {
1: "multi_hop",
2: "temporal",
3: "open_domain",
4: "single_hop",
5: "adversarial",
}
_STOP = {
"the", "a", "an", "of", "in", "on", "at", "to", "and", "or", "for", "with", "was", "were",
"is", "are", "did", "does", "do", "her", "his", "their", "she", "he", "they", "it", "that",
"this", "what", "when", "where", "who", "how", "why", "which", "caroline", "melanie",
}
def answer_tokens(answer: str) -> list[str]:
"""Content tokens of the expected answer, for a paraphrase-tolerant check."""
words = re.findall(r"[a-z0-9']+", str(answer).lower())
return [w for w in words if w not in _STOP and len(w) > 1]
def turn_index(conversation: dict) -> dict[str, dict]:
"""dia_id -> turn, across every session of one conversation."""
index: dict[str, dict] = {}
for key, value in conversation.items():
if not key.startswith("session_") or key.endswith("date_time"):
continue
if not isinstance(value, list):
continue
for turn in value:
dia_id = turn.get("dia_id")
if dia_id:
index[dia_id] = {"dia_id": dia_id, "session": key, **turn}
return index
def turn_text(turn: dict) -> str:
text = str(turn.get("text") or "").strip()
if not text:
# some turns are images; the caption is the only textual evidence
text = str(turn.get("blip_caption") or turn.get("caption") or "").strip()
speaker = str(turn.get("speaker") or "").strip()
return f"{speaker}: {text}" if speaker else text
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", type=Path,
default=Path(r"data/net_locomo/locomo10.json"))
parser.add_argument("--output-dir", type=Path, default=Path("data/net_locomo"))
parser.add_argument("--per-category", type=int, default=40,
help="cap questions per category (0 = all 1,986)")
parser.add_argument("--pool", type=int, default=16,
help="candidate turns per question, evidence always included")
parser.add_argument("--seed", type=int, default=20260915)
args = parser.parse_args()
conversations = json.loads(args.source.read_text(encoding="utf-8"))
rng = random.Random(args.seed)
# Gather every question first so per-category sampling is stratified across
# conversations rather than exhausted from the first one.
by_category: dict[int, list[tuple[int, dict, dict, dict]]] = defaultdict(list)
for conv_index, item in enumerate(conversations):
conversation = item["conversation"]
index = turn_index(conversation)
for qa in item["qa"]:
category = int(qa.get("category") or 0)
if category not in CATEGORY_NAMES:
continue
by_category[category].append((conv_index, item, index, qa))
episodes = []
stats = Counter()
skipped = Counter()
for category in sorted(by_category):
bucket = by_category[category]
rng.shuffle(bucket)
chosen = bucket if args.per_category <= 0 else bucket[:args.per_category]
for conv_index, item, index, qa in chosen:
conversation = item["conversation"]
evidence_ids = [str(e) for e in (qa.get("evidence") or [])]
evidence_turns = [index[e] for e in evidence_ids if e in index]
if category != 5 and not evidence_turns:
skipped["no_evidence_in_index"] += 1
continue
pool: list[dict] = list(evidence_turns)
seen_ids = {t["dia_id"] for t in pool}
same_session = [t for t in index.values()
if evidence_turns and t["session"] == evidence_turns[0]["session"]
and t["dia_id"] not in seen_ids]
rng.shuffle(same_session)
for turn in same_session:
if len(pool) >= args.pool:
break
if turn["dia_id"] not in seen_ids:
pool.append(turn)
seen_ids.add(turn["dia_id"])
others = [t for t in index.values() if t["dia_id"] not in seen_ids]
rng.shuffle(others)
for turn in others:
if len(pool) >= args.pool:
break
pool.append(turn)
seen_ids.add(turn["dia_id"])
# dedupe identical surface text: duplicate records would let a
# positive be "found" without reading the evidence
texts, positives = [], []
evidence_text = {t["dia_id"] for t in evidence_turns}
for turn in pool:
text = turn_text(turn)
if not text:
continue
if text not in texts:
texts.append(text)
# Adversarial questions DO carry evidence, but no answer: the cited
# turn discusses the topic without stating the answer, so it is the
# trap, not the answer. It stays in the pool (as a hard distractor)
# and is deliberately NOT marked positive -- otherwise the harness
# reads `bool(positives)` as "answerable" and scores a correct
# refusal as a wrong answer.
if category != 5 and turn["dia_id"] in evidence_text:
if text not in positives:
positives.append(text)
if category != 5 and not positives:
skipped["positives_deduped_away"] += 1
continue
if len(texts) < 4:
skipped["pool_too_small"] += 1
continue
if category == 5:
acceptable = []
answerable = False
else:
answer = str(qa.get("answer") or "").strip()
if not answer:
skipped["empty_answer"] += 1
continue
acceptable = [answer]
answerable = True
episodes.append({
"query": str(qa["question"]).strip(),
"candidates": [{"text": t} for t in texts],
# empty for adversarial on purpose: the harness infers
# "answerable" from bool(positives)
"positive_indices": [texts.index(p) for p in positives],
"metadata": {
"category": CATEGORY_NAMES[category],
"acceptable": acceptable,
"answerable": answerable,
"answer_tokens": answer_tokens(qa.get("answer") or ""),
"adversarial_answer": qa.get("adversarial_answer"),
"evidence": evidence_ids,
"sample_id": item.get("sample_id"),
"locomo_category": category,
"source": "locomo10",
},
})
stats[CATEGORY_NAMES[category]] += 1
out = args.output_dir / "eval.jsonl"
out.parent.mkdir(parents=True, exist_ok=True)
with out.open("w", encoding="utf-8") as handle:
for episode in episodes:
handle.write(json.dumps(episode, ensure_ascii=False) + "\n")
manifest = {
"source": str(args.source),
"dataset": "LoCoMo (snap-research/locomo) locomo10.json",
"provenance": "real published benchmark; dialogue, questions and answers are the dataset's own",
"conversations": len(conversations),
"written": len(episodes),
"pool_cap": args.pool,
"per_category_cap": args.per_category,
"seed": args.seed,
"by_category": dict(stats),
"skipped": dict(skipped),
"caveats": [
"candidate pool is bounded per question, so this is not a full 5,882-turn haystack run",
"category adversarial is scored as 'must refuse': the dataset provides a plausible wrong answer",
"answers phrased differently from the benchmark string fail exact-anchor matching; "
"metadata.answer_tokens supports a paraphrase-tolerant secondary metric",
],
}
(args.output_dir / "manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
print(json.dumps({k: manifest[k] for k in
("written", "by_category", "pool_cap", "per_category_cap")}, ensure_ascii=False))
if skipped:
print("skipped:", dict(skipped))
print(f"wrote {out}")
return 0
if __name__ == "__main__":
raise SystemExit(main())