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
+49
View File
@@ -0,0 +1,49 @@
"""Validate the refusal detector in both directions.
Over-firing would inflate the known-question false-refusal rate; under-firing
would inflate the unknown-refusal rate. Both directions are printed with the
verbatim reply so the judgement can be checked by eye.
"""
from __future__ import annotations
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from V2_dpskw.eval_scoring import is_refusal, squash
from V2_dpskw.rescore_e2e import load_corpus, load_run
def main(corpus, run):
cases = load_corpus(Path(corpus), 25)
rows = load_run(Path(run))
unknown_refused, unknown_asserted = [], []
ans_refused, ans_refused_but_answered = [], []
for case, row in zip(cases, rows):
reply = row.get("reply", "")
refused = is_refusal(reply)
hit = [v for v in case["acceptable"] if v and squash(v) in squash(reply)]
if not case["answerable"]:
(unknown_refused if refused else unknown_asserted).append((case, row))
elif refused:
(ans_refused_but_answered if hit else ans_refused).append((case, row))
print(f"UNANSWERABLE cases where the detector says REFUSED ({len(unknown_refused)})")
for case, row in unknown_refused:
print(f" Q {case['query']}\n R {row.get('reply','')[:110]}")
print(f"\nUNANSWERABLE cases where the detector says ASSERTED ({len(unknown_asserted)})")
for case, row in unknown_asserted:
print(f" Q {case['query']}\n R {row.get('reply','')[:110]}")
print(f"\nANSWERABLE cases flagged as refusals ({len(ans_refused) + len(ans_refused_but_answered)})")
for tag, group in (("ALSO-MATCHED-ANCHOR (harmless)", ans_refused_but_answered), ("COUNTED AS FALSE REFUSAL", ans_refused)):
print(f" -- {tag}: {len(group)}")
for case, row in group:
print(f" Q {case['query']}\n R {row.get('reply','')[:110]}")
if __name__ == "__main__":
main(sys.argv[1], sys.argv[2])