Compare commits

...

43 Commits

Author SHA1 Message Date
iomgaa 630a9c4636 层3 E2: estimator.py——式(4) 几何均值先验 chunk_prior + 式(5) 贝叶斯目标 bayesian_target,17 单测
对应 docs/04 §4 E2。detach 双防线(chunk_prior 内为本质防线、
bayesian_target 末尾为防御性第二道,分别对应参考 trainer:2196/2205);
log 域抗下溢;clamp 下限守定理 4.1(b);k_sem∈[0,N] 口径校验;
方差收缩 toy 模拟对拍 validate_chunk_mc_estimator.py 精神
(MSE_bayes < MSE_freq,含劣先验 ±0.1 稳健性)。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-22 08:01:33 -04:00
iomgaa e06e5ed7a2 层3 E1: similarity.py——语义相似度 φ(rouge1 多重集/edit_distance 词级)+ 式(3) k_sem 聚合,21 单测
对应 docs/04 §4 E1。三处对参考实现的替代:edit 吃 str 内部按词切(不再比
token id,回归 tokenizer 无关)、rouge1 集合改多重集(ROUGE-1 标准定义,
数学文本重复词多)、去掉 1e-8 分母平滑(全同串精确得 1)。φ 默认
edit_distance 对齐论文 §5.1。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-22 07:58:05 -04:00
iomgaa e7cf27cecc docs: roadmap 存档点更新——docs/04 已精读,下一步开写 E1(压缩前存档)
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 09:48:40 -04:00
iomgaa 142aeb8ab1 docs: 层 3 章节文档 docs/04(语义相似度 φ + MC 估计器,式3/4/5)
论文精读(§3.2.1-3.2.2 式3/4/5)+ 参考实现解剖(带行号)+ 保留/替代/删除 + E1-E3
重构任务 + 验证方式。核心:logit-free 支点(比文本非比 logits)、Dirichlet 贝叶斯
π̂ 反塌缩、detach 命门(层 2 §4.1 姊妹篇)。解剖抓出参考实现三处不一致(φ 比文本
vs 比 token id、rouge1 集合 vs 多重集)作重构清理依据。roadmap 存档点/索引同步。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 09:40:58 -04:00
iomgaa fac8e0dcc5 docs: 层 2 关账——§6.1 远程实证勘误"预期见毛刺" + roadmap 存档点 + 接口回看
- docs/03 §6.1:两次远程跑(sanity/noclip)均平稳,勘误当初"预期毛刺"的错误
  预测;三条实证结论(爆炸真机制但本区间高度阻尼、grad_norm 是裁剪前值、层 5
  真正杀手锏是 logit-free)
- roadmap 存档点:层 2  关账,判据全过、关键发现、接口回看(全判"深",
  三条不返工小注);章节索引 03 标已关账;下一步层 3

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 09:30:45 -04:00
iomgaa 8b362eae09 层2: max_grad_norm 提进 DistillConfig(显式化静默稳定器)+ noclip 对照模式
首冒烟发现:sanity 的 loss 平滑、无预期毛刺,因 HF 默认 max_grad_norm=1.0 把
反向 KL 的梯度爆炸(§4.1,实测 grad_norm 14→2 是裁剪前范数)默默压平了——正是
本项目要堵的"静默行为"。

- configs.py: DistillConfig 加 max_grad_norm=1.0(默认=原 HF 行为),docstring 讲清
  它是 §4.1 爆炸的隐形稳定器、日志 grad_norm 是裁剪前值;__post_init__ 校验 >0
- train_whitebox.py: FULL 显式写出、TrainingArguments 传入;build_config 加 noclip
  模式(max_grad_norm=1e9≈关裁剪 + lr 5× + 15 步)暴露原始爆炸供教学对照
- .sh: 用法加 noclip 模式说明

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 08:34:14 -04:00
iomgaa f1b6d1f668 修复: 禁用 HF Xet(hf-mirror 不代理 CAS,teacher 4B 下载 401)
远程冒烟实撞:hf-mirror 供小文件正常,但大权重走 Xet 协议会绕过镜像直连
cas-server.xethub.hf.co 并返 401 Unauthorized。加 HF_HUB_DISABLE_XET=1 退回
经典 HTTP/LFS 下载路径(镜像支持)。两个训练脚本都加,防层 1 换机重下也撞。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:58:16 -04:00
iomgaa 2f4780b69f docs: roadmap 存档点更新——层 2 代码关账(U1-U5 全绿),列远程冒烟三盯待办
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:24:25 -04:00
iomgaa b7f24d635e 层2/U5: white-box OPD 自包含训练脚本(train_whitebox.py + .sh)
对应 docs/03 §5 U5,对齐层 1 train_sft 骨架,换成层 2 装配:
- 双模型:student Qwen3-0.6B(fp32+bf16混训) + teacher Qwen3-4B(bf16 推理)
- prompt_only collator(无 teacher 缓存,现场 on-policy 生成)
- DistillTrainer 装配:teacher_model/teacher_tokenizer + beta/温度/生成参数
- FULL = §5 默认(beta=1、lr=1e-6、B=4×GA=4×4卡=全局64、max_new_tokens=1024)
- 首 prompt 自检(末尾须为生成引导符);sanity=50步冒烟
- .sh 前置清单强调白盒扛两份全词表 logits、OOM 阶梯、首跑下载 4B teacher

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:23:42 -04:00
iomgaa 404abc22bf 重构: load_sft_dataset 改吃散装参数(磨平接口回看记录的毛刺)
深模块修正:本函数只用 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>
2026-07-19 05:21:16 -04:00
iomgaa 0ca60ea93f 层2/U4: DistillTrainer 编排(on-policy 生成→teacher no_grad 前向→反向KL)
trainer.py(对应 docs/03 §5 U4):
- build_generated_batch: 纯函数,generate 输出重建 ids/attention/labels;
  "首个 eos 及之前有效"用 cumsum-self==0 实现,稳健对付 pad==eos / pad!=eos / 无eos
- DistillTrainer.compute_loss 四步:生成(no_grad,unwrap)→重建→双前向→移位+divergence
- 损失几何复用 compute_prompt_length(与 sft_loss 同款);梯度只经 student
- 构造时校验 teacher/student 同 tokenizer(白盒前提,比 DT:2876 更早)
- teacher eval+冻结、设备迁移推迟到 compute_loss;同层1退出梯度累积新式契约

test_distill.py(对应 docs/03 §2.5):
- build_generated_batch 6 例:eos居中屏蔽/无eos全监督/batch混长/pad==eos/立即eos/prompt段-100

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:16:57 -04:00
iomgaa e5a28e8e77 层2/U3: SFTCollator 放开 prompt-only 模式(供 on-policy 生成)
data.py(对应 docs/03 §5 U3):
- 加 prompt_only 开关:True 时输出 prompts/prompt_attention_mask(不产 labels,
  由 U4 生成后重建);False 时 SFT 双预算路径逐字不变
- max_length 改可选:prompt-only 无总预算;SFT 模式缺它构造即报错
- 兑现 T3 为 on-policy 生成预留的口子;生成用左 padding(右边界对齐)

test_data.py:
- 新增 prompt_only 模式:返回 prompt 张量/不报错、左 padding、截断、剥末轮 assistant
- 回归守卫:SFT 模式仍拒绝 prompt-only 行("SFT 路径行为不变")

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:02:24 -04:00
iomgaa de7f36828a 层2/U2: 纯逻辑 token_divergence(式(2) 全词表 KL/JSD)+ 梯度爆炸演示单测
trainer.py(对应 docs/03 §5 U2,与 sft_loss 同为纯张量损失函数):
- token_divergence: 全词表精确 KL(β=1 反向=式(2) / β=0 前向 / (0,1) JSD)
- 只吃两组已对齐 logits + labels 掩码,不做移位(复用 T4 几何,留给 U4)
- 删参考实现 top-k/尾桶/nan_to_num(本地全词表恒有限);按 β 分支省一份 probs
- per-token mean 与 sft_loss 同尺度;全掩码/beta 越界显式报错

test_divergence.py(对应 docs/03 §6.1):
- 对拍 PyTorch torch.distributions.kl_divergence(独立 oracle,非同式自证)
- 方向性: 反向罚越界(mode-seeking) / 前向罚漏覆盖(mode-covering)
- §4.1 梯度爆炸: teacher 概率 1e-1→1e-6 时 student 梯度范数单调暴涨(层5对照桩)

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 04:00:25 -04:00
iomgaa e42af5256f 层2/U1: 新增 DistillConfig(式(2) 白盒蒸馏参数)+ docs/03 表述打磨
configs.py(对应 docs/03 §5 U1):
- DistillConfig 自包含、不继承 SFTConfig;teacher_model 进 config、student 留脚本
- 三处刻意缺席: 无 teacher_completions_path/max_length/top_k(现场生成+全词表)
- 两温度分名: kl_temperature(散度 softmax)vs gen_temperature(on-policy 采样)
- 默认即式(2): beta=1 反向KL、温度1、纯采样 top_p=1;lr=1e-6(论文§5.1蒸馏)
- __post_init__ 8 分支构造即校验(gen_temperature>0 护 on-policy 语义)

docs/03:
- §2.3 β 三副面孔表: 记号统一 π、补 mode-seeking 对称、附全词表/稀疏双镜像说明
- §2.3 三条实现约定(温度/log域/batchmean)由一句话拆成可扫读表格
- §3 偏差清单上方补统领抉择原则

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 03:37:44 -04:00
iomgaa db392d9c81 存档点:压缩前补记层 2 开写前两件待办(默认参数确认+显存重审、自查题对答案)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:47:08 -04:00
iomgaa eb56883267 层1 正式关账:判据全过、接口回看完成、疤痕档案与诊断工具箱入存档点
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:44:58 -04:00
iomgaa 872d4bd6a6 层1: diag_generate ruff 修复
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:37:28 -04:00
iomgaa e472a66959 层1: checkpoint 生成测试脚本(关账判据 3)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:37:13 -04:00
iomgaa 4f13365ffa 层1: 损失缩放终解——显式退出 HF 新式梯度累积契约(model_accepts_loss_kwargs=False)
第二幕根因:新式契约的 ×world_size 补偿在基类 compute_loss 尾部(v5
trainer.py:2028),整体重写会绕过它 → loss 与梯度 ÷4(sanity 0.244≈0.85/4)。
按 HF 文档建议(trainer.py:1977)退出新式契约回经典行为;docs/02 勘误改为
两幕全记录。

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:25:11 -04:00
iomgaa ee16e29846 层1: 修复 HF 梯度累积契约坑——compute_loss 按新式契约返回 sum/num_items_in_batch
根因(探针定案):预训练模型真实 CE≈0.85,训练日志 7.5≈0.94×8(累积步数)。
Qwen3 forward 接受 loss_kwargs → Trainer 走新式契约不再除以累积步数,我们
返回裸 mean 导致日志与梯度同放大 8 倍。sft_loss 纯函数不动,适配收口在
compute_loss;docs/02 §5 勘误起点预期(~0.85)并记录此坑。

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:18:09 -04:00
iomgaa 1976230250 层1: 损失探针脚本——预训练模型走管线逐行算 CE,区分数据问题与训练环节问题
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:14:31 -04:00
iomgaa ff572bf4b9 层1: collator 对齐诊断脚本——排查远程 loss 7.5 与 \50118 解码异常(逐环检验缓存/模板/分词/切片)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 10:05:46 -04:00
iomgaa a0faec0df7 层1: 修复远程 sanity OOM——per_device batch 8→2、累积 2→8(全局 64 不变)
根因:150k 大词表下显存大头是 (B,T,V) logits 链(fp32 ~20G@B=8)与逐层激活,
均正比于 B 而与 0.6B 参数量无关。docs/02 §2.6 旧显存估算勘误入档;
train_sft.sh 加 expandable_segments 防碎片。

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 09:21:38 -04:00
iomgaa 68216ace22 docs: 第三章勘误——off-policy 切片是数据轨迹上的 KL 蒸馏而非混 SFT;补 prompt-only 下静默空转的说明
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 09:02:41 -04:00
iomgaa 1d272b6d6c docs: roadmap 索引同步(第三章已写待读)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 08:34:45 -04:00
iomgaa 3c1530a1db docs: 第三章——层 2 white-box OPD(式(2) 精读、standard 路径解剖、U1-U5 重构任务)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 08:34:29 -04:00
iomgaa 4621ebae31 层1/T2: max_tokens 放大至 16384(防截断,上限不计费);并发提到 16;进度行加速度与预计剩余时间
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 08:20:57 -04:00
iomgaa c5a3b7d0bb 层1/T5: 自包含训练脚本 train_sft.sh + train_sft.py;层 1 代码收口
- train_sft.py: FSDP 环境变量前置块(import 前,DDP 下无害);fail-fast 加载
  顺序(数据→tokenizer→模型);首样本自检打印(真 tokenizer 掩码边界肉眼核对);
  remove_unused_columns=False 等非显然约束逐条注释;sanity 模式 = replace 覆盖
- train_sft.sh: 显式 CUDA_VISIBLE_DEVICES 4 卡、PYTHONUNBUFFERED、HF 镜像/缓存
  改道 /data,前置检查清单(含 scp 数据命令)
- 本地验证:fail-fast 到 teacher 缓存缺失处显式报错(941/1000,59 条真实命中
  反向证明 prompt_key 契约端到端成立)
- roadmap 存档点:层 1 代码完成,进入运行阶段

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 08:16:24 -04:00
iomgaa 42a4349e69 层1/T2: 生成脚本接入本地已有的 DAPO 去重版 parquet(17917 行),补试跑说明
数据管线已用真实数据验证(加载→归一→seed 抽子集,5 条无报错)

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 08:11:01 -04:00
iomgaa e72d07aced 层1: ars_opd 改为 editable 安装,修复脚本/调试器 import 失败
- pyproject.toml: 最小打包配置(依赖仍统一走 requirements*.txt)
- .vscode/launch.json: 调试当前文件 + teacher 生成两个配置,cwd 锚定仓库根
- setup_remote.sh 与 CLAUDE.md 常用命令同步 pip install -e

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 08:06:40 -04:00
iomgaa 20e6d97427 层1/T2: teacher.py 批量生成 + sha256 JSONL 缓存;teacher 改定 MiniMax-M3
- teacher.py: 通用 OpenAI 兼容客户端(配置驱动 base_url,替代 OpenRouter 专用);
  缓存即断点(逐条落盘+flush,重跑自动续传);单条失败先落盘其余、结束汇总显式报错;
  M3 思考段 <think>...</think> 入库前剥离(只剥开头一段)
- configs.py: 新增 TeacherGenConfig(采样参数显式化;连接三元组走 .env)
- scripts/generate_teacher_completions.py: 自包含生成脚本(本地跑,与训练侧
  同 seed 同子集约束已注明)
- teacher 决策变更同步:.env.example / docs/00 关键设定与存档点 / docs/02
- tests/test_teacher.py: 10 个单测(假客户端注入),含与 attach 的端到端契约闭环

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 07:52:53 -04:00
iomgaa f5bb852fde 层1/T4: trainer.py 掩码 SFT 损失 + 最小 HF Trainer 子类(docs/02 §2.4)
- sft_loss 纯张量函数:batch-min prompt_length、移位切片、labels 重掩码;
  空 batch/无 prompt 行显式报错(参考实现静默归零,差异已标注)
- SFTTrainer 只重写 compute_loss 与 log;不向模型传 labels(防内部损失
  绕过重掩码);日志并入每步有效 token 数供 sanity 监控
- tests/test_trainer.py: 7 个单测,含手算对拍、移位对齐、窗口外无关性、
  漏切重掩码

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 05:02:41 -04:00
iomgaa 5ea58ddf59 层1/T3: data.py 数据管线——加载、messages 归一、teacher 缓存挂接、双预算 collator(docs/02 §2.2-2.3)
- to_messages 三分支归一,except:pass 改显式报错(差异标注在注释)
- prompt_key: sha256 内容寻址,作为与 teacher.py 的缓存契约单点定义
- attach_teacher_completions: 缺键一次性报全,绝不静默跳过
- SFTCollator: 双预算截断 + 未截断长度定边界(坑二)+ -100 掩码 + 左 padding;
  prompt-only 行显式报错(层 2 接 on-policy 再放开)
- tests/test_data.py: 19 个单测,玩具字符级 tokenizer 覆盖超长解答/超长题目/
  enable_thinking 双取值/左 padding/缓存契约

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 04:51:19 -04:00
iomgaa 58d75cc56a 层1/T1: SFTConfig dataclass 与构造校验(docs/02 §4)
- ars_opd/configs.py: 冻结 dataclass,机器路径无默认强制显式传入;
  双预算/enable_thinking/lr 差异均按 §7 规范标注
- __post_init__ 构造即校验,防"completion 预算为零→loss 恒 0"静默空训练
- tests/test_configs.py: 6 个校验测试

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 04:40:45 -04:00
iomgaa 240f404416 工作模式变更入档:层 1 起 Claude 编写代码、用户精读提问;存档点同步
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 04:30:55 -04:00
iomgaa 6c128ad6f5 docs/02:FSDP 决策定案——DDP 到 1.7B,触发点写死(4B 或 seq>8K),env 坑今日消除(脚本头部前置块)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 04:21:30 -04:00
iomgaa 0bcfc6efd4 层 0 关账:两端环境验收通过(本地 4 passed + 远程 4 passed),存档点进入层 1
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 04:05:17 -04:00
iomgaa 12e2a8b0ae docs: 第二章——层1 SFT 基线开工文档(参考实现解剖+保留/替代/删除清单+任务表);CLAUDE.md 映射表补 data.py
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 03:51:06 -04:00
iomgaa 77f0419bad roadmap:训练数据定为 DAPO-Math-17K(对齐论文§5.1),登记 φ/β 两处论文-代码默认值背离
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 03:47:30 -04:00
iomgaa 9c0ec15716 setup_remote.sh 自清 pip 残留;存档点记录远程根分区结构性问题(/root/zym 507G)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 03:39:14 -04:00
iomgaa 991e277316 CLAUDE.md:日志不缓存升级为全局硬规则(含 conda run 事故伤疤标注)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 03:30:34 -04:00
iomgaa b679bcb24f 修复 setup_remote.sh:弃用 conda run(输出整体缓冲),改为直调环境内 pip/python 实现实时日志
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 03:28:13 -04:00
iomgaa f255d5613a docs: roadmap 增加存档点小节(断点恢复用),登记层 0 当前状态
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 03:20:12 -04:00
33 changed files with 3849 additions and 15 deletions
+3 -3
View File
@@ -1,8 +1,8 @@
# 复制为 .env 并填入真实值(.env 已被 gitignore,严禁提交密钥)
# teacher APIOpenAI 兼容格式DeepSeek / MiniMax 二选一填
TEACHER_API_BASE=https://api.deepseek.com/v1
# teacher APIOpenAI 兼容格式2026-07-18 定:自建 new-api 网关 + MiniMax-M3
TEACHER_API_BASE=https://newapi.iomgaa.online/v1
TEACHER_API_KEY=
TEACHER_MODEL=deepseek-chat
TEACHER_MODEL=MiniMax-M3
# W&B(仅远程训练需要)
WANDB_API_KEY=
+26
View File
@@ -0,0 +1,26 @@
{
// VSCode 调试配置。前提:本地环境已 `pip install -e .`(见 README/CLAUDE.md),
// 且右下角解释器已选 ars-opd 环境。cwd 固定为仓库根:脚本里的相对路径
// data/...、outputs/...)都以仓库根为基准。
"version": "0.2.0",
"configurations": [
{
"name": "调试当前文件",
"type": "debugpy",
"request": "launch",
"program": "${file}",
"console": "integratedTerminal",
"cwd": "${workspaceFolder}",
"justMyCode": false
},
{
"name": "teacher 批量生成(1k 子集)",
"type": "debugpy",
"request": "launch",
"program": "${workspaceFolder}/scripts/generate_teacher_completions.py",
"console": "integratedTerminal",
"cwd": "${workspaceFolder}",
"justMyCode": false
}
]
}
+4
View File
@@ -0,0 +1,4 @@
{
"python-envs.defaultEnvManager": "ms-python.python:conda",
"python-envs.defaultPackageManager": "ms-python.python:conda"
}
+5 -3
View File
@@ -1,7 +1,7 @@
# CLAUDE.md
> [!URGENT]
> 1. 本项目是**学习驱动的科研重构项目**:通过分层重建 OmniOPD 参考实现来学习论文方法,同时产出可作为后续研究基座的干净代码库。学习优先——用户要逐章理解每个模块,**不要一次性替用户写完所有代码**;每章先讲清楚,重构任务由用户主导、你配合
> 1. 本项目是**学习驱动的科研重构项目**:通过分层重建 OmniOPD 参考实现来学习论文方法,同时产出可作为后续研究基座的干净代码库。学习优先——每章先讲清楚原理再动代码。**分工(2026-07-18 用户定):代码由 Claude 编写,用户逐行精读并提问**;因此代码必须严格教学导向(遵守 §7 注释规范),每个模块写完后主动讲解设计要点与易错点,用户的提问优先于推进进度
> 2. 所有思考过程和回复必须使用**简体中文**。
## 1. 项目元数据
@@ -25,7 +25,8 @@
| `ars_opd/similarity.py` | §3.2.1 式(3) | 语义相似度 φ(rouge1 / edit_distance |
| `ars_opd/estimator.py` | §3.2.2 式(4)(5) | MC 估计 + Dirichlet 贝叶斯平滑 π̂ |
| `ars_opd/chunking.py` | §3.2.3 式(6)(7) | 熵计算 + peak-entropy chunk 选择 |
| `ars_opd/teacher.py` | §3.2.1 | teacher rollout 采样(OpenAI 兼容 API + 缓存 / vLLM |
| `ars_opd/data.py` | —(IO 边缘,无论文锚点) | 数据加载(DAPO parquet/HF+ messages 归一 + 双预算掩码 collator |
| `ars_opd/teacher.py` | §3.2.1 | teacher rollout 采样(OpenAI 兼容 API + 缓存 / vLLM);层 1 起步能力:批量生成+sha256 缓存 |
| `ars_opd/trainer.py` | §3.2.4 式(8) | chunk 蒸馏损失 + trust-region KL 锚定,编排全流程 |
## 4. 常用命令
@@ -33,6 +34,7 @@
```bash
conda activate ars-opd && pytest tests/ -x -q # 单元测试(本地 CPU
conda activate ars-opd && ruff check ars_opd/ --fix && ruff format ars_opd/
pip install -e . --no-build-isolation --no-deps # 新环境一次性:注册 ars_opd 包(否则脚本 import 报错)
```
## 5. 本地-远程规则
@@ -40,7 +42,7 @@ conda activate ars-opd && ruff check ars_opd/ --fix && ruff format ars_opd/
> [!CRITICAL]
> - 远程机 gpu-a800-0608×A800-80G**只允许用其中 4 块**;所有 GPU 命令必须显式 `CUDA_VISIBLE_DEVICES=<idx>`,先 `nvidia-smi` 确认空闲卡,严禁自动选卡。
> - 远程**根分区仅剩 12G**conda 环境、HF 缓存(`HF_HOME`)、模型、checkpoint、数据集一律放 `/data/zym/` 下。
> - 长任务用 tmux 跑,日志缓存`python -u` / `PYTHONUNBUFFERED=1`),确保可实时检查
> - **任何长时间运行的命令(训练、pip/conda 安装、脚本)禁止日志缓存**,宁可承担延时也要实时可查:python 加 `-u` / `PYTHONUNBUFFERED=1`**不用 `conda run` 包裹长命令**(它整体缓冲输出直到结束——2026-07 曾因此把正常安装误判为卡死),改为直调 `<env>/bin/pip`、`<env>/bin/python`;远程长任务一律 tmux
> - 远程机器上不改代码,只 `git pull` + 跑脚本;本地不跑训练。
## 6. 学习工作流
+4
View File
@@ -0,0 +1,4 @@
"""ars_opdOmniOPD (arXiv:2606.01476v2) 的分层重构实现。
模块与论文的对应关系见 CLAUDE.md §3(单一事实源),此处不重复。
"""
+310
View File
@@ -0,0 +1,310 @@
"""实验配置(层 1SFTConfig / TeacherGenConfig;层 2DistillConfig)。
设计约定(对应 CLAUDE.md §2"配置显式化"):
- 所有实验参数必须是这里某个 dataclass 的字段;代码里出现魔法数字/路径即违规。
- 密钥(API key 等)不进配置类,走 `.env`(见 teacher.py)。
- 机器相关路径(数据集、输出目录)不给默认值,强制调用方显式传入——
防止参考实现里 `/fsx` 硬编码那类"在别人机器上必炸"的坑。
- 各层的 config 自包含、不互相继承:层与层是不同实验,共享基类会把它们耦合,
违背"从上读到下看懂全部流程"CLAUDE.md §2)。字段重复是有意接受的成本。
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True)
class SFTConfig:
"""层 1 SFT 基线的全部实验参数。
论文锚点:§3.1 式(1) 的标准交叉熵 SFT;但按 §5.1 的基线定义,
训练数据是 teacher rollout(离线蒸馏),不是人写答案——所以有
`teacher_completions_path` 字段:DAPO 是 prompt-only 数据集,
解答一律来自 teacher 生成的缓存文件。
frozen=True:配置一旦构造即只读。训练中途被悄悄改掉的配置是最难
排查的 bug 来源之一;要换参数就构造一个新实例,留下明确的代码痕迹。
"""
# ---- 机器相关路径(无默认值,必须显式传入)----
dataset_path: str
"""DAPO-Math-17K 的本地 parquet 路径(文件或目录),或 HF Hub 数据集名。"""
output_dir: str
"""checkpoint 与日志输出目录(远程必须落在 /data/zym 下)。"""
# ---- 数据 ----
teacher_completions_path: str | None = None
"""teacher 解答缓存(teacher.py 生成的 JSONL)。None 表示数据集自带
assistant 轮次;若数据实际是 prompt-only 又没给此路径,data.py 会显式报错,
不做静默兜底。"""
dataset_split: str = "train"
subset_size: int | None = 1000
"""随机抽取的子集大小(控制 teacher API 成本,roadmap 定为 ~1k);None = 全量。"""
# ---- 序列双预算 ----
# 非显然约束:prompt 与 completion 必须各有独立预算。若只用一个 max_length
# 从右截断,超长解答会把 prompt 挤空,模型在"没有题目"的样本上学习解答
# ——这是参考实现 collator(trainer:267-292) 的头号正确性卖点,此处继承。
max_length: int = 4096
"""prompt + completion 的总 token 预算。"""
max_prompt_length: int = 1024
"""prompt 单独预算;completion 实际预算 = max_length - len(截断后 prompt)。"""
enable_thinking: bool = False
"""Qwen3 chat 模板的思考开关。False 时模板注入空 `<think>\\n\\n</think>`。
非显然约束:此开关改变渲染后的 prompt 文本,从而改变 prompt/completion
的 token 边界——训练与推理必须取同一值,否则掩码整体错位。"""
# ---- 优化 ----
# 差异标注:论文 §5.1 的蒸馏训练用 lr=1e-6,参考实现 SFT 默认 2e-5
# (train_distillation.py:73)。SFT 有真实 token 监督、信号密集,从参考实现取
# 2e-5;层 5 的蒸馏配置再回到论文的 1e-6。
learning_rate: float = 2e-5
per_device_train_batch_size: int = 2
"""非显然约束:别看 0.6B 小就调大它——显存大头是 (B,T,V) 的 logits 链
fp32 一份 ~20G@B=8)与逐层激活,都正比于 B 而与参数量无关;B=8 实测
爆 80G 卡(2026-07-18 远程 sanity)。"""
gradient_accumulation_steps: int = 8
"""全局 batch = 2(per_device) × 4(卡) × 8(累积) = 64,与参考实现注释的
训练规模(trainer 配置注释"global batch 64")对齐。"""
num_train_epochs: int = 1
max_steps: int = -1
""">0 时覆盖 num_train_epochs,只跑这么多步——远程 50 步 sanity 用;-1 = 按 epoch。"""
lr_scheduler_type: str = "linear"
warmup_ratio: float = 0.0
gradient_checkpointing: bool = False
"""0.6B 学生显存富余,不开(省 ~40% 显存、慢 ~30%)。注意:FSDP 下此开关
是 no-op,真正的开关是 FSDP_ACTIVATION_CHECKPOINTING 环境变量(见
scripts/ 训练脚本头部的前置块,docs/02 §2.6)。"""
bf16: bool = True
seed: int = 42
# ---- 日志与保存 ----
logging_steps: int = 1
save_steps: int = 100
save_total_limit: int = 2
report_to: str = "none"
""""none""wandb"。默认 none:本地调试不该悄悄往外发数据,远程脚本显式开。"""
def __post_init__(self) -> None:
"""构造即校验:配置错误必须在训练开始前炸,而不是跑到第一个超长样本才炸。"""
if self.max_prompt_length >= self.max_length:
raise ValueError(
f"max_prompt_length({self.max_prompt_length}) 必须小于 "
f"max_length({self.max_length}),否则 completion 预算为零,"
f"所有样本的 labels 将全为 -100,loss 恒为 0 且无报错——静默空训练。"
)
if self.learning_rate <= 0:
raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}")
if self.subset_size is not None and self.subset_size <= 0:
raise ValueError(
f"subset_size 必须为正整数或 None(全量),收到 {self.subset_size}"
)
if self.max_steps == 0 or self.max_steps < -1:
raise ValueError(
f"max_steps 只接受 -1(按 epoch)或正整数,收到 {self.max_steps}"
)
@dataclass(frozen=True)
class TeacherGenConfig:
"""teacher 批量生成(层 1 能力)的采样与执行参数。
连接信息(API 地址/密钥/模型名)不在这里——那是部署环境的事实,走 `.env`
teacher.py 读取);这里只放"换一组值就是换一个实验"的采样参数。
"""
temperature: float = 1.0
top_p: float = 0.95
"""MiniMax M 系官方推荐采样参数:temperature=1.0, top_p=0.95。"""
max_tokens: int = 16384
"""teacher 单条回复的 token 上限。这是上限不是目标——按实际生成量计费,
放大它不增加正常解答的成本,只给最难的题留出写完的空间(8192 时 59 条实测
截断 2 条)。非显然约束:M3 的思考段也计入此额度,设太小会把解答挤没。"""
strip_think: bool = True
"""剥离 content 开头的 <think>...</think> 思考段。SFT 的监督目标是最终
解答;student 以 enable_thinking=False 训练,学思考段会与模板约定矛盾。"""
concurrency: int = 16
"""并发请求数(线程池大小)。上限看网关的承受力,报 429 就调小。"""
max_retries: int = 3
"""单请求的网络级重试次数(openai 客户端内建指数退避)。"""
system_prompt: str | None = None
"""None = 不加 system 轮(DAPO 题面自带作答指令,不需要额外指挥)。"""
def __post_init__(self) -> None:
if self.max_tokens <= 0:
raise ValueError(f"max_tokens 必须为正,收到 {self.max_tokens}")
if self.concurrency < 1:
raise ValueError(f"concurrency 必须 ≥1,收到 {self.concurrency}")
if self.temperature < 0:
raise ValueError(f"temperature 必须 ≥0,收到 {self.temperature}")
@dataclass(frozen=True)
class DistillConfig:
"""层 2 white-box OPD 基线(token 级反向 KL 蒸馏)的全部实验参数。
论文锚点:§3.1 式(2) 的 on-policy 白盒蒸馏 L = E_{y~π_θ}[Σ_t KL(π_θ ‖ π_T)]。
与层 1 SFTConfig 的三处结构性差异(docs/03 §3 偏差清单):
- 无 teacher_completions_pathteacher 现场前向给出全词表 logits、student 现场
on-policy 生成轨迹,两者都不落盘缓存,故层 2 不需要 teacher 解答文件。
- 无 max_length(总预算):completion 不再来自数据,而是 model.generate 生成,
序列总长 = prompt(≤max_prompt_length) + 生成(≤max_new_tokens),由两个预算
各自界定,不需要一个总的右截断预算。
- 无 top_k:§4 删除清单——本地同 tokenizer teacher 放得下全词表,恒走精确
全词表 KL,不做参考实现默认的 top-1 稀疏近似(那是 API 传输妥协,非论文成分)。
teacher_model 在此、student 在脚本(U5 的常量,同层 1 的 STUDENT_MODEL):
student 是被训练的固定基线,teacher 是"换一个就是换一个实验"的旋钮,故归 config。
"""
# ---- 机器相关路径(无默认值,必须显式传入)----
dataset_path: str
"""DAPO-Math-17K 的本地 parquet 路径(文件或目录),或 HF Hub 数据集名。
层 2 只用题面(prompt-only),不读数据自带的任何 completion。"""
output_dir: str
"""checkpoint 与日志输出目录(远程必须落在 /data/zym 下)。"""
# ---- teacher(层 2 的核心旋钮)----
teacher_model: str = "Qwen/Qwen3-4B"
"""本地 HF teacher 模型名。非机器路径(HF Hub 名各机可复现),故给默认值。
非显然约束:必须与 student **同 tokenizer**——KL 是逐词表位对齐求和,词表不
一致则第 v 个分量对不上、相除无意义(docs/03 §1)。此约束在 U4 构造 Trainer 时
比对 get_vocab() 显式校验,不匹配即报错,不静默。"""
# ---- 数据(复用层 1 的抽取逻辑,同 seed 同子集)----
dataset_split: str = "train"
subset_size: int | None = 1000
"""随机抽取的子集大小;None = 全量。非显然约束:与层 1 同 seed 同 size 才能
在同一批题上对比 SFT 与蒸馏,否则两层看的是不同题、曲线不可比。"""
# ---- 序列预算(prompt 截断 + 生成上限,见类 docstring 为何无 max_length----
max_prompt_length: int = 1024
"""prompt 单独预算(prompt-only collator 按此左截断)。"""
max_new_tokens: int = 1024
"""student on-policy 生成的 token 上限。与 max_prompt_length 之和即序列总长 T
显存账(§5)按 T=2048 估算。"""
enable_thinking: bool = False
"""Qwen3 chat 模板思考开关,喂给 student 生成。非显然约束:与层 1 取同值,
否则 prompt 渲染文本变、生成分布与 SFT 基线不可比(docs/02 §2.3 边界契约)。"""
# ---- 蒸馏损失(式(2) 与 docs/03 §2.3 三副面孔)----
beta: float = 1.0
"""KL 方向系数。0=前向 KL(π_T‖π_θ, mode-covering)1=反向 KL(π_θ‖π_T,
mode-seeking)=**式(2)**(0,1)=JSD 插值。默认 1 即论文式(2);参数保留是因为
前向/反向/JSD 是同一公式(U2 顺手覆盖),且层 5 的 KL 锚要用前向。"""
kl_temperature: float = 1.0
"""散度内 softmax 前除进两侧 logits 的温度(docs/03 §2.3)。升温放大尾部
"暗知识"排序。非显然约束:它与下面的 gen_temperature 是**两个不同**的温度
——这个调的是 loss 里分布的软硬,那个调的是采样的随机性;恰好都默认 1.0,
但改一个不影响另一个。默认 1.0 即式(2)(不做温度缩放)。"""
# ---- on-policy 生成采样(式(2) 的 y~π_θ 期望)----
gen_temperature: float = 1.0
gen_top_p: float = 1.0
"""student 生成轨迹的采样参数。默认 temperature=1.0/top_p=1.0 = 纯采样自 π_θ,
最忠实于式(2) 的 on-policy 期望(docs/03 §3 抉择原则:本质忠于论文)。
调低是拿保真度换"少生成垃圾",卡了再动。"""
# ---- 优化 ----
# 差异标注:层 1 SFT 用 2e-5(信号密集的真 token 监督);层 2 是蒸馏,从论文
# §5.1 的蒸馏 lr=1e-6。小 lr 在这里还有额外好处:式(2) 的反向 KL 会梯度爆炸
# (§4.1on-policy 采到 teacher 眼中的烂 token 时 log(π_θ/π_T)→∞),小步长
# 帮训练在毛刺中存活——这毛刺本身是层 2 要观察的教学目标(docs/03 §6.3)。
learning_rate: float = 1e-6
per_device_train_batch_size: int = 4
"""非显然约束:白盒蒸馏的显存大头是 **两份**全词表 logitsstudent+teacher
(B,T,V) bf16 各 ~2.5G@B=4/T=2048+ log_softmax 中间量,比层 1 更紧。B=4 是
§5 估算值(student 训练全套 ~10G + teacher 推理副本 ~9G + 两份 logits ~15G
A800-80G 起步安全),但**必须**在首次远程冒烟用 nvidia-smi 实测确认,OOM 阶梯:
先降 B 到 2、仍不够再开 gradient_checkpointing。"""
gradient_accumulation_steps: int = 4
"""全局 batch = 4(per_device) × 4(卡) × 4(累积) = 64,与层 1 保持一致。"""
num_train_epochs: int = 1
max_steps: int = -1
""">0 时覆盖 num_train_epochs——远程 50 步冒烟用(docs/03 §6.3);-1 = 按 epoch。"""
lr_scheduler_type: str = "linear"
warmup_ratio: float = 0.0
max_grad_norm: float = 1.0
"""梯度裁剪阈值。此前是 HF Trainer 的静默默认(1.0),现显式化——它是式(2)
反向 KL 梯度爆炸(§4.1)的**隐形稳定器**on-policy 采到 teacher 眼中烂 token
时单步梯度范数可炸到十几(2026-07-19 首冒烟实测 grad_norm 14→2),HF 默认
裁到 1.0 才让 loss 曲线平稳。把它设得远大于实测范数(≈关闭裁剪)可暴露原始
爆炸,供教学对照(train_whitebox.py 的 noclip 模式)。非显然约束:日志里的
grad_norm 是**裁剪前**范数,故 14→2 那串本身就是爆炸证据,只是被裁剪掩盖了。"""
gradient_checkpointing: bool = False
"""默认不开(§5 显存账 B=4 富余);OOM 时作为降 batch 之后的第二道降显存手段。
注意 FSDP 下此开关是 no-opdocs/02 §2.6),但层 2 坚持 DDP 故此处有效。"""
bf16: bool = True
seed: int = 42
# ---- 日志与保存 ----
logging_steps: int = 1
save_steps: int = 100
save_total_limit: int = 2
report_to: str = "none"
""""none""wandb"。默认 none:本地调试不该悄悄往外发数据,远程脚本显式开。"""
def __post_init__(self) -> None:
"""构造即校验:配置错误必须在加载 4B teacher(GB 级下载)之前炸。"""
if not 0.0 <= self.beta <= 1.0:
raise ValueError(
f"beta 必须在 [0,1]0=前向/1=反向/中间=JSD),收到 {self.beta}"
)
if self.kl_temperature <= 0:
# 温度除进 logits,≤0 会翻转或炸掉分布
raise ValueError(f"kl_temperature 必须为正,收到 {self.kl_temperature}")
if self.gen_temperature <= 0:
# 非显然约束:0 在 HF 里是 greedy,会退化 on-policy 采样为确定性解码,
# 破坏式(2) 的 y~π_θ 期望;要纯 on-policy 就必须 >0
raise ValueError(
f"gen_temperature 必须为正(0=greedy 破坏 on-policy),"
f"收到 {self.gen_temperature}"
)
if not 0.0 < self.gen_top_p <= 1.0:
raise ValueError(f"gen_top_p 必须在 (0,1],收到 {self.gen_top_p}")
if self.max_new_tokens <= 0:
raise ValueError(f"max_new_tokens 必须为正,收到 {self.max_new_tokens}")
if self.learning_rate <= 0:
raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}")
if self.max_grad_norm <= 0:
# 用远大于实测范数的值≈关闭裁剪;≤0 无意义(0 会把梯度裁没)
raise ValueError(f"max_grad_norm 必须为正,收到 {self.max_grad_norm}")
if self.subset_size is not None and self.subset_size <= 0:
raise ValueError(
f"subset_size 必须为正整数或 None(全量),收到 {self.subset_size}"
)
if self.max_steps == 0 or self.max_steps < -1:
raise ValueError(
f"max_steps 只接受 -1(按 epoch)或正整数,收到 {self.max_steps}"
)
+407
View File
@@ -0,0 +1,407 @@
"""数据管线(IO 边缘,无论文锚点):加载 → messages 归一 → 挂接 teacher 解答 → collator。
层 1 的数据流(对应 docs/02 §1 的基线定义:SFT = 在 teacher rollout 上的离线蒸馏):
DAPO parquetprompt-only
→ to_messages 归一成 [{"role","content"}] 列表
→ 按 seed 抽子集
→ attach_teacher_completions 从 JSONL 缓存挂上 teacher 解答(assistant 轮)
→ SFTCollator 分词、双预算截断、-100 掩码、左 padding
本模块与 teacher.py 的缓存契约由 `prompt_key` 单点定义:teacher.py 生成缓存、
本模块消费缓存,双方必须用同一个函数算键。
"""
from __future__ import annotations
import ast
import hashlib
import json
from pathlib import Path
from typing import Any
import torch
from datasets import Dataset, load_dataset
# F.cross_entropy 的 ignore_index 默认值;标了它的位置不产生 loss
IGNORE_INDEX = -100
# ---------------------------------------------------------------------------
# messages 归一
# ---------------------------------------------------------------------------
def _parse_stringified_list(value: str, column: str) -> list:
"""parquet 有时把 list 存成其字符串形态,用 ast 还原。
差异标注:参考实现(train_distillation.py:299-305)在这里 `except: pass` 静默吞错,
坏行会以原始字符串流进 collator,在 apply_chat_template 处以难懂的方式炸;
我们显式报错,错误信息直接指向坏数据本身。
"""
try:
parsed = ast.literal_eval(value)
except (ValueError, SyntaxError) as e:
raise ValueError(
f"{column!r} 是字符串但无法解析为 Python 字面量(坏数据行):"
f"{value[:200]!r}"
) from e
if not isinstance(parsed, (list, tuple)):
raise ValueError(f"{column!r} 解析结果不是列表:{type(parsed).__name__}")
return list(parsed)
def to_messages(example: dict[str, Any]) -> dict[str, list[dict[str, str]]]:
"""把三种来源格式归一成 messages 列:[{"role": ..., "content": ...}, ...]。
支持(与参考实现 train_distillation.py:295-321 相同的三分支):
- ``messages`` 列:直取;
- ``prompt`` 列(DAPO parquet,列名不副实——装的是完整 chat 列表):改名;
- ``question`` 列(gsm8k 风格纯文本):包成单 user 轮。
差异标注:参考实现对不认识的行 `return x` 静默放行,我们显式报错。
"""
if "messages" in example:
msgs = example["messages"]
column = "messages"
elif "prompt" in example:
msgs = example["prompt"]
column = "prompt"
elif "question" in example:
return {"messages": [{"role": "user", "content": example["question"]}]}
else:
raise ValueError(
f"无法识别的数据行:既无 messages/prompt 也无 question 列,"
f"实有列 {sorted(example.keys())}"
)
if isinstance(msgs, str):
msgs = _parse_stringified_list(msgs, column)
msgs = list(msgs)
if not msgs:
raise ValueError(f"{column!r} 是空列表(坏数据行)")
for m in msgs:
if not isinstance(m, dict) or "role" not in m or "content" not in m:
raise ValueError(
f"{column!r} 中存在非 {{role, content}} 结构的元素:{m!r}"
)
return {"messages": [{"role": m["role"], "content": m["content"]} for m in msgs]}
# ---------------------------------------------------------------------------
# teacher 解答缓存(与 teacher.py 的契约)
# ---------------------------------------------------------------------------
def prompt_key(messages: list[dict[str, str]]) -> str:
"""teacher 缓存的键:对 messages 的规范化 JSON 取 sha256。
差异标注:参考实现(distillation_trainer.py:984)用 `str(hash(prompt))`——
Python 对 str 的 hash 默认加盐,跨进程/跨次运行不稳定,缓存必然失效重生成。
sha256 内容寻址:同一道题永远同一个键。
只取 role/content 两个字段参与哈希:DAPO 行里其余元数据(data_source 等)
变了不应导致缓存失效。
"""
canon = [{"role": m["role"], "content": m["content"]} for m in messages]
return hashlib.sha256(json.dumps(canon, ensure_ascii=False).encode()).hexdigest()
def attach_teacher_completions(dataset: Dataset, jsonl_path: str) -> Dataset:
"""把 teacher 解答缓存(JSONL,每行 {"key", "completion"})挂到数据集上。
- 末轮已是 assistant 的行保持原样(数据自带解答,不覆盖);
- 任何 prompt-only 行在缓存中查不到键 → 收集齐所有缺失后一次性报错,
提示先运行 teacher 生成——绝不静默跳过(跳过 = 悄悄改变训练集组成)。
"""
path = Path(jsonl_path)
if not path.exists():
raise FileNotFoundError(
f"teacher 解答缓存不存在:{jsonl_path}。先运行 teacher.py 的批量生成。"
)
cache: dict[str, str] = {}
with open(path, encoding="utf-8") as f:
for line_no, line in enumerate(f, 1):
if not line.strip():
continue
rec = json.loads(line) # 坏行直接炸,带行号
if "key" not in rec or "completion" not in rec:
raise ValueError(f"{jsonl_path}:{line_no} 缺少 key/completion 字段")
cache[rec["key"]] = rec["completion"]
# 先整体扫描缺失,一次性报全——比在 .map 里炸第一条更省来回
missing = [
i
for i, ex in enumerate(dataset)
if ex["messages"][-1]["role"] != "assistant"
and prompt_key(ex["messages"]) not in cache
]
if missing:
raise KeyError(
f"{len(missing)}/{len(dataset)} 行在 teacher 缓存中查不到解答"
f"(首个缺失行 index={missing[0]})。检查:teacher 生成是否用了同一"
f"子集与同一 seed?(子集抽取在 load_sft_dataset 中先于挂接发生,"
f"两侧 seed 不同则键集合不同)"
)
def _attach(ex: dict[str, Any]) -> dict[str, Any]:
msgs = ex["messages"]
if msgs[-1]["role"] == "assistant":
return ex
completion = cache[prompt_key(msgs)]
return {"messages": msgs + [{"role": "assistant", "content": completion}]}
return dataset.map(_attach)
# ---------------------------------------------------------------------------
# 数据集加载(入口)
# ---------------------------------------------------------------------------
def load_sft_dataset(
dataset_path: str,
dataset_split: str = "train",
subset_size: int | None = None,
seed: int = 42,
teacher_completions_path: str | None = None,
) -> Dataset:
"""数据管线入口:加载 → 归一 → 抽子集 →(可选)挂 teacher 解答。
收散装参数而非整个 config(深模块:本函数只用这 5 个字段,不该索要一整个
SFTConfig)。这样层 1SFTConfig)、层 2DistillConfig,无 teacher 缓存)、
诊断脚本都能直接调,无需伪造无关字段。teacher_completions_path=None 时
返回 prompt-only 数据集(末轮 user,供 on-policy 生成);给了则挂 teacher
解答(末轮 assistant,供 SFT)。
返回只含 ``messages`` 一列的 Dataset。
"""
ds = _load_raw(dataset_path, dataset_split)
ds = ds.map(
to_messages,
remove_columns=[c for c in ds.column_names if c != "messages"],
)
if subset_size is not None and subset_size < len(ds):
# 非显然约束:抽子集必须在挂接 teacher 解答之前、且由 seed 完全确定——
# teacher.py 生成缓存时会走完全相同的"加载→归一→抽子集"路径,两侧 seed
# 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。层 2 与层 1
# 用同 seed 同 subset_size,才能在同一批题上对比 SFT 与蒸馏。
ds = ds.shuffle(seed=seed).select(range(subset_size))
if teacher_completions_path is not None:
ds = attach_teacher_completions(ds, teacher_completions_path)
return ds
def _load_raw(dataset_path: str, split: str) -> Dataset:
"""三分支加载:parquet 目录 / 单 parquet 文件 / HF Hub 数据集名。
差异标注:参考实现(train_distillation.py:292)对 Hub 分支硬编码 config 名
"main"(gsm8k 专用);我们不硬编码——需要特定 config 的数据集请下载成
parquet 本地加载。
"""
p = Path(dataset_path)
if p.is_dir():
return load_dataset("parquet", data_dir=dataset_path, split=split)
if dataset_path.endswith(".parquet"):
return load_dataset("parquet", data_files=dataset_path, split=split)
return load_dataset(dataset_path, split=split)
# ---------------------------------------------------------------------------
# Collator:本层最核心的一段(对拍 distillation_trainer.py:210-343
# ---------------------------------------------------------------------------
class SFTCollator:
"""把一个 batch 的 messages 变成训练/生成所需的定长张量,两种模式二选一。
prompt_only=False(层 1 SFT,默认)——输出 input_ids/attention_mask/labels
核心设计(继承参考实现的双预算方案,docs/02 §2.3):prompt 与 completion
各自独立预算——prompt 用 max_prompt_length 截断,completion 上限是
max_length - len(截断后 prompt)。若只用一个总预算从右截断,超长解答会把
prompt 挤空,模型在"没有题目"的样本上学解答。要求每行末轮是 assistant,
否则报错(prompt-only 行在纯 SFT 下只产生零 loss = 静默空训练)。
prompt_only=True(层 2 white-box OPDdocs/03 §5 U3)——只渲染 prompt、
输出 prompts/prompt_attention_mask 供 model.generate 做 on-policy 生成;
completion 由生成产生、labels 由 U4 的 DistillTrainer 在生成后重建,故此模式
不产 labels、也不吃 max_length。这兑现了参考实现为 on-policy 生成留的口子
(层 1 曾故意关掉,见此前 git 历史)。
与参考实现的其余差异:空 <think> 的一次性诊断打印改为单元测试断言(契约进
测试,不进运行时日志)。
"""
def __init__(
self,
tokenizer: "Any",
max_prompt_length: int,
max_length: int | None = None,
enable_thinking: bool = False,
prompt_only: bool = False,
) -> None:
"""tokenizer 需实现 HF 接口:apply_chat_template / __call__ / pad_token_id。
max_length 仅 SFT 模式需要(completion 预算依赖它);prompt_only 模式下
completion 是生成的、无总预算,故 max_length 可为 None。
"""
if not prompt_only and max_length is None:
raise ValueError(
"SFT 模式(prompt_only=False)必须提供 max_length——completion "
"预算 = max_length - len(prompt),缺它无法确定解答截断点。"
)
self.tokenizer = tokenizer
self.max_length = max_length
self.max_prompt_length = max_prompt_length
self.enable_thinking = enable_thinking
self.prompt_only = prompt_only
# pad→eos 回退:左 padding 位置的 attention_mask 恒为 0pad 值不参与
# 任何计算,只需要一个合法 token id 占位,借用 eos 即可
if tokenizer.pad_token_id is not None:
self.pad_token_id: int = tokenizer.pad_token_id
elif tokenizer.eos_token_id is not None:
self.pad_token_id = tokenizer.eos_token_id
else:
raise ValueError("tokenizer 既无 pad_token 也无 eos_token,无法 padding")
def __call__(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
"""按模式分派:prompt_only 走生成用 prompt 张量,否则走 SFT 双预算。"""
if self.prompt_only:
return self._collate_prompt_only(examples)
return self._collate_sft(examples)
def _collate_prompt_only(
self, examples: list[dict[str, Any]]
) -> dict[str, torch.Tensor]:
"""层 2:只渲染 prompt 供 on-policy 生成,不产 completion/labels。
返回(B = batch 大小,P = batch 内最长 prompt 长度):
- prompts: (B, P) 左 padding
- prompt_attention_mask: (B, P) padding 位置为 0
非显然约束:生成必须左 padding——所有 prompt 右对齐到同一右边界,
model.generate 从该边界统一续写;右 padding 会让短 prompt 的生成从 pad
中间开始,全乱。这也是层 1 SFT 就选左 padding 的原因(全项目一种约定)。
"""
all_prompt_ids: list[list[int]] = []
for example in examples:
messages = example["messages"]
# prompt-only 数据末轮是 user;若末轮已是 assistant 则剥掉,取生成前上下文
prompt_msgs = (
messages[:-1] if messages[-1]["role"] == "assistant" else messages
)
if not prompt_msgs:
raise ValueError(
"prompt_only collator 收到空 prompt(无可生成的上下文)"
)
# 与 SFT 模式同样带生成引导符渲染(add_generation_prompt=True):
# prompt 末尾就是 "<|im_start|>assistant\n...",生成从此续写
formatted_prompt = self.tokenizer.apply_chat_template(
prompt_msgs,
tokenize=False,
add_generation_prompt=True,
enable_thinking=self.enable_thinking,
)
prompt_ids: list[int] = self.tokenizer(
formatted_prompt,
truncation=True,
max_length=self.max_prompt_length,
add_special_tokens=False,
)["input_ids"]
all_prompt_ids.append(prompt_ids)
return {
"prompts": _left_pad(all_prompt_ids, self.pad_token_id), # (B, P)
"prompt_attention_mask": _left_pad(
[[1] * len(ids) for ids in all_prompt_ids], 0
), # (B, P)
}
def _collate_sft(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
"""层 1 SFTmessages(末轮 assistant)→ 定长张量。
返回(B = batch 大小,T = batch 内最长序列长度):
- input_ids: (B, T) 左 padding
- attention_mask: (B, T) padding 位置为 0
- labels: (B, T) padding 与 prompt 位置为 -100completion 位置为 token id
"""
all_input_ids: list[list[int]] = []
all_labels: list[list[int]] = []
for example in examples:
messages = example["messages"]
if len(messages) < 2 or messages[-1]["role"] != "assistant":
raise ValueError(
"SFTCollator 收到 prompt-only 行(末轮不是 assistant)。纯 SFT 下"
"它只会产生全 -100 的零 loss 样本——静默空训练。检查 teacher "
"解答是否挂接成功。"
)
# prompt = 末轮 assistant 之前的全部轮次,渲染时带生成引导符
# "<|im_start|>assistant\n..."),这样 completion 是纯解答文本的分词
formatted_prompt = self.tokenizer.apply_chat_template(
messages[:-1],
tokenize=False,
add_generation_prompt=True,
enable_thinking=self.enable_thinking,
)
# prompt 自己的预算内截断。沿用 tokenizer 默认右截断(与参考实现一致):
# 超预算的题目被截掉尾部(含生成引导符)——1024 预算下 DAPO 极少触发,
# 触发时该样本退化但不会污染边界(边界用未截断长度算,见下)
prompt_ids: list[int] = self.tokenizer(
formatted_prompt,
truncation=True,
max_length=self.max_prompt_length,
add_special_tokens=False,
)["input_ids"]
# 非显然约束(docs/02 坑一/坑二):completion 边界必须用"未截断 prompt
# 的分词长度"从整段渲染中切出。BPE 分词不满足拼接稳定性,分开渲染
# prompt 和 completion 再拼接 ≠ 整段渲染后分词;而若用截断后长度当切分
# 点,会把 prompt 尾部的 token 误标成 completion——静默的语义错误。
formatted_full = self.tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=False,
enable_thinking=self.enable_thinking,
)
full_ids: list[int] = self.tokenizer(
formatted_full, truncation=False, add_special_tokens=False
)["input_ids"]
untruncated_prompt_len = len(
self.tokenizer(
formatted_prompt, truncation=False, add_special_tokens=False
)["input_ids"]
)
completion_ids = full_ids[untruncated_prompt_len:]
# completion 预算 = 总预算 - 截断后 prompt 实长。配置校验
# max_prompt_length < max_length)保证它恒 > 0
completion_budget = self.max_length - len(prompt_ids)
completion_ids = completion_ids[:completion_budget]
all_input_ids.append(prompt_ids + completion_ids)
# prompt 位置标 -100:题目不产生 loss,只学解答
all_labels.append([IGNORE_INDEX] * len(prompt_ids) + completion_ids)
# 左 paddingbatch 内所有序列右对齐。纯 SFT 用右 padding 也行,但左 padding
# 让 trainer 能用一个标量 prompt_length 切 batchdocs/02 §2.4),且与
# 层 2+ 的生成场景(生成必须左 padding)统一,全项目只有一种 padding 约定
return {
"input_ids": _left_pad(all_input_ids, self.pad_token_id), # (B, T)
"attention_mask": _left_pad(
[[1] * len(ids) for ids in all_input_ids], 0
), # (B, T)
"labels": _left_pad(all_labels, IGNORE_INDEX), # (B, T)
}
def _left_pad(seqs: list[list[int]], pad_value: int) -> torch.Tensor:
"""把变长序列在左侧补齐成 (B, T) 张量,T = batch 内最大长度。"""
t_max = max(len(s) for s in seqs)
return torch.tensor(
[[pad_value] * (t_max - len(s)) + s for s in seqs], dtype=torch.long
)
+89
View File
@@ -0,0 +1,89 @@
"""MC 估计与 Dirichlet 贝叶斯平滑——论文 §3.2.2 式(4)(5)。
把 similarity.py 产出的软计数 k_sem(外部 teacher 信号)与学生自身的
chunk 置信度 π̄(内部先验)融合成有界目标 π̂ ∈ (0, 1],供层 5 的 chunk
损失当乘子:loss_c = −π̂ · mean(log p)。定理 4.1 三性质由此获得:
(a) π̂ 有界 → 无白盒式(2) 的梯度爆炸;(b) π̂ > 0 → k_sem=0 也不塌缩;
(c) 先验收缩 → 方差小于频率估计 k/N。
纯逻辑模块(CLAUDE.md §2):只依赖 torch,toy 张量本地 CPU 可测。
"""
import torch
def chunk_prior(log_probs: torch.Tensor) -> torch.Tensor:
"""式(4):π̄ = exp((1/C)·Σ_t log p_t)——学生对整个 chunk 的几何均值置信度。
C 个 token 概率的几何均值,充当式(5) 的贝叶斯先验:teacher 采样(k_sem)
是主信号,π̄ 只是"学生自己觉得这段有多稳"的地板,防 k_sem=0 时目标归零。
参数:
log_probs: 学生对 chunk 内各 token 的对数概率,shape (C,),值 ≤ 0。
(调用方从 log_softmax 后 gather 标签位置所得,层 5 负责。)
返回:
π̄,标量张量 shape ()、值 ∈ [1e-8, 1],**不带梯度**。
实现细节:
- log 域先均值再 exp:直接连乘 C=50 个小概率会下溢
(50 个 0.01 → 1e-100,超出 fp32 下限 ~1e-38),log 域安全。
- detach 命门(参考实现 distillation_trainer.py:2196 同):π̄ 是学生
自身概率的函数,若保留梯度,优化器会发现"压低自己的 chunk 概率
→ π̄→0 → π̂ 变小 → 损失权重变小"这条逃逸路径——恰在 k_sem=0
(teacher 否定)的 chunk 上最有利可图,这些 chunk 最先塌缩。
π̄ 只能当常数先验,不能当优化变量。锁死断言见
tests/test_estimator_detach.py。
- clamp 下限 1e-8:极端负的均值 exp 后可能下溢为 0,而定理 4.1(b)
的反塌缩要求 π̄ 严格为正。差异标注:参考实现 clamp(1e-8, 1.0),
上限实为冗余——log p ≤ 0 ⇒ mean ≤ 0 ⇒ exp ≤ 1,此处省去。
"""
if log_probs.numel() == 0:
raise ValueError("log_probs 为空:chunk 至少要含 1 个 token")
log_pi_bar = log_probs.detach().mean() # (C,) -> ()
return log_pi_bar.exp().clamp(min=1e-8)
def bayesian_target(
k_sem: float,
pi_bar: torch.Tensor,
n_rollouts: int,
alpha: float,
) -> torch.Tensor:
"""式(5):π̂ = (k_sem + α·π̄) / (N + α)——chunk 接受概率的贝叶斯估计。
等价凸组合视角(论文式10):
π̂ = N/(N+α) · (k_sem/N) + α/(N+α) · π̄
"teacher 频率估计""学生先验"的加权平均;默认 N=10、α=1 时权重
约 91% : 9%,teacher 主导,先验只兜底。
参数:
k_sem: 式(3) 的软匹配计数,∈ [0, N](aggregate_similarity 产出)。
pi_bar: 式(4) 的先验 π̄,标量张量(chunk_prior 产出)。
n_rollouts: teacher rollout 数 N。**必须等于算 k_sem 时的
len(teacher_rollouts)**——分子分母口径不一致会系统性偏移 π̂。
alpha: 先验强度 α ≥ 0。α=0 退化为频率估计 k/N(层 6 的
no_bayesian 消融,参考 config.py:315),失去定理 4.1(b) 保护。
返回:
π̂,标量张量 shape ()、值 ∈ [1e-8, 1],**不带梯度**。
实现细节:
- 差异标注:参考实现(distillation_trainer.py:2205)在此对 π̂ 整体
detach;我们的 π̄ 在 chunk_prior 内已 detach,此处的 detach 是
第二道防线——防止将来有人把带梯度的张量传进 pi_bar。
- clamp(1e-8, 1.0):下限防 α=0 且 k_sem=0 时 π̂=0(乘子归零则该
chunk 完全失去监督);上限防 pi_bar 越界传入时 π̂ 溢出概率语义。
"""
if n_rollouts < 1:
raise ValueError(f"n_rollouts 必须 ≥ 1,得到 {n_rollouts}")
if alpha < 0:
raise ValueError(f"alpha 必须 ≥ 0,得到 {alpha}")
if not 0.0 <= k_sem <= n_rollouts:
raise ValueError(
f"k_sem={k_sem} 越界 [0, {n_rollouts}]:检查是否与"
f" len(teacher_rollouts) 口径一致"
)
# 式(5): π̂ = (k_sem + α·π̄) / (N + α)
pi_hat = (k_sem + alpha * pi_bar) / (n_rollouts + alpha) # () -> ()
return pi_hat.clamp(1e-8, 1.0).detach()
+131
View File
@@ -0,0 +1,131 @@
"""语义相似度 φ 与 chunk 级聚合 k_sem——论文 §3.2.1 式(3)。
logit-free 的支点:学生 chunk 对不对,不再比 token 概率(层 2 白盒式(2)),
改比"学生 chunk 文本" vs "teacher rollout 文本"的语义相似度。φ 只依赖
文本本身,与两侧 tokenizer 无关,teacher 只需能吐文本(任何 API 均可)。
纯逻辑模块(CLAUDE.md §2):只依赖标准库,可脱离 torch 在本地 CPU 测试。
teacher rollout 怎么采出来是层 5 teacher.py 的事,本模块只吃现成字符串。
"""
from collections import Counter
def rouge1(hypothesis: str, reference: str) -> float:
"""ROUGE-1 F1(unigram 重叠率),φ 的候选度量之一(论文 §3.2.1)。
以词为单位(空白切分)统计两串的 unigram 重叠,算 F1。
词袋语义:只看"用了哪些词",不看词序——"a b" vs "b a" 得 1.0。
参数:
hypothesis: 学生 chunk 文本。
reference: teacher rollout 文本。
返回:
F1 ∈ [0, 1];任一侧无词(空串/纯空白)时为 0.0。
实现细节:
- 差异标注:参考实现(distillation_trainer.py:1670)用 set 去重后求交,
会把 "x x x x" vs "x" 判成满分 1.0;此处用 Counter 多重集
(ROUGE-1 标准定义),重复词按 min 计数配对,同例只得 0.4。
数学推理文本里重复 token(数字、"="、变量名)极常见,去重会失真。
- 差异标注:参考实现分母加 1e-8 防零除,代价是全同串 F1≈0.99999998
而非精确 1;此处 overlap==0 时提前返回,分母恒正,无需平滑。
"""
hyp_counts = Counter(hypothesis.split())
ref_counts = Counter(reference.split())
if not hyp_counts or not ref_counts:
return 0.0
# 多重集交:每个词按两侧出现次数的 min 配对
overlap = sum((hyp_counts & ref_counts).values())
if overlap == 0:
return 0.0
precision = overlap / sum(hyp_counts.values())
recall = overlap / sum(ref_counts.values())
return 2 * precision * recall / (precision + recall)
def edit_similarity(hypothesis: str, reference: str) -> float:
"""归一化编辑相似度 1 Levenshtein/max(m,n),论文 §5.1 的默认 φ。
以词为单位(空白切分)算 Levenshtein 距离(插入/删除/替换各计 1),
再归一化到 [0, 1] 取反。顺序敏感:"a b" vs "b a" 距离 2,相似度 0——
与 rouge1 的词袋语义形成互补。
参数:
hypothesis: 学生 chunk 文本。
reference: teacher rollout 文本。
返回:
相似度 ∈ [0, 1];两侧均空为 1.0(零距离),仅一侧空为 0.0(全删/全插)。
实现细节:
- 差异标注:参考实现(distillation_trainer.py:1682)吃 token id 列表,
相似度随 tokenizer 切法漂移,违背本层"文本是公共语言"的初衷
(docs/04 §2.1 坑①);此处吃 str、内部按词切,与 rouge1 统一口径。
- 两行滚动 DP(同参考实现):空间 O(n) 而非 O(m·n)。
"""
hyp_words = hypothesis.split()
ref_words = reference.split()
m, n = len(hyp_words), len(ref_words)
if m == 0 and n == 0:
return 1.0
if m == 0 or n == 0:
return 0.0
# prev[j] = 前一行的 dist(hyp[:i-1], ref[:j]);curr 原地滚动复用
prev = list(range(n + 1))
curr = [0] * (n + 1)
for i in range(1, m + 1):
curr[0] = i
for j in range(1, n + 1):
cost = 0 if hyp_words[i - 1] == ref_words[j - 1] else 1
curr[j] = min(
prev[j] + 1, # 删除 hyp[i-1]
curr[j - 1] + 1, # 插入 ref[j-1]
prev[j - 1] + cost, # 替换(相同则免费)
)
prev, curr = curr, prev
return 1.0 - prev[n] / max(m, n)
def phi(hypothesis: str, reference: str, metric: str = "edit_distance") -> float:
"""语义相似度 φ(y_c, ŷ_c) ∈ [0, 1],论文 §3.2.1 式(3) 的原子度量。
参数:
hypothesis: 学生 chunk 文本。
reference: teacher rollout 文本。
metric: "edit_distance"(默认)或 "rouge1"
差异标注:参考实现配置默认 rouge1(config.py:299),与论文 §5.1
的 edit_distance 背离;此处从论文。
返回:
相似度 ∈ [0, 1]。
"""
if metric == "edit_distance":
return edit_similarity(hypothesis, reference)
if metric == "rouge1":
return rouge1(hypothesis, reference)
raise ValueError(f"未知相似度度量: {metric!r}(可选 'edit_distance' / 'rouge1')")
def aggregate_similarity(
student_chunk: str,
teacher_rollouts: list[str],
metric: str = "edit_distance",
) -> float:
"""式(3):k_sem = Σ_{i=1}^{N} φ(y_c, ŷ_c^{(i)}),chunk 的语义匹配计数。
学生 chunk 与 N 个 teacher rollout 逐一算 φ 后求和。φ 连续,故 k_sem 是
[0, N] 上的实数——"软计数":k_sem≈N 意为学生这段与 teacher 高度一致,
k_sem≈0 意为 teacher 从不这么写。它是 π̂(式5)里唯一的外部 teacher 信号。
参数:
student_chunk: 学生 chunk 文本(C 个 token 解码所得)。
teacher_rollouts: N 段 teacher 续写文本,与学生 chunk 共享同一前缀
y_<c(对应关系由"同一前缀现场生成"保证,无需搜索匹配,docs/04 §1)。
metric: 传给 phi,默认 "edit_distance"
返回:
k_sem ∈ [0, N],N = len(teacher_rollouts)。空列表得 0.0(空和)。
"""
return sum(phi(student_chunk, r, metric) for r in teacher_rollouts)
+197
View File
@@ -0,0 +1,197 @@
"""teacher rollout 采样(IO 边缘,论文 §3.2.1)。
层 1 起步能力:给一批 prompt 批量生成解答,落盘 sha256 键的 JSONL 缓存
(键契约在 data.prompt_key 单点定义,本模块与 data.attach_teacher_completions
共用)。层 5 在此长出 chunk 前缀续写的 MC rollout 能力。
连接信息从 `.env` 读取(TEACHER_API_BASE / TEACHER_API_KEY / TEACHER_MODEL),
密钥永不出现在代码与配置类里。
缓存即断点:生成过程逐条追加写盘,任何中断(网络、Ctrl-C、单条失败)后
重跑同一命令,已完成的条目自动跳过——API 花的钱不会白花。
"""
from __future__ import annotations
import json
import os
import re
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
from dotenv import load_dotenv
from openai import OpenAI
from ars_opd.configs import TeacherGenConfig
from ars_opd.data import prompt_key
Messages = list[dict[str, str]]
def _load_teacher_env(env_file: str | None = None) -> tuple[str, str, str]:
"""从 .env(及进程环境)读取 API 连接三元组,缺一项都显式报错。"""
load_dotenv(env_file)
values = {}
for name in ("TEACHER_API_BASE", "TEACHER_API_KEY", "TEACHER_MODEL"):
value = os.environ.get(name, "").strip()
if not value:
raise ValueError(
f"环境变量 {name} 未设置。复制 .env.example 为 .env 并填入真实值。"
)
values[name] = value
return (
values["TEACHER_API_BASE"],
values["TEACHER_API_KEY"],
values["TEACHER_MODEL"],
)
def _strip_leading_think(text: str) -> str:
"""剥离 content 开头的 <think>...</think> 段(M3 等 reasoning 模型会内联思考)。
只剥开头一段:解答正文里若出现字面 "<think>" 字样(例如题目在讨论标签本身),
不应被误删。
"""
return re.sub(r"^\s*<think>.*?</think>\s*", "", text, count=1, flags=re.DOTALL)
class TeacherClient:
"""OpenAI 兼容的 teacher 客户端:单条生成 + 采样参数收口。
差异标注:参考实现是 OpenRouter 专用客户端(带其私有请求头与站点字段);
我们用通用 OpenAI 客户端 + base_url 配置驱动,任何兼容网关(new-api、
vLLM serve、官方 API)都无需改代码。
测试注入口:传入 client/model 可绕过 .env 与真实网络(见 tests/test_teacher.py)。
"""
def __init__(
self,
gen_config: TeacherGenConfig,
client: OpenAI | None = None,
model: str | None = None,
) -> None:
self.cfg = gen_config
if client is None:
base, key, env_model = _load_teacher_env()
client = OpenAI(
base_url=base, api_key=key, max_retries=gen_config.max_retries
)
model = model or env_model
if model is None:
raise ValueError("注入 client 时必须同时指定 model")
self.client = client
self.model = model
def generate(self, messages: Messages) -> str:
"""对单条 promptmessages 列表,末轮为 user)生成解答文本。
返回剥离思考段、去首尾空白后的解答。空解答直接报错——空字符串写进
缓存会在训练时变成全 -100 的空样本(trainer 会炸,但应在这里更早炸)。
"""
if self.cfg.system_prompt is not None:
messages = [
{"role": "system", "content": self.cfg.system_prompt}
] + messages
resp = self.client.chat.completions.create(
model=self.model,
messages=messages,
temperature=self.cfg.temperature,
top_p=self.cfg.top_p,
max_tokens=self.cfg.max_tokens,
)
content = resp.choices[0].message.content or ""
if self.cfg.strip_think:
content = _strip_leading_think(content)
content = content.strip()
if not content:
raise ValueError(
"teacher 返回空解答(可能:max_tokens 太小把思考截断在半途,"
"或模型拒答)。该条不会入缓存。"
)
return content
def generate_completions(
prompts: list[Messages],
cache_path: str,
teacher: TeacherClient,
) -> None:
"""批量生成解答并追加写入 JSONL 缓存(每行 {"key", "completion", "preview"})。
- 已在缓存中的键直接跳过(断点续传);
- 并发线程池执行,每完成一条立即写盘并 flush(中断不丢已完成的结果);
- 单条失败不中断其余任务(并发中的兄弟请求已经花了钱,先让它们落盘),
全部结束后若有失败则汇总显式报错——重跑即续传,绝不静默缺数据。
"""
path = Path(cache_path)
path.parent.mkdir(parents=True, exist_ok=True)
done_keys = _cached_keys(path)
todo = [(prompt_key(p), p) for p in prompts]
todo = [(k, p) for k, p in todo if k not in done_keys]
print(
f"[teacher] 共 {len(prompts)} 条:缓存命中 {len(prompts) - len(todo)}"
f"待生成 {len(todo)},并发 {teacher.cfg.concurrency}",
flush=True,
)
if not todo:
return
failures: list[tuple[str, str]] = []
finished = 0
start = time.monotonic()
# 写盘收口在主线程(as_completed 消费端),工作线程只跑网络请求——
# 多线程同写一个文件句柄会交错损坏 JSONL
with open(path, "a", encoding="utf-8") as f:
with ThreadPoolExecutor(max_workers=teacher.cfg.concurrency) as pool:
futures = {pool.submit(teacher.generate, p): (k, p) for k, p in todo}
for fut in as_completed(futures):
key, p = futures[fut]
try:
completion = fut.result()
except Exception as e: # noqa: BLE001 —— 收集后统一显式报错,非静默吞错
failures.append((key, repr(e)))
continue
finally:
finished += 1
if finished % 20 == 0 or finished == len(todo):
elapsed = time.monotonic() - start
rate = finished / elapsed * 60 # 条/分
eta = (len(todo) - finished) / rate if rate > 0 else 0
print(
f"[teacher] {finished}/{len(todo)} 完成 | "
f"{rate:.1f} 条/分 | 已用 {elapsed / 60:.1f} 分 | "
f"预计剩余 {eta:.0f}",
flush=True,
)
record = {
"key": key,
"completion": completion,
# preview 仅供人工抽查缓存文件,消费端(attach)只认 key/completion
"preview": p[-1]["content"][:80],
}
f.write(json.dumps(record, ensure_ascii=False) + "\n")
f.flush()
if failures:
examples = "; ".join(f"{k[:12]}…: {err}" for k, err in failures[:3])
raise RuntimeError(
f"{len(failures)}/{len(todo)} 条生成失败(成功的已入缓存,重跑本命令"
f"即断点续传)。前几条错误:{examples}"
)
def _cached_keys(path: Path) -> set[str]:
"""读取缓存中已有的键集合;文件不存在视为空缓存(首跑)。"""
if not path.exists():
return set()
keys = set()
with open(path, encoding="utf-8") as f:
for line_no, line in enumerate(f, 1):
if not line.strip():
continue
rec = json.loads(line) # 坏行直接炸:缓存损坏必须暴露,不能悄悄重新生成
keys.add(rec["key"])
return keys
+423
View File
@@ -0,0 +1,423 @@
"""训练编排(IO 边缘)。层 1 形态:掩码 SFT 损失 + 最小 HF Trainer 子类。
论文锚点:§3.1 式(1) 的标准交叉熵 SFT(监督目标是 teacher rollout,见 docs/02 §1)。
层 5 会在此模块长出式(8) 的 chunk 蒸馏损失与 KL 锚定;届时式(1) 路径保留为基线。
结构:损失算法收在纯张量函数 `sft_loss`(可用 toy 张量在 CPU 单测),
`SFTTrainer` 只做接线——模型前向、调 `sft_loss`、把指标并进日志,
其余一切(优化器、调度、DDP、checkpoint 保存)原样继承 HF Trainer。
"""
from __future__ import annotations
from typing import Any
import torch
import torch.nn.functional as F
from transformers import Trainer
from ars_opd.data import IGNORE_INDEX
# ---------------------------------------------------------------------------
# 损失算法(纯张量函数,对拍 distillation_trainer.py:751-759 + 2776-2874
# ---------------------------------------------------------------------------
def compute_prompt_length(attention_mask: torch.Tensor, labels: torch.Tensor) -> int:
"""batch 级 prompt 边界 = batch 内最短的"有效长度 - completion 长度"
参数(B = batchT = padding 后长度):
- attention_mask: (B, T),左 padding 位置为 0
- labels: (B, T)completion 位置为 token id,其余为 -100
返回标量 pl。非显然约束:取 batch **最小值**是为了不切掉任何行的
completion token——左 padding 下所有序列右对齐,第 r 行的 completion 起点
索引是 T - comp_r ≥ prompt_r ≥ min,故切片 [pl:] 必然包含全部 completion
代价是长 prompt 行会漏进一些 prompt token,由 sft_loss 用 labels 重掩码兜住。
"""
full_lengths = attention_mask.sum(dim=1) # (B,) 每行非 padding 的 token 数
completion_lengths = (labels != IGNORE_INDEX).sum(dim=1) # (B,)
return int((full_lengths - completion_lengths).min().item())
def sft_loss(
logits: torch.Tensor,
input_ids: torch.Tensor,
labels: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, int]:
"""式(1)completion 位置上的移位交叉熵。
参数(B = batchT = padding 后长度,V = 词表大小):
- logits: (B, T, V) 模型对全序列的输出
- input_ids / labels / attention_mask: (B, T)SFTCollator 的产物
返回 (标量 loss, 本 batch 有效 completion token 数)。
切片几何(docs/02 §2.4 的四行,此处为权威实现):
位置 t-1 的 logit 预测位置 t 的 token,故 logits 取 [pl-1, T-1) 、
targets 取 [pl, T),两段长度同为 T-pl,逐位对齐。
"""
pl = compute_prompt_length(attention_mask, labels)
if pl < 1:
# 需要位置 pl-1 的 logit 存在;pl=0 意味着某行完全没有 prompt token
# 数据管线出了问题(collator 保证 prompt 至少含模板 token
raise ValueError(f"prompt_length={pl} < 1,存在无 prompt 的数据行")
shifted_logits = logits[:, pl - 1 : -1, :] # (B, T, V) -> (B, T-pl, V)
targets = input_ids[:, pl:].clone() # (B, T-pl)clone: 下面要原地改写
# 重掩码:切片里漏进的 prompt token(长 prompt 行)与 padding 全部置 -100。
# labels 是权威掩码,切片几何只是省算力——正确性完全由这一步保证
invalid = labels[:, pl:] == IGNORE_INDEX # (B, T-pl)
targets[invalid] = IGNORE_INDEX
num_valid = int((~invalid).sum().item())
if num_valid == 0:
# 差异标注:参考实现(trainer:2810-2812)对非有限 loss 返回零梯度标量静默
# 继续;我们显式报错——collator 已挡下 prompt-only 行,走到这里仍全被
# 掩码只可能是数据坏了(如 teacher 返回空解答),必须暴露而非跳过
raise ValueError(
"本 batch 没有任何有效 completion token(全被 -100 掩码)。"
"检查 teacher 解答是否为空、completion 预算是否被截光。"
)
vocab = shifted_logits.shape[-1]
loss = F.cross_entropy(
shifted_logits.reshape(-1, vocab), # (B*(T-pl), V)
targets.reshape(-1), # (B*(T-pl),)
ignore_index=IGNORE_INDEX,
)
return loss, num_valid
def token_divergence(
student_logits: torch.Tensor,
teacher_logits: torch.Tensor,
labels: torch.Tensor,
beta: float = 1.0,
temperature: float = 1.0,
) -> tuple[torch.Tensor, int]:
"""式(2)completion 位置上的 token 级(广义)KL 散度,全词表精确。
论文锚点:§3.1 式(2) L = E_{y~π_θ}[Σ_t KL(π_θ(·|y_<t,x) ‖ π_T(·|y_<t,x))]。
本函数只管"给定两组**已对齐**的 logits,算散度标量"——on-policy 采样(y~π_θ)
与移位对齐([pl-1:-1])由调用方(DistillTrainer, U4)负责,复用 sft_loss 同一套
切片几何。故这里不吃 input_ids:散度是分布对分布,不需要目标 token,labels
仅用于定位有效位置。
参数(B=batch,T=已移位对齐长度,V=词表大小):
- student_logits: (B, T, V)student 前向输出(带梯度)
- teacher_logits: (B, T, V)teacher no_grad 前向输出(无梯度)
- labels: (B, T)completion 位为 token id、其余为 IGNORE_INDEX;只做有效位掩码
- beta: KL 方向(docs/03 §2.3 三副面孔)。0=前向 KL(π_T‖π_θ)、1=反向 KL(π_θ‖π_T)=
式(2)、(0,1)=JSD 插值
- temperature: softmax 前除进两侧 logits 的温度(§2.3),软化/锐化分布
返回 (per-token mean 散度标量, 有效 token 数)。
差异标注:参考实现(DT:2408-2491)含 top-k 稀疏 + 尾桶快路,我们只保留全词表这
一条精确路径(docs/03 §4 删除清单:本地同 tokenizer teacher 放得下)。因全词表
log_softmax 对有限 logits 恒有限,也不需要参考在 -inf 支持集上的 nan_to_num 兜底。
"""
if not 0.0 <= beta <= 1.0:
raise ValueError(f"beta 必须在 [0,1]0=前向/1=反向/中间=JSD),收到 {beta}")
# 温度除进 logits、softmax 之前(§2.3):调分布形状,不是等比缩概率
student_logits = student_logits / temperature
teacher_logits = teacher_logits / temperature
# 全程 log 域(§2.3 数值稳定性)。log_probs 两侧都要;probs 按分支只算用得上的
# 那一份——(B,T,V) 在真实规模下每份 ~2.5G(§5 显存账),默认 β=1 热路径只需
# student_probs,不materialize teacher_probs
student_log_probs = F.log_softmax(student_logits, dim=-1) # (B, T, V)
teacher_log_probs = F.log_softmax(teacher_logits, dim=-1) # (B, T, V)
if beta == 1.0:
# 反向 KL(π_θ‖π_T) = Σ_v π_θ (log π_θ log π_T) —— 式(2)
# 非显然约束(§4.1 梯度爆炸源):对 student logit 的梯度含 π_θ·log(π_θ/π_T),
# on-policy 采到 teacher 眼中烂 token(π_T→0)时 log 比值→∞,单 token 梯度可
# 炸掉整个 batch。这正是层 2 要亲眼观察、层 5 用有界乘子 π̂ 替换的病灶
per_token = (
student_log_probs.exp() * (student_log_probs - teacher_log_probs)
).sum(-1) # (B, T, V) -> (B, T)
elif beta == 0.0:
# 前向 KL(π_T‖π_θ) = Σ_v π_T (log π_T log π_θ)
per_token = (
teacher_log_probs.exp() * (teacher_log_probs - student_log_probs)
).sum(-1)
else:
# JSD 插值:m = (1−β)π_θ + β π_T;β·KL(π_T‖m) + (1−β)·KL(π_θ‖m)
student_probs = student_log_probs.exp()
teacher_probs = teacher_log_probs.exp()
mixture = (1.0 - beta) * student_probs + beta * teacher_probs
# clamp_min(tiny) 防 log0(§2.3):混合概率理论上恒正,此处是浮点下溢兜底
log_mixture = mixture.clamp_min(torch.finfo(mixture.dtype).tiny).log()
kl_teacher = (teacher_probs * (teacher_log_probs - log_mixture)).sum(-1)
kl_student = (student_probs * (student_log_probs - log_mixture)).sum(-1)
per_token = beta * kl_teacher + (1.0 - beta) * kl_student
# 掩码 + per-token mean(docs/03 §2.3reduction 实义为 sum/有效token数,
# 与层 1 sft_loss 同尺度,两条 loss 曲线才可比)
mask = labels != IGNORE_INDEX # (B, T)
num_valid = int(mask.sum().item())
if num_valid == 0:
# 与 sft_loss 同纪律:走到这里全被掩码只可能是数据/生成坏了,显式报错不静默
raise ValueError(
"本 batch 没有任何有效 completion token(全被 -100 掩码)。"
"检查 on-policy 生成是否产出了空 completion。"
)
loss = per_token[mask].sum() / num_valid
return loss, num_valid
def build_generated_batch(
prompts: torch.Tensor,
prompt_attention_mask: torch.Tensor,
gen_output: torch.Tensor,
eos_token_id: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""把 model.generate 的输出重建成 (input_ids, attention_mask, labels)。
对应 docs/03 §2.5"生成结果重建 input_ids/labels 写回"。on-policy 下 completion
不来自数据、而是 student 现场生成,故 labels 也在生成后现造:prompt 段全 -100、
生成段有效处填 token id,供 token_divergence 掩码。
参数(B=batchP=prompt padding 长度,G=本 batch 最大生成长度):
- prompts: (B, P) 左 padding 的 promptSFTCollator prompt_only 产物)
- prompt_attention_mask: (B, P) prompt 左 padding 位为 0
- gen_output: (B, P+G) model.generate 输出(前 P 列即 prompts,后 G 列是生成)
- eos_token_id: 生成终止符 id
返回 (input_ids (B,P+G), attention_mask (B,P+G), labels (B,P+G))。
非显然约束(生成段的右 padding 掩码):generate 对提前结束的序列在右侧补
padding 到 batch 最大长度。"首个 eos(含)之前有效"用 `cumsum - self == 0`
实现——它精确保留到首个 eos、屏蔽其后一切(无论其后是 eos 还是 pad,也无论
pad_token 是否等于 eos),避免"pad==eos 时把补位当解答""pad!=eos 时漏掉补位"
两种静默错误。无 eos(撞 max_new_tokens)则整段生成全有效。
"""
b, p = prompts.shape
gen_tokens = gen_output[:, p:] # (B, G) 纯生成段
is_eos = gen_tokens == eos_token_id # (B, G)
# cumsum - self:截至本位、其**之前**出现过的 eos 数;==0 即"首个 eos 及之前"
gen_valid = (is_eos.cumsum(dim=1) - is_eos.long()) == 0 # (B, G) bool
input_ids = gen_output
attention_mask = torch.cat(
[prompt_attention_mask, gen_valid.long()], dim=1
) # (B, P+G)
prompt_labels = torch.full(
(b, p), IGNORE_INDEX, dtype=torch.long, device=prompts.device
)
gen_labels = torch.where(
gen_valid, gen_tokens, torch.full_like(gen_tokens, IGNORE_INDEX)
) # 生成段:有效处填 token id,其余 -100
labels = torch.cat([prompt_labels, gen_labels], dim=1) # (B, P+G)
return input_ids, attention_mask, labels
# ---------------------------------------------------------------------------
# Trainer 接线
# ---------------------------------------------------------------------------
class SFTTrainer(Trainer):
"""最小 SFT Trainer:只重写 compute_loss 与 log,其余全部继承 HF Trainer。
用法(见 scripts/ 训练脚本):与 HF Trainer 完全同参构造,
data_collator 传 SFTCollatortrain_dataset 传 load_sft_dataset 的产物。
"""
def __init__(self, *args: Any, **kwargs: Any) -> None:
super().__init__(*args, **kwargs)
# 关 dropout(参考实现同款,docs/02 §3 保留清单):层 2+ 的蒸馏要求
# student/ref 两次前向可比,前向必须确定;Qwen3 默认无 dropout
# 此处是无害的统一前置
for module in self.model.modules():
if isinstance(module, torch.nn.Dropout):
module.p = 0.0
# 非显然约束:显式退出 HF 的新式梯度累积契约(v5 trainer.py:1977 文档
# 原话:"If you are not using num_items_in_batch ... overwrite
# self.model_accepts_loss_kwargs to False")。新式契约要求返回
# sum/全局token数,且依赖基类 compute_loss 尾部的 ×world_size 补偿
# trainer.py:2028)——我们整体重写了 compute_loss,那段补偿不会执行,
# 曾致 loss 与梯度 ÷4(2026-07-18,第二幕;第一幕是返回裸 mean 被 ×8,
# 全程记录见 docs/02 §5)。退出后回到经典契约:返回本微批 mean,
# Trainer 负责 ÷累积步数,日志跨卡平均,版本稳定
self.model_accepts_loss_kwargs = False
self._token_counts: list[int] = []
def compute_loss(
self,
model: Any,
inputs: dict[str, torch.Tensor],
return_outputs: bool = False,
num_items_in_batch: int | None = None,
) -> torch.Tensor | tuple[torch.Tensor, Any]:
# 非显然约束:不把 labels 传给模型前向。HF 模型收到 labels 会自己算
# "全序列移位 CE"并放进 outputs.loss,那会绕过 batch-min 切片与重掩码;
# 损失的权威实现只能有 sft_loss 一处
outputs = model(
input_ids=inputs["input_ids"],
attention_mask=inputs["attention_mask"],
)
loss, num_tokens = sft_loss(
outputs.logits,
inputs["input_ids"],
inputs["labels"],
inputs["attention_mask"],
)
# num_items_in_batch 有意忽略:已在 __init__ 退出新式契约(见彼处注释),
# 本函数返回微批 mean,÷累积步数由 Trainer.training_step 负责。
# 代价是微批按相同权重而非 token 数加权(百分之几的偏差,与参考实现同行为)
self._token_counts.append(num_tokens)
return (loss, outputs) if return_outputs else loss
def log(self, logs: dict[str, float], start_time: float | None = None) -> None:
"""在 HF 的常规日志里并入每步有效 token 数均值。
sanity run 时盯这个数:若它远小于预期(≈batch 内解答总长),说明掩码
把 completion 也吞了——loss 曲线看不出这种错,token 数看得出。
"""
if self._token_counts:
logs["sft/num_tokens_per_step"] = sum(self._token_counts) / len(
self._token_counts
)
self._token_counts = []
super().log(logs, start_time)
class DistillTrainer(Trainer):
"""层 2 white-box OPDon-policy 生成 → teacher no_grad 前向 → token 级反向 KL。
论文锚点:§3.1 式(2)。每个微批现场编排三步(docs/03 §2.2 主流程的精简版,
已按 §4 删除 buffer/稀疏路径/off-policy 抽签):
1. student 采样生成轨迹 y~π_θ(on-policyno_grad——只采样不回传);
2. student 带梯度前向 + teacher no_grad 前向,得两份全词表 logits;
3. token_divergence 算 KL,梯度只经 student 那一路。
与 SFTTrainer 的关系:损失几何(移位对齐、batch-min prompt_length、labels 重
掩码)完全同款,直接复用 compute_prompt_length;唯一差别是"completion 从哪来"
——SFT 读数据缓存,这里 student 现场生成。
构造(见 scripts/train_whitebox.py):像 HF Trainer 一样传 model(student)/args/
train_dataset/data_collator(prompt_only 的 SFTCollator),另用关键字传 teacher_model
与 teacher_tokenizer,以及蒸馏超参。
"""
def __init__(
self,
*args: Any,
teacher_model: Any,
teacher_tokenizer: Any,
beta: float = 1.0,
kl_temperature: float = 1.0,
gen_temperature: float = 1.0,
gen_top_p: float = 1.0,
max_new_tokens: int = 1024,
**kwargs: Any,
) -> None:
super().__init__(*args, **kwargs)
# student tokenizer 从 collator 取(prompt_only collator 必持有它)
student_tokenizer = self.data_collator.tokenizer
# 白盒前提校验(docs/03 §1、§2.6):KL 逐词表位对齐,同 tokenizer 才有意义。
# 构造时就炸——不像参考实现(DT:2876)拖到第一步前向才炸
if teacher_tokenizer.get_vocab() != student_tokenizer.get_vocab():
raise ValueError(
"teacher 与 student 的 tokenizer 词表不一致——白盒 KL 要求逐词表位"
"对应(docs/03 §1)。请换用与 student 同 tokenizer 的 teacher。"
)
self._student_tokenizer = student_tokenizer
# teachereval + 冻结参数 + 每进程一份副本(DDP 每卡一份,docs/03 §2.6)。
# 设备迁移推迟到 compute_loss——此刻 student 还没被 Trainer 放到卡上
self.teacher = teacher_model.eval()
for param in self.teacher.parameters():
param.requires_grad_(False)
self.beta = beta
self.kl_temperature = kl_temperature
self.gen_temperature = gen_temperature
self.gen_top_p = gen_top_p
self.max_new_tokens = max_new_tokens
# 关 dropout:生成、student 前向、teacher 前向三者须确定可比(同 SFTTrainer)
for module in self.model.modules():
if isinstance(module, torch.nn.Dropout):
module.p = 0.0
# 同层 1:显式退出 HF 新式梯度累积契约(docs/02 §5trainer.py:1977
self.model_accepts_loss_kwargs = False
self._token_counts: list[int] = []
def compute_loss(
self,
model: Any,
inputs: dict[str, torch.Tensor],
return_outputs: bool = False,
num_items_in_batch: int | None = None,
) -> torch.Tensor | tuple[torch.Tensor, Any]:
prompts = inputs["prompts"]
prompt_attention_mask = inputs["prompt_attention_mask"]
# teacher 迁到 student 所在卡(一次性;.to 幂等,后续步是 no-op)
if self.teacher.device != prompts.device:
self.teacher = self.teacher.to(prompts.device)
# 1. on-policy 生成。no_gradGKD 标准做法是"采样一次、再 teacher-forcing
# 前向算分布",梯度经第 3 步的前向回传,不经采样本身。DDP 下须用 unwrap
# 后的模型(DDP 包装体不暴露 generate
unwrapped_model = self.accelerator.unwrap_model(model)
with torch.no_grad():
gen_output = unwrapped_model.generate(
input_ids=prompts,
attention_mask=prompt_attention_mask,
max_new_tokens=self.max_new_tokens,
do_sample=True,
temperature=self.gen_temperature,
top_p=self.gen_top_p,
pad_token_id=self._student_tokenizer.pad_token_id,
eos_token_id=self._student_tokenizer.eos_token_id,
)
# 2. 重建 input_ids/attention_mask/labelsprompt 段 -100、生成段有效处填 id)
input_ids, attention_mask, labels = build_generated_batch(
prompts,
prompt_attention_mask,
gen_output,
self._student_tokenizer.eos_token_id,
)
# 3. student 带梯度前向 + teacher no_grad 前向(两份全词表 logits)。
# 非显然约束:teacher no_grad 免掉的是它内部几十层激活的反向图(②),
# 但输出 logits(①)仍占满 (B,L,V) 显存——两份都要算进 §5 显存账
student_outputs = model(input_ids=input_ids, attention_mask=attention_mask)
with torch.no_grad():
teacher_logits = self.teacher(
input_ids=input_ids, attention_mask=attention_mask
).logits
# 4. 移位对齐(与 sft_loss 同款几何,docs/02 §2.4)→ token_divergence。
# 切片漏进的 prompt token 与生成段右 padding 由 labels 重掩码兜住
pl = compute_prompt_length(attention_mask, labels)
if pl < 1:
raise ValueError(f"prompt_length={pl} < 1,生成批次存在无 prompt 的行")
loss, num_tokens = token_divergence(
student_outputs.logits[:, pl - 1 : -1, :],
teacher_logits[:, pl - 1 : -1, :],
labels[:, pl:],
beta=self.beta,
temperature=self.kl_temperature,
)
self._token_counts.append(num_tokens)
return (loss, student_outputs) if return_outputs else loss
def log(self, logs: dict[str, float], start_time: float | None = None) -> None:
"""并入每步的生成 token 数均值:远程冒烟盯它,骤降=生成塌成空串。"""
if self._token_counts:
logs["distill/num_gen_tokens_per_step"] = sum(self._token_counts) / len(
self._token_counts
)
self._token_counts = []
super().log(logs, start_time)
+28 -6
View File
@@ -1,6 +1,25 @@
# 00 · 分层重构路线图
> 原则:按论文概念的依赖顺序逐层重建,每层完成后代码可运行、可验证。学习路径 = 提交历史。
> 后面各层不提前细化——细节在进入该层时随章节文档长出来(依据见 `appendix-claudemd-decisions.md` 的延迟接入哲学)。
## 当前进度(存档点)
> 每次断点(层完成/工作暂停)更新此节。恢复上下文时:读 CLAUDE.md → 本节 → 对应章节文档。
- **日期**: 2026-07-19
- **当前层**: 层 3(相似度 φ + MC 估计器),**docs/04 已精读**(用户已理解:logit-free 支点=比文本非比 logits、k_sem 是外部 teacher 信号 π̄ 只是防塌缩地板、detach 命门、chunk-vs-prefix 对应靠"同一前缀现场生成 teacher 续写"非搜索匹配);**下一步开写 E1**。层 3 全本地 CPU 纯逻辑,不碰 GPU/远程。E1 similarity.py(式3 φ+k_sem,两度量统一到词级文本、默认 edit_distance 对齐论文§5.1)、E2 estimator.py(式4 π̄ 几何均值+detach、式5 π̂ 贝叶斯凸组合)、E3 test_estimator_detach.py 接真实现。设计取舍已定见 docs/04 §3。层 2 ✅ 已关账
- **层 2 代码构成**: U1 DistillConfigconfigs.py,两温度分名/三处刻意缺席/max_grad_norm 显式化);U2 token_divergence + 梯度爆炸单测(trainer.py,全词表 KL/JSD);U3 SFTCollator prompt_only 模式(data.py);U4 DistillTrainer + build_generated_batchtrainer.py,生成→双前向→divergence);U5 train_whitebox.py/.shfull/sanity/noclip 三模式)。附带:load_sft_dataset 毛刺已磨平(改吃散装参数);hf-mirror 不代理 Xet CAS → HF_HUB_DISABLE_XET=1
- **层 2 关账判据全过**(详见 docs/03 §6.1 实证): 66 单测全绿;B=4/T=2048 **实测不 OOM**(§5 估算成立);两次远程跑(sanity 裁到 1.0 / noclip ≈关裁剪+lr5×)均平稳、不 NaN、生成不塌(num_gen ~1900-4096)、loss 0.35→0.21 下降;checkpoint 存下
- **层 2 关键发现(勘误"预期见毛刺")**: §4.1 梯度爆炸是真机制(U2 单测坐实单 token π_T→0 暴涨),但真实训练**高度阻尼**——同门 teacher + per-token mean 摊平,batch 级 grad_norm 峰值仅 ~14 且只降不升,关裁剪也不炸。启示:`grad_norm` 日志是**裁剪前**值(14→2 那串即爆炸证据,被 HF 默认 max_grad_norm=1.0 静默压平,现已显式化);层 5 有界乘子 π̂ 真正杀手锏是 **logit-free**(白盒的同 tokenizer 约束把你锁在温和区间)
- **层 2 接口回看**(§6.5 每层必做,全部判"深",无需返工的毛刺): DistillConfig/token_divergence/build_generated_batch/load_sft_dataset(已修) 接口均远简于实现。三条**记录不返工**的小注:① DistillTrainer 从 self.data_collator.tokenizer 取 student tokenizer(隐式耦合,但省一个冗余参数,可接受);② SFTCollator 名字略超范(现含 prompt_only 非 SFT 模式),rename 的 churn 不值;③ token_divergence 的 labels 仅作掩码非目标(已在 docstring 标注)
- **层 0**: ✅ 已关账(2026-07-18
- **层 1**: ✅ 已关账(2026-07-18)。判据全过:正本缓存 sha `33deb18c…`(1000 条,键唯一,think 残留 0,仅 2 条硬题截断);正式 1 epoch loss 0.94→0.6056s/16 步);checkpoint 生成通顺(`/data/zym/outputs/sft_qwen3-0.6b_dapo1k`);接口回看完成(全部模块判"深";毛刺记录:load_sft_dataset 吃整个 SFTConfig 迫使诊断脚本填假 output_dir,层 2 第二消费方出现时定夺)
- **层 1 疤痕档案**(详见 docs/02 §2.6/§5 勘误): ① 显存大头是 (B,T,V) logits 链与激活(正比 B×T,与参数量无关),B=8 曾爆 80G;② HF 梯度累积契约两幕剧(×8 → ÷4),终解 = model_accepts_loss_kwargs=False 退出新式契约;③ 缓存正本纪律:本地生成一次、单向 scp 分发、sha256 对账,两侧独立生成曾花双份钱且内容漂移
- **诊断工具箱**: scripts/diag_collator.py(对齐链逐环)、diag_loss_probe.py(预训练 CE 基准 0.85)、diag_generate.py(生成质量)——层 2+ 数值异常照此三板斧
- **已完成学习**: 第一/二/三章全部精讲(第三章含 standard 路径三大反直觉点、两温度、抉择原则、§4.1 梯度爆炸机制+实证)
- **未精讲的文档账**: docs/01 的 §3.7(KL 锚三处实现差异)、§3.8(论文外稳定器)、§4(训练步流程走读)
- **远程磁盘备忘**: 根分区 100% 的结构性原因是 `/root/zym`(507G 历史工作区)压在根分区,建议择期整体搬迁 `/data`;临时缓解 = 清 `/tmp/pip-unpack-*`、旧 tar.gz、journal。所有新增写盘已改道 `/data/zym`
## 分层计划
@@ -20,9 +39,9 @@
|------|------|------|
| `00-roadmap.md` | 本文 | ✅ |
| `01-paper-code-map.md` | 论文 §3-§4 精读 + 参考实现全景解剖 + 概念↔代码对照表 | ✅ |
| `02-sft-baseline.md` | 层 1SFT 与数据管线 | |
| `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性 | ⬜ |
| `04-mc-estimator.md` | 层 3:MC 估计 + 贝叶斯平滑 | ⬜ |
| `02-sft-baseline.md` | 层 1SFT 与数据管线 | ✅ 待读 |
| `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性(含 §6.1 远程实证) | ✅ 已关账 |
| `04-mc-estimator.md` | 层 3语义相似度 φ + MC 估计 + 贝叶斯平滑(式3/4/5 | ✅ 待读 |
| `05-entropy-chunking.md` | 层 4:熵调度 | ⬜ |
| `06-omniopd-full.md` | 层 5:完整损失与 teacher 客户端 | ⬜ |
| `07-eval-ablation.md` | 层 6:评测与消融 | ⬜ |
@@ -35,6 +54,9 @@
| rollout 数 N | 10 | 10 | 每个 chunk 的 teacher MC 采样数(§4.2 证明 N=10 是甜点) |
| chunk 长度 C | 50 | 50 | token 数 |
| 先验强度 α | 1.0 | 1.0 | `chunk_alpha`Dirichlet 平滑 |
| 相似度 φ | ROUGE-1 | ROUGE-1 | 备选 edit_distance |
| Student | Qwen3-8B 级 | Qwen3-0.6B | 跑通优先 |
| Teacher | Qwen3.5-397B / Claude / Gemini | DeepSeek 或 MiniMaxOpenAI 兼容) | logit-free 主路径 |
| 相似度 φ | **edit_distance**(§5.1 | edit_distance | ⚠️ 代码默认 rouge1config L298)与论文默认背离,须显式指定 |
| KL 锚权重 β | **0.1**(§5.1) | 0.1 | ⚠️ 代码默认 `mc_kl_weight=0` 与论文背离,须显式指定 |
| 训练数据 | DAPO-Math-17Kprompt-only | 同(层 1 先抽 ~1k 子集控制 API 成本) | 一份数据服务层 1-6;学生升到 1.7B 后可直接对表论文 Table 1 |
| Student | Qwen3-1.7B / 4B | Qwen3-0.6B | 跑通优先;升级 1.7B 即可与论文对比 |
| Teacher | Qwen3-32B / Claude-4.5-Haiku / Gemini-2.5-Flash | MiniMax-M3(自建 new-api 网关,OpenAI 兼容;2026-07-18 由 DeepSeek 改定) | logit-free 主路径;M3 是 reasoning 模型,思考段入库前剥离(teacher.py strip_think |
| SFT 基线定义 | teacher rollout 上的离线蒸馏(非人写答案) | 同 | 对应参考实现 `_generate_teacher_completions` + JSONL 缓存路径 |
+87
View File
@@ -0,0 +1,87 @@
# 02 · 层 1:SFT 基线与数据管线
> 本章目标:搞懂"SFT 基线"在本项目中的确切含义,解剖参考实现的数据管线与掩码交叉熵,然后动手建起 `ars_opd` 的第一批模块。行号缩写:`trainer:` = `references/ars-opd/trl/trl/experimental/distillation/distillation_trainer.py``train_dist:` = `references/ars-opd/train_distillation.py`。
## 1. 论文侧:SFT 基线的确切定义
式(1)是标准交叉熵,但注意论文 §5.1 对基线的定义:**SFT = 在 teacher rollout 上的离线蒸馏**Kim & Rush 2016 式 sequence-level distillation),不是"在人写答案上训练"。流程:拿 DAPO-Math-17K 的题目 → teacher 生成解答 → 学生对解答做掩码交叉熵。§4.4 的 Thm 4.4 顺带证明了这种 SFT **不具有** tokenizer/风格不变性(损失绑死 teacher 的具体 token 选择)——这是它后面被 OmniOPD 超越的理论伏笔。
本层设定(roadmap 已定):student Qwen3-0.6Bteacher MiniMax-M3(自建 new-api 网关,OpenAI 兼容;2026-07-18 由 DeepSeek 改定);数据抽 DAPO ~1k 子集控制 API 成本;目标是**管线跑通 + loss 正常下降**,不追分数。
## 2. 参考实现解剖
### 2.1 两套 SFT,别认错
| | `train_sft_sanity.py`129 行) | trainer 内 `lmbda=0` 路径(真基线) |
|---|---|---|
| 用途 | 数据质检:绕过全部自定义机制,拿官方 SFTTrainer 验证"数据本身可学" | 论文的 SFT 基线 |
| prompt 掩码 | **无**——整段对话(含用户提问)都算 loss | 有——prompt 位置标 -100 |
| 路由条件 | 独立脚本 | `no_teacher and lmbda==0.0`trainer:2855 |
教学点:sanity 脚本是**工具**不是基线;但"先用最笨的官方管线验证数据可学,再上自定义机制"这个调试策略本身值得继承。
### 2.2 数据格式与加载(train_dist:284-326
- DAPO parquet 列:`['data_source','prompt','ability','reward_model','extra_info']`——**`prompt` 列名不副实**,装的是完整 chat 列表 `[{user},{assistant}]`(若已有解答)或仅 `[{user}]`
- `format_messages`train_dist:295-321)把三种来源归一成 `messages` 列:`messages` 直取 / `prompt` 改名 / `question` 包成单 user 轮。parquet 会把 list 存成字符串,故有 `ast.literal_eval` 修复——**外面套着 `except: pass` 静默吞错(train_dist:299-305),我们 CLAUDE.md 明令禁止,重构时改为显式报错**。
- chat 模板**不在**数据阶段应用,推迟到 collator 逐 batch 应用(与 sanity 脚本相反)——好设计:`enable_thinking` 等模板决策收口一处。
### 2.3 Collator——本层最核心的一段(trainer:210-343
职责一句话:把 `messages` 变成 `input_ids/attention_mask/labels`,且**长解答永远不能把题目挤没**。行为清单:
| 行为 | 位置 | 为什么 load-bearing |
|------|------|---------------------|
| prompt/completion **各自独立预算** | prompt 用 `max_prompt_length` 截断(trainer:267-273);completion 上限 = `max_length - len(prompt)`trainer:292 | 单一 `max_length` 截断时,超长解答会把 prompt 截成空——参考实现的头号正确性卖点 |
| 标签构造 | `labels = [-100]*len(prompt) + completion`trainer:299 | -100 = 交叉熵的 ignore_indexprompt 不产生 loss |
| **左** padding | trainer:309-335 | 让整个 batch 能用一个标量 `prompt_length` 切分(见 2.4 |
| 边界确定 | 用**未截断**的 prompt 重分词长度切出 completiontrainer:286-289 | 模板渲染后 prompt+completion 的拼接分词 ≠ 分开分词,必须用同一渲染再切 |
| `enable_thinking` 透传 + 空 `<think>` 诊断 | trainer:252-266 | Qwen3 模板 no-think 时注入空 `<think>\n\n</think>`;此开关变了,prompt/completion 边界跟着变——错一次全错 |
| prompt-only 行 | 全 -100trainer:300-303 | 纯 SFT 下这种行 loss=0(有专门空 batch 兜护 trainer:2810 |
### 2.4 损失路径(trainer:2776-2874
```
prompt_length = batch 内 (总长-完成长) 的最小值 # trainer:751-759
logits[:, prompt_length-1 : -1] vs ids[:, prompt_length:] # 移位对齐
targets[labels==-100 处] = -100 # 重掩码,trainer:2803-2805
F.cross_entropy(..., ignore_index=-100)
```
两个精妙点:① `prompt_length` 取 **batch 最小值**保证不切掉任何 completion token,代价是长 prompt 行会有 prompt token 漏进"completion 切片"——由第 ③ 步用 labels **重掩码**兜住(labels 是权威掩码,切片几何只是加速);② 移位 `-1`:位置 t-1 的 logit 预测位置 t 的 token,SFT/蒸馏所有损失都踩这条对齐线。
### 2.5 teacher 生成与缓存(trainer:934-1068
触发条件 `lmbda==0 and use_teacher_server`;已有解答的行跳过;缓存 = `output_dir/openrouter_completions_rank{rank}.jsonl`,键为 `str(hash(prompt))`trainer:984——`hash()` 跨进程不稳定,我们换 sha256)。我们的重构把这段独立成 `teacher.py` 的第一个能力:「给一批 prompt 生成解答,带落盘缓存」——层 5 再给它长出 MC rollout 能力。
### 2.6 基础设施笔记(0.6B 用不上但要知道)
参考实现用 torchrun + FSDP,有个著名 trick`FSDP_ACTIVATION_CHECKPOINTING` 环境变量必须在 import accelerate/transformers **之前**设置(train_dist:15-21),且 `TrainingArguments.gradient_checkpointing` 在 FSDP 下是 no-optrain_dist:390-393)。
**我们的决策(2026-07-18 讨论定)**0.6B 乃至 1.7B 学生 4×A800 都用 **DDP**FSDP 触发点 = **换 4B 学生或序列 >8K**。⚠️ 显存账勘误(2026-07-18 sanity 实爆教训):"1.7B 全套 ~27G/卡"的旧估算只算了参数系(参数+梯度+Adam),漏了两个与参数量无关、正比于 batch 的大头:(B,T,V) logits 链(fp32 一份即 B×T×151936×4 字节,B=8/T=4096 时 ~20Gcross_entropy 内部 log_softmax 再来一份)与逐层激活(~30-40G@B=8)。150k 大词表下**显存瓶颈是 B×T,不是模型大小**;对策 = 压 per_device batch 用梯度累积补(B=2×4卡×8累积=64 不变)。但"import 前设环境变量"这个坑**今天就消除**:T5 训练脚本头部从第一天起内置 `os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", ...)` 前置块——DDP 下无害 no-op,未来启用 FSDP 时只改 TrainingArguments 字段,核心模块零改动。启用后必须 `nvidia-smi` 实测显存验证生效(分布式配置"设置了但静默无效"是常态,不可信配置)。另:层 2-5 坚持 DDP 还有调试纯度考量——on-policy 生成/ref model 搬运与 FSDP 的交界是参考实现最毛的地方,先保证"出错必是算法错"。
## 3. 保留 / 替代 / 删除清单(CLAUDE.md §6.2 规定动作)
| 决策 | 项目 |
|------|------|
| **保留** | 独立双预算、左 padding、-100 掩码、batch-min prompt_length + 重掩码、移位对齐、pad→eos 回退、disable_dropout、enable_thinking 收口 collator |
| **替代** | `hash()` → sha256`ast.literal_eval``except:pass` → 显式报错;OpenRouter 专用客户端 → 通用 OpenAI 兼容客户端(配置驱动);FSDP → DDP(0.6B |
| **删除** | `_RepeatBatchDataLoader` + RepeatSampler + 整套 buffer 机制(`lmbda=0` 下是空转,trainer:858-932)、vLLM 学生生成、Liger、on/off-policy 指标、ebopd 配置群 |
## 4. 重构任务(你主导,我配合)
| # | 任务 | 落点 | 备注 |
|---|------|------|------|
| T1 | `SFTConfig` dataclass(模型/数据路径、双预算、lr、enable_thinking…) | `ars_opd/configs.py` | 全部显式,禁默认藏参 |
| T2 | teacher 批量生成 + sha256 JSONL 缓存 | `ars_opd/teacher.py`(首个能力) | 读 `.env`;先对 ~1k 子集生成 |
| T3 | 数据加载(parquet/HF 双支持)+ `to_messages` + collator | `ars_opd/data.py`(新模块,IO 边缘) | CLAUDE.md 映射表需同步加行 |
| T4 | 掩码 SFT 损失 + 最小训练循环(HF Trainer 子类) | `ars_opd/trainer.py`(最小形态) | 只做 2.4 那四行的事 |
| T5 | 自包含实验脚本(写死全参数,零参数复现) | `scripts/train_sft.sh` | 触发 Video-Tree §2.5 规则接入;显式 CUDA_VISIBLE_DEVICES 4 卡 |
建议顺序 T1→T3→T4(本地可测)→T2(要 API key)→T5(远程)。**默认参数**(已通过;teacher 2026-07-18 改定):teacher 用 `MiniMax-M3`(自建 new-api 网关;M 系是 reasoning 模型,content 可能内联 `<think>` 思考段,入库前由 teacher.py 剥离)、`enable_thinking=False``max_length=4096 / max_prompt_length=1024`、子集 1000 题。
## 5. 验证方式
1. **本地(CPU**:collator 单测对拍参考行为——构造超长解答样本断言 prompt 未被截空;断言 -100 位置分布;断言 enable_thinking 两种取值下边界正确。数据加载单测:断言 `except:pass` 已变显式报错。
2. **远程**:先 sanity(官方 SFTTrainer 或我们管线跑 50 步)确认 loss 稳定下降;再全量 1 epoch。W&B 或实时日志盯 `sft/loss` 与有效 token 数。⚠️ 两条实测勘误(2026-07-18):① 起点不是想象的 ~2-3——预训练 Qwen3-0.6B 对 M3 风格数学文本的真实 CE ≈ **0.85**(探针 scripts/diag_loss_probe.py 实测),健康曲线 ≈ 0.9→0.4;② **HF 梯度累积契约坑(两幕剧)**:模型 forward 接受 loss_kwargs 时(Qwen3 是),Trainer 默认按"新式契约"对待自定义 compute_loss——第一幕:返回裸 mean 会被 ×累积步数(首跑 loss 7.5 ≈ 0.94×8);第二幕:改成 sum/num_items 后又 ÷world_size0.244 ≈ 0.85÷4),因为新式契约的 ×num_processes 补偿在**基类** compute_loss 尾部(v5 trainer.py:2028),整体重写 compute_loss 会绕过它。终解 = 按 HF 文档(trainer.py:1977)显式 `self.model_accepts_loss_kwargs = False` 退出新式契约,回到"返回 mean、Trainer ÷累积步数"的经典行为(代价:微批等权而非 token 加权,偏差百分之几,与参考实现同行为)。教训:**整体重写框架方法时,必须检查基类同名方法里除了你替换的逻辑还捎带了什么**。
3. 关账判据:1k 子集 1 epoch 跑完,loss 曲线正常,checkpoint 能被 `from_pretrained` 加载并生成通顺文本。
+140
View File
@@ -0,0 +1,140 @@
# 03 · 层 2White-box OPD 基线(token 级反向 KL
> 本章目标:吃透论文式(2) 的 on-policy 白盒蒸馏及其梯度爆炸脆弱性(§4.1),解剖参考实现 `distillation_mode="standard"` 路径,定出层 2 的重构任务。行号缩写:`DT:` = `references/ars-opd/trl/trl/experimental/distillation/distillation_trainer.py``TR:` = `references/ars-opd/train_distillation.py``CFG:` = `.../distillation/distillation_config.py`。
## 1. 论文侧:式(2) 精读
$$\mathcal{L}_{\text{OPD}} = \mathbb{E}_{y \sim \pi_\theta(\cdot|x)}\Big[\sum_t D_{KL}\big(\pi_\theta(\cdot|y_{<t},x)\,\|\,\pi_T(\cdot|y_{<t},x)\big)\Big]$$
三个成分逐个看:
| 成分 | 含义 | 为什么 |
|------|------|--------|
| $y \sim \pi_\theta$ | **on-policy**:轨迹由 student 自己采样 | 治 SFT 的曝光偏差——式(1) 只在 teacher 轨迹上监督,student 推理时一旦走出熟悉区域就没见过纠正信号;on-policy 让 teacher 在"student 实际会犯错的地方"给监督 |
| $D_{KL}(\pi_\theta\|\pi_T)$ | **反向 KL**student 在前)| mode-seekingstudent 容量小,与其平摊质量模仿 teacher 全分布(前向 KL 的 mode-covering),不如集中质量学好 teacher 的主模式 |
| 逐 token 求和 | token 级分布对齐 | 这就是"白盒":需要 teacher 每个位置的**完整 logits**,因此 teacher 必须本地可跑、且与 student **同 tokenizer**(词表逐位对应才能算 KL |
**§4.1 梯度爆炸(本层的理论主课,OmniOPD 的出发点)**:反向 KL 对 student logit 的梯度含 $\log\frac{\pi_\theta(v)}{\pi_T(v)}$ 项。on-policy 下 student 会采到 teacher 认为极差的 token$\pi_T \to 0$),此时 log 比值 $\to \infty$,单个 token 的梯度可以炸掉整个 batch。第一章已给过一句话版本,本层用代码把它钉死:重构任务里有一个"梯度范数随 $\pi_T$ 衰减而暴涨"的单元测试(对应层 3 detach 测试的姊妹篇——一个证明旧方案为什么坏,一个证明新方案为什么稳)。
顺带记住对比锚点:层 5 的式(8) 用"有界乘子 π̂"替换这里的 log 比值,这正是两代方法的分水岭。
## 2. 参考实现解剖(standard 路径)
### 2.1 一个重要的事实先行
**训练脚本从不传本地 teacher**TR:420-424 只传 model/args/dataset):仓库实际跑的 standard 蒸馏全走 **teacher server 路径**(环境变量 `TEACHER_URL` 触发,TR:329-330 → DT:2891-2920);**本地 teacher 路径(DT:2921-2944)只有直接构造 `DistillationTrainer(teacher_model=...)` 才会触发**。我们层 2 恰恰要走后者(4 卡 A800 放得下 4B teacher),所以两条路都要看懂,但以本地路为重构蓝本。
### 2.2 compute_loss 主流程(DT:2841-2946
```
守卫: no_teacher and lmbda<1 → 显式报错 (DT:2869)
student 前向(带梯度) (DT:2883)
prompt_length 切齐 + [pl-1:-1] 移位 (DT:2887-2889) ← 与 SFT 完全同款,T4 已实现
teacher logits:
server 路 → top-k logprobs 传输 (DT:2898)
本地路 → teacher.eval() + no_grad 前向 (DT:2923, 2578-2594)
divergence: generalized_jsd_loss / 稀疏快路 (DT:2926-2942)
```
on/off-policy 抽签**不在这里**——在 `_prepare_inputs``_fill_buffer`(见 2.5)。
### 2.3 generalized_jsd_lossDT:2408-2491)——β 的三副面孔
记号:$\pi_\theta$ = student$\pi_T$ = teacherKL 里"在前"的那个分布是被求期望的一方。
| β | 数学 | 语义 | 代码 |
|---|------|------|------|
| 0 | $KL(\pi_T\|\pi_\theta)$ | 前向,**mode-covering**teacher 在前,student 被迫摊平质量去覆盖 teacher 的全分布 | DT:150-151 |
| 1 | $KL(\pi_\theta\|\pi_T)$ | 反向,**mode-seeking****式(2) 用这个**student 在前,集中质量学 teacher 主模式 | DT:152-153 |
| (0,1) | $\beta KL(\pi_T\|m)+(1{-}\beta)KL(\pi_\theta\|m)$$m=(1{-}\beta)\pi_\theta{+}\beta\pi_T$ | JSD 插值(β 同时是混合权重与两项权重) | DT:154-162 |
> 行号说明:上表指向 **全词表** 路径(`F.kl_div`,我们 U2 走这条)。参考实现**默认**走 top-1 稀疏(§2.4),对应 DT:133-148 的 masked 镜像分支——同样三支 β、同样数学,只是在截断支持集上手算而非调 `F.kl_div`。
另有三条与 β 语义无关、但读代码时容易卡住的实现约定。它们各自独立,只是恰好都在这个函数里;前两条是正确性/稳定性刚需(我们保留),第三条是历史包袱(我们纠名):
| 实现约定 | 位置 | 是什么 / 为什么 load-bearing | 我们 U2 |
|----------|------|------------------------------|---------|
| 温度除进 logits | DT:2439-2440 | `logits / τ` 必须在 softmax **之前**做——这是在调分布形状(升温 τ>1 放大尾部的"暗知识"排序),不是等比缩概率。`softmax(z/τ) ≠ softmax(z)/τ`,位置错了就不再是合法分布 | 保留(默认 τ=1,此步为恒等) |
| 全程 log 域运算 | DT:143,156-159 | 15 万词表下单个概率小到 1e-8,而 KL 全是乘除,直接算会下溢成 0 → NaN。对策:全程存 log-prob(乘变加、除变减)。两个衍生 trick:混合分布 $m$ 的**加法**在 log 域要用 `logsumexp`(log 里的加法天然是乘法);masked 位概率为 0,取 log 前先 `clamp_min(tiny)``log(0)=-inf` | 保留(仅全词表这一路径) |
| `batchmean` 名不副实 | DT:2386-2399 | 先滤 `labels≠-100` 留下 completion 位,再 `jsd.sum() / 有效 token 数`——实义是 **per-token mean**,不是 PyTorch `batchmean` 那个"÷ 序列条数"。这样量纲与 SFT 的 per-token 交叉熵一致,两条 loss 曲线才可比 | 纠名为 `per_token_mean`,名字即文档 |
### 2.4 默认配置不是全词表 KL!(本章最大陷阱)
默认 `loss_top_k=1, loss_add_tail=True`CFG:353,368)→ 走 **top-1 稀疏快路**DT:2507-2549):支持集 = {实际采样 token} {teacher top-1} + **尾桶**(第 K+1 个桶收纳截断外全部概率质量:$\log(1-\sum e^{\text{top}k})$DT:108-118,防 top-1 时 loss 平凡为 0)。要严格对齐式(2) 的全词表反向 KL,必须 `loss_top_k=0`。top-k>1 时支持集依 β 取 teacher top-k / student top-k / 两者并集去重(DT:2449-2467)。
server 路径更粗:teacher 只回传 top-k logprobs 三张表(DT:2676-2679),β>0 时 CFG 强制 top-1CFG:540-544)。这些截断全是**传输/显存工程妥协**(第一章 §3.7 讨论过),不是论文成分。
### 2.5 buffer 与 on-policy 生成(DT:846-932, 1071-1263
- `_RepeatBatchDataLoader` 把同一 collated batch 重复 `gradient_accumulation_steps` 次(DT:346-371),`_fill_buffer` 按**切片级**伯努利抽签 on/off-policy`random() <= lmbda`DT:879,主进程抽签后广播)。
- **off-policy 切片 = 在数据集自带 completion 的轨迹上做 KL 蒸馏**DT:893-894 原样保留数据轨迹;损失仍是 KL,不是交叉熵)——GKD 的 λ 插值本义:λ=0 离线蒸馏、λ=1 纯 on-policy。两个易误解处:轨迹不是 teacher 现场采样的;DAPO prompt-only 下这些切片 labels 全 -100KL 被掩码归零 = **静默空转的算力浪费**(唯一例外:`lmbda=0`+server 触发 teacher 生成,DT:901-906,即层 1 SFT 的数据来源)。
- on-policy 切片:vLLM colocate 生成(按 `vllm_sync_frequency` 同步权重,DT:1094-1106)或 `model.generate`DT:1113-1168);生成结果重建 input_ids/labels 写回 bufferDT:1170-1263labels 只在 completion 段有效)。
- loss 前向是对已生成序列的 teacher-forcingDT:2883)——"采样一次、前向算分布",GKD 标准做法。
### 2.6 本地 teacher 的基础设施(DT:504-586
加载后 `accelerator.prepare_model(teacher, evaluation_mode=True)`(DT:584,随 DDP 每卡一份副本);同 tokenizer 校验比较 `get_vocab()`DT:734-740),不匹配在 compute_loss 处显式报错(DT:2876);前向 `eval() + no_grad`DT:2581-2586)。
### 2.7 顺带发现的坑(解剖副产物)
| 坑 | 位置 | 说明 |
|----|------|------|
| `lmbda=1 + no_teacher` 穿过守卫后在深处崩 | DT:2872 vs DT:2593 | 报错文案宣称合法,实际必崩——守卫条件写错 |
| off-policy ≠ "teacher 采样的离线数据" | DT:893-894 | 轨迹来自数据集固有 completion(损失仍是 KL);prompt-only 数据下整个切片被掩码归零,静默空转 |
| `num_generations>1 且 lmbda<1` 会造重复样本 | CFG:593-596 | 官方注释自己承认 |
## 3. 与式(2) 的偏差清单(默认配置下)
抉择原则(下表每一行都由它推出,而非逐条拍板):**式(2) 的算法本质忠于论文(on-policy + 反向 KL + 全词表分布对齐);参考实现的默认近似是为它的处境——API 传输 + 大 teacher——妥协出来的,换了我们的处境(本地同 tokenizer 的 4B teacher + 4×A800)就不继承;不改优化方向的表面差异,选对诊断/教学最有利的;一般性凡免费且未来有用则保留、凡昂贵且当前数据上空转则删除。** 一个反直觉推论:正因处境不同,我们回归论文本质反而比参考默认更贴式(2)(支持集那行)。逐行推导见下,"我们层 2"列即结论。
| 项 | 参考实现默认 | 严格式(2) | 我们层 2 |
|----|--------------|-----------|----------|
| 支持集 | top-1 稀疏 + 尾桶 | 全词表 | **全词表**`top_k=0` 等价;同 tokenizer 本地 teacher 使我们能比参考默认更贴论文) |
| on-policy 比例 | lmbda=1.0 ✓ | 纯 on-policy | lmbda=1(不实现混合抽签) |
| KL 方向 | beta=1.0 ✓ | 反向 | β 作为参数保留(前向/反向/JSD 同一公式,纯逻辑函数顺手覆盖,also 层 5 KL 锚要用前向) |
| 温度 | 1.0 ✓ | 1 | 1.0 |
| reduction | per-token mean | 论文 token 求和 | per-token mean(与 SFT loss 同尺度,才能对比曲线;差一个常数因子不改优化方向) |
## 4. 保留 / 替代 / 删除
| 决策 | 项目 |
|------|------|
| **保留** | "生成一次 + teacher-forcing 前向"结构;prompt 边界切齐/重掩码(直接复用 T4 的 `compute_prompt_length`);teacher `eval+no_grad`;同 tokenizer 校验(显式报错);β 语义与温度;per-token mean |
| **替代** | vLLM 学生生成 → `model.generate`0.6B 生成不慢,省掉 colocate+权重同步整套复杂度;卡了吞吐再回来接 vLLM,触发条件记录于此);teacher server → 本地 Qwen3-4B HF 前向(全词表精确);buffer/RepeatBatchDataLoader/切片抽签 → 每个 batch 现场生成现场用(lmbda=1 下 buffer 是纯开销) |
| **删除** | top-k/尾桶/top-1 稀疏快路(全词表放得下就不近似;显存账见 §5)、`reverse_kl_top_1_mode`、teacher server 客户端、Liger、off-policy 混训、on/off-policy 指标群 |
## 5. 重构任务(Claude 写码、你精读提问)
| # | 任务 | 落点 | 备注 |
|---|------|------|------|
| U1 | `DistillConfig` dataclassteacher 模型、β、温度、生成参数) | `ars_opd/configs.py` | 追加,不动 SFTConfig |
| U2 | 纯逻辑 divergence:全词表 masked token-KL/JSD(β 参数化)+ 单测 | `ars_opd/trainer.py`(纯张量函数,同 `sft_loss` 地位) | 单测含**梯度爆炸演示**:$\pi_T\to0$ 时梯度范数暴涨的断言(§4.1 的可执行版本) |
| U3 | collator 放开 prompt-only(返回 prompts/prompt_attention_mask 供生成) | `ars_opd/data.py` | 兑现 T3 预留的口子;SFT 路径行为不变(回归测试盯住) |
| U4 | `DistillTrainer`:生成 → teacher no_grad 前向 → divergence | `ars_opd/trainer.py` | on-policy 生成用 `model.generate`;tokenizer 一致性构造时校验 |
| U5 | 自包含脚本 | `scripts/train_whitebox.sh` + `.py` | teacher Qwen/Qwen3-4B;数据复用同一 DAPO 子集(prompt-only,无需 teacher 缓存) |
**显存账(A800-80G,验证 U 系列前算给自己看)**:全词表 logits (B,T,V) 是大头——B=4、T=2048、V≈151k 的 bf16 logits 单张 ≈2.5Gstudent+teacher 两份 + log_softmax 中间量 ≈15G 级,加 0.6B 训练全套 ~10G 与 4B teacher 推理副本 ~9GB=4/T=2048 起步安全;生成长度先压 1024。参数进 `DistillConfig` 显式化。
**默认参数提案(可否决)**:β=1.0、temperature=1.0、lmbda 固定 1(不做参数)、`max_new_tokens=1024`、per_device batch 4、teacher `Qwen/Qwen3-4B`、数据同 seed 同 1k 子集(复用层 1 的抽取逻辑,不需要 teacher 解答缓存)。
## 6. 验证方式
1. **本地 CPU 单测**divergence 手算对拍(V=5 玩具分布,β∈{0,1,0.5} 三点各一);β=1 与 β=0 的方向性断言(teacher 置信/弥散两种分布下 loss 排序);**梯度爆炸测试**:固定 studentteacher 对采样 token 的概率从 1e-1 衰减到 1e-6,断言梯度范数单调暴涨且超阈值——为层 5 的"有界乘子"对照埋桩。
2. **collator 回归**:放开 prompt-only 后,全部既有 SFT 测试必须原样通过。
3. **远程冒烟**:盯 KL loss 曲线、`grad_norm``distill/num_gen_tokens_per_step`
4. 关账判据:白盒蒸馏跑通不 NaN,生成不塌空,接口回看完成。
### 6.1 远程实证(2026-07-19,两次跑,勘误当初的"预期见毛刺")
**当初预测错了**docs 原写"预期能看到 loss 毛刺(梯度爆炸实况)"。实跑**没有毛刺**,两次都平稳。诚实记录 + 解释:
| 跑 | 配置 | loss | grad_norm |
|----|------|------|-----------|
| sanity | 裁到 1.0、lr 1e-6、50 步 | 平滑 0.35→0.22 | 14.42 → ~2(单调降) |
| noclip | ≈关裁剪、lr 5e-6、15 步 | 平滑 0.35→0.21 | 14.42 → ~2(无尖峰,且更快收敛) |
三条实证结论:
- **爆炸是真机制、但本区间高度阻尼**。§4.1 在单测里坐实(单个 π_T→1e-6 的 token 梯度暴涨),但真实训练里:① **同门 teacher**Qwen3 0.6B↔4B)使 student 很少采到 teacher 真恨的 token;② **per-token mean 把每步 ~4000 token 的梯度尖峰摊平**(单测看单 token 机制,真实看上千 token 平均后果)。故 batch 级 grad_norm 峰值只 ~14(比健康 ~2 高 5-7 倍,但远非几十上百),且只降不升。
- **关键坑:`grad_norm` 日志是裁剪前值**。sanity 的 14→2 那串本身就是爆炸证据,只是 HF 默认 `max_grad_norm=1.0` 把**步长**裁掉了、loss 才平滑——这个静默稳定器现已提进 `DistillConfig`(见其注释)。noclip 关掉它,grad_norm 曲线几乎不变(第 1 步两跑完全相同=14.42,验证确定性),但大步长反而**加速收敛**、仍不炸。
- **重构层 5 动机的认知**:白盒的"同 tokenizer"硬约束把你锁在相对温和的区间(换跨家族 teacher 会先被词表校验拦下),所以层 5 有界乘子 π̂ 的真正杀手锏不在"防这个温和爆炸",而在 **logit-free**(teacher 只给文本、拿不到 logits,白盒根本跑不了)。爆炸的干净见证留在 U2 单测。
+154
View File
@@ -0,0 +1,154 @@
# 04 · 层 3:语义相似度 φ + MC 估计器(式 3/4/5
> 本章目标:吃透 OmniOPD 如何把"要 teacher logits"(层 2 白盒的硬约束)换成"比 teacher 文本"logit-free),并用 Dirichlet 贝叶斯平滑把稀疏的相似度信号变成稳定、非零的监督乘子 π̂。然后建 `ars_opd/similarity.py`(式3)与 `ars_opd/estimator.py`(式4/5)两个**纯逻辑**模块,本地 CPU 对拍参考实现。
> 行号缩写:`DT:` = `references/ars-opd/trl/trl/experimental/distillation/distillation_trainer.py``CFG:` = `.../distillation_config.py``VAL:` = `references/ars-opd/validate_chunk_mc_estimator.py`。
## 1. 论文侧:从"要 logits"到"比文本"(§3.2.1-3.2.2
### 1.1 大局:这一层是 logit-free 的支点
层 2 白盒 OPD(式2)要 teacher 每个位置的完整分布,还要同 tokenizer——我们实测这把方法锁死在同门 teacher。§3.2.1 用一句话拆掉它:**别在 token 概率上匹配,改在文本语义上匹配**。teacher 只需生成文本 rollout;学生某段 chunk 对不对,由"学生这段文本" vs "teacher 那几段 rollout 文本"的语义相似度判定。
这一步同时解三个问题(§3.2.1 首段):① teacher 逐 token 查询 O(T) 不可行 → chunk 化降到 O(T/C);② tokenizer 不一致导致学生的精确 token 在 teacher rollout 里根本不出现、产生稀疏零梯度 → 改比语义;③ 由此得到跨架构可用的信号。结构上灵感来自 Speculative Decoding 的验证阶段——把一个 C-token chunk 当作单个"验证单元"。
### 1.2 式(3):语义相似度聚合 k_sem
学生生成 on-policy 轨迹 y,从中选 M 个长度 C 的 chunk(选法是层 4 的熵调度)。对某个 chunk c,把它之前的前缀 y_<c 喂给 teacherteacher 生成 N 个 rollout,逐个与学生 chunk 比相似度、求和:
$$k_{\text{sem}}^{(c)} = \sum_{i=1}^{N} \phi\big(y_c,\, y_{\text{teacher}}^{(i)}\big),\qquad \phi:(y_c, y_{\text{teacher}})\mapsto[0,1]$$
| 要点 | 说明 |
|------|------|
| φ 是**连续**相似度 | ROUGE-1 单词重叠 / 编辑距离,∈[0,1],非 0/1 硬票 |
| k_sem ∈ **[0, N]** 是实数 | N 个 [0,1] 相似度之和;`test_estimator_detach.py` 口语叫"票数",实质是连续和 |
| 比的是**文本语义**不是 token | 词级重叠对 tokenizer/风格不变——teacher 换词、换词表,只要意思对 φ 就高(§3.2.1 末:不惩罚风格偏差与词表不匹配) |
### 1.3 式(4):学生先验 π̄(贝叶斯先验)
$$\bar\pi_\theta^{(c)} = \Big(\prod_{t\in c}\pi_\theta(y_t\mid x,y_{<t})\Big)^{1/C} = \exp\!\Big(\tfrac1C\sum_{t\in c}\log\pi_\theta(y_t\mid\cdot)\Big)$$
学生对这段 chunk 每 token 概率的**几何均值** = exp(平均 log 概率)。它是学生**自己**"平均每 token 有多自信",∈(0,1],拿来当贝叶斯先验。几何均值(非算术)才是自然的 chunk 级概率:整段联合概率开 C 次方。
### 1.4 式(5)Dirichlet-Multinomial 贝叶斯平滑 π̂(本层灵魂)
$$\hat\pi_{\text{teacher}}^{(c)} = \frac{k_{\text{sem}}^{(c)} + \alpha\,\bar\pi_\theta^{(c)}}{N + \alpha}
\;\overset{式(10)}{=}\; \underbrace{\tfrac{N}{N+\alpha}}_{}\underbrace{\tfrac{k_{\text{sem}}}{N}}_{\hat\pi_{\text{freq}}\;\text{频率估计}} + \underbrace{\tfrac{\alpha}{N+\alpha}}_{}\underbrace{\bar\pi_\theta}_{\text{学生先验}}$$
π̂ 是"频率估计 k_sem/N"与"学生先验 π̄"的**凸组合**,权重 N 对 α。α = 先验强度(`chunk_alpha`,默认 1.0)。
### 1.5 为什么这样设计:定理 4.1 的三性质 + detach 命门
π̂ 最终进式(8)$\mathcal{L}_{\text{chunk}} = -\hat\pi^{(c)}\sum_{t\in c}\log\pi_\theta(y_t\mid\cdot)$(层 5 的事,此处只需知道 π̂ 是**乘子**)。定理 4.1 证明这套设计同时关掉两个失败模式:
| 性质 | 内容 | 对照层 2 |
|------|------|---------|
| (a) 不爆炸 | π̂∈[0,1] 是**乘子**,不在分母、不在 log 里;每 chunk 梯度被学生 score function 卡住有界(式11 | 层 2 反向 KL 的 log(π_θ/π_T) 在 π_T→0 时**无界爆炸** |
| (b) 不塌缩 | 先验保证 π̂ ≥ α·π̄/(N+α) **> 0 恒成立**(式12),即便 k_sem=0teacher 全否定) | 裸频率估计 k_sem/N 在 k_sem=0 时**塌成 0**、梯度死在最该纠正处 |
| (c) 方差收缩 | 贝叶斯 MSE 有闭式(式13),方差比频率估计严格缩小 (N/(N+α))²<1 | —— |
**detach 命门**`tests/test_estimator_detach.py` 守的,docs/01 §3.4 推导):π̄ 是学生**自己**的概率。若 π̄ 在式(8) 里不 detach,梯度会经它回传——优化器发现**压低学生对自己 token 的概率**能把乘子 π̂ 推向 0、从而逃避 -π̂·log π_θ 的惩罚(p·ln(1/p)→0,指数快过对数),而且**恰好在 k_sem=0 的 chunk 上塌缩**(那里 π̂=α·π̄/(N+α) 纯由 π̄ 驱动)——最需要学的地方最先崩。detach 把 π̄ 变成常量乘子,损失退化为"按 π̂ 权重强化学生 token",梯度方向恒为增大 p。这正是 (a)(b) 得以成立的机制根源,也是层 2 §4.1 的姊妹篇:一个证明旧方案为什么炸,一个建新方案为什么稳。
## 2. 参考实现解剖(带行号)
### 2.1 φ 的两个实现——与三处不一致(重构要抹平)
| 度量 | 位置 | 输入表示 | 算法 |
|------|------|---------|------|
| rouge1 | DT:1670 `_compute_rouge1` | **词集合**`.split()` 去重)于 decode 后**文本** | set 重叠的 F1 = 2PR/(P+R) |
| edit | DT:1682 `_compute_edit_similarity` | **token id 列表**(顺序敏感) | 1 归一化 Levenshtein / max(m,n) |
⚠️ 三处不一致(都是重构要统一的):
1. **rouge1 比文本、edit 比 token id**DT:1850 vs 1846)。edit 用 token id **重新耦合了 tokenizer**——直接违背 §3.2.1 "跨 tokenizer" 的立身之本;只在 teacher/student 同 tokenizer 的 vLLM 路径侥幸能跑。
2. **rouge1 两种算法**trainer 用**集合**DT:1672),validate 脚本用**多重集计数** `min(ref_cnt, hyp_cnt)`VAL:63)——同名不同义。
3. edit 的归一化用 max(m,n),标准 ROUGE-1 其实是多重集——参考实现里 rouge/edit/bleu 各行其是(VAL:55-153 有 7 种度量的大杂烩)。
### 2.2 k_sem 聚合(DT:1841-1852
```
for teacher_chunk in teacher_chunks (N 个):
sim = edit(student_ids, teacher_ids) 或 rouge1(student_text, teacher_text)
k_score += sim # 连续求和 = 式(3) 的 k_sem
```
### 2.3 chunk 级 π̄ / π̂ / detachDT:2195-2205)——正主
```python
# 式(4) 学生先验:几何均值,detach(命门①,DT:2196
log_pi_bar = chunk_lps.detach().sum() / chunk_len
pi_bar = log_pi_bar.exp().clamp(min=1e-8, max=1.0)
# 式(5) 贝叶斯目标:再 detach(命门②双保险,DT:2201
pi_hat = (k + self.chunk_alpha * pi_bar) / (self.chunk_mc_samples + self.chunk_alpha)
pi_hat = pi_hat.clamp(min=1e-8, max=1.0).detach()
chunk_loss = -pi_hat * chunk_lps.mean() # 式(8) chunk 项
```
两处 detach2196 的 `chunk_lps.detach()` 与 2201 的 `pi_hat.detach()`)**任一都足以**切断逃逸路(k 是 python float,π̄ detach 后 π̂ 已无梯度);参考实现两处都留是防御。clamp 的 1e-8 下限防 log0/精确零;上限 1.0 其实自然满足(k≤N、π̄≤1 ⇒ π̂≤1),是防御。**差异标注**:式(8) 论文是 Σlog π_θ,参考用 `.mean()`(除以 chunk 长)——per-token mean,此处是层 5 的事,先记下。
### 2.4 token 级基线(DT:1465)——论文说"不可行"的朴素版,作对照
```python
pi_hat = ((k_counts.float() + alpha * student_probs_at_token.detach()) / (N + alpha)).detach()
mc_loss = -pi_hat * student_log_probs_at_token
```
同一个式(5),但落在**单 token** 上(C=1):teacher 每步查询、数学生精确 token 的经验频率。这就是 §3.2.1 开头说的 O(T) 不可行、且 tokenizer 不一致下 k 恒 0 的朴素方案。`test_estimator_detach.py` 现在的 C=1 简化正对应这个基线。我们重构做 **chunk 级**2.3)。
### 2.5 validate 脚本的对拍策略(VAL,本层验证方式的蓝本)
`validate_chunk_mc_estimator.py` 回答"廉价的文本相似度 π̂ 能否逼近昂贵的真值":
- **ground truth**VAL:13,348):teacher 在**学生 chunk 的 token 上**的几何均值概率 π̄_teacher = exp(mean(log P_teacher))——这需要 teacher logprobs(白盒),是 π̂ 想廉价逼近的对象。
- **估计**VAL:417-419):k_continuous = mean(sim)·Nbayes = (k + α·prior)/(N+α)。
- **判据**VAL:433-437):对 7 种度量各算 MSE_freq vs MSE_bayes、Spearman 相关;验证**贝叶斯平滑降 MSE**(定理 4.1c)、哪种 φ 最相关。
我们重构的纯逻辑单测照此精神,但用 **toy 数据**(不连真 teacher):手构 k_sem/π̄ 断言 π̂ 公式与性质,再用 toy 模拟验证"MSE_bayes < MSE_freq"与"k=0 时 π̂>0"。
### 2.6 配置默认(CFG
| 参数 | 符号 | 默认 | 论文 |
|------|------|------|------|
| `chunk_mc_samples` | N | 10CFG:277 | 10(§4.2 甜点) |
| `chunk_alpha` | α | 1.0CFG:281 | 1.0 |
| `chunk_length` | C | 50CFG:273 附近) | 50 |
| `chunk_similarity` | φ | **rouge1**CFG:299 | **edit_distance**(§5.1)——⚠️ 背离,须显式指定 |
| `no_bayesian` | — | FalseCFG:315 | 消融:直接用 k/N(频率),验证塌缩 |
## 3. 保留 / 替代 / 删除
| 决策 | 项目 |
|------|------|
| **保留** | 式(4) 几何均值先验 + detach(命门);式(5) 贝叶斯凸组合;连续 φ 求和成 k_sem;clamp 下限防零;no_bayesian 消融(留作层 6 消融开关) |
| **替代** | φ 统一到**词级文本**`.split()`)——edit 也比 words,不再比 token id(抹平 2.1 坑①,回归 tokenizer 无关);rouge1 集合/多重集二选一并注明;默认 φ 显式设 edit_distance(对齐论文 §5.1,不用 code 的 rouge1 默认) |
| **删除** | token 级基线路径(DT:1465,朴素不可行版);validate 里 bleu/jaccard/exact_match 等多余度量(只留 rouge1 + edit);vLLM/API 采样(那是层 5 teacher.py 的事,层 3 纯逻辑只吃已算好的 k_sem 与 log 概率) |
## 4. 重构任务(Claude 写码、你精读提问)
| # | 任务 | 落点 | 备注 |
|---|------|------|------|
| E1 | `rouge1(hyp, ref)` + `edit_similarity(hyp, ref)` + `phi(hyp, ref, metric)` + `aggregate_similarity(student_chunk, teacher_rollouts, metric)` | `ars_opd/similarity.py`(纯逻辑,只依赖标准库) | 两度量都吃 str、内部 `.split()` 比 words;φ 默认 edit_distancek_sem = Σφ(式3 |
| E2 | `chunk_prior(log_probs)`(式4 几何均值,**detach**+ `bayesian_target(k_sem, pi_bar, n, alpha)`(式5 | `ars_opd/estimator.py`(纯逻辑,只依赖 torch | detach 在 `chunk_prior` 内;`bayesian_target` 再 detach 防御;clamp 下限 |
| E3 | 把 `test_estimator_detach.py` 接到**真实现** | `tests/test_estimator_detach.py` | 现为独立 C=1 数学测试;追加对 `chunk_prior`/`bayesian_target` 的同名断言,锁死 detach 不被误删 |
**接口预想**(E1/E2 公共函数,全类型注解):
```
# similarity.py(纯 str/list,可脱离 torch 测)
rouge1(hypothesis: str, reference: str) -> float # 词集合 F1 ∈[0,1]
edit_similarity(hypothesis: str, reference: str) -> float # 1 归一化 Levenshtein ∈[0,1]
phi(hypothesis: str, reference: str, metric: str = "edit_distance") -> float
aggregate_similarity(student_chunk: str, teacher_rollouts: list[str], metric: str) -> float # 式(3) k_sem
# estimator.pytorchtoy 张量可测)
chunk_prior(log_probs: torch.Tensor) -> torch.Tensor # 式(4) π̄,detach(C,)->标量
bayesian_target(k_sem: float, pi_bar: torch.Tensor, n_rollouts: int, alpha: float) -> torch.Tensor # 式(5) π̂
```
## 5. 验证方式
1. **similarity.py 单测**:手构字符串断言 φ 值(如全同 chunk→edit=1、rouge1=1;不相交→0;部分重叠手算);k_sem = Σφ 的连续性;空串边界。
2. **estimator.py 单测(对拍参考精神)**
- **公式对拍**:手构 log_probs/k_sem,断言 π̄=exp(mean log p)、π̂=(k+απ̄)/(N+α) 与手算一致;凸组合式(10) 恒等。
- **性质对拍**:π̂ ∈(0,1] 恒;**k=0 时 π̂ = α·π̄/(N+α) > 0**(定理 4.1b 反塌缩,= detach 测试的正面);
- **方差收缩(toy 模拟)**:固定真值 μ、采样 N 个 [0,1] 相似度多次,断言 π̂ 的 MSE < 频率估计 k/N 的 MSE(定理 4.1c)。
- **detach 命门**:对真 `chunk_prior`/`bayesian_target` 复刻 `test_estimator_detach.py` 的两个世界断言(detach→k=0 仍增大 p;不 detach→k=0 且 p<1/e 时逃逸)。
3. 关账判据:similarity/estimator 全单测本地 CPU 通过;`test_estimator_detach.py` 已接真实现;接口回看完成(两模块判"深")。
+15
View File
@@ -0,0 +1,15 @@
# 最小打包配置:只为让 `pip install -e .` 把 ars_opd 注册进环境,
# 使脚本/调试器/远程从任意目录都能 import ars_opd(不再依赖 cwd 恰好是仓库根)。
# 依赖不在这里声明——统一走 requirements*.txt,避免两处清单漂移。
[build-system]
requires = ["setuptools>=64"]
build-backend = "setuptools.build_meta"
[project]
name = "ars-opd"
version = "0.1.0"
description = "OmniOPD (arXiv:2606.01476v2) 分层重构实现"
requires-python = ">=3.11"
[tool.setuptools]
packages = ["ars_opd"]
+125
View File
@@ -0,0 +1,125 @@
"""诊断脚本:逐环检验 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)
# 环 5collator 全流程后,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())
+49
View File
@@ -0,0 +1,49 @@
"""层 1 关账判据 3:训练后 checkpoint 能被 from_pretrained 加载并生成通顺解答。
远程运行(CPU 即可,0.6B 生成 512 token 约 1-2 分钟):
python -u scripts/diag_generate.py
"""
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from ars_opd.data import load_sft_dataset
MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物
# 不挂 teacher 解答(teacher_completions_path 缺省):只取题目做推理输入
ds = load_sft_dataset(
"data/dapo-math-17k-unique.parquet", subset_size=1000, seed=42
)
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, dtype=torch.float32)
model.eval()
# 取子集第 900+ 行附近的题(训练时见过,此处只验"会不会说话"不验泛化)
for i in (900, 950):
prompt = tok.apply_chat_template(
ds[i]["messages"],
tokenize=False,
add_generation_prompt=True,
enable_thinking=False, # 必须与训练取值一致(docs/02 §2.3 边界契约)
)
inputs = tok(prompt, return_tensors="pt", add_special_tokens=False)
with torch.no_grad():
out = model.generate(
**inputs, max_new_tokens=512, do_sample=False, temperature=None, top_p=None
)
completion = tok.decode(
out[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True
)
print(f"===== 样本 {i} 题目 =====")
print(ds[i]["messages"][-1]["content"][120:280], "")
print("----- 生成(前 600 字符)-----")
print(completion[:600])
print()
print(
"判读:应为步骤化数学解答(markdown 风格、以 Answer: 行收尾的倾向);"
"乱码/复读/空输出 = 不通过。",
flush=True,
)
+68
View File
@@ -0,0 +1,68 @@
"""损失探针:用未训练的预训练模型走完整管线,逐行算 loss(层 1 疑点排查第二步)。
判读(训练日志初始 loss ≈ 7.5):
- 探针也 ≈ 7:管线一致,loss 高是数据/模型现实 → 去查数据(垃圾长文、乱码占比);
- 探针 ≈ 2-4:管线(本脚本与训练共用)没问题但训练环节另有妖 → 查训练循环差异。
同时打印 HF 模型内建 CE(同一数学的独立实现)交叉验证 sft_loss。
远程运行(CPU 即可,约 1-2 分钟):
python -u scripts/diag_loss_probe.py
"""
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from ars_opd.configs import SFTConfig
from ars_opd.data import SFTCollator, load_sft_dataset
from ars_opd.trainer import sft_loss
MODEL = "Qwen/Qwen3-0.6B"
cfg = SFTConfig(
dataset_path="data/dapo-math-17k-unique.parquet",
output_dir="/tmp/diag",
subset_size=8,
seed=42,
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
)
ds = load_sft_dataset(
cfg.dataset_path,
cfg.dataset_split,
cfg.subset_size,
cfg.seed,
cfg.teacher_completions_path,
)
tok = AutoTokenizer.from_pretrained(MODEL)
collator = SFTCollator(
tok,
max_length=cfg.max_length,
max_prompt_length=cfg.max_prompt_length,
enable_thinking=False,
)
model = AutoModelForCausalLM.from_pretrained(MODEL, dtype=torch.float32)
model.eval()
print(f"{'':>3} {'sft_loss':>9} {'HF内建CE':>9} {'监督tok':>7} 解答开头")
total, total_n = 0.0, 0
for i in range(len(ds)):
batch = collator([ds[i]])
with torch.no_grad():
out = model(
input_ids=batch["input_ids"], attention_mask=batch["attention_mask"]
)
ours, n = sft_loss(
out.logits, batch["input_ids"], batch["labels"], batch["attention_mask"]
)
# 交叉验证:HF 内建损失(labels 传入模型,内部自动移位)与 sft_loss
# 是同一数学的两个独立实现,单行 batch 下应当几乎相等
hf = model(
input_ids=batch["input_ids"],
attention_mask=batch["attention_mask"],
labels=batch["labels"],
).loss
head = ds[i]["messages"][-1]["content"][:40].replace("\n", " ")
print(f"{i:>3} {ours.item():>9.3f} {hf.item():>9.3f} {n:>7} {head}", flush=True)
total += ours.item() * n
total_n += n
print(f"\n按 token 加权平均: {total / total_n:.3f}(对照训练日志初始 loss ≈ 7.5)", flush=True)
+36
View File
@@ -0,0 +1,36 @@
"""层 1:为 DAPO 1k 子集生成 teacherMiniMax-M3)解答缓存。
自包含实验脚本:全部参数写死在此,零参数复现。在**本地**运行(纯 API 调用,
不需要 GPU;本机可直连自建网关):
conda activate ars-opd
python -u scripts/generate_teacher_completions.py
前置:
1. .env 已填 TEACHER_API_BASE / TEACHER_API_KEY / TEACHER_MODEL
2. DAPO parquet 已下载到 DATASET_PATH(见 docs/02 §4)。
中断安全:缓存逐条落盘,重跑本脚本自动跳过已完成条目(断点续传)。
"""
from ars_opd.configs import TeacherGenConfig
from ars_opd.data import load_sft_dataset
from ars_opd.teacher import TeacherClient, generate_completions
# 非显然约束:这里的 dataset/subset_size/seed 必须与 T5 训练脚本完全一致——
# 两侧各自走"加载→归一→抽子集",seed 相同才是同一批题(data.py 有详注)
DATASET_PATH = "data/dapo-math-17k-unique.parquet" # DAPO 官方去重版,17917 行
CACHE_PATH = "data/teacher_completions_dapo1k_minimax-m3.jsonl"
# 试跑说明:首次建议把下面 subset_size 临时改成 5,跑通并人工抽查缓存里的解答
# 质量(think 是否剥净、格式是否正常)后再改回 1000 重跑。放心改:subset 是对
# 同一 seed 的洗牌序列取前缀,前 5 条与前 1000 条的头 5 条完全相同,试跑写入的
# 缓存在正式跑时全部命中,一分钱不浪费。
# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
dataset = load_sft_dataset(DATASET_PATH, subset_size=1000, seed=42)
prompts = [row["messages"] for row in dataset]
teacher = TeacherClient(TeacherGenConfig()) # 采样参数全用 configs.py 的显式默认
generate_completions(prompts, CACHE_PATH, teacher)
print(f"完成。缓存文件:{CACHE_PATH}")
+8 -3
View File
@@ -15,6 +15,8 @@ export CONDA_PKGS_DIRS=$DATA_ROOT/conda_pkgs # conda 包缓存默认在根
export PIP_CACHE_DIR=$DATA_ROOT/pip_cache # pip 缓存默认在根分区
export TMPDIR=$DATA_ROOT/tmp # 大 wheel 解压临时目录
mkdir -p "$DATA_ROOT"/{envs,hf_cache,conda_pkgs,pip_cache,tmp}
# 清理上次中断可能残留的 pip 临时目录(含曾被 conda run 吞掉 TMPDIR 而落在 /tmp 的)
rm -rf "$DATA_ROOT"/tmp/pip-* /tmp/pip-unpack-* 2>/dev/null || true
# ---- 1. conda 环境(建在 /data,不建在 ~----
if [ ! -d "$ENV_PATH" ]; then
@@ -29,7 +31,10 @@ else
fi
# ---- 3. 依赖 ----
conda run -p "$ENV_PATH" pip install -r "$REPO_DIR/requirements.txt" -r "$REPO_DIR/requirements-remote.txt"
# 直接调环境内 pip:conda run 会整体缓冲子命令输出(违反"日志实时可查"规矩),弃用
"$ENV_PATH/bin/pip" install -r "$REPO_DIR/requirements.txt" -r "$REPO_DIR/requirements-remote.txt"
# ars_opd 以 editable 方式注册进环境:脚本从任意目录都能 import,不依赖 cwd
"$ENV_PATH/bin/pip" install -e "$REPO_DIR" --no-build-isolation --no-deps
# ---- 4. 环境变量持久化(写入 ~/.bashrc,幂等)----
if ! grep -q "ars-opd-rebuild env" ~/.bashrc; then
@@ -46,9 +51,9 @@ fi
# ---- 5. 验证 ----
echo "=== 验证 torch/CUDA ==="
conda run -p "$ENV_PATH" python -c "import torch; print('torch', torch.__version__, '| cuda可用:', torch.cuda.is_available(), '| 卡数:', torch.cuda.device_count())"
"$ENV_PATH/bin/python" -c "import torch; print('torch', torch.__version__, '| cuda可用:', torch.cuda.is_available(), '| 卡数:', torch.cuda.device_count())"
echo "=== 验证单元测试 ==="
conda run -p "$ENV_PATH" python -m pytest "$REPO_DIR/tests" -q
"$ENV_PATH/bin/python" -m pytest "$REPO_DIR/tests" -q
echo "=== 磁盘检查(根分区不应有明显增长)==="
df -h / /data | tail -2
echo "全部完成。日常使用:输入 opd 进入环境与目录。"
+154
View File
@@ -0,0 +1,154 @@
"""层 1:SFT 基线训练入口(由 train_sft.sh 经 torchrun 启动,勿直接 python 运行)。
自包含实验脚本:全部参数写死在下方 FULL 配置里,零参数复现;sanity 模式只是
对 FULL 的两处显式覆盖(50 步 + 独立输出目录)。
"""
# ---- FSDP 前置块(必须在一切 transformers/accelerate import 之前)----
# 非显然约束:FSDP 的激活检查点开关是 accelerate 在 import 时读取的环境变量
# (参考实现 train_distillation.py:15-21 的著名坑);写在 import 后会静默无效。
# DDP 下本变量是无害 no-op——现在就位是为了未来换 4B 学生/FSDP 时只改此处一行,
# 且改完必须 nvidia-smi 实测显存验证生效(docs/02 §2.6)。
import os
os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", "false")
import dataclasses
import sys
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments
from ars_opd.configs import SFTConfig
from ars_opd.data import IGNORE_INDEX, SFTCollator, load_sft_dataset
from ars_opd.trainer import SFTTrainer
STUDENT_MODEL = "Qwen/Qwen3-0.6B"
FULL = SFTConfig(
dataset_path="data/dapo-math-17k-unique.parquet",
output_dir="/data/zym/outputs/sft_qwen3-0.6b_dapo1k",
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
subset_size=1000,
seed=42, # 非显然约束:与 generate_teacher_completions.py 一致,否则缓存大面积 miss
max_length=4096,
max_prompt_length=1024,
enable_thinking=False,
learning_rate=2e-5,
per_device_train_batch_size=2, # B=8 曾爆 80G:大头是 (B,T,V) logits 链与激活,见 SFTConfig 注释
gradient_accumulation_steps=8, # 全局 batch = 2 × 4 卡 × 8 = 64
num_train_epochs=1,
max_steps=-1,
lr_scheduler_type="linear",
warmup_ratio=0.0,
gradient_checkpointing=False,
bf16=True,
logging_steps=1,
save_steps=100,
save_total_limit=2,
report_to="none", # 层 1 先靠 tmux 实时日志;W&B 触发条件见 appendix C 表
)
def build_config() -> SFTConfig:
"""按命令行模式产出配置。frozen dataclass 的换参方式:replace 构造新实例。"""
mode = sys.argv[1] if len(sys.argv) > 1 else "full"
if mode == "full":
return FULL
if mode == "sanity":
return dataclasses.replace(
FULL, max_steps=50, output_dir=FULL.output_dir + "-sanity"
)
raise ValueError(f"未知模式 {mode!r},只接受 full / sanity")
def smoke_check_first_batch(dataset, collator, tokenizer) -> None:
"""训练前解码第一个 batch 供肉眼核对(只在 rank0 打印一次)。
单测用玩具 tokenizer 钉死了预算/边界的算法(tests/test_data.py),但真
tokenizer 的模板渲染只能在这里肉眼验证:掩码边界是否落在 assistant 起点、
no-think 时空 <think> 块是否在 prompt 侧。这是参考实现"一次性诊断打印"
的合理化版本(docs/02 §2.3)。
"""
batch = collator([dataset[0]])
ids, labels = batch["input_ids"][0], batch["labels"][0]
masked = labels == IGNORE_INDEX
prompt_text = tokenizer.decode(ids[masked], skip_special_tokens=False)
completion_text = tokenizer.decode(ids[~masked], skip_special_tokens=False)
print(
"=" * 30
+ " 首样本自检(人工核对掩码边界)"
+ "=" * 30
+ f"\n[prompt 段 | {int(masked.sum())} tok | 不产生 loss]\n"
+ f"{prompt_text[-300:]}\n"
+ f"\n[completion 段 | {int((~masked).sum())} tok | 监督目标]\n"
+ f"{completion_text[:300]}\n"
+ "=" * 80,
flush=True,
)
def main() -> None:
cfg = build_config()
rank0 = int(os.environ.get("RANK", "0")) == 0
# 加载顺序刻意 fail-fast:数据(毫秒级,最易配错)→ tokenizer(几 MB)→
# 模型(GB 级下载)。teacher 缓存缺失要在下模型之前炸出来
dataset = load_sft_dataset(
cfg.dataset_path,
cfg.dataset_split,
cfg.subset_size,
cfg.seed,
cfg.teacher_completions_path,
)
tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
collator = SFTCollator(
tokenizer,
max_length=cfg.max_length,
max_prompt_length=cfg.max_prompt_length,
enable_thinking=cfg.enable_thinking,
)
if rank0:
smoke_check_first_batch(dataset, collator, tokenizer)
model = AutoModelForCausalLM.from_pretrained(STUDENT_MODEL, dtype=torch.float32)
args = TrainingArguments(
output_dir=cfg.output_dir,
# 非显然约束:必须关掉列裁剪。HF Trainer 默认删除模型 forward 签名里
# 没有的数据列——"messages" 会被整列删光,collator 收到空字典且不报错
remove_unused_columns=False,
learning_rate=cfg.learning_rate,
per_device_train_batch_size=cfg.per_device_train_batch_size,
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
num_train_epochs=cfg.num_train_epochs,
max_steps=cfg.max_steps,
lr_scheduler_type=cfg.lr_scheduler_type,
warmup_ratio=cfg.warmup_ratio,
gradient_checkpointing=cfg.gradient_checkpointing,
bf16=cfg.bf16,
seed=cfg.seed,
logging_steps=cfg.logging_steps,
logging_first_step=True,
save_strategy="steps",
save_steps=cfg.save_steps,
save_total_limit=cfg.save_total_limit,
report_to=cfg.report_to,
ddp_find_unused_parameters=False, # 全参训练无闲置参数,省一次全模型扫描
dataloader_num_workers=2, # collator 逐 batch 分词在 CPU,双 worker 与 GPU 重叠
)
trainer = SFTTrainer(
model=model,
args=args,
train_dataset=dataset,
data_collator=collator,
)
trainer.train()
trainer.save_model() # 终态模型(save_pretrained 格式,含 config
if rank0:
tokenizer.save_pretrained(cfg.output_dir)
print(f"训练完成,模型已存至 {cfg.output_dir}", flush=True)
if __name__ == "__main__":
main()
+28
View File
@@ -0,0 +1,28 @@
#!/usr/bin/env bash
# 层 1SFT 基线训练(远程 gpu-a800-060 专用;本地不跑训练)。
#
# 用法(tmux 内执行,日志实时可查):
# bash scripts/train_sft.sh sanity # 50 步冒烟:看首样本自检 + loss 是否从 ~2-3 下降
# bash scripts/train_sft.sh # 正式:1k 子集 1 epoch
#
# 前置检查清单:
# 1. nvidia-smi 确认下方 GPUS 四张卡空闲(只许用 8 卡中的 4 张,严禁自动选卡)
# 2. data/ 下已有两个文件(gitignore 不随 git 走,本地 scp 上来):
# scp data/dapo-math-17k-unique.parquet data/teacher_completions_dapo1k_minimax-m3.jsonl \
# <远程>:/data/zym/ars-opd-rebuild/data/
# 3. 代码是最新:git -C /data/zym/ars-opd-rebuild pull
set -euo pipefail
cd "$(dirname "$0")/.." # 锚定仓库根:py 内 data/... 相对路径以此为基准
GPUS=0,1,2,3 # ⚠️ 改这里前先 nvidia-smi
MODE=${1:-full}
export CUDA_VISIBLE_DEVICES=$GPUS
export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5
export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长序列 batch 易碎片化,按需扩段
export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com}
export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满
# hf-mirror 不代理 HF Xet CAS(大权重走 Xet 会 401,见 train_whitebox.sh 详注)
export HF_HUB_DISABLE_XET=1
torchrun --nproc_per_node=4 --master_port=29571 scripts/train_sft.py "$MODE"
+168
View File
@@ -0,0 +1,168 @@
"""层 2white-box OPD 训练入口(由 train_whitebox.sh 经 torchrun 启动)。
自包含实验脚本:全部参数写死在下方 FULL 配置里,零参数复现;sanity 模式只是
对 FULL 的两处显式覆盖(50 步 + 独立输出目录)。
与层 1 train_sft.py 的结构差异:双模型(student + 本地 teacher)、prompt-only
数据(无 teacher 缓存,现场 on-policy 生成)、DistillTrainer 编排。
"""
# ---- FSDP 前置块(必须在一切 transformers/accelerate import 之前,同 train_sft.py----
import os
os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", "false")
import dataclasses
import sys
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments
from ars_opd.configs import DistillConfig
from ars_opd.data import SFTCollator, load_sft_dataset
from ars_opd.trainer import DistillTrainer
STUDENT_MODEL = "Qwen/Qwen3-0.6B" # 被训练的固定基线(同层 1,脚本级常量)
FULL = DistillConfig(
dataset_path="data/dapo-math-17k-unique.parquet",
output_dir="/data/zym/outputs/whitebox_qwen3-0.6b_dapo1k",
teacher_model="Qwen/Qwen3-4B", # 本地全词表 teacher(须与 student 同 tokenizer
subset_size=1000,
seed=42, # 与层 1 一致:同一批题上对比 SFT 与蒸馏
max_prompt_length=1024,
max_new_tokens=1024, # 与 max_prompt_length 之和 = 序列总长 T≈2048(§5 显存账)
enable_thinking=False,
beta=1.0, # 反向 KL = 式(2)
kl_temperature=1.0,
gen_temperature=1.0, # 纯采样自 π_θ(忠实 on-policy
gen_top_p=1.0,
learning_rate=1e-6, # 论文 §5.1 蒸馏 lr;小步长也帮训练在梯度爆炸毛刺中存活
per_device_train_batch_size=4, # §5 估算,首次远程必须 nvidia-smi 核实不 OOM
gradient_accumulation_steps=4, # 全局 batch = 4 × 4 卡 × 4 = 64(同层 1)
num_train_epochs=1,
max_steps=-1,
max_grad_norm=1.0, # 显式写出这个此前静默的稳定器(§4.1 爆炸靠它压平,见 config 注释)
bf16=True,
logging_steps=1,
save_steps=100,
save_total_limit=2,
report_to="none",
)
def build_config() -> DistillConfig:
"""按命令行模式产出配置。frozen dataclass 换参方式:replace 构造新实例。"""
mode = sys.argv[1] if len(sys.argv) > 1 else "full"
if mode == "full":
return FULL
if mode == "sanity":
return dataclasses.replace(
FULL, max_steps=50, output_dir=FULL.output_dir + "-sanity"
)
if mode == "noclip":
# §4.1 教学对照:关闭裁剪 + 稍抬 lr,暴露反向 KL 原始爆炸。max_grad_norm
# 设远高于实测范数(~14)故永不触发≈无裁剪;lr 5×放大让爆炸在 loss 上可见。
# 与 sanity(裁到 1.0、lr 1e-6 的平滑曲线)并排 = 白盒脆弱性活教材,层 5 对照
return dataclasses.replace(
FULL,
max_steps=15,
max_grad_norm=1e9,
learning_rate=5e-6,
output_dir=FULL.output_dir + "-noclip",
)
raise ValueError(f"未知模式 {mode!r},只接受 full / sanity / noclip")
def smoke_check_first_prompt(dataset, collator, tokenizer) -> None:
"""训练前解码第一个 prompt 供肉眼核对(只在 rank0 打印一次)。
prompt-only 模式的自检重点:prompt 末尾应是生成引导符("...assistant\\n" +
no-think 时的空 <think>),student 将从此续写。若末尾不对,生成的分布与
训练目标会错位。
"""
batch = collator([dataset[0]])
prompt_ids = batch["prompts"][0]
mask = batch["prompt_attention_mask"][0].bool()
text = tokenizer.decode(prompt_ids[mask], skip_special_tokens=False)
print(
"=" * 30
+ " 首个 prompt 自检(供 on-policy 生成)"
+ "=" * 30
+ f"\n[{int(mask.sum())} tok,末尾应为生成引导符]\n{text[-400:]}\n"
+ "=" * 80,
flush=True,
)
def main() -> None:
cfg = build_config()
rank0 = int(os.environ.get("RANK", "0")) == 0
# 加载顺序 fail-fast(同 train_sft.py):数据(毫秒级)→ tokenizer(几 MB)→
# 模型(GB 级)。层 2 无 teacher 缓存,数据是 prompt-only 子集
dataset = load_sft_dataset(
cfg.dataset_path, cfg.dataset_split, cfg.subset_size, cfg.seed
)
student_tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
teacher_tokenizer = AutoTokenizer.from_pretrained(cfg.teacher_model)
collator = SFTCollator(
student_tokenizer,
max_prompt_length=cfg.max_prompt_length,
enable_thinking=cfg.enable_thinking,
prompt_only=True, # 层 2:只出 prompt 张量,completion 靠生成
)
if rank0:
smoke_check_first_prompt(dataset, collator, student_tokenizer)
# student fp32 + bf16 混合精度(同层 1);teacher 直接 bf16(只推理,省显存)
student = AutoModelForCausalLM.from_pretrained(STUDENT_MODEL, dtype=torch.float32)
teacher = AutoModelForCausalLM.from_pretrained(
cfg.teacher_model, dtype=torch.bfloat16
)
args = TrainingArguments(
output_dir=cfg.output_dir,
remove_unused_columns=False, # 保住 messages 列供 collator(同层 1 注释)
learning_rate=cfg.learning_rate,
per_device_train_batch_size=cfg.per_device_train_batch_size,
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
num_train_epochs=cfg.num_train_epochs,
max_steps=cfg.max_steps,
lr_scheduler_type=cfg.lr_scheduler_type,
warmup_ratio=cfg.warmup_ratio,
max_grad_norm=cfg.max_grad_norm,
gradient_checkpointing=cfg.gradient_checkpointing,
bf16=cfg.bf16,
seed=cfg.seed,
logging_steps=cfg.logging_steps,
logging_first_step=True,
save_strategy="steps",
save_steps=cfg.save_steps,
save_total_limit=cfg.save_total_limit,
report_to=cfg.report_to,
ddp_find_unused_parameters=False,
dataloader_num_workers=2,
)
trainer = DistillTrainer(
model=student,
args=args,
train_dataset=dataset,
data_collator=collator,
teacher_model=teacher,
teacher_tokenizer=teacher_tokenizer, # 构造时校验与 student 同词表
beta=cfg.beta,
kl_temperature=cfg.kl_temperature,
gen_temperature=cfg.gen_temperature,
gen_top_p=cfg.gen_top_p,
max_new_tokens=cfg.max_new_tokens,
)
trainer.train()
trainer.save_model()
if rank0:
student_tokenizer.save_pretrained(cfg.output_dir)
print(f"训练完成,模型已存至 {cfg.output_dir}", flush=True)
if __name__ == "__main__":
main()
+36
View File
@@ -0,0 +1,36 @@
#!/usr/bin/env bash
# 层 2white-box OPD 训练(远程 gpu-a800-060 专用;本地不跑训练)。
#
# 用法(tmux 内执行,日志实时可查):
# bash scripts/train_whitebox.sh sanity # 50 步冒烟:首 prompt 自检 + KL loss + 生成数
# bash scripts/train_whitebox.sh noclip # §4.1 对照:15 步,关裁剪+抬 lr,暴露原始
# # 梯度爆炸(loss 毛刺);与 sanity 平滑曲线并排
# bash scripts/train_whitebox.sh # 正式:1k 子集 1 epoch
#
# 前置检查清单:
# 1. nvidia-smi 确认下方 GPUS 四张卡空闲(只许用 8 卡中的 4 张,严禁自动选卡)。
# ⚠️ 白盒显存比层 1 紧:student 训练全套 + teacher(4B) 推理副本 + **两份**全词表
# logits(student/teacher),§5 估算 B=4/T=2048 起步安全,但首跑必须盯 nvidia-smi;
# 若 OOM,降 per_device_train_batch_size 到 2,仍不够再开 gradient_checkpointing
# (改 DistillConfig,注意 checkpointing 与 generate 的 use_cache 交互)。
# 2. data/dapo-math-17k-unique.parquet 已在(层 2 无需 teacher 缓存,纯 prompt-only):
# scp data/dapo-math-17k-unique.parquet <远程>:/data/zym/ars-opd-rebuild/data/
# 3. 代码最新:git -C /data/zym/ars-opd-rebuild pull
# 4. 首跑会下载 teacher Qwen3-4BGB 级)到 HF_HOME,确保 /data 有空间
set -euo pipefail
cd "$(dirname "$0")/.." # 锚定仓库根
GPUS=0,1,2,3 # ⚠️ 改这里前先 nvidia-smi
MODE=${1:-full}
export CUDA_VISIBLE_DEVICES=$GPUS
export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5
export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长生成序列易碎片化,按需扩段
export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com}
export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满
# 非显然坑:hf-mirror 不代理 HF 的 Xet CAS——大权重走 Xet 会直连
# cas-server.xethub.hf.co 并返 4012026-07-19 teacher 4B 下载实撞)。禁用 Xet
# 退回经典 HTTP/LFS 下载(镜像支持)。若仍不行:pip uninstall hf_xet
export HF_HUB_DISABLE_XET=1
torchrun --nproc_per_node=4 --master_port=29572 scripts/train_whitebox.py "$MODE"
+49
View File
@@ -0,0 +1,49 @@
"""SFTConfig 的构造校验测试(层 1 / T1)。
只测"配置错误必须在构造时炸"这一条约定;参数语义本身没有逻辑可测。
"""
import dataclasses
import pytest
from ars_opd.configs import SFTConfig
def make(**overrides):
"""最小合法配置;单测只关心被覆盖的那个字段。"""
base = dict(dataset_path="dummy.parquet", output_dir="/tmp/dummy")
base.update(overrides)
return SFTConfig(**base)
def test_合法配置可构造():
cfg = make()
assert cfg.max_length > cfg.max_prompt_length
def test_prompt预算吞掉总预算时报错():
# 这是最危险的静默失败:completion 预算为 0 → labels 全 -100 → loss 恒 0
with pytest.raises(ValueError, match="max_prompt_length"):
make(max_prompt_length=4096, max_length=4096)
def test_非法学习率报错():
with pytest.raises(ValueError, match="learning_rate"):
make(learning_rate=0.0)
def test_非法子集大小报错():
with pytest.raises(ValueError, match="subset_size"):
make(subset_size=0)
def test_非法max_steps报错():
with pytest.raises(ValueError, match="max_steps"):
make(max_steps=0)
def test_配置冻结不可变():
cfg = make()
with pytest.raises(dataclasses.FrozenInstanceError):
cfg.learning_rate = 1e-3
+318
View File
@@ -0,0 +1,318 @@
"""层 1 / T3:数据管线单测(docs/02 §5.1 规定的验证项)。
用字符级玩具 tokenizer 在 CPU 上对拍 collator 行为,不依赖网络下载真模型。
玩具模板刻意模仿 Qwen3 的关键结构:生成引导符 + no-think 时注入空思考块。
"""
import json
import pytest
from datasets import Dataset
from ars_opd.data import (
IGNORE_INDEX,
SFTCollator,
attach_teacher_completions,
prompt_key,
to_messages,
)
class ToyTokenizer:
"""字符级 tokenizer:一个字符一个 tokenid = 码点)。
模板契约与真 chat 模板同构:
- 每轮渲染成 "[role]content"
- assistant 轮(或生成引导符后)在 no-think 模式下注入 "<T></T>"(模仿
Qwen3 的空 <think>\\n\\n</think>);
- 完整渲染 == prompt 渲染 + 解答文本,保证边界可精确断言。
"""
def __init__(self, pad_token_id=0, eos_token_id=1):
self.pad_token_id = pad_token_id
self.eos_token_id = eos_token_id
def apply_chat_template(
self,
messages,
tokenize=False,
add_generation_prompt=False,
enable_thinking=False,
):
think = "" if enable_thinking else "<T></T>"
parts = []
for m in messages:
prefix = think if m["role"] == "assistant" else ""
parts.append(f"[{m['role']}]{prefix}{m['content']}")
text = "".join(parts)
if add_generation_prompt:
text += f"[assistant]{think}"
return text
def __call__(
self,
text,
truncation=False,
max_length=None,
padding=False,
add_special_tokens=False,
):
ids = [ord(c) for c in text]
if truncation and max_length is not None:
ids = ids[:max_length]
return {"input_ids": ids}
def ids_of(text):
return [ord(c) for c in text]
def row(question, answer=None):
msgs = [{"role": "user", "content": question}]
if answer is not None:
msgs.append({"role": "assistant", "content": answer})
return {"messages": msgs}
# ---------------------------------------------------------------------------
# to_messages
# ---------------------------------------------------------------------------
def test_dapo_prompt列直接归一():
ex = {"prompt": [{"role": "user", "content": "1+1=?"}], "data_source": "dapo"}
assert to_messages(ex) == {"messages": [{"role": "user", "content": "1+1=?"}]}
def test_字符串化的列表被还原():
ex = {"prompt": "[{'role': 'user', 'content': 'hi'}]"}
assert to_messages(ex)["messages"] == [{"role": "user", "content": "hi"}]
def test_坏字符串显式报错而非静默放行():
# 参考实现 except:pass 会让这行以字符串形态流进 collator
with pytest.raises(ValueError, match="无法解析"):
to_messages({"prompt": "[{'role': broken"})
def test_question列包成单user轮():
assert to_messages({"question": "2+2=?"}) == {
"messages": [{"role": "user", "content": "2+2=?"}]
}
def test_无法识别的行报错():
with pytest.raises(ValueError, match="无法识别"):
to_messages({"foo": 1})
# ---------------------------------------------------------------------------
# prompt_key(与 teacher.py 的缓存契约)
# ---------------------------------------------------------------------------
def test_同题同键_不同题不同键():
m1 = [{"role": "user", "content": "q"}]
m2 = [{"role": "user", "content": "q'"}]
assert prompt_key(m1) == prompt_key(m1)
assert prompt_key(m1) != prompt_key(m2)
def test_额外元数据字段不影响键():
plain = [{"role": "user", "content": "q"}]
noisy = [{"role": "user", "content": "q", "source": "dapo"}]
assert prompt_key(plain) == prompt_key(noisy)
# ---------------------------------------------------------------------------
# attach_teacher_completions
# ---------------------------------------------------------------------------
def write_cache(path, entries):
with open(path, "w", encoding="utf-8") as f:
for msgs, completion in entries:
f.write(
json.dumps({"key": prompt_key(msgs), "completion": completion}) + "\n"
)
def test_挂接teacher解答(tmp_path):
q = [{"role": "user", "content": "1+1=?"}]
cache = tmp_path / "cache.jsonl"
write_cache(cache, [(q, "答案是 2")])
ds = Dataset.from_list([{"messages": q}])
out = attach_teacher_completions(ds, str(cache))
assert out[0]["messages"][-1] == {"role": "assistant", "content": "答案是 2"}
def test_缓存缺键一次性报全部缺失(tmp_path):
cache = tmp_path / "cache.jsonl"
write_cache(cache, [])
ds = Dataset.from_list([row("q1"), row("q2")])
with pytest.raises(KeyError, match="2/2"):
attach_teacher_completions(ds, str(cache))
def test_自带解答的行不被覆盖(tmp_path):
cache = tmp_path / "cache.jsonl"
write_cache(cache, [])
ds = Dataset.from_list([row("q", "人写的答案")])
out = attach_teacher_completions(ds, str(cache))
assert out[0]["messages"][-1]["content"] == "人写的答案"
def test_缓存文件不存在报错():
ds = Dataset.from_list([row("q")])
with pytest.raises(FileNotFoundError):
attach_teacher_completions(ds, "/不存在/cache.jsonl")
# ---------------------------------------------------------------------------
# SFTCollator
# ---------------------------------------------------------------------------
def make_collator(**kw):
defaults = dict(max_length=1000, max_prompt_length=100, enable_thinking=False)
defaults.update(kw)
return SFTCollator(ToyTokenizer(), **defaults)
def test_基本形态_掩码与边界():
collator = make_collator()
batch = collator([row("ab", "cd")])
prompt_text = "[user]ab[assistant]<T></T>"
labels = batch["labels"][0].tolist()
# prompt 全 -100completion 位置是解答的 token
assert labels[: len(prompt_text)] == [IGNORE_INDEX] * len(prompt_text)
assert labels[len(prompt_text) :] == ids_of("cd")
assert batch["input_ids"][0].tolist() == ids_of(prompt_text + "cd")
assert batch["attention_mask"][0].tolist() == [1] * (len(prompt_text) + 2)
def test_超长解答不挤占prompt():
# 头号正确性卖点:completion 被截,prompt 一个 token 不少
prompt_text = "[user]ab[assistant]<T></T>"
collator = make_collator(max_length=len(prompt_text) + 3)
batch = collator([row("ab", "x" * 50)])
input_ids = batch["input_ids"][0].tolist()
assert input_ids[: len(prompt_text)] == ids_of(prompt_text) # prompt 完整
assert len(input_ids) == len(prompt_text) + 3 # completion 只剩预算内 3 个
def test_超长题目截断但边界不错位():
# 坑二场景:prompt 超预算被截断,completion 的 token 必须仍然精确
# (切分点用未截断长度,而非截断后长度)
collator = make_collator(max_prompt_length=10)
batch = collator([row("很长的题目" * 20, "答案")])
labels = batch["labels"][0].tolist()
non_masked = [t for t in labels if t != IGNORE_INDEX]
assert non_masked == ids_of("答案") # 解答 token 一个不错
assert sum(t == IGNORE_INDEX for t in labels) == 10 # prompt 恰被截到预算
def test_enable_thinking两种取值边界都正确():
for thinking in (False, True):
collator = make_collator(enable_thinking=thinking)
batch = collator([row("q", "ans")])
non_masked = [t for t in batch["labels"][0].tolist() if t != IGNORE_INDEX]
assert non_masked == ids_of("ans"), f"enable_thinking={thinking} 时边界错位"
def test_nothink模板注入空思考块():
# 参考实现的一次性诊断打印,在这里变成永久契约
text = ToyTokenizer().apply_chat_template(
[{"role": "user", "content": "q"}],
add_generation_prompt=True,
enable_thinking=False,
)
assert text.endswith("<T></T>")
def test_sft模式仍拒绝prompt_only行():
# 回归守卫(docs/03 §5 U3"SFT 路径行为不变"):SFT 模式下 prompt-only 行
# 仍是静默空训练闸门,必须报错——放开只发生在显式 prompt_only=True 模式
with pytest.raises(ValueError, match="prompt-only"):
make_collator()([row("没有答案的题")])
def test_sft模式缺max_length构造即报错():
with pytest.raises(ValueError, match="max_length"):
SFTCollator(
ToyTokenizer(), max_prompt_length=50
) # 非 prompt_only 却无 max_length
def test_左padding对齐():
collator = make_collator()
batch = collator([row("ab", "cd"), row("a", "c")])
t = batch["input_ids"].shape[1]
short_mask = batch["attention_mask"][1].tolist()
n_pad = t - short_mask.count(1)
assert n_pad > 0
assert short_mask[:n_pad] == [0] * n_pad # padding 在左
assert batch["labels"][1].tolist()[:n_pad] == [IGNORE_INDEX] * n_pad
assert batch["input_ids"][1].tolist()[:n_pad] == [0] * n_pad # pad_token_id=0
def test_pad回退到eos():
collator = SFTCollator(
ToyTokenizer(pad_token_id=None, eos_token_id=7),
max_length=100,
max_prompt_length=50,
)
assert collator.pad_token_id == 7
# ---------------------------------------------------------------------------
# SFTCollatorprompt_only 模式(层 2 / U3
# ---------------------------------------------------------------------------
def make_prompt_collator(**kw):
defaults = dict(max_prompt_length=100, enable_thinking=False, prompt_only=True)
defaults.update(kw)
return SFTCollator(ToyTokenizer(), **defaults)
def test_prompt_only模式返回prompt张量且不报错():
# 层 2prompt-only 行是常态,不再报错;渲染带生成引导符供 generate 续写
collator = make_prompt_collator()
batch = collator([row("ab")])
expected = "[user]ab[assistant]<T></T>"
assert set(batch.keys()) == {"prompts", "prompt_attention_mask"}
assert batch["prompts"][0].tolist() == ids_of(expected)
assert batch["prompt_attention_mask"][0].tolist() == [1] * len(expected)
def test_prompt_only模式左padding对齐():
# 生成要求左 padding:短 prompt 在左侧补 pad,右边界对齐
collator = make_prompt_collator()
batch = collator([row("abc"), row("a")])
t = batch["prompts"].shape[1]
short_mask = batch["prompt_attention_mask"][1].tolist()
n_pad = t - short_mask.count(1)
assert n_pad > 0
assert short_mask[:n_pad] == [0] * n_pad # padding 在左
assert batch["prompts"][1].tolist()[:n_pad] == [0] * n_pad # pad_token_id=0
def test_prompt_only模式截断到prompt预算():
collator = make_prompt_collator(max_prompt_length=8)
batch = collator([row("很长的题目" * 20)])
assert batch["prompts"].shape[1] == 8 # 恰截到 max_prompt_length
def test_prompt_only模式剥掉末轮assistant():
# 若数据碰巧带了 assistant 轮,取生成前上下文(剥掉它再加生成引导符)
collator = make_prompt_collator()
batch = collator([row("q", "已有答案")])
expected = "[user]q[assistant]<T></T>" # 不含"已有答案"
assert batch["prompts"][0].tolist() == ids_of(expected)
+91
View File
@@ -0,0 +1,91 @@
"""层 2 / U4:生成重建纯逻辑单测(docs/03 §2.5)。
只测 build_generated_batch——把 model.generate 输出重建成 input_ids/attention/labels
的张量逻辑。DistillTrainer 的编排(生成→双前向→divergence)需真模型,由远程冒烟
验证。重点盯"首个 eos 之后一律屏蔽"这条最易 off-by-one 的规则。
"""
import torch
from ars_opd.data import IGNORE_INDEX
from ars_opd.trainer import build_generated_batch
EOS = 99
PAD = 0
def test_eos在中间_其后全屏蔽():
# 行内:prompt=[pad,u,u],生成=[a, EOS, pad]a 与 eos 有效,eos 后的 pad 无效
prompts = torch.tensor([[PAD, 5, 5]])
prompt_mask = torch.tensor([[0, 1, 1]])
gen_output = torch.tensor([[PAD, 5, 5, 7, EOS, PAD]]) # (1, P+G)=(1,6)
ids, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
assert ids.tolist() == gen_output.tolist() # input_ids 即生成全序列
# prompt 段全 -100;生成段 [7, EOS, -100]
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX]
# attentionprompt 左 pad=0,生成段 eos 及之前=1、其后=0
assert attn[0].tolist() == [0, 1, 1, 1, 1, 0]
def test_无eos撞max时整段生成有效():
prompts = torch.tensor([[5, 5, 5]])
prompt_mask = torch.tensor([[1, 1, 1]])
gen_output = torch.tensor([[5, 5, 5, 8, 9, 10]]) # 生成三 token,无 eos
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [8, 9, 10] # 全监督
assert attn[0].tolist() == [1, 1, 1, 1, 1, 1]
def test_batch内不同生成长度_各自正确对齐():
# 行0 提前 eos(右侧被补 pad 到 batch 宽度);行1 撞 max。二者共用同一 (B,P+G)
prompts = torch.tensor([[PAD, 5, 5], [5, 5, 5]])
prompt_mask = torch.tensor([[0, 1, 1], [1, 1, 1]])
gen_output = torch.tensor(
[
[PAD, 5, 5, 7, EOS, PAD], # 行0:生成 [7, EOS],末位 pad 补齐
[5, 5, 5, 8, 9, 10], # 行1:生成 [8, 9, 10]
]
)
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX]
assert labels[1].tolist() == [IGNORE_INDEX] * 3 + [8, 9, 10]
assert attn[0].tolist() == [0, 1, 1, 1, 1, 0]
assert attn[1].tolist() == [1, 1, 1, 1, 1, 1]
def test_pad等于eos也不误判():
# 关键 cornerpad_token == eos_token。首个 eos 有效、其后补位的 eos 全屏蔽
prompts = torch.tensor([[5, 5, 5]])
prompt_mask = torch.tensor([[1, 1, 1]])
gen_output = torch.tensor([[5, 5, 5, 7, EOS, EOS]]) # 末位补的 pad 恰好==eos
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX]
assert attn[0].tolist() == [1, 1, 1, 1, 1, 0] # 第二个 eos 被当补位屏蔽
def test_立即eos_只留一个token():
prompts = torch.tensor([[5, 5, 5]])
prompt_mask = torch.tensor([[1, 1, 1]])
gen_output = torch.tensor([[5, 5, 5, EOS, PAD, PAD]]) # 第一步就 eos
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [EOS, IGNORE_INDEX, IGNORE_INDEX]
assert attn[0].tolist() == [1, 1, 1, 1, 0, 0]
def test_prompt段恒为负100():
prompts = torch.tensor([[PAD, PAD, 5, 5]])
prompt_mask = torch.tensor([[0, 0, 1, 1]])
gen_output = torch.tensor([[PAD, PAD, 5, 5, 8, 9]])
_, _, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
assert labels[0, :4].tolist() == [IGNORE_INDEX] * 4 # prompt 段(含左 pad)全 -100
+131
View File
@@ -0,0 +1,131 @@
"""层 2 / U2token 级散度单测(docs/03 §2、§4.1、§6.1)。
只测纯张量函数 token_divergenceDistillTrainer 类由远程冒烟验证。
三块:
1. 对拍 PyTorch 自带 KL(独立 oracle,非同式自证);
2. 反向/前向 KL 的方向性(mode-seeking vs mode-covering);
3. §4.1 梯度爆炸演示——teacher 概率趋 0 时 student 梯度暴涨(层 5 有界乘子的对照桩)。
构造技巧:softmax(log p) = pp 已归一),故用 `probs.log()` 当 logits 即可精确
控制两侧分布,让手算/对拍成为可能。
"""
import pytest
import torch
from torch.distributions import Categorical, kl_divergence
from ars_opd.data import IGNORE_INDEX
from ars_opd.trainer import token_divergence
V = 5 # 玩具词表
def logits_of(probs: list[float]) -> torch.Tensor:
"""概率向量 -> (1, 1, V) logits,使 log_softmax 后精确还原该分布。"""
return torch.tensor(probs).log().reshape(1, 1, V)
ONE_VALID = torch.zeros(1, 1, dtype=torch.long) # 单个有效 tokenid 0 ≠ -100
# ---- 1. 对拍 PyTorch KL ----
@pytest.mark.parametrize("beta", [0.0, 1.0, 0.5])
def test_散度对拍pytorch_kl(beta):
p_s = [0.10, 0.20, 0.30, 0.25, 0.15]
p_t = [0.05, 0.05, 0.40, 0.40, 0.10]
loss, num = token_divergence(logits_of(p_s), logits_of(p_t), ONE_VALID, beta=beta)
assert num == 1
cs, ct = Categorical(torch.tensor(p_s)), Categorical(torch.tensor(p_t))
if beta == 1.0: # 反向 KL(π_θ‖π_T)
oracle = kl_divergence(cs, ct)
elif beta == 0.0: # 前向 KL(π_T‖π_θ)
oracle = kl_divergence(ct, cs)
else: # JSD:对混合分布的两支 KL 加权
m = Categorical((1 - beta) * torch.tensor(p_s) + beta * torch.tensor(p_t))
oracle = beta * kl_divergence(ct, m) + (1 - beta) * kl_divergence(cs, m)
assert torch.allclose(loss, oracle, atol=1e-6)
def test_同分布散度为零():
p = [0.1, 0.2, 0.3, 0.25, 0.15]
for beta in (0.0, 1.0, 0.5):
loss, _ = token_divergence(logits_of(p), logits_of(p), ONE_VALID, beta=beta)
assert torch.allclose(loss, torch.zeros(()), atol=1e-6)
# ---- 2. 方向性:反向罚"越界",前向罚"漏覆盖" ----
def test_kl方向性():
peaked = [0.90, 0.025, 0.025, 0.025, 0.025]
diffuse = [1 / V] * V
# 情形 Ateacher 尖、student 弥散——student 把质量放到 teacher≈0 处。
# 反向 KL(π_θ‖π_T) 因 log(π_θ/π_T) 在越界 token 上爆大而重罚;前向相对轻。
rev_A = token_divergence(
logits_of(diffuse), logits_of(peaked), ONE_VALID, beta=1.0
)[0]
fwd_A = token_divergence(
logits_of(diffuse), logits_of(peaked), ONE_VALID, beta=0.0
)[0]
assert rev_A > fwd_A # 反向惩罚 student 越出 teacher 支持集(mode-seeking
# 情形 Bstudent 尖、teacher 弥散——teacher 的质量落在 student≈0 处。
# 前向 KL(π_T‖π_θ) 重罚"漏覆盖";反向相对轻。
rev_B = token_divergence(
logits_of(peaked), logits_of(diffuse), ONE_VALID, beta=1.0
)[0]
fwd_B = token_divergence(
logits_of(peaked), logits_of(diffuse), ONE_VALID, beta=0.0
)[0]
assert fwd_B > rev_B # 前向惩罚 student 没覆盖 teacher 的质量(mode-covering
def test_温度升高软化分布降低反向kl():
# student 与 teacher 都尖但尖在不同 token;升温软化两侧 → 反向 KL 下降
s, t = [0.90, 0.025, 0.025, 0.025, 0.025], [0.025, 0.90, 0.025, 0.025, 0.025]
cold = token_divergence(logits_of(s), logits_of(t), ONE_VALID, temperature=1.0)[0]
hot = token_divergence(logits_of(s), logits_of(t), ONE_VALID, temperature=4.0)[0]
assert hot < cold
# ---- 3. §4.1 梯度爆炸演示 ----
def test_梯度爆炸_teacher概率趋0时student梯度暴涨():
# student 固定:对"采样 token"(id 0) 给最高 logit(模拟 on-policy 采到它)
base = [1.5, 0.5, 0.3, 0.2, 0.1]
epsilons = [1e-1, 1e-2, 1e-3, 1e-4, 1e-5, 1e-6]
grad_norms = []
for eps in epsilons:
student_logits = torch.tensor(base).reshape(1, 1, V).requires_grad_(True)
# teachertoken0 概率 = eps(越来越"厌恶"它),其余 (1-eps) 均分
t_probs = [(1 - eps) / (V - 1)] * V
t_probs[0] = eps
teacher_logits = torch.tensor(t_probs).log().reshape(1, 1, V)
loss, _ = token_divergence(student_logits, teacher_logits, ONE_VALID, beta=1.0)
loss.backward()
grad_norms.append(student_logits.grad.norm().item())
# 单调暴涨:teacher 越否定采样 tokenstudent 梯度范数越大
for lo, hi in zip(grad_norms, grad_norms[1:]):
assert hi > lo
# 末端(π_T=1e-6)远超首端(π_T=1e-1)——§4.1 的可执行证据,
# 为层 5"有界乘子 π̂"的稳定性对照埋桩
assert grad_norms[-1] > 5 * grad_norms[0]
def test_全掩码batch显式报错():
all_masked = torch.full((1, 1), IGNORE_INDEX, dtype=torch.long)
with pytest.raises(ValueError, match="有效 completion"):
token_divergence(logits_of([0.2] * V), logits_of([0.2] * V), all_masked)
def test_beta越界报错():
with pytest.raises(ValueError, match="beta"):
token_divergence(
logits_of([0.2] * V), logits_of([0.2] * V), ONE_VALID, beta=2.0
)
+162
View File
@@ -0,0 +1,162 @@
"""estimator.py 单测——docs/04 §5.2:公式对拍 + 定理 4.1 性质 + 方差收缩。
对拍精神源自参考实现 validate_chunk_mc_estimator.py(比 MSE_freq vs
MSE_bayes),但全用 toy 数据本地 CPU 跑,不连真 teacher。
detach 命门的两个世界断言在 tests/test_estimator_detach.py(E3)。
"""
import math
import pytest
import torch
from ars_opd.estimator import bayesian_target, chunk_prior
# ------------------------------------------------------------- chunk_prior
def test_prior_is_geometric_mean():
# 式(4) 手算:p = [0.9, 0.1] → π̄ = exp((log .9 + log .1)/2) = √0.09 = 0.3
log_probs = torch.log(torch.tensor([0.9, 0.1]))
assert math.isclose(chunk_prior(log_probs).item(), 0.3, rel_tol=1e-6)
def test_prior_uniform_probs():
# 全同概率的几何均值 = 该概率本身
log_probs = torch.full((50,), math.log(0.5))
assert math.isclose(chunk_prior(log_probs).item(), 0.5, rel_tol=1e-6)
def test_prior_shape_and_range():
pi_bar = chunk_prior(torch.log(torch.rand(50).clamp(1e-6, 1.0)))
assert pi_bar.shape == () # (C,) -> 标量
assert 0.0 < pi_bar.item() <= 1.0
def test_prior_log_domain_survives_underflow():
# 50 个 p=0.01 直接连乘 = 1e-100(fp32 下溢为 0);log 域算出 0.01
log_probs = torch.full((50,), math.log(0.01))
assert math.isclose(chunk_prior(log_probs).item(), 0.01, rel_tol=1e-4)
def test_prior_clamp_floor():
# 极端负 log 均值 → exp 下溢,clamp 兜到 1e-8 保持严格为正(定理 4.1b 前提)
log_probs = torch.full((5,), -1e9)
assert chunk_prior(log_probs).item() == pytest.approx(1e-8)
def test_prior_is_detached():
# detach 命门:π̄ 不带梯度(逃逸机制的完整断言在 test_estimator_detach.py)
log_probs = torch.log(torch.tensor([0.5, 0.5], requires_grad=True))
pi_bar = chunk_prior(log_probs)
assert not pi_bar.requires_grad
def test_prior_empty_raises():
with pytest.raises(ValueError, match="为空"):
chunk_prior(torch.tensor([]))
# --------------------------------------------------------- bayesian_target
def test_target_formula_hand_computed():
# 式(5) 手算:k=8, π̄=0.5, N=10, α=1 → π̂ = (8 + 0.5)/11 = 0.77272…
pi_hat = bayesian_target(8.0, torch.tensor(0.5), n_rollouts=10, alpha=1.0)
assert math.isclose(pi_hat.item(), 8.5 / 11, rel_tol=1e-6)
def test_target_convex_combination_identity():
# 式(10) 恒等:π̂ = N/(N+α)·(k/N) + α/(N+α)·π̄,任取参数逐点核对
k, pi_bar, n, alpha = 3.7, torch.tensor(0.42), 10, 1.5
direct = bayesian_target(k, pi_bar, n, alpha).item()
convex = (n / (n + alpha)) * (k / n) + (alpha / (n + alpha)) * pi_bar.item()
assert math.isclose(direct, convex, rel_tol=1e-6)
def test_target_anti_collapse_at_k_zero():
# 定理 4.1(b):k=0(teacher 全否定)时 π̂ = α·π̄/(N+α) > 0,监督不归零
pi_hat = bayesian_target(0.0, torch.tensor(0.3), n_rollouts=10, alpha=1.0)
assert math.isclose(pi_hat.item(), 0.3 / 11, rel_tol=1e-6)
assert pi_hat.item() > 0
def test_target_full_score_shrinks_below_one():
# 贝叶斯收缩:k=N 满分时 π̂ = (N+α·π̄)/(N+α) < 1(只要 π̄<1)——
# 先验把估计从两端往中间拉,这正是方差收缩的来源
pi_hat = bayesian_target(10.0, torch.tensor(0.5), n_rollouts=10, alpha=1.0)
assert math.isclose(pi_hat.item(), 10.5 / 11, rel_tol=1e-6)
assert pi_hat.item() < 1.0
def test_target_bounded_in_unit_interval():
# 定理 4.1(a):任意合法参数下 π̂ ∈ (0, 1]
for k in [0.0, 2.5, 10.0]:
for p in [1e-8, 0.5, 1.0]:
v = bayesian_target(k, torch.tensor(p), 10, 1.0).item()
assert 0.0 < v <= 1.0
def test_target_alpha_zero_is_frequency_estimate():
# α=0 退化为 k/N(no_bayesian 消融);k=0 时被 clamp 兜到 1e-8 而非 0
assert math.isclose(
bayesian_target(7.0, torch.tensor(0.5), 10, 0.0).item(), 0.7, rel_tol=1e-6
)
assert bayesian_target(0.0, torch.tensor(0.5), 10, 0.0).item() == pytest.approx(
1e-8
)
def test_target_is_detached_even_with_grad_input():
# 第二道防线:pi_bar 带梯度传入,π̂ 仍必须 detach
pi_bar = torch.tensor(0.5, requires_grad=True)
pi_hat = bayesian_target(5.0, pi_bar, 10, 1.0)
assert not pi_hat.requires_grad
def test_target_validation_raises():
pi_bar = torch.tensor(0.5)
with pytest.raises(ValueError, match="n_rollouts"):
bayesian_target(0.0, pi_bar, 0, 1.0)
with pytest.raises(ValueError, match="alpha"):
bayesian_target(0.0, pi_bar, 10, -0.1)
with pytest.raises(ValueError, match="越界"):
bayesian_target(11.0, pi_bar, 10, 1.0) # k_sem > N:口径不一致
with pytest.raises(ValueError, match="越界"):
bayesian_target(-0.5, pi_bar, 10, 1.0)
# --------------------------------------------- 方差收缩(定理 4.1c,toy 模拟)
def test_variance_shrinkage_beats_frequency_estimate():
"""toy 模拟对拍 validate_chunk_mc_estimator.py 的 MSE_freq vs MSE_bayes。
设真值 μ:每次试验采 N=10 个相似度 sim_i(均值 μ 的噪声),
频率估计 = mean(sim) = k/N,贝叶斯估计 = (k + α·π̄)/(N+α)。
先验 π̄ = μ(理想先验)时,收缩纯降方差、零偏差代价,MSE 必更小。
"""
torch.manual_seed(0)
mu, n, alpha = 0.7, 10, 1.0
trials = 2000
# (trials, N) 的相似度样本:均值 μ、截断到 [0,1]
sims = (mu + 0.25 * torch.randn(trials, n)).clamp(0.0, 1.0)
k = sims.sum(dim=1) # (trials,) 每次试验的 k_sem
freq = k / n
bayes = (k + alpha * mu) / (n + alpha)
mse_freq = ((freq - mu) ** 2).mean().item()
mse_bayes = ((bayes - mu) ** 2).mean().item()
assert mse_bayes < mse_freq
def test_variance_shrinkage_robust_to_imperfect_prior():
# 先验偏离真值(π̄ = μ±0.1)仍应赢:α=1、N=10 时先验权重仅 1/11,
# 引入的偏差平方远小于省下的方差(定理 4.1c 在论文设定下的稳健性)
torch.manual_seed(1)
mu, n, alpha, trials = 0.6, 10, 1.0, 2000
sims = (mu + 0.25 * torch.randn(trials, n)).clamp(0.0, 1.0)
k = sims.sum(dim=1)
mse_freq = ((k / n - mu) ** 2).mean().item()
for prior in [mu - 0.1, mu + 0.1]:
bayes = (k + alpha * prior) / (n + alpha)
assert ((bayes - mu) ** 2).mean().item() < mse_freq
+142
View File
@@ -0,0 +1,142 @@
"""similarity.py 单测——docs/04 §5.1:手构字符串钉死 φ 与 k_sem(式3)。
全部本地 CPU、纯标准库,不依赖 torch/transformers。
关键手算用例在各测试的注释里逐步展开,方便对着验算。
"""
import math
import pytest
from ars_opd.similarity import aggregate_similarity, edit_similarity, phi, rouge1
# ---------------------------------------------------------------- rouge1
def test_rouge1_identical_is_exact_one():
# 全同串必须精确 = 1.0(参考实现因分母 +1e-8 只能得 ≈0.99999998)
s = "so x = 5 and y = 12"
assert rouge1(s, s) == 1.0
def test_rouge1_disjoint_is_zero():
assert rouge1("a b c", "x y z") == 0.0
def test_rouge1_partial_overlap_hand_computed():
# hyp = {a, b, c}, ref = {a, b, d}:overlap = 2
# precision = 2/3, recall = 2/3, F1 = 2·(2/3)(2/3) / (4/3) = 2/3
assert math.isclose(rouge1("a b c", "a b d"), 2 / 3)
def test_rouge1_multiset_counts_repeats():
# 多重集语义:hyp = [x,x,x,x], ref = [x] → overlap = min(4,1) = 1
# precision = 1/4, recall = 1/1, F1 = 2·(1/4)/(5/4) = 0.4
# (参考实现的 set 版会给满分 1.0——数学文本重复词多,这是关键失真点)
assert math.isclose(rouge1("x x x x", "x"), 0.4)
def test_rouge1_is_bag_of_words_order_blind():
# 词袋:只看用了哪些词,不看顺序
assert rouge1("a b", "b a") == 1.0
def test_rouge1_empty_sides():
assert rouge1("", "a b") == 0.0
assert rouge1("a b", "") == 0.0
assert rouge1("", "") == 0.0
assert rouge1(" ", "a") == 0.0 # 纯空白 split 后无词
# ---------------------------------------------------------- edit_similarity
def test_edit_identical_is_one():
s = "so x = 5 and y = 12"
assert edit_similarity(s, s) == 1.0
def test_edit_totally_different_is_zero():
# ["a","b"] vs ["c","d"]:2 次替换,dist=2, max(m,n)=2 → 1 1 = 0
assert edit_similarity("a b", "c d") == 0.0
def test_edit_single_substitution_hand_computed():
# ["a","b","c"] vs ["a","x","c"]:1 次替换,dist=1, max=3 → 2/3
assert math.isclose(edit_similarity("a b c", "a x c"), 2 / 3)
def test_edit_insertion_hand_computed():
# ["a","b"] vs ["a","x","b"]:1 次插入,dist=1, max=3 → 2/3
assert math.isclose(edit_similarity("a b", "a x b"), 2 / 3)
def test_edit_is_order_sensitive():
# ["a","b"] vs ["b","a"]:两次替换 dist=2 → 0.0;与 rouge1 的 1.0 互补
assert edit_similarity("a b", "b a") == 0.0
assert rouge1("a b", "b a") == 1.0
def test_edit_empty_sides():
assert edit_similarity("", "") == 1.0 # 零距离
assert edit_similarity("a b", "") == 0.0 # 全删
assert edit_similarity("", "a b") == 0.0 # 全插
def test_edit_asymmetric_lengths():
# ["a"] vs ["a","b","c","d"]:3 次插入,dist=3, max=4 → 1/4
assert math.isclose(edit_similarity("a", "a b c d"), 1 / 4)
# --------------------------------------------------------------------- phi
def test_phi_default_is_edit_distance():
# 论文 §5.1 默认;"a b" vs "b a" 恰能区分两度量(edit=0, rouge1=1)
assert phi("a b", "b a") == edit_similarity("a b", "b a") == 0.0
def test_phi_dispatch():
h, r = "a b c", "a b d"
assert phi(h, r, metric="rouge1") == rouge1(h, r)
assert phi(h, r, metric="edit_distance") == edit_similarity(h, r)
def test_phi_unknown_metric_raises():
with pytest.raises(ValueError, match="bleu"):
phi("a", "a", metric="bleu")
# ----------------------------------------------------- aggregate_similarity
def test_aggregate_is_sum_of_phi():
# 式(3) 手算:rollouts 与 "a b c" 的 edit 相似度分别为 1.0, 2/3, 0.0
chunk = "a b c"
rollouts = ["a b c", "a x c", "x y z"]
expected = 1.0 + 2 / 3 + 0.0
assert math.isclose(aggregate_similarity(chunk, rollouts), expected)
def test_aggregate_bounds():
# k_sem ∈ [0, N]:全同 → N,全不同 → 0
n = 5
assert aggregate_similarity("a b", ["a b"] * n) == float(n)
assert aggregate_similarity("a b", ["x y"] * n) == 0.0
def test_aggregate_is_continuous_soft_count():
# φ 连续 ⇒ k_sem 非整数是常态(区别于 token 精确匹配的硬计数)
k = aggregate_similarity("a b c", ["a b c", "a x c"])
assert 1.0 < k < 2.0
def test_aggregate_empty_rollouts():
assert aggregate_similarity("a b", []) == 0.0
def test_aggregate_metric_passthrough():
# "a b" vs "b a":edit 全零,rouge1 全满——验证 metric 真的传下去了
chunk, rollouts = "a b", ["b a", "b a"]
assert aggregate_similarity(chunk, rollouts, metric="edit_distance") == 0.0
assert aggregate_similarity(chunk, rollouts, metric="rouge1") == 2.0
+143
View File
@@ -0,0 +1,143 @@
"""层 1 / T2:teacher 批量生成与缓存单测。
用假 OpenAI 客户端注入(TeacherClient 的测试口),验证思考段剥离、缓存契约
(与 data.attach_teacher_completions 的端到端闭环)、断点续传、失败汇总。
"""
from types import SimpleNamespace
import pytest
from datasets import Dataset
from ars_opd.configs import TeacherGenConfig
from ars_opd.data import attach_teacher_completions, prompt_key
from ars_opd.teacher import TeacherClient, _load_teacher_env, generate_completions
class FakeClient:
"""最小 OpenAI 客户端替身:chat.completions.create 按 responder 出内容。"""
def __init__(self, responder):
self.calls = []
self._responder = responder
self.chat = SimpleNamespace(completions=SimpleNamespace(create=self._create))
def _create(self, model, messages, **kwargs):
self.calls.append(messages)
content = self._responder(messages)
return SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content=content))]
)
def make_teacher(responder, **cfg_overrides):
cfg = TeacherGenConfig(**cfg_overrides)
return TeacherClient(cfg, client=FakeClient(responder), model="fake-m3")
def user(q):
return [{"role": "user", "content": q}]
# ---------------------------------------------------------------------------
# TeacherClient.generate
# ---------------------------------------------------------------------------
def test_剥离开头思考段():
teacher = make_teacher(lambda m: "<think>心算一下</think>\n答案是 42")
assert teacher.generate(user("q")) == "答案是 42"
def test_正文中的think字样不误删():
teacher = make_teacher(lambda m: "<think>x</think>正文提到 <think> 标签本身")
assert teacher.generate(user("q")) == "正文提到 <think> 标签本身"
def test_只剩思考段等于空解答_报错():
teacher = make_teacher(lambda m: "<think>思考被截断在半途")
# 未闭合的 think 段剥不掉,但闭合后为空的要报错
teacher_empty = make_teacher(lambda m: "<think>只有思考</think> ")
with pytest.raises(ValueError, match="空解答"):
teacher_empty.generate(user("q"))
# 未闭合时保留原文(宁可保留可疑内容也不静默删成空)
assert "<think>" in teacher.generate(user("q"))
def test_关闭strip_think则原样保留():
teacher = make_teacher(lambda m: "<think>a</think>b", strip_think=False)
assert teacher.generate(user("q")) == "<think>a</think>b"
def test_system_prompt前置():
teacher = make_teacher(lambda m: "ok", system_prompt="你是数学助教")
teacher.generate(user("q"))
sent = teacher.client.calls[0]
assert sent[0] == {"role": "system", "content": "你是数学助教"}
assert sent[1]["role"] == "user"
def test_注入client但不给model报错():
with pytest.raises(ValueError, match="model"):
TeacherClient(TeacherGenConfig(), client=FakeClient(lambda m: "x"), model=None)
# ---------------------------------------------------------------------------
# generate_completions:缓存契约与断点续传
# ---------------------------------------------------------------------------
def test_端到端契约_生成的缓存能被attach消费(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("1+1=?"), user("2+2=?")]
teacher = make_teacher(lambda m: f"对「{m[-1]['content']}」的解答")
generate_completions(prompts, cache, teacher)
ds = Dataset.from_list([{"messages": p} for p in prompts])
out = attach_teacher_completions(ds, cache)
assert out[0]["messages"][-1]["content"] == "对「1+1=?」的解答"
assert out[1]["messages"][-1]["content"] == "对「2+2=?」的解答"
def test_断点续传_已缓存的不重新生成(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("q1"), user("q2")]
teacher = make_teacher(lambda m: "a")
generate_completions([prompts[0]], cache, teacher)
assert len(teacher.client.calls) == 1
generate_completions(prompts, cache, teacher) # q1 命中缓存
assert len(teacher.client.calls) == 2 # 只多了 q2 一次调用
def test_单条失败_其余落盘_结束时汇总报错(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("好题"), user("坏题")]
def responder(m):
if m[-1]["content"] == "坏题":
raise RuntimeError("网关 500")
return "解答"
teacher = make_teacher(responder)
with pytest.raises(RuntimeError, match="1/2"):
generate_completions(prompts, cache, teacher)
# 成功的那条已经在缓存里,重跑只会补坏题
from ars_opd.teacher import _cached_keys
from pathlib import Path
assert _cached_keys(Path(cache)) == {prompt_key(prompts[0])}
# ---------------------------------------------------------------------------
# .env 读取
# ---------------------------------------------------------------------------
def test_env缺失显式报错(monkeypatch):
for name in ("TEACHER_API_BASE", "TEACHER_API_KEY", "TEACHER_MODEL"):
monkeypatch.delenv(name, raising=False)
with pytest.raises(ValueError, match="TEACHER_API_BASE"):
_load_teacher_env(env_file="/不存在的路径/.env")
+118
View File
@@ -0,0 +1,118 @@
"""层 1 / T4:掩码 SFT 损失单测(docs/02 §2.4 的切片几何与重掩码)。
只测纯张量函数 sft_loss / compute_prompt_lengthSFTTrainer 类是 HF 接线,
由远程 sanity run 验证。张量全部手工构造,期望值可手算。
"""
import pytest
import torch
import torch.nn.functional as F
from ars_opd.data import IGNORE_INDEX
from ars_opd.trainer import compute_prompt_length, sft_loss
V = 7 # 玩具词表大小
def onehot_logits(target_ids, scale=10.0):
"""构造在 target_ids 处放尖峰的 logits。(T,) -> (T, V)"""
t = torch.tensor(target_ids)
return F.one_hot(t, V).float() * scale
def batch_of_one(input_ids, labels, attention=None):
"""单行 batch 的三件套,logits 另配。"""
ids = torch.tensor([input_ids])
lab = torch.tensor([labels])
att = torch.ones_like(ids) if attention is None else torch.tensor([attention])
return ids, lab, att
def test_prompt_length_取batch最小且不数padding():
# pad p p c c c p p p c c c
labels = torch.tensor(
[
[IGNORE_INDEX] * 3 + [5, 6, 5],
[IGNORE_INDEX] * 3 + [6, 5, 6],
]
)
attention = torch.tensor([[0, 1, 1, 1, 1, 1], [1, 1, 1, 1, 1, 1]])
# 行 0 有效长 5、completion 3 → prompt 2;行 1 是 6-3=3batch 最小 = 2
assert compute_prompt_length(attention, labels) == 2
def test_损失与手算交叉熵一致():
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX, IGNORE_INDEX, 5, 6])
torch.manual_seed(0)
logits = torch.randn(1, 4, V)
loss, num = sft_loss(logits, ids, lab, att)
# pl=2:位置 1、2 的 logit 分别预测位置 2、3 的 token(5 和 6
expected = F.cross_entropy(logits[0, 1:3], torch.tensor([5, 6]))
assert torch.allclose(loss, expected)
assert num == 2
def test_移位对齐_预测下一个token而非当前():
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX, IGNORE_INDEX, 5, 6])
# 位置 t 的尖峰指向位置 t+1 的 token(正确的"预测下一个")→ loss ≈ 0
next_logits = onehot_logits([4, 5, 6, 0]).unsqueeze(0)
loss_next, _ = sft_loss(next_logits, ids, lab, att)
# 位置 t 的尖峰指向位置 t 自己的 token(错误的"复读当前")→ loss 大
self_logits = onehot_logits([3, 4, 5, 6]).unsqueeze(0)
loss_self, _ = sft_loss(self_logits, ids, lab, att)
assert loss_next.item() < 0.01
assert loss_self.item() > 5.0
def test_窗口外的logits不影响损失():
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX, IGNORE_INDEX, 5, 6])
torch.manual_seed(0)
logits = torch.randn(1, 4, V)
loss_base, _ = sft_loss(logits, ids, lab, att)
# pl=2 → 用到的窗口是位置 [1, 3);位置 0(prompt 内部)和 3(末位)不参与
perturbed = logits.clone()
perturbed[0, 0] += 100.0
perturbed[0, 3] -= 100.0
loss_pert, _ = sft_loss(perturbed, ids, lab, att)
assert torch.allclose(loss_base, loss_pert)
def test_batch_min切片漏进的prompt_token被重掩码():
# 行 0pad1 + prompt2 + comp3;行 1prompt3 + comp3 → pl = min(2,3) = 2
ids = torch.tensor([[0, 3, 4, 5, 6, 5], [3, 4, 3, 6, 5, 6]])
lab = torch.tensor(
[
[IGNORE_INDEX] * 3 + [5, 6, 5],
[IGNORE_INDEX] * 3 + [6, 5, 6],
]
)
att = torch.tensor([[0, 1, 1, 1, 1, 1], [1, 1, 1, 1, 1, 1]])
torch.manual_seed(1)
logits = torch.randn(2, 6, V)
loss_base, num = sft_loss(logits, ids, lab, att)
assert num == 6 # 两行各 3 个 completion token,漏进切片的 prompt 位不计数
# 位置 1 的 logit 预测位置 2——两行的位置 2 都在切片内但都是 -100(行 0 是
# prompt 尾、行 1 是漏进来的 prompt token)。改它不该动 loss
perturbed = logits.clone()
perturbed[:, 1] += 100.0
loss_pert, _ = sft_loss(perturbed, ids, lab, att)
assert torch.allclose(loss_base, loss_pert)
def test_全掩码batch显式报错():
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX] * 4)
with pytest.raises(ValueError, match="有效 completion"):
sft_loss(torch.randn(1, 4, V), ids, lab, att)
def test_无prompt行显式报错():
ids, lab, att = batch_of_one([3, 4], [3, 4]) # labels 全有效 → prompt 长 0
with pytest.raises(ValueError, match="prompt_length"):
sft_loss(torch.randn(1, 2, V), ids, lab, att)