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:
@@ -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])
|
||||
Reference in New Issue
Block a user