Files
ars-opd-rebuild/ars_opd/configs.py
T
iomgaa e42af5256f 层2/U1: 新增 DistillConfig(式(2) 白盒蒸馏参数)+ docs/03 表述打磨
configs.py(对应 docs/03 §5 U1):
- DistillConfig 自包含、不继承 SFTConfig;teacher_model 进 config、student 留脚本
- 三处刻意缺席: 无 teacher_completions_path/max_length/top_k(现场生成+全词表)
- 两温度分名: kl_temperature(散度 softmax)vs gen_temperature(on-policy 采样)
- 默认即式(2): beta=1 反向KL、温度1、纯采样 top_p=1;lr=1e-6(论文§5.1蒸馏)
- __post_init__ 8 分支构造即校验(gen_temperature>0 护 on-policy 语义)

docs/03:
- §2.3 β 三副面孔表: 记号统一 π、补 mode-seeking 对称、附全词表/稀疏双镜像说明
- §2.3 三条实现约定(温度/log域/batchmean)由一句话拆成可扫读表格
- §3 偏差清单上方补统领抉择原则

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 03:37:44 -04:00

300 lines
15 KiB
Python
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.
"""实验配置(层 1SFTConfig / TeacherGenConfig;层 2DistillConfig)。
设计约定(对应 CLAUDE.md §2"配置显式化"):
- 所有实验参数必须是这里某个 dataclass 的字段;代码里出现魔法数字/路径即违规。
- 密钥(API key 等)不进配置类,走 `.env`(见 teacher.py)。
- 机器相关路径(数据集、输出目录)不给默认值,强制调用方显式传入——
防止参考实现里 `/fsx` 硬编码那类"在别人机器上必炸"的坑。
- 各层的 config 自包含、不互相继承:层与层是不同实验,共享基类会把它们耦合,
违背"从上读到下看懂全部流程"(CLAUDE.md §2)。字段重复是有意接受的成本。
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True)
class SFTConfig:
"""层 1 SFT 基线的全部实验参数。
论文锚点:§3.1 式(1) 的标准交叉熵 SFT;但按 §5.1 的基线定义,
训练数据是 teacher rollout(离线蒸馏),不是人写答案——所以有
`teacher_completions_path` 字段:DAPO 是 prompt-only 数据集,
解答一律来自 teacher 生成的缓存文件。
frozen=True:配置一旦构造即只读。训练中途被悄悄改掉的配置是最难
排查的 bug 来源之一;要换参数就构造一个新实例,留下明确的代码痕迹。
"""
# ---- 机器相关路径(无默认值,必须显式传入)----
dataset_path: str
"""DAPO-Math-17K 的本地 parquet 路径(文件或目录),或 HF Hub 数据集名。"""
output_dir: str
"""checkpoint 与日志输出目录(远程必须落在 /data/zym 下)。"""
# ---- 数据 ----
teacher_completions_path: str | None = None
"""teacher 解答缓存(teacher.py 生成的 JSONL)。None 表示数据集自带
assistant 轮次;若数据实际是 prompt-only 又没给此路径,data.py 会显式报错,
不做静默兜底。"""
dataset_split: str = "train"
subset_size: int | None = 1000
"""随机抽取的子集大小(控制 teacher API 成本,roadmap 定为 ~1k);None = 全量。"""
# ---- 序列双预算 ----
# 非显然约束:prompt 与 completion 必须各有独立预算。若只用一个 max_length
# 从右截断,超长解答会把 prompt 挤空,模型在"没有题目"的样本上学习解答
# ——这是参考实现 collator(trainer:267-292) 的头号正确性卖点,此处继承。
max_length: int = 4096
"""prompt + completion 的总 token 预算。"""
max_prompt_length: int = 1024
"""prompt 单独预算;completion 实际预算 = max_length - len(截断后 prompt)。"""
enable_thinking: bool = False
"""Qwen3 chat 模板的思考开关。False 时模板注入空 `<think>\\n\\n</think>`。
非显然约束:此开关改变渲染后的 prompt 文本,从而改变 prompt/completion
的 token 边界——训练与推理必须取同一值,否则掩码整体错位。"""
# ---- 优化 ----
# 差异标注:论文 §5.1 的蒸馏训练用 lr=1e-6,参考实现 SFT 默认 2e-5
# (train_distillation.py:73)。SFT 有真实 token 监督、信号密集,从参考实现取
# 2e-5;层 5 的蒸馏配置再回到论文的 1e-6。
learning_rate: float = 2e-5
per_device_train_batch_size: int = 2
"""非显然约束:别看 0.6B 小就调大它——显存大头是 (B,T,V) 的 logits 链
fp32 一份 ~20G@B=8)与逐层激活,都正比于 B 而与参数量无关;B=8 实测
爆 80G 卡(2026-07-18 远程 sanity)。"""
gradient_accumulation_steps: int = 8
"""全局 batch = 2(per_device) × 4(卡) × 8(累积) = 64,与参考实现注释的
训练规模(trainer 配置注释"global batch 64")对齐。"""
num_train_epochs: int = 1
max_steps: int = -1
""">0 时覆盖 num_train_epochs,只跑这么多步——远程 50 步 sanity 用;-1 = 按 epoch。"""
lr_scheduler_type: str = "linear"
warmup_ratio: float = 0.0
gradient_checkpointing: bool = False
"""0.6B 学生显存富余,不开(省 ~40% 显存、慢 ~30%)。注意:FSDP 下此开关
是 no-op,真正的开关是 FSDP_ACTIVATION_CHECKPOINTING 环境变量(见
scripts/ 训练脚本头部的前置块,docs/02 §2.6)。"""
bf16: bool = True
seed: int = 42
# ---- 日志与保存 ----
logging_steps: int = 1
save_steps: int = 100
save_total_limit: int = 2
report_to: str = "none"
""""none" 或 "wandb"。默认 none:本地调试不该悄悄往外发数据,远程脚本显式开。"""
def __post_init__(self) -> None:
"""构造即校验:配置错误必须在训练开始前炸,而不是跑到第一个超长样本才炸。"""
if self.max_prompt_length >= self.max_length:
raise ValueError(
f"max_prompt_length({self.max_prompt_length}) 必须小于 "
f"max_length({self.max_length}),否则 completion 预算为零,"
f"所有样本的 labels 将全为 -100,loss 恒为 0 且无报错——静默空训练。"
)
if self.learning_rate <= 0:
raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}")
if self.subset_size is not None and self.subset_size <= 0:
raise ValueError(
f"subset_size 必须为正整数或 None(全量),收到 {self.subset_size}"
)
if self.max_steps == 0 or self.max_steps < -1:
raise ValueError(
f"max_steps 只接受 -1(按 epoch)或正整数,收到 {self.max_steps}"
)
@dataclass(frozen=True)
class TeacherGenConfig:
"""teacher 批量生成(层 1 能力)的采样与执行参数。
连接信息(API 地址/密钥/模型名)不在这里——那是部署环境的事实,走 `.env`
teacher.py 读取);这里只放"换一组值就是换一个实验"的采样参数。
"""
temperature: float = 1.0
top_p: float = 0.95
"""MiniMax M 系官方推荐采样参数:temperature=1.0, top_p=0.95。"""
max_tokens: int = 16384
"""teacher 单条回复的 token 上限。这是上限不是目标——按实际生成量计费,
放大它不增加正常解答的成本,只给最难的题留出写完的空间(8192 时 59 条实测
截断 2 条)。非显然约束:M3 的思考段也计入此额度,设太小会把解答挤没。"""
strip_think: bool = True
"""剥离 content 开头的 <think>...</think> 思考段。SFT 的监督目标是最终
解答;student 以 enable_thinking=False 训练,学思考段会与模板约定矛盾。"""
concurrency: int = 16
"""并发请求数(线程池大小)。上限看网关的承受力,报 429 就调小。"""
max_retries: int = 3
"""单请求的网络级重试次数(openai 客户端内建指数退避)。"""
system_prompt: str | None = None
"""None = 不加 system 轮(DAPO 题面自带作答指令,不需要额外指挥)。"""
def __post_init__(self) -> None:
if self.max_tokens <= 0:
raise ValueError(f"max_tokens 必须为正,收到 {self.max_tokens}")
if self.concurrency < 1:
raise ValueError(f"concurrency 必须 ≥1,收到 {self.concurrency}")
if self.temperature < 0:
raise ValueError(f"temperature 必须 ≥0,收到 {self.temperature}")
@dataclass(frozen=True)
class DistillConfig:
"""层 2 white-box OPD 基线(token 级反向 KL 蒸馏)的全部实验参数。
论文锚点:§3.1 式(2) 的 on-policy 白盒蒸馏 L = E_{y~π_θ}[Σ_t KL(π_θ ‖ π_T)]。
与层 1 SFTConfig 的三处结构性差异(docs/03 §3 偏差清单):
- 无 teacher_completions_pathteacher 现场前向给出全词表 logits、student 现场
on-policy 生成轨迹,两者都不落盘缓存,故层 2 不需要 teacher 解答文件。
- 无 max_length(总预算):completion 不再来自数据,而是 model.generate 生成,
序列总长 = prompt(≤max_prompt_length) + 生成(≤max_new_tokens),由两个预算
各自界定,不需要一个总的右截断预算。
- 无 top_k:§4 删除清单——本地同 tokenizer teacher 放得下全词表,恒走精确
全词表 KL,不做参考实现默认的 top-1 稀疏近似(那是 API 传输妥协,非论文成分)。
teacher_model 在此、student 在脚本(U5 的常量,同层 1 的 STUDENT_MODEL):
student 是被训练的固定基线,teacher 是"换一个就是换一个实验"的旋钮,故归 config。
"""
# ---- 机器相关路径(无默认值,必须显式传入)----
dataset_path: str
"""DAPO-Math-17K 的本地 parquet 路径(文件或目录),或 HF Hub 数据集名。
层 2 只用题面(prompt-only),不读数据自带的任何 completion。"""
output_dir: str
"""checkpoint 与日志输出目录(远程必须落在 /data/zym 下)。"""
# ---- teacher(层 2 的核心旋钮)----
teacher_model: str = "Qwen/Qwen3-4B"
"""本地 HF teacher 模型名。非机器路径(HF Hub 名各机可复现),故给默认值。
非显然约束:必须与 student **同 tokenizer**——KL 是逐词表位对齐求和,词表不
一致则第 v 个分量对不上、相除无意义(docs/03 §1)。此约束在 U4 构造 Trainer 时
比对 get_vocab() 显式校验,不匹配即报错,不静默。"""
# ---- 数据(复用层 1 的抽取逻辑,同 seed 同子集)----
dataset_split: str = "train"
subset_size: int | None = 1000
"""随机抽取的子集大小;None = 全量。非显然约束:与层 1 同 seed 同 size 才能
在同一批题上对比 SFT 与蒸馏,否则两层看的是不同题、曲线不可比。"""
# ---- 序列预算(prompt 截断 + 生成上限,见类 docstring 为何无 max_length----
max_prompt_length: int = 1024
"""prompt 单独预算(prompt-only collator 按此左截断)。"""
max_new_tokens: int = 1024
"""student on-policy 生成的 token 上限。与 max_prompt_length 之和即序列总长 T
显存账(§5)按 T=2048 估算。"""
enable_thinking: bool = False
"""Qwen3 chat 模板思考开关,喂给 student 生成。非显然约束:与层 1 取同值,
否则 prompt 渲染文本变、生成分布与 SFT 基线不可比(docs/02 §2.3 边界契约)。"""
# ---- 蒸馏损失(式(2) 与 docs/03 §2.3 三副面孔)----
beta: float = 1.0
"""KL 方向系数。0=前向 KL(π_T‖π_θ, mode-covering)1=反向 KL(π_θ‖π_T,
mode-seeking)=**式(2)**(0,1)=JSD 插值。默认 1 即论文式(2);参数保留是因为
前向/反向/JSD 是同一公式(U2 顺手覆盖),且层 5 的 KL 锚要用前向。"""
kl_temperature: float = 1.0
"""散度内 softmax 前除进两侧 logits 的温度(docs/03 §2.3)。升温放大尾部
"暗知识"排序。非显然约束:它与下面的 gen_temperature 是**两个不同**的温度
——这个调的是 loss 里分布的软硬,那个调的是采样的随机性;恰好都默认 1.0,
但改一个不影响另一个。默认 1.0 即式(2)(不做温度缩放)。"""
# ---- on-policy 生成采样(式(2) 的 y~π_θ 期望)----
gen_temperature: float = 1.0
gen_top_p: float = 1.0
"""student 生成轨迹的采样参数。默认 temperature=1.0/top_p=1.0 = 纯采样自 π_θ,
最忠实于式(2) 的 on-policy 期望(docs/03 §3 抉择原则:本质忠于论文)。
调低是拿保真度换"少生成垃圾",卡了再动。"""
# ---- 优化 ----
# 差异标注:层 1 SFT 用 2e-5(信号密集的真 token 监督);层 2 是蒸馏,从论文
# §5.1 的蒸馏 lr=1e-6。小 lr 在这里还有额外好处:式(2) 的反向 KL 会梯度爆炸
# (§4.1on-policy 采到 teacher 眼中的烂 token 时 log(π_θ/π_T)→∞),小步长
# 帮训练在毛刺中存活——这毛刺本身是层 2 要观察的教学目标(docs/03 §6.3)。
learning_rate: float = 1e-6
per_device_train_batch_size: int = 4
"""非显然约束:白盒蒸馏的显存大头是 **两份**全词表 logitsstudent+teacher
(B,T,V) bf16 各 ~2.5G@B=4/T=2048+ log_softmax 中间量,比层 1 更紧。B=4 是
§5 估算值(student 训练全套 ~10G + teacher 推理副本 ~9G + 两份 logits ~15G
A800-80G 起步安全),但**必须**在首次远程冒烟用 nvidia-smi 实测确认,OOM 阶梯:
先降 B 到 2、仍不够再开 gradient_checkpointing。"""
gradient_accumulation_steps: int = 4
"""全局 batch = 4(per_device) × 4(卡) × 4(累积) = 64,与层 1 保持一致。"""
num_train_epochs: int = 1
max_steps: int = -1
""">0 时覆盖 num_train_epochs——远程 50 步冒烟用(docs/03 §6.3);-1 = 按 epoch。"""
lr_scheduler_type: str = "linear"
warmup_ratio: float = 0.0
gradient_checkpointing: bool = False
"""默认不开(§5 显存账 B=4 富余);OOM 时作为降 batch 之后的第二道降显存手段。
注意 FSDP 下此开关是 no-opdocs/02 §2.6),但层 2 坚持 DDP 故此处有效。"""
bf16: bool = True
seed: int = 42
# ---- 日志与保存 ----
logging_steps: int = 1
save_steps: int = 100
save_total_limit: int = 2
report_to: str = "none"
""""none" 或 "wandb"。默认 none:本地调试不该悄悄往外发数据,远程脚本显式开。"""
def __post_init__(self) -> None:
"""构造即校验:配置错误必须在加载 4B teacherGB 级下载)之前炸。"""
if not 0.0 <= self.beta <= 1.0:
raise ValueError(
f"beta 必须在 [0,1]0=前向/1=反向/中间=JSD),收到 {self.beta}"
)
if self.kl_temperature <= 0:
# 温度除进 logits,≤0 会翻转或炸掉分布
raise ValueError(f"kl_temperature 必须为正,收到 {self.kl_temperature}")
if self.gen_temperature <= 0:
# 非显然约束:0 在 HF 里是 greedy,会退化 on-policy 采样为确定性解码,
# 破坏式(2) 的 y~π_θ 期望;要纯 on-policy 就必须 >0
raise ValueError(
f"gen_temperature 必须为正(0=greedy 破坏 on-policy),"
f"收到 {self.gen_temperature}"
)
if not 0.0 < self.gen_top_p <= 1.0:
raise ValueError(f"gen_top_p 必须在 (0,1],收到 {self.gen_top_p}")
if self.max_new_tokens <= 0:
raise ValueError(f"max_new_tokens 必须为正,收到 {self.max_new_tokens}")
if self.learning_rate <= 0:
raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}")
if self.subset_size is not None and self.subset_size <= 0:
raise ValueError(
f"subset_size 必须为正整数或 None(全量),收到 {self.subset_size}"
)
if self.max_steps == 0 or self.max_steps < -1:
raise ValueError(
f"max_steps 只接受 -1(按 epoch)或正整数,收到 {self.max_steps}"
)