From e42af5256fb22fc92b8967c16ba395a59e019ebd Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sun, 19 Jul 2026 03:37:44 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=822/U1:=20=E6=96=B0=E5=A2=9E=20DistillCo?= =?UTF-8?q?nfig=EF=BC=88=E5=BC=8F(2)=20=E7=99=BD=E7=9B=92=E8=92=B8?= =?UTF-8?q?=E9=A6=8F=E5=8F=82=E6=95=B0=EF=BC=89+=20docs/03=20=E8=A1=A8?= =?UTF-8?q?=E8=BF=B0=E6=89=93=E7=A3=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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) --- ars_opd/configs.py | 146 +++++++++++++++++++++++++++++++++++++++- docs/03-whitebox-opd.md | 20 ++++-- 2 files changed, 161 insertions(+), 5 deletions(-) diff --git a/ars_opd/configs.py b/ars_opd/configs.py index 0c74c0c..f508873 100644 --- a/ars_opd/configs.py +++ b/ars_opd/configs.py @@ -1,10 +1,12 @@ -"""实验配置(当前只含层 1 的 SFTConfig,后续层在此文件追加各自的 dataclass)。 +"""实验配置(层 1:SFTConfig / TeacherGenConfig;层 2:DistillConfig)。 设计约定(对应 CLAUDE.md §2"配置显式化"): - 所有实验参数必须是这里某个 dataclass 的字段;代码里出现魔法数字/路径即违规。 - 密钥(API key 等)不进配置类,走 `.env`(见 teacher.py)。 - 机器相关路径(数据集、输出目录)不给默认值,强制调用方显式传入—— 防止参考实现里 `/fsx` 硬编码那类"在别人机器上必炸"的坑。 +- 各层的 config 自包含、不互相继承:层与层是不同实验,共享基类会把它们耦合, + 违背"从上读到下看懂全部流程"(CLAUDE.md §2)。字段重复是有意接受的成本。 """ from __future__ import annotations @@ -153,3 +155,145 @@ class TeacherGenConfig: 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_path:teacher 现场前向给出全词表 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.1,on-policy 采到 teacher 眼中的烂 token 时 log(π_θ/π_T)→∞),小步长 + # 帮训练在毛刺中存活——这毛刺本身是层 2 要观察的教学目标(docs/03 §6.3)。 + learning_rate: float = 1e-6 + + per_device_train_batch_size: int = 4 + """非显然约束:白盒蒸馏的显存大头是 **两份**全词表 logits(student+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-op(docs/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 teacher(GB 级下载)之前炸。""" + 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}" + ) diff --git a/docs/03-whitebox-opd.md b/docs/03-whitebox-opd.md index 87060a5..e2a30b8 100644 --- a/docs/03-whitebox-opd.md +++ b/docs/03-whitebox-opd.md @@ -40,13 +40,23 @@ on/off-policy 抽签**不在这里**——在 `_prepare_inputs`→`_fill_buffer` ### 2.3 generalized_jsd_loss(DT:2408-2491)——β 的三副面孔 +记号:$\pi_\theta$ = student,$\pi_T$ = teacher;KL 里"在前"的那个分布是被求期望的一方。 + | β | 数学 | 语义 | 代码 | |---|------|------|------| -| 0 | $KL(\pi_T\|\pi_\theta)$ | 前向,mode-covering | DT:150-151 | -| 1 | $KL(\pi_\theta\|\pi_T)$ | **反向,式(2) 用这个** | DT:152-153 | -| (0,1) | $\beta KL(p_T\|m)+(1{-}\beta)KL(p_\theta\|m)$,$m=(1{-}\beta)p_\theta{+}\beta p_T$ | JSD 插值 | DT:154-162 | +| 0 | $KL(\pi_T\|\pi_\theta)$ | 前向,**mode-covering**:teacher 在前,student 被迫摊平质量去覆盖 teacher 的全分布 | DT:150-151 | +| 1 | $KL(\pi_\theta\|\pi_T)$ | 反向,**mode-seeking**,**式(2) 用这个**:student 在前,集中质量学 teacher 主模式 | DT:152-153 | +| (0,1) | $\beta KL(\pi_T\|m)+(1{-}\beta)KL(\pi_\theta\|m)$,$m=(1{-}\beta)\pi_\theta{+}\beta\pi_T$ | JSD 插值(β 同时是混合权重与两项权重) | DT:154-162 | -细节:温度在 softmax 前除进两侧 logits(DT:2439-2440);全程 log 域运算(logsumexp 混合、`clamp_min(tiny)` 防 log0,DT:143,156-159);reduction=`batchmean` 实为 **sum / 有效 token 数**(labels≠-100 先滤,DT:2386-2399)——名字叫 batchmean,实义是 per-token mean,重构时按实义命名。 +> 行号说明:上表指向 **全词表** 路径(`F.kl_div`,我们 U2 走这条)。参考实现**默认**走 top-1 稀疏(§2.4),对应 DT:133-148 的 masked 镜像分支——同样三支 β、同样数学,只是在截断支持集上手算而非调 `F.kl_div`。 + +另有三条与 β 语义无关、但读代码时容易卡住的实现约定。它们各自独立,只是恰好都在这个函数里;前两条是正确性/稳定性刚需(我们保留),第三条是历史包袱(我们纠名): + +| 实现约定 | 位置 | 是什么 / 为什么 load-bearing | 我们 U2 | +|----------|------|------------------------------|---------| +| 温度除进 logits | DT:2439-2440 | `logits / τ` 必须在 softmax **之前**做——这是在调分布形状(升温 τ>1 放大尾部的"暗知识"排序),不是等比缩概率。`softmax(z/τ) ≠ softmax(z)/τ`,位置错了就不再是合法分布 | 保留(默认 τ=1,此步为恒等) | +| 全程 log 域运算 | DT:143,156-159 | 15 万词表下单个概率小到 1e-8,而 KL 全是乘除,直接算会下溢成 0 → NaN。对策:全程存 log-prob(乘变加、除变减)。两个衍生 trick:混合分布 $m$ 的**加法**在 log 域要用 `logsumexp`(log 里的加法天然是乘法);masked 位概率为 0,取 log 前先 `clamp_min(tiny)` 防 `log(0)=-inf` | 保留(仅全词表这一路径) | +| `batchmean` 名不副实 | DT:2386-2399 | 先滤 `labels≠-100` 留下 completion 位,再 `jsd.sum() / 有效 token 数`——实义是 **per-token mean**,不是 PyTorch `batchmean` 那个"÷ 序列条数"。这样量纲与 SFT 的 per-token 交叉熵一致,两条 loss 曲线才可比 | 纠名为 `per_token_mean`,名字即文档 | ### 2.4 默认配置不是全词表 KL!(本章最大陷阱) @@ -75,6 +85,8 @@ server 路径更粗:teacher 只回传 top-k logprobs 三张表(DT:2676-2679 ## 3. 与式(2) 的偏差清单(默认配置下) +抉择原则(下表每一行都由它推出,而非逐条拍板):**式(2) 的算法本质忠于论文(on-policy + 反向 KL + 全词表分布对齐);参考实现的默认近似是为它的处境——API 传输 + 大 teacher——妥协出来的,换了我们的处境(本地同 tokenizer 的 4B teacher + 4×A800)就不继承;不改优化方向的表面差异,选对诊断/教学最有利的;一般性凡免费且未来有用则保留、凡昂贵且当前数据上空转则删除。** 一个反直觉推论:正因处境不同,我们回归论文本质反而比参考默认更贴式(2)(支持集那行)。逐行推导见下,"我们层 2"列即结论。 + | 项 | 参考实现默认 | 严格式(2) | 我们层 2 | |----|--------------|-----------|----------| | 支持集 | top-1 稀疏 + 尾桶 | 全词表 | **全词表**(`top_k=0` 等价;同 tokenizer 本地 teacher 使我们能比参考默认更贴论文) |