- 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>
8.7 KiB
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.6B;teacher 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_index,prompt 不产生 loss |
| 左 padding | trainer:309-335 | 让整个 batch 能用一个标量 prompt_length 切分(见 2.4) |
| 边界确定 | 用未截断的 prompt 重分词长度切出 completion(trainer:286-289) | 模板渲染后 prompt+completion 的拼接分词 ≠ 分开分词,必须用同一渲染再切 |
enable_thinking 透传 + 空 <think> 诊断 |
trainer:252-266 | Qwen3 模板 no-think 时注入空 <think>\n\n</think>;此开关变了,prompt/completion 边界跟着变——错一次全错 |
| prompt-only 行 | 全 -100(trainer: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-op(train_dist:390-393)。
我们的决策(2026-07-18 讨论定):0.6B 乃至 1.7B 学生 4×A800 都用 DDP(显存账:1.7B 全套 ~27G/卡,80G 舒适);FSDP 触发点 = 换 4B 学生或序列 >8K。但"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. 验证方式
- 本地(CPU):collator 单测对拍参考行为——构造超长解答样本断言 prompt 未被截空;断言 -100 位置分布;断言 enable_thinking 两种取值下边界正确。数据加载单测:断言
except:pass已变显式报错。 - 远程:先 sanity(官方 SFTTrainer 或我们管线跑 50 步)确认 loss 从 ~2-3 稳定下降;再全量 1 epoch。W&B 或实时日志盯
sft/loss与有效 token 数。 - 关账判据:1k 子集 1 epoch 跑完,loss 曲线正常,checkpoint 能被
from_pretrained加载并生成通顺文本。