- 引入 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,读写关闭时与原生模型逐位相同
50 lines
1.9 KiB
Python
50 lines
1.9 KiB
Python
"""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])
|