20e6d97427
- 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>
88 lines
8.7 KiB
Markdown
88 lines
8.7 KiB
Markdown
# 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. 验证方式
|
||
|
||
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` 加载并生成通顺文本。
|