Files
ars-opd-rebuild/docs/02-sft-baseline.md
T
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

88 lines
9.6 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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 是),自定义 compute_loss 必须返回 `sum/num_items_in_batch` 而非裸 mean,否则日志与梯度都放大"累积步数"倍(首跑 loss 7.5 ≈ 0.94×8 即此坑;修复在 trainer.py compute_loss)。
3. 关账判据:1k 子集 1 epoch 跑完,loss 曲线正常,checkpoint 能被 `from_pretrained` 加载并生成通顺文本。