"""诊断脚本:逐环检验 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.dataset_path, cfg.dataset_split, cfg.subset_size, cfg.seed, cfg.teacher_completions_path, ) 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())