Files
ars-opd-rebuild/docs/02-sft-baseline.md
T
iomgaa a0faec0df7 层1: 修复远程 sanity OOM——per_device batch 8→2、累积 2→8(全局 64 不变)
根因:150k 大词表下显存大头是 (B,T,V) logits 链(fp32 ~20G@B=8)与逐层激活,
均正比于 B 而与 0.6B 参数量无关。docs/02 §2.6 旧显存估算勘误入档;
train_sft.sh 加 expandable_segments 防碎片。

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 09:21:38 -04:00

9.1 KiB
Raw Blame History

02 · 层 1:SFT 基线与数据管线

本章目标:搞懂"SFT 基线"在本项目中的确切含义,解剖参考实现的数据管线与掩码交叉熵,然后动手建起 ars_opd 的第一批模块。行号缩写:trainer: = references/ars-opd/trl/trl/experimental/distillation/distillation_trainer.pytrain_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.py129 行) trainer 内 lmbda=0 路径(真基线)
用途 数据质检:绕过全部自定义机制,拿官方 SFTTrainer 验证"数据本身可学" 论文的 SFT 基线
prompt 掩码 ——整段对话(含用户提问)都算 loss 有——prompt 位置标 -100
路由条件 独立脚本 no_teacher and lmbda==0.0trainer:2855

教学点:sanity 脚本是工具不是基线;但"先用最笨的官方管线验证数据可学,再上自定义机制"这个调试策略本身值得继承。

2.2 数据格式与加载(train_dist:284-326

  • DAPO parquet 列:['data_source','prompt','ability','reward_model','extra_info']——prompt 列名不副实,装的是完整 chat 列表 [{user},{assistant}](若已有解答)或仅 [{user}]
  • format_messagestrain_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) + completiontrainer: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_lengthbatch 最小值保证不切掉任何 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,有个著名 trickFSDP_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 都用 DDPFSDP 触发点 = 换 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() → sha256ast.literal_evalexcept: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=Falsemax_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 加载并生成通顺文本。