404abc22bf
深模块修正:本函数只用 5 个字段,却索要整个 SFTConfig——层 1 无痛,但诊断脚本 被迫伪造 output_dir(4 处 /tmp/diag、outputs/_unused),层 2 更因 DistillConfig 无 teacher_completions_path 而无法复用。改收 dataset_path/split/subset_size/seed/ teacher_completions_path 五个散装参数(接口终于比实现轻)。 - data.py: 签名改散装参数;移除 TYPE_CHECKING 的 SFTConfig 依赖 - train_sft / diag_loss_probe / diag_collator: 仍持 SFTConfig(喂 collator),改调用点 - diag_generate / generate_teacher_completions: 只为 load 而造 config,直接丢弃、 去掉伪造 output_dir,改传字面量 - 为 U5 层 2 训练脚本能直接 load_sft_dataset(distill_cfg 的字段) 铺路 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
126 lines
4.8 KiB
Python
126 lines
4.8 KiB
Python
"""诊断脚本:逐环检验 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())
|