docs: 第二章——层1 SFT 基线开工文档(参考实现解剖+保留/替代/删除清单+任务表);CLAUDE.md 映射表补 data.py

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-18 03:51:06 -04:00
parent 77f0419bad
commit 12e2a8b0ae
3 changed files with 91 additions and 5 deletions
+4 -4
View File
@@ -10,9 +10,9 @@
- **日期**: 2026-07-18
- **当前层**: 层 0(环境与骨架),收尾中
- **已完成**: 文档第一章精讲完毕(式 1-8 + §3.1-3.6 逐节过);detach 命门测试预置并通过;依赖清单/远程脚本/密钥模板入库;gitea 双端打通(SSH 222 端口)
- **进行中**: 本地 conda env `ars-opd` 在装依赖;远程 `setup_remote.sh` 在跑
- **层 0 关账判据**: 两端 `pytest tests/ -q` 全绿
- **下一步**: 层 1SFT 基线):解剖参考实现 `train_sft_sanity.py` 与 collator → 写 `docs/02-sft-baseline.md` → 数据管线 + 最小训练脚本(Qwen3-0.6B,远程 4 卡)
- **进行中**: 本地 env 依赖安装移交用户手动执行(CPU torch + 清华源;后台安装因网络超时反复失败);远程已验收 ✅(torch+8 卡、pytest 4/4
- **层 0 关账判据**: 两端 `pytest tests/ -q` 全绿(远程 ✅,本地待装完)
- **下一步**: 层 1 开工文档 `docs/02-sft-baseline.md` 已写好(含解剖、保留/替代/删除清单、任务表 T1-T5、默认参数提案),用户阅读后按 T1→T3→T4→T2→T5 动手
- **未精讲的文档账**: docs/01 的 §3.7(KL 锚三处实现差异)、§3.8(论文外稳定器)、§4(训练步流程走读)
- **远程磁盘备忘**: 根分区 100% 的结构性原因是 `/root/zym`(507G 历史工作区)压在根分区,建议择期整体搬迁 `/data`;临时缓解 = 清 `/tmp/pip-unpack-*`、旧 tar.gz、journal。所有新增写盘已改道 `/data/zym`
@@ -34,7 +34,7 @@
|------|------|------|
| `00-roadmap.md` | 本文 | ✅ |
| `01-paper-code-map.md` | 论文 §3-§4 精读 + 参考实现全景解剖 + 概念↔代码对照表 | ✅ |
| `02-sft-baseline.md` | 层 1SFT 与数据管线 | |
| `02-sft-baseline.md` | 层 1SFT 与数据管线 | ✅ 待读 |
| `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性 | ⬜ |
| `04-mc-estimator.md` | 层 3:MC 估计 + 贝叶斯平滑 | ⬜ |
| `05-entropy-chunking.md` | 层 4:熵调度 | ⬜ |
+85
View File
@@ -0,0 +1,85 @@
# 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 DeepSeekOpenAI 兼容);数据抽 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)。**0.6B 学生 4×A800 用 DDP 即可**(每卡放得下整模型),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` 加载并生成通顺文本。