From ff572bf4b9d420a1b7b027e0830ad2c549ceb8d1 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 10:05:46 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=821:=20collator=20=E5=AF=B9=E9=BD=90?= =?UTF-8?q?=E8=AF=8A=E6=96=AD=E8=84=9A=E6=9C=AC=E2=80=94=E2=80=94=E6=8E=92?= =?UTF-8?q?=E6=9F=A5=E8=BF=9C=E7=A8=8B=20loss=207.5=20=E4=B8=8E=20\50118?= =?UTF-8?q?=20=E8=A7=A3=E7=A0=81=E5=BC=82=E5=B8=B8=EF=BC=88=E9=80=90?= =?UTF-8?q?=E7=8E=AF=E6=A3=80=E9=AA=8C=E7=BC=93=E5=AD=98/=E6=A8=A1?= =?UTF-8?q?=E6=9D=BF/=E5=88=86=E8=AF=8D/=E5=88=87=E7=89=87=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 --- scripts/diag_collator.py | 119 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 119 insertions(+) create mode 100644 scripts/diag_collator.py diff --git a/scripts/diag_collator.py b/scripts/diag_collator.py new file mode 100644 index 0000000..e5565f2 --- /dev/null +++ b/scripts/diag_collator.py @@ -0,0 +1,119 @@ +"""诊断脚本:逐环检验 collator 对齐链(层 1 关账前的疑点排查)。 + +背景:远程 sanity 中初始 loss ~7.5(预期 ~2-3),且首样本自检的 completion 段 +出现 `$k \\50118$`(teacher 原文是 `$k \\leq 2018$`)。本脚本把 +"缓存文本 → 模板渲染 → 分词 → 边界切片 → 解码"逐环单测,定位腐坏点。 + +远程运行(CPU 即可,tokenizer 用已有 HF 缓存): + python -u scripts/diag_collator.py +""" + +import hashlib +import sys + +from transformers import AutoTokenizer + +from ars_opd.configs import SFTConfig +from ars_opd.data import IGNORE_INDEX, SFTCollator, load_sft_dataset + +CACHE = "data/teacher_completions_dapo1k_minimax-m3.jsonl" +DATASET = "data/dapo-math-17k-unique.parquet" +MODEL = "Qwen/Qwen3-0.6B" + + +def check(name: str, ok: bool, detail: str = "") -> bool: + print(f"[{'通过' if ok else '失败'}] {name}" + (f" —— {detail}" if detail else ""), flush=True) + return ok + + +def first_diff(a: str, b: str) -> int: + n = min(len(a), len(b)) + for i in range(n): + if a[i] != b[i]: + return i + return -1 if len(a) == len(b) else n + + +def main() -> None: + # 环 0:缓存文件指纹(与本地对比,排除 scp 传坏/版本不一致) + digest = hashlib.sha256(open(CACHE, "rb").read()).hexdigest() + print(f"缓存文件 sha256: {digest[:16]}… (与本地对比)", flush=True) + + cfg = SFTConfig( + dataset_path=DATASET, + output_dir="/tmp/diag", + subset_size=5, + seed=42, + teacher_completions_path=CACHE, + ) + ds = load_sft_dataset(cfg) + msgs = ds[0]["messages"] + completion_text = msgs[-1]["content"] + + # 环 1:本机数据管线出来的 teacher 文本是否干净 + check( + "环1 缓存→数据集文本干净", + "\\leq 2018" in completion_text and "\\50118" not in completion_text, + f"开头: {completion_text[:60]!r}", + ) + + tok = AutoTokenizer.from_pretrained(MODEL) + fp = tok.apply_chat_template( + msgs[:-1], tokenize=False, add_generation_prompt=True, enable_thinking=False + ) + ff = tok.apply_chat_template( + msgs, tokenize=False, add_generation_prompt=False, enable_thinking=False + ) + + # 环 2:完整渲染必须以 prompt 渲染为前缀(collator 边界法的前提假设!) + prefix_ok = ff.startswith(fp) + check("环2 完整渲染以 prompt 渲染为前缀", prefix_ok) + if not prefix_ok: + i = first_diff(ff, fp) + print(f" 首个分歧在第 {i} 字符:\n" + f" prompt 渲染: …{fp[max(0, i - 60) : i + 60]!r}\n" + f" 完整渲染: …{ff[max(0, i - 60) : i + 60]!r}", flush=True) + + # 环 3:完整渲染中 teacher 文本是否原样存在(模板会不会改写 content) + check( + "环3 完整渲染保留 teacher 原文", + "\\leq 2018" in ff and "\\50118" not in ff, + "" if "\\leq 2018" in ff else "模板改写了 assistant content!", + ) + + # 环 4:token 级前缀(坑一:拼接稳定性) + full_ids = tok(ff, add_special_tokens=False)["input_ids"] + fp_ids = tok(fp, add_special_tokens=False)["input_ids"] + tok_prefix_ok = full_ids[: len(fp_ids)] == fp_ids + check("环4 token 级前缀一致(无跨界合并)", tok_prefix_ok) + if not tok_prefix_ok: + i = next(k for k in range(len(fp_ids)) if full_ids[k] != fp_ids[k]) + lo, hi = max(0, i - 3), i + 4 + print(f" 首个分歧在 token {i}/{len(fp_ids)}:\n" + f" prompt 侧: {[tok.decode([t]) for t in fp_ids[lo:hi]]}\n" + f" 完整侧: {[tok.decode([t]) for t in full_ids[lo:hi]]}", flush=True) + + # 环 5:collator 全流程后,completion 解码应等于完整渲染去掉 prompt 的尾段前缀 + collator = SFTCollator( + tok, + max_length=cfg.max_length, + max_prompt_length=cfg.max_prompt_length, + enable_thinking=False, + ) + batch = collator([ds[0]]) + ids, labels = batch["input_ids"][0], batch["labels"][0] + comp_decoded = tok.decode(ids[labels != IGNORE_INDEX], skip_special_tokens=False) + expected_tail = ff[len(fp) :] if prefix_ok else "(环2 已失败,无期望值)" + tail_ok = prefix_ok and expected_tail.startswith(comp_decoded[:200]) + check("环5 completion 解码 == 渲染尾段", tail_ok) + if prefix_ok and not tail_ok: + i = first_diff(comp_decoded, expected_tail) + print(f" 首个分歧在第 {i} 字符:\n" + f" 解码: …{comp_decoded[max(0, i - 50) : i + 50]!r}\n" + f" 期望: …{expected_tail[max(0, i - 50) : i + 50]!r}", flush=True) + + print("\n诊断完成。把全部输出贴回对话。", flush=True) + + +if __name__ == "__main__": + sys.exit(main())