# 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 DeepSeek(OpenAI 兼容);数据抽 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` 透传 + 空 `` 诊断 | trainer:252-266 | Qwen3 模板 no-think 时注入空 `\n\n`;此开关变了,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 用 `deepseek-chat`(非 reasoner,短答案省钱)、`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 从 ~2-3 稳定下降;再全量 1 epoch。W&B 或实时日志盯 `sft/loss` 与有效 token 数。 3. 关账判据:1k 子集 1 epoch 跑完,loss 曲线正常,checkpoint 能被 `from_pretrained` 加载并生成通顺文本。