Compare commits
48 Commits
8647c5a89d
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 630a9c4636 | |||
| e06e5ed7a2 | |||
| e7cf27cecc | |||
| 142aeb8ab1 | |||
| fac8e0dcc5 | |||
| 8b362eae09 | |||
| f1b6d1f668 | |||
| 2f4780b69f | |||
| b7f24d635e | |||
| 404abc22bf | |||
| 0ca60ea93f | |||
| e5a28e8e77 | |||
| de7f36828a | |||
| e42af5256f | |||
| db392d9c81 | |||
| eb56883267 | |||
| 872d4bd6a6 | |||
| e472a66959 | |||
| 4f13365ffa | |||
| ee16e29846 | |||
| 1976230250 | |||
| ff572bf4b9 | |||
| a0faec0df7 | |||
| 68216ace22 | |||
| 1d272b6d6c | |||
| 3c1530a1db | |||
| 4621ebae31 | |||
| c5a3b7d0bb | |||
| 42a4349e69 | |||
| e72d07aced | |||
| 20e6d97427 | |||
| f5bb852fde | |||
| 5ea58ddf59 | |||
| 58d75cc56a | |||
| 240f404416 | |||
| 6c128ad6f5 | |||
| 0bcfc6efd4 | |||
| 12e2a8b0ae | |||
| 77f0419bad | |||
| 9c0ec15716 | |||
| 991e277316 | |||
| b679bcb24f | |||
| f255d5613a | |||
| 270a9ab291 | |||
| cccccc0ce6 | |||
| 72799a34dc | |||
| c9eddd5be8 | |||
| d0bedd564c |
@@ -0,0 +1,8 @@
|
||||
# 复制为 .env 并填入真实值(.env 已被 gitignore,严禁提交密钥)
|
||||
# teacher API(OpenAI 兼容格式;2026-07-18 定:自建 new-api 网关 + MiniMax-M3)
|
||||
TEACHER_API_BASE=https://newapi.iomgaa.online/v1
|
||||
TEACHER_API_KEY=
|
||||
TEACHER_MODEL=MiniMax-M3
|
||||
|
||||
# W&B(仅远程训练需要)
|
||||
WANDB_API_KEY=
|
||||
Vendored
+26
@@ -0,0 +1,26 @@
|
||||
{
|
||||
// VSCode 调试配置。前提:本地环境已 `pip install -e .`(见 README/CLAUDE.md),
|
||||
// 且右下角解释器已选 ars-opd 环境。cwd 固定为仓库根:脚本里的相对路径
|
||||
// (data/...、outputs/...)都以仓库根为基准。
|
||||
"version": "0.2.0",
|
||||
"configurations": [
|
||||
{
|
||||
"name": "调试当前文件",
|
||||
"type": "debugpy",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal",
|
||||
"cwd": "${workspaceFolder}",
|
||||
"justMyCode": false
|
||||
},
|
||||
{
|
||||
"name": "teacher 批量生成(1k 子集)",
|
||||
"type": "debugpy",
|
||||
"request": "launch",
|
||||
"program": "${workspaceFolder}/scripts/generate_teacher_completions.py",
|
||||
"console": "integratedTerminal",
|
||||
"cwd": "${workspaceFolder}",
|
||||
"justMyCode": false
|
||||
}
|
||||
]
|
||||
}
|
||||
Vendored
+4
@@ -0,0 +1,4 @@
|
||||
{
|
||||
"python-envs.defaultEnvManager": "ms-python.python:conda",
|
||||
"python-envs.defaultPackageManager": "ms-python.python:conda"
|
||||
}
|
||||
@@ -1,7 +1,7 @@
|
||||
# CLAUDE.md
|
||||
|
||||
> [!URGENT]
|
||||
> 1. 本项目是**学习驱动的科研重构项目**:通过分层重建 OmniOPD 参考实现来学习论文方法,同时产出可作为后续研究基座的干净代码库。学习优先——用户要逐章理解每个模块,**不要一次性替用户写完所有代码**;每章先讲清楚,重构任务由用户主导、你配合。
|
||||
> 1. 本项目是**学习驱动的科研重构项目**:通过分层重建 OmniOPD 参考实现来学习论文方法,同时产出可作为后续研究基座的干净代码库。学习优先——每章先讲清楚原理再动代码。**分工(2026-07-18 用户定):代码由 Claude 编写,用户逐行精读并提问**;因此代码必须严格教学导向(遵守 §7 注释规范),每个模块写完后主动讲解设计要点与易错点,用户的提问优先于推进进度。
|
||||
> 2. 所有思考过程和回复必须使用**简体中文**。
|
||||
|
||||
## 1. 项目元数据
|
||||
@@ -25,7 +25,8 @@
|
||||
| `ars_opd/similarity.py` | §3.2.1 式(3) | 语义相似度 φ(rouge1 / edit_distance) |
|
||||
| `ars_opd/estimator.py` | §3.2.2 式(4)(5) | MC 估计 + Dirichlet 贝叶斯平滑 π̂ |
|
||||
| `ars_opd/chunking.py` | §3.2.3 式(6)(7) | 熵计算 + peak-entropy chunk 选择 |
|
||||
| `ars_opd/teacher.py` | §3.2.1 | teacher rollout 采样(OpenAI 兼容 API + 缓存 / vLLM) |
|
||||
| `ars_opd/data.py` | —(IO 边缘,无论文锚点) | 数据加载(DAPO parquet/HF)+ messages 归一 + 双预算掩码 collator |
|
||||
| `ars_opd/teacher.py` | §3.2.1 | teacher rollout 采样(OpenAI 兼容 API + 缓存 / vLLM);层 1 起步能力:批量生成+sha256 缓存 |
|
||||
| `ars_opd/trainer.py` | §3.2.4 式(8) | chunk 蒸馏损失 + trust-region KL 锚定,编排全流程 |
|
||||
|
||||
## 4. 常用命令
|
||||
@@ -33,6 +34,7 @@
|
||||
```bash
|
||||
conda activate ars-opd && pytest tests/ -x -q # 单元测试(本地 CPU)
|
||||
conda activate ars-opd && ruff check ars_opd/ --fix && ruff format ars_opd/
|
||||
pip install -e . --no-build-isolation --no-deps # 新环境一次性:注册 ars_opd 包(否则脚本 import 报错)
|
||||
```
|
||||
|
||||
## 5. 本地-远程规则
|
||||
@@ -40,7 +42,7 @@ conda activate ars-opd && ruff check ars_opd/ --fix && ruff format ars_opd/
|
||||
> [!CRITICAL]
|
||||
> - 远程机 gpu-a800-060:8×A800-80G,**只允许用其中 4 块**;所有 GPU 命令必须显式 `CUDA_VISIBLE_DEVICES=<idx>`,先 `nvidia-smi` 确认空闲卡,严禁自动选卡。
|
||||
> - 远程**根分区仅剩 12G**:conda 环境、HF 缓存(`HF_HOME`)、模型、checkpoint、数据集一律放 `/data/zym/` 下。
|
||||
> - 长任务用 tmux 跑,日志不缓存(`python -u` / `PYTHONUNBUFFERED=1`),确保可实时检查。
|
||||
> - **任何长时间运行的命令(训练、pip/conda 安装、脚本)禁止日志缓存**,宁可承担延时也要实时可查:python 加 `-u` / `PYTHONUNBUFFERED=1`;**不用 `conda run` 包裹长命令**(它整体缓冲输出直到结束——2026-07 曾因此把正常安装误判为卡死),改为直调 `<env>/bin/pip`、`<env>/bin/python`;远程长任务一律 tmux。
|
||||
> - 远程机器上不改代码,只 `git pull` + 跑脚本;本地不跑训练。
|
||||
|
||||
## 6. 学习工作流
|
||||
@@ -49,10 +51,23 @@ conda activate ars-opd && ruff check ars_opd/ --fix && ruff format ars_opd/
|
||||
2. 重构一个模块前,先对照参考实现列出其全部行为(含 trick 和 workaround),逐一确认保留/替代/删除。
|
||||
3. 每个纯逻辑模块完成后,用 toy 数据对拍参考实现的对应逻辑(参考其 `validate_mc_estimator.py` / `validate_chunk_mc_estimator.py`)。
|
||||
4. 文档规范:优先表格与公式,代码块 ≤15 行(展示思路用伪代码,完整代码引用文件路径),引用参考实现必须带 `文件:行号`。
|
||||
5. **每层完成后做接口回看**:逐模块自问"接口是否比实现简单得多"(深模块判据);若某接口的参数/约定复杂到接近实现本身,先记录并重构,再进入下一层。规则来源与哲学对照见 `docs/appendix-claudemd-decisions.md`。
|
||||
|
||||
## 7. 代码规范
|
||||
## 7. 代码规范(教学导向)
|
||||
|
||||
- 公共函数完整类型注解;模块/类/函数写中文 docstring(功能、参数、返回、关键实现细节)。
|
||||
**注释分工**:`docs/` 章节负责讲原理,代码注释负责做索引,两者不重复。代码注释只写三类内容:
|
||||
|
||||
| 类型 | 要求 | 示例 |
|
||||
|------|------|------|
|
||||
| 论文锚点 | 实现论文公式/机制的函数,docstring 首行标出处;关键行旁给公式本体 | `# 式(5): π̂ = (k_sem + α·π̄) / (N + α)` |
|
||||
| 非显然约束 | 只解释"为什么必须这样"及违反后果,不解释"这行在干什么";load-bearing 的反直觉点必须写 | `# π̂ 必须 detach:否则优化器会压低学生自身概率把乘子 π̂ 推向 0 以逃逸惩罚,teacher 否定的 chunk 最先塌缩` |
|
||||
| 差异标注 | 凡有意偏离论文或参考实现处,注明对方做法与我们的理由 | `# 参考实现(trainer:2205)对 chunk 内取 mean,论文式(8)为 sum,此处从论文` |
|
||||
|
||||
**类型与 shape**:
|
||||
- 模块间公共接口(`ars_opd/` 各模块导出的函数/类)强制完整类型注解——接口注解本身就是教学信息;模块内私有 helper 从宽。
|
||||
- 类型注解表达不了张量 shape,故 shape 是硬要求:docstring 注明参数/返回的 shape,函数体内关键变换旁加行注释(如 `# (B, T, V) -> (B, T)`)。
|
||||
- docstring 用中文,含功能、参数、返回、关键实现细节。
|
||||
|
||||
**其余硬规则**:
|
||||
- **严禁** `except Exception: pass`;出错直接报错,不用默认值兜底。
|
||||
- 张量函数在 docstring 中注明各参数的 shape。
|
||||
- 提交信息用中文,说明"这一步对应哪一章/哪个模块"。
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
"""ars_opd:OmniOPD (arXiv:2606.01476v2) 的分层重构实现。
|
||||
|
||||
模块与论文的对应关系见 CLAUDE.md §3(单一事实源),此处不重复。
|
||||
"""
|
||||
@@ -0,0 +1,310 @@
|
||||
"""实验配置(层 1:SFTConfig / TeacherGenConfig;层 2:DistillConfig)。
|
||||
|
||||
设计约定(对应 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_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
|
||||
|
||||
max_grad_norm: float = 1.0
|
||||
"""梯度裁剪阈值。此前是 HF Trainer 的静默默认(1.0),现显式化——它是式(2)
|
||||
反向 KL 梯度爆炸(§4.1)的**隐形稳定器**:on-policy 采到 teacher 眼中烂 token
|
||||
时单步梯度范数可炸到十几(2026-07-19 首冒烟实测 grad_norm 14→2),HF 默认
|
||||
裁到 1.0 才让 loss 曲线平稳。把它设得远大于实测范数(≈关闭裁剪)可暴露原始
|
||||
爆炸,供教学对照(train_whitebox.py 的 noclip 模式)。非显然约束:日志里的
|
||||
grad_norm 是**裁剪前**范数,故 14→2 那串本身就是爆炸证据,只是被裁剪掩盖了。"""
|
||||
|
||||
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.max_grad_norm <= 0:
|
||||
# 用远大于实测范数的值≈关闭裁剪;≤0 无意义(0 会把梯度裁没)
|
||||
raise ValueError(f"max_grad_norm 必须为正,收到 {self.max_grad_norm}")
|
||||
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}"
|
||||
)
|
||||
+407
@@ -0,0 +1,407 @@
|
||||
"""数据管线(IO 边缘,无论文锚点):加载 → messages 归一 → 挂接 teacher 解答 → collator。
|
||||
|
||||
层 1 的数据流(对应 docs/02 §1 的基线定义:SFT = 在 teacher rollout 上的离线蒸馏):
|
||||
|
||||
DAPO parquet(prompt-only)
|
||||
→ to_messages 归一成 [{"role","content"}] 列表
|
||||
→ 按 seed 抽子集
|
||||
→ attach_teacher_completions 从 JSONL 缓存挂上 teacher 解答(assistant 轮)
|
||||
→ SFTCollator 分词、双预算截断、-100 掩码、左 padding
|
||||
|
||||
本模块与 teacher.py 的缓存契约由 `prompt_key` 单点定义:teacher.py 生成缓存、
|
||||
本模块消费缓存,双方必须用同一个函数算键。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from datasets import Dataset, load_dataset
|
||||
|
||||
# F.cross_entropy 的 ignore_index 默认值;标了它的位置不产生 loss
|
||||
IGNORE_INDEX = -100
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# messages 归一
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _parse_stringified_list(value: str, column: str) -> list:
|
||||
"""parquet 有时把 list 存成其字符串形态,用 ast 还原。
|
||||
|
||||
差异标注:参考实现(train_distillation.py:299-305)在这里 `except: pass` 静默吞错,
|
||||
坏行会以原始字符串流进 collator,在 apply_chat_template 处以难懂的方式炸;
|
||||
我们显式报错,错误信息直接指向坏数据本身。
|
||||
"""
|
||||
try:
|
||||
parsed = ast.literal_eval(value)
|
||||
except (ValueError, SyntaxError) as e:
|
||||
raise ValueError(
|
||||
f"列 {column!r} 是字符串但无法解析为 Python 字面量(坏数据行):"
|
||||
f"{value[:200]!r}"
|
||||
) from e
|
||||
if not isinstance(parsed, (list, tuple)):
|
||||
raise ValueError(f"列 {column!r} 解析结果不是列表:{type(parsed).__name__}")
|
||||
return list(parsed)
|
||||
|
||||
|
||||
def to_messages(example: dict[str, Any]) -> dict[str, list[dict[str, str]]]:
|
||||
"""把三种来源格式归一成 messages 列:[{"role": ..., "content": ...}, ...]。
|
||||
|
||||
支持(与参考实现 train_distillation.py:295-321 相同的三分支):
|
||||
- ``messages`` 列:直取;
|
||||
- ``prompt`` 列(DAPO parquet,列名不副实——装的是完整 chat 列表):改名;
|
||||
- ``question`` 列(gsm8k 风格纯文本):包成单 user 轮。
|
||||
|
||||
差异标注:参考实现对不认识的行 `return x` 静默放行,我们显式报错。
|
||||
"""
|
||||
if "messages" in example:
|
||||
msgs = example["messages"]
|
||||
column = "messages"
|
||||
elif "prompt" in example:
|
||||
msgs = example["prompt"]
|
||||
column = "prompt"
|
||||
elif "question" in example:
|
||||
return {"messages": [{"role": "user", "content": example["question"]}]}
|
||||
else:
|
||||
raise ValueError(
|
||||
f"无法识别的数据行:既无 messages/prompt 也无 question 列,"
|
||||
f"实有列 {sorted(example.keys())}"
|
||||
)
|
||||
|
||||
if isinstance(msgs, str):
|
||||
msgs = _parse_stringified_list(msgs, column)
|
||||
msgs = list(msgs)
|
||||
|
||||
if not msgs:
|
||||
raise ValueError(f"列 {column!r} 是空列表(坏数据行)")
|
||||
for m in msgs:
|
||||
if not isinstance(m, dict) or "role" not in m or "content" not in m:
|
||||
raise ValueError(
|
||||
f"列 {column!r} 中存在非 {{role, content}} 结构的元素:{m!r}"
|
||||
)
|
||||
|
||||
return {"messages": [{"role": m["role"], "content": m["content"]} for m in msgs]}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# teacher 解答缓存(与 teacher.py 的契约)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def prompt_key(messages: list[dict[str, str]]) -> str:
|
||||
"""teacher 缓存的键:对 messages 的规范化 JSON 取 sha256。
|
||||
|
||||
差异标注:参考实现(distillation_trainer.py:984)用 `str(hash(prompt))`——
|
||||
Python 对 str 的 hash 默认加盐,跨进程/跨次运行不稳定,缓存必然失效重生成。
|
||||
sha256 内容寻址:同一道题永远同一个键。
|
||||
|
||||
只取 role/content 两个字段参与哈希:DAPO 行里其余元数据(data_source 等)
|
||||
变了不应导致缓存失效。
|
||||
"""
|
||||
canon = [{"role": m["role"], "content": m["content"]} for m in messages]
|
||||
return hashlib.sha256(json.dumps(canon, ensure_ascii=False).encode()).hexdigest()
|
||||
|
||||
|
||||
def attach_teacher_completions(dataset: Dataset, jsonl_path: str) -> Dataset:
|
||||
"""把 teacher 解答缓存(JSONL,每行 {"key", "completion"})挂到数据集上。
|
||||
|
||||
- 末轮已是 assistant 的行保持原样(数据自带解答,不覆盖);
|
||||
- 任何 prompt-only 行在缓存中查不到键 → 收集齐所有缺失后一次性报错,
|
||||
提示先运行 teacher 生成——绝不静默跳过(跳过 = 悄悄改变训练集组成)。
|
||||
"""
|
||||
path = Path(jsonl_path)
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"teacher 解答缓存不存在:{jsonl_path}。先运行 teacher.py 的批量生成。"
|
||||
)
|
||||
|
||||
cache: dict[str, str] = {}
|
||||
with open(path, encoding="utf-8") as f:
|
||||
for line_no, line in enumerate(f, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
rec = json.loads(line) # 坏行直接炸,带行号
|
||||
if "key" not in rec or "completion" not in rec:
|
||||
raise ValueError(f"{jsonl_path}:{line_no} 缺少 key/completion 字段")
|
||||
cache[rec["key"]] = rec["completion"]
|
||||
|
||||
# 先整体扫描缺失,一次性报全——比在 .map 里炸第一条更省来回
|
||||
missing = [
|
||||
i
|
||||
for i, ex in enumerate(dataset)
|
||||
if ex["messages"][-1]["role"] != "assistant"
|
||||
and prompt_key(ex["messages"]) not in cache
|
||||
]
|
||||
if missing:
|
||||
raise KeyError(
|
||||
f"{len(missing)}/{len(dataset)} 行在 teacher 缓存中查不到解答"
|
||||
f"(首个缺失行 index={missing[0]})。检查:teacher 生成是否用了同一"
|
||||
f"子集与同一 seed?(子集抽取在 load_sft_dataset 中先于挂接发生,"
|
||||
f"两侧 seed 不同则键集合不同)"
|
||||
)
|
||||
|
||||
def _attach(ex: dict[str, Any]) -> dict[str, Any]:
|
||||
msgs = ex["messages"]
|
||||
if msgs[-1]["role"] == "assistant":
|
||||
return ex
|
||||
completion = cache[prompt_key(msgs)]
|
||||
return {"messages": msgs + [{"role": "assistant", "content": completion}]}
|
||||
|
||||
return dataset.map(_attach)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 数据集加载(入口)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def load_sft_dataset(
|
||||
dataset_path: str,
|
||||
dataset_split: str = "train",
|
||||
subset_size: int | None = None,
|
||||
seed: int = 42,
|
||||
teacher_completions_path: str | None = None,
|
||||
) -> Dataset:
|
||||
"""数据管线入口:加载 → 归一 → 抽子集 →(可选)挂 teacher 解答。
|
||||
|
||||
收散装参数而非整个 config(深模块:本函数只用这 5 个字段,不该索要一整个
|
||||
SFTConfig)。这样层 1(SFTConfig)、层 2(DistillConfig,无 teacher 缓存)、
|
||||
诊断脚本都能直接调,无需伪造无关字段。teacher_completions_path=None 时
|
||||
返回 prompt-only 数据集(末轮 user,供 on-policy 生成);给了则挂 teacher
|
||||
解答(末轮 assistant,供 SFT)。
|
||||
|
||||
返回只含 ``messages`` 一列的 Dataset。
|
||||
"""
|
||||
ds = _load_raw(dataset_path, dataset_split)
|
||||
ds = ds.map(
|
||||
to_messages,
|
||||
remove_columns=[c for c in ds.column_names if c != "messages"],
|
||||
)
|
||||
if subset_size is not None and subset_size < len(ds):
|
||||
# 非显然约束:抽子集必须在挂接 teacher 解答之前、且由 seed 完全确定——
|
||||
# teacher.py 生成缓存时会走完全相同的"加载→归一→抽子集"路径,两侧 seed
|
||||
# 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。层 2 与层 1
|
||||
# 用同 seed 同 subset_size,才能在同一批题上对比 SFT 与蒸馏。
|
||||
ds = ds.shuffle(seed=seed).select(range(subset_size))
|
||||
if teacher_completions_path is not None:
|
||||
ds = attach_teacher_completions(ds, teacher_completions_path)
|
||||
return ds
|
||||
|
||||
|
||||
def _load_raw(dataset_path: str, split: str) -> Dataset:
|
||||
"""三分支加载:parquet 目录 / 单 parquet 文件 / HF Hub 数据集名。
|
||||
|
||||
差异标注:参考实现(train_distillation.py:292)对 Hub 分支硬编码 config 名
|
||||
"main"(gsm8k 专用);我们不硬编码——需要特定 config 的数据集请下载成
|
||||
parquet 本地加载。
|
||||
"""
|
||||
p = Path(dataset_path)
|
||||
if p.is_dir():
|
||||
return load_dataset("parquet", data_dir=dataset_path, split=split)
|
||||
if dataset_path.endswith(".parquet"):
|
||||
return load_dataset("parquet", data_files=dataset_path, split=split)
|
||||
return load_dataset(dataset_path, split=split)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Collator:本层最核心的一段(对拍 distillation_trainer.py:210-343)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SFTCollator:
|
||||
"""把一个 batch 的 messages 变成训练/生成所需的定长张量,两种模式二选一。
|
||||
|
||||
prompt_only=False(层 1 SFT,默认)——输出 input_ids/attention_mask/labels:
|
||||
核心设计(继承参考实现的双预算方案,docs/02 §2.3):prompt 与 completion
|
||||
各自独立预算——prompt 用 max_prompt_length 截断,completion 上限是
|
||||
max_length - len(截断后 prompt)。若只用一个总预算从右截断,超长解答会把
|
||||
prompt 挤空,模型在"没有题目"的样本上学解答。要求每行末轮是 assistant,
|
||||
否则报错(prompt-only 行在纯 SFT 下只产生零 loss = 静默空训练)。
|
||||
|
||||
prompt_only=True(层 2 white-box OPD,docs/03 §5 U3)——只渲染 prompt、
|
||||
输出 prompts/prompt_attention_mask 供 model.generate 做 on-policy 生成;
|
||||
completion 由生成产生、labels 由 U4 的 DistillTrainer 在生成后重建,故此模式
|
||||
不产 labels、也不吃 max_length。这兑现了参考实现为 on-policy 生成留的口子
|
||||
(层 1 曾故意关掉,见此前 git 历史)。
|
||||
|
||||
与参考实现的其余差异:空 <think> 的一次性诊断打印改为单元测试断言(契约进
|
||||
测试,不进运行时日志)。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: "Any",
|
||||
max_prompt_length: int,
|
||||
max_length: int | None = None,
|
||||
enable_thinking: bool = False,
|
||||
prompt_only: bool = False,
|
||||
) -> None:
|
||||
"""tokenizer 需实现 HF 接口:apply_chat_template / __call__ / pad_token_id。
|
||||
|
||||
max_length 仅 SFT 模式需要(completion 预算依赖它);prompt_only 模式下
|
||||
completion 是生成的、无总预算,故 max_length 可为 None。
|
||||
"""
|
||||
if not prompt_only and max_length is None:
|
||||
raise ValueError(
|
||||
"SFT 模式(prompt_only=False)必须提供 max_length——completion "
|
||||
"预算 = max_length - len(prompt),缺它无法确定解答截断点。"
|
||||
)
|
||||
self.tokenizer = tokenizer
|
||||
self.max_length = max_length
|
||||
self.max_prompt_length = max_prompt_length
|
||||
self.enable_thinking = enable_thinking
|
||||
self.prompt_only = prompt_only
|
||||
# pad→eos 回退:左 padding 位置的 attention_mask 恒为 0,pad 值不参与
|
||||
# 任何计算,只需要一个合法 token id 占位,借用 eos 即可
|
||||
if tokenizer.pad_token_id is not None:
|
||||
self.pad_token_id: int = tokenizer.pad_token_id
|
||||
elif tokenizer.eos_token_id is not None:
|
||||
self.pad_token_id = tokenizer.eos_token_id
|
||||
else:
|
||||
raise ValueError("tokenizer 既无 pad_token 也无 eos_token,无法 padding")
|
||||
|
||||
def __call__(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
|
||||
"""按模式分派:prompt_only 走生成用 prompt 张量,否则走 SFT 双预算。"""
|
||||
if self.prompt_only:
|
||||
return self._collate_prompt_only(examples)
|
||||
return self._collate_sft(examples)
|
||||
|
||||
def _collate_prompt_only(
|
||||
self, examples: list[dict[str, Any]]
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""层 2:只渲染 prompt 供 on-policy 生成,不产 completion/labels。
|
||||
|
||||
返回(B = batch 大小,P = batch 内最长 prompt 长度):
|
||||
- prompts: (B, P) 左 padding
|
||||
- prompt_attention_mask: (B, P) padding 位置为 0
|
||||
|
||||
非显然约束:生成必须左 padding——所有 prompt 右对齐到同一右边界,
|
||||
model.generate 从该边界统一续写;右 padding 会让短 prompt 的生成从 pad
|
||||
中间开始,全乱。这也是层 1 SFT 就选左 padding 的原因(全项目一种约定)。
|
||||
"""
|
||||
all_prompt_ids: list[list[int]] = []
|
||||
for example in examples:
|
||||
messages = example["messages"]
|
||||
# prompt-only 数据末轮是 user;若末轮已是 assistant 则剥掉,取生成前上下文
|
||||
prompt_msgs = (
|
||||
messages[:-1] if messages[-1]["role"] == "assistant" else messages
|
||||
)
|
||||
if not prompt_msgs:
|
||||
raise ValueError(
|
||||
"prompt_only collator 收到空 prompt(无可生成的上下文)"
|
||||
)
|
||||
# 与 SFT 模式同样带生成引导符渲染(add_generation_prompt=True):
|
||||
# prompt 末尾就是 "<|im_start|>assistant\n...",生成从此续写
|
||||
formatted_prompt = self.tokenizer.apply_chat_template(
|
||||
prompt_msgs,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=self.enable_thinking,
|
||||
)
|
||||
prompt_ids: list[int] = self.tokenizer(
|
||||
formatted_prompt,
|
||||
truncation=True,
|
||||
max_length=self.max_prompt_length,
|
||||
add_special_tokens=False,
|
||||
)["input_ids"]
|
||||
all_prompt_ids.append(prompt_ids)
|
||||
|
||||
return {
|
||||
"prompts": _left_pad(all_prompt_ids, self.pad_token_id), # (B, P)
|
||||
"prompt_attention_mask": _left_pad(
|
||||
[[1] * len(ids) for ids in all_prompt_ids], 0
|
||||
), # (B, P)
|
||||
}
|
||||
|
||||
def _collate_sft(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
|
||||
"""层 1 SFT:messages(末轮 assistant)→ 定长张量。
|
||||
|
||||
返回(B = batch 大小,T = batch 内最长序列长度):
|
||||
- input_ids: (B, T) 左 padding
|
||||
- attention_mask: (B, T) padding 位置为 0
|
||||
- labels: (B, T) padding 与 prompt 位置为 -100,completion 位置为 token id
|
||||
"""
|
||||
all_input_ids: list[list[int]] = []
|
||||
all_labels: list[list[int]] = []
|
||||
|
||||
for example in examples:
|
||||
messages = example["messages"]
|
||||
if len(messages) < 2 or messages[-1]["role"] != "assistant":
|
||||
raise ValueError(
|
||||
"SFTCollator 收到 prompt-only 行(末轮不是 assistant)。纯 SFT 下"
|
||||
"它只会产生全 -100 的零 loss 样本——静默空训练。检查 teacher "
|
||||
"解答是否挂接成功。"
|
||||
)
|
||||
|
||||
# prompt = 末轮 assistant 之前的全部轮次,渲染时带生成引导符
|
||||
# ("<|im_start|>assistant\n..."),这样 completion 是纯解答文本的分词
|
||||
formatted_prompt = self.tokenizer.apply_chat_template(
|
||||
messages[:-1],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=self.enable_thinking,
|
||||
)
|
||||
# prompt 自己的预算内截断。沿用 tokenizer 默认右截断(与参考实现一致):
|
||||
# 超预算的题目被截掉尾部(含生成引导符)——1024 预算下 DAPO 极少触发,
|
||||
# 触发时该样本退化但不会污染边界(边界用未截断长度算,见下)
|
||||
prompt_ids: list[int] = self.tokenizer(
|
||||
formatted_prompt,
|
||||
truncation=True,
|
||||
max_length=self.max_prompt_length,
|
||||
add_special_tokens=False,
|
||||
)["input_ids"]
|
||||
|
||||
# 非显然约束(docs/02 坑一/坑二):completion 边界必须用"未截断 prompt
|
||||
# 的分词长度"从整段渲染中切出。BPE 分词不满足拼接稳定性,分开渲染
|
||||
# prompt 和 completion 再拼接 ≠ 整段渲染后分词;而若用截断后长度当切分
|
||||
# 点,会把 prompt 尾部的 token 误标成 completion——静默的语义错误。
|
||||
formatted_full = self.tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=False,
|
||||
enable_thinking=self.enable_thinking,
|
||||
)
|
||||
full_ids: list[int] = self.tokenizer(
|
||||
formatted_full, truncation=False, add_special_tokens=False
|
||||
)["input_ids"]
|
||||
untruncated_prompt_len = len(
|
||||
self.tokenizer(
|
||||
formatted_prompt, truncation=False, add_special_tokens=False
|
||||
)["input_ids"]
|
||||
)
|
||||
completion_ids = full_ids[untruncated_prompt_len:]
|
||||
|
||||
# completion 预算 = 总预算 - 截断后 prompt 实长。配置校验
|
||||
# (max_prompt_length < max_length)保证它恒 > 0
|
||||
completion_budget = self.max_length - len(prompt_ids)
|
||||
completion_ids = completion_ids[:completion_budget]
|
||||
|
||||
all_input_ids.append(prompt_ids + completion_ids)
|
||||
# prompt 位置标 -100:题目不产生 loss,只学解答
|
||||
all_labels.append([IGNORE_INDEX] * len(prompt_ids) + completion_ids)
|
||||
|
||||
# 左 padding:batch 内所有序列右对齐。纯 SFT 用右 padding 也行,但左 padding
|
||||
# 让 trainer 能用一个标量 prompt_length 切 batch(docs/02 §2.4),且与
|
||||
# 层 2+ 的生成场景(生成必须左 padding)统一,全项目只有一种 padding 约定
|
||||
return {
|
||||
"input_ids": _left_pad(all_input_ids, self.pad_token_id), # (B, T)
|
||||
"attention_mask": _left_pad(
|
||||
[[1] * len(ids) for ids in all_input_ids], 0
|
||||
), # (B, T)
|
||||
"labels": _left_pad(all_labels, IGNORE_INDEX), # (B, T)
|
||||
}
|
||||
|
||||
|
||||
def _left_pad(seqs: list[list[int]], pad_value: int) -> torch.Tensor:
|
||||
"""把变长序列在左侧补齐成 (B, T) 张量,T = batch 内最大长度。"""
|
||||
t_max = max(len(s) for s in seqs)
|
||||
return torch.tensor(
|
||||
[[pad_value] * (t_max - len(s)) + s for s in seqs], dtype=torch.long
|
||||
)
|
||||
@@ -0,0 +1,89 @@
|
||||
"""MC 估计与 Dirichlet 贝叶斯平滑——论文 §3.2.2 式(4)(5)。
|
||||
|
||||
把 similarity.py 产出的软计数 k_sem(外部 teacher 信号)与学生自身的
|
||||
chunk 置信度 π̄(内部先验)融合成有界目标 π̂ ∈ (0, 1],供层 5 的 chunk
|
||||
损失当乘子:loss_c = −π̂ · mean(log p)。定理 4.1 三性质由此获得:
|
||||
(a) π̂ 有界 → 无白盒式(2) 的梯度爆炸;(b) π̂ > 0 → k_sem=0 也不塌缩;
|
||||
(c) 先验收缩 → 方差小于频率估计 k/N。
|
||||
|
||||
纯逻辑模块(CLAUDE.md §2):只依赖 torch,toy 张量本地 CPU 可测。
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def chunk_prior(log_probs: torch.Tensor) -> torch.Tensor:
|
||||
"""式(4):π̄ = exp((1/C)·Σ_t log p_t)——学生对整个 chunk 的几何均值置信度。
|
||||
|
||||
C 个 token 概率的几何均值,充当式(5) 的贝叶斯先验:teacher 采样(k_sem)
|
||||
是主信号,π̄ 只是"学生自己觉得这段有多稳"的地板,防 k_sem=0 时目标归零。
|
||||
|
||||
参数:
|
||||
log_probs: 学生对 chunk 内各 token 的对数概率,shape (C,),值 ≤ 0。
|
||||
(调用方从 log_softmax 后 gather 标签位置所得,层 5 负责。)
|
||||
|
||||
返回:
|
||||
π̄,标量张量 shape ()、值 ∈ [1e-8, 1],**不带梯度**。
|
||||
|
||||
实现细节:
|
||||
- log 域先均值再 exp:直接连乘 C=50 个小概率会下溢
|
||||
(50 个 0.01 → 1e-100,超出 fp32 下限 ~1e-38),log 域安全。
|
||||
- detach 命门(参考实现 distillation_trainer.py:2196 同):π̄ 是学生
|
||||
自身概率的函数,若保留梯度,优化器会发现"压低自己的 chunk 概率
|
||||
→ π̄→0 → π̂ 变小 → 损失权重变小"这条逃逸路径——恰在 k_sem=0
|
||||
(teacher 否定)的 chunk 上最有利可图,这些 chunk 最先塌缩。
|
||||
π̄ 只能当常数先验,不能当优化变量。锁死断言见
|
||||
tests/test_estimator_detach.py。
|
||||
- clamp 下限 1e-8:极端负的均值 exp 后可能下溢为 0,而定理 4.1(b)
|
||||
的反塌缩要求 π̄ 严格为正。差异标注:参考实现 clamp(1e-8, 1.0),
|
||||
上限实为冗余——log p ≤ 0 ⇒ mean ≤ 0 ⇒ exp ≤ 1,此处省去。
|
||||
"""
|
||||
if log_probs.numel() == 0:
|
||||
raise ValueError("log_probs 为空:chunk 至少要含 1 个 token")
|
||||
log_pi_bar = log_probs.detach().mean() # (C,) -> ()
|
||||
return log_pi_bar.exp().clamp(min=1e-8)
|
||||
|
||||
|
||||
def bayesian_target(
|
||||
k_sem: float,
|
||||
pi_bar: torch.Tensor,
|
||||
n_rollouts: int,
|
||||
alpha: float,
|
||||
) -> torch.Tensor:
|
||||
"""式(5):π̂ = (k_sem + α·π̄) / (N + α)——chunk 接受概率的贝叶斯估计。
|
||||
|
||||
等价凸组合视角(论文式10):
|
||||
π̂ = N/(N+α) · (k_sem/N) + α/(N+α) · π̄
|
||||
即"teacher 频率估计"与"学生先验"的加权平均;默认 N=10、α=1 时权重
|
||||
约 91% : 9%,teacher 主导,先验只兜底。
|
||||
|
||||
参数:
|
||||
k_sem: 式(3) 的软匹配计数,∈ [0, N](aggregate_similarity 产出)。
|
||||
pi_bar: 式(4) 的先验 π̄,标量张量(chunk_prior 产出)。
|
||||
n_rollouts: teacher rollout 数 N。**必须等于算 k_sem 时的
|
||||
len(teacher_rollouts)**——分子分母口径不一致会系统性偏移 π̂。
|
||||
alpha: 先验强度 α ≥ 0。α=0 退化为频率估计 k/N(层 6 的
|
||||
no_bayesian 消融,参考 config.py:315),失去定理 4.1(b) 保护。
|
||||
|
||||
返回:
|
||||
π̂,标量张量 shape ()、值 ∈ [1e-8, 1],**不带梯度**。
|
||||
|
||||
实现细节:
|
||||
- 差异标注:参考实现(distillation_trainer.py:2205)在此对 π̂ 整体
|
||||
detach;我们的 π̄ 在 chunk_prior 内已 detach,此处的 detach 是
|
||||
第二道防线——防止将来有人把带梯度的张量传进 pi_bar。
|
||||
- clamp(1e-8, 1.0):下限防 α=0 且 k_sem=0 时 π̂=0(乘子归零则该
|
||||
chunk 完全失去监督);上限防 pi_bar 越界传入时 π̂ 溢出概率语义。
|
||||
"""
|
||||
if n_rollouts < 1:
|
||||
raise ValueError(f"n_rollouts 必须 ≥ 1,得到 {n_rollouts}")
|
||||
if alpha < 0:
|
||||
raise ValueError(f"alpha 必须 ≥ 0,得到 {alpha}")
|
||||
if not 0.0 <= k_sem <= n_rollouts:
|
||||
raise ValueError(
|
||||
f"k_sem={k_sem} 越界 [0, {n_rollouts}]:检查是否与"
|
||||
f" len(teacher_rollouts) 口径一致"
|
||||
)
|
||||
# 式(5): π̂ = (k_sem + α·π̄) / (N + α)
|
||||
pi_hat = (k_sem + alpha * pi_bar) / (n_rollouts + alpha) # () -> ()
|
||||
return pi_hat.clamp(1e-8, 1.0).detach()
|
||||
@@ -0,0 +1,131 @@
|
||||
"""语义相似度 φ 与 chunk 级聚合 k_sem——论文 §3.2.1 式(3)。
|
||||
|
||||
logit-free 的支点:学生 chunk 对不对,不再比 token 概率(层 2 白盒式(2)),
|
||||
改比"学生 chunk 文本" vs "teacher rollout 文本"的语义相似度。φ 只依赖
|
||||
文本本身,与两侧 tokenizer 无关,teacher 只需能吐文本(任何 API 均可)。
|
||||
|
||||
纯逻辑模块(CLAUDE.md §2):只依赖标准库,可脱离 torch 在本地 CPU 测试。
|
||||
teacher rollout 怎么采出来是层 5 teacher.py 的事,本模块只吃现成字符串。
|
||||
"""
|
||||
|
||||
from collections import Counter
|
||||
|
||||
|
||||
def rouge1(hypothesis: str, reference: str) -> float:
|
||||
"""ROUGE-1 F1(unigram 重叠率),φ 的候选度量之一(论文 §3.2.1)。
|
||||
|
||||
以词为单位(空白切分)统计两串的 unigram 重叠,算 F1。
|
||||
词袋语义:只看"用了哪些词",不看词序——"a b" vs "b a" 得 1.0。
|
||||
|
||||
参数:
|
||||
hypothesis: 学生 chunk 文本。
|
||||
reference: teacher rollout 文本。
|
||||
|
||||
返回:
|
||||
F1 ∈ [0, 1];任一侧无词(空串/纯空白)时为 0.0。
|
||||
|
||||
实现细节:
|
||||
- 差异标注:参考实现(distillation_trainer.py:1670)用 set 去重后求交,
|
||||
会把 "x x x x" vs "x" 判成满分 1.0;此处用 Counter 多重集
|
||||
(ROUGE-1 标准定义),重复词按 min 计数配对,同例只得 0.4。
|
||||
数学推理文本里重复 token(数字、"="、变量名)极常见,去重会失真。
|
||||
- 差异标注:参考实现分母加 1e-8 防零除,代价是全同串 F1≈0.99999998
|
||||
而非精确 1;此处 overlap==0 时提前返回,分母恒正,无需平滑。
|
||||
"""
|
||||
hyp_counts = Counter(hypothesis.split())
|
||||
ref_counts = Counter(reference.split())
|
||||
if not hyp_counts or not ref_counts:
|
||||
return 0.0
|
||||
# 多重集交:每个词按两侧出现次数的 min 配对
|
||||
overlap = sum((hyp_counts & ref_counts).values())
|
||||
if overlap == 0:
|
||||
return 0.0
|
||||
precision = overlap / sum(hyp_counts.values())
|
||||
recall = overlap / sum(ref_counts.values())
|
||||
return 2 * precision * recall / (precision + recall)
|
||||
|
||||
|
||||
def edit_similarity(hypothesis: str, reference: str) -> float:
|
||||
"""归一化编辑相似度 1 − Levenshtein/max(m,n),论文 §5.1 的默认 φ。
|
||||
|
||||
以词为单位(空白切分)算 Levenshtein 距离(插入/删除/替换各计 1),
|
||||
再归一化到 [0, 1] 取反。顺序敏感:"a b" vs "b a" 距离 2,相似度 0——
|
||||
与 rouge1 的词袋语义形成互补。
|
||||
|
||||
参数:
|
||||
hypothesis: 学生 chunk 文本。
|
||||
reference: teacher rollout 文本。
|
||||
|
||||
返回:
|
||||
相似度 ∈ [0, 1];两侧均空为 1.0(零距离),仅一侧空为 0.0(全删/全插)。
|
||||
|
||||
实现细节:
|
||||
- 差异标注:参考实现(distillation_trainer.py:1682)吃 token id 列表,
|
||||
相似度随 tokenizer 切法漂移,违背本层"文本是公共语言"的初衷
|
||||
(docs/04 §2.1 坑①);此处吃 str、内部按词切,与 rouge1 统一口径。
|
||||
- 两行滚动 DP(同参考实现):空间 O(n) 而非 O(m·n)。
|
||||
"""
|
||||
hyp_words = hypothesis.split()
|
||||
ref_words = reference.split()
|
||||
m, n = len(hyp_words), len(ref_words)
|
||||
if m == 0 and n == 0:
|
||||
return 1.0
|
||||
if m == 0 or n == 0:
|
||||
return 0.0
|
||||
# prev[j] = 前一行的 dist(hyp[:i-1], ref[:j]);curr 原地滚动复用
|
||||
prev = list(range(n + 1))
|
||||
curr = [0] * (n + 1)
|
||||
for i in range(1, m + 1):
|
||||
curr[0] = i
|
||||
for j in range(1, n + 1):
|
||||
cost = 0 if hyp_words[i - 1] == ref_words[j - 1] else 1
|
||||
curr[j] = min(
|
||||
prev[j] + 1, # 删除 hyp[i-1]
|
||||
curr[j - 1] + 1, # 插入 ref[j-1]
|
||||
prev[j - 1] + cost, # 替换(相同则免费)
|
||||
)
|
||||
prev, curr = curr, prev
|
||||
return 1.0 - prev[n] / max(m, n)
|
||||
|
||||
|
||||
def phi(hypothesis: str, reference: str, metric: str = "edit_distance") -> float:
|
||||
"""语义相似度 φ(y_c, ŷ_c) ∈ [0, 1],论文 §3.2.1 式(3) 的原子度量。
|
||||
|
||||
参数:
|
||||
hypothesis: 学生 chunk 文本。
|
||||
reference: teacher rollout 文本。
|
||||
metric: "edit_distance"(默认)或 "rouge1"。
|
||||
差异标注:参考实现配置默认 rouge1(config.py:299),与论文 §5.1
|
||||
的 edit_distance 背离;此处从论文。
|
||||
|
||||
返回:
|
||||
相似度 ∈ [0, 1]。
|
||||
"""
|
||||
if metric == "edit_distance":
|
||||
return edit_similarity(hypothesis, reference)
|
||||
if metric == "rouge1":
|
||||
return rouge1(hypothesis, reference)
|
||||
raise ValueError(f"未知相似度度量: {metric!r}(可选 'edit_distance' / 'rouge1')")
|
||||
|
||||
|
||||
def aggregate_similarity(
|
||||
student_chunk: str,
|
||||
teacher_rollouts: list[str],
|
||||
metric: str = "edit_distance",
|
||||
) -> float:
|
||||
"""式(3):k_sem = Σ_{i=1}^{N} φ(y_c, ŷ_c^{(i)}),chunk 的语义匹配计数。
|
||||
|
||||
学生 chunk 与 N 个 teacher rollout 逐一算 φ 后求和。φ 连续,故 k_sem 是
|
||||
[0, N] 上的实数——"软计数":k_sem≈N 意为学生这段与 teacher 高度一致,
|
||||
k_sem≈0 意为 teacher 从不这么写。它是 π̂(式5)里唯一的外部 teacher 信号。
|
||||
|
||||
参数:
|
||||
student_chunk: 学生 chunk 文本(C 个 token 解码所得)。
|
||||
teacher_rollouts: N 段 teacher 续写文本,与学生 chunk 共享同一前缀
|
||||
y_<c(对应关系由"同一前缀现场生成"保证,无需搜索匹配,docs/04 §1)。
|
||||
metric: 传给 phi,默认 "edit_distance"。
|
||||
|
||||
返回:
|
||||
k_sem ∈ [0, N],N = len(teacher_rollouts)。空列表得 0.0(空和)。
|
||||
"""
|
||||
return sum(phi(student_chunk, r, metric) for r in teacher_rollouts)
|
||||
@@ -0,0 +1,197 @@
|
||||
"""teacher rollout 采样(IO 边缘,论文 §3.2.1)。
|
||||
|
||||
层 1 起步能力:给一批 prompt 批量生成解答,落盘 sha256 键的 JSONL 缓存
|
||||
(键契约在 data.prompt_key 单点定义,本模块与 data.attach_teacher_completions
|
||||
共用)。层 5 在此长出 chunk 前缀续写的 MC rollout 能力。
|
||||
|
||||
连接信息从 `.env` 读取(TEACHER_API_BASE / TEACHER_API_KEY / TEACHER_MODEL),
|
||||
密钥永不出现在代码与配置类里。
|
||||
|
||||
缓存即断点:生成过程逐条追加写盘,任何中断(网络、Ctrl-C、单条失败)后
|
||||
重跑同一命令,已完成的条目自动跳过——API 花的钱不会白花。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from openai import OpenAI
|
||||
|
||||
from ars_opd.configs import TeacherGenConfig
|
||||
from ars_opd.data import prompt_key
|
||||
|
||||
Messages = list[dict[str, str]]
|
||||
|
||||
|
||||
def _load_teacher_env(env_file: str | None = None) -> tuple[str, str, str]:
|
||||
"""从 .env(及进程环境)读取 API 连接三元组,缺一项都显式报错。"""
|
||||
load_dotenv(env_file)
|
||||
values = {}
|
||||
for name in ("TEACHER_API_BASE", "TEACHER_API_KEY", "TEACHER_MODEL"):
|
||||
value = os.environ.get(name, "").strip()
|
||||
if not value:
|
||||
raise ValueError(
|
||||
f"环境变量 {name} 未设置。复制 .env.example 为 .env 并填入真实值。"
|
||||
)
|
||||
values[name] = value
|
||||
return (
|
||||
values["TEACHER_API_BASE"],
|
||||
values["TEACHER_API_KEY"],
|
||||
values["TEACHER_MODEL"],
|
||||
)
|
||||
|
||||
|
||||
def _strip_leading_think(text: str) -> str:
|
||||
"""剥离 content 开头的 <think>...</think> 段(M3 等 reasoning 模型会内联思考)。
|
||||
|
||||
只剥开头一段:解答正文里若出现字面 "<think>" 字样(例如题目在讨论标签本身),
|
||||
不应被误删。
|
||||
"""
|
||||
return re.sub(r"^\s*<think>.*?</think>\s*", "", text, count=1, flags=re.DOTALL)
|
||||
|
||||
|
||||
class TeacherClient:
|
||||
"""OpenAI 兼容的 teacher 客户端:单条生成 + 采样参数收口。
|
||||
|
||||
差异标注:参考实现是 OpenRouter 专用客户端(带其私有请求头与站点字段);
|
||||
我们用通用 OpenAI 客户端 + base_url 配置驱动,任何兼容网关(new-api、
|
||||
vLLM serve、官方 API)都无需改代码。
|
||||
|
||||
测试注入口:传入 client/model 可绕过 .env 与真实网络(见 tests/test_teacher.py)。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
gen_config: TeacherGenConfig,
|
||||
client: OpenAI | None = None,
|
||||
model: str | None = None,
|
||||
) -> None:
|
||||
self.cfg = gen_config
|
||||
if client is None:
|
||||
base, key, env_model = _load_teacher_env()
|
||||
client = OpenAI(
|
||||
base_url=base, api_key=key, max_retries=gen_config.max_retries
|
||||
)
|
||||
model = model or env_model
|
||||
if model is None:
|
||||
raise ValueError("注入 client 时必须同时指定 model")
|
||||
self.client = client
|
||||
self.model = model
|
||||
|
||||
def generate(self, messages: Messages) -> str:
|
||||
"""对单条 prompt(messages 列表,末轮为 user)生成解答文本。
|
||||
|
||||
返回剥离思考段、去首尾空白后的解答。空解答直接报错——空字符串写进
|
||||
缓存会在训练时变成全 -100 的空样本(trainer 会炸,但应在这里更早炸)。
|
||||
"""
|
||||
if self.cfg.system_prompt is not None:
|
||||
messages = [
|
||||
{"role": "system", "content": self.cfg.system_prompt}
|
||||
] + messages
|
||||
resp = self.client.chat.completions.create(
|
||||
model=self.model,
|
||||
messages=messages,
|
||||
temperature=self.cfg.temperature,
|
||||
top_p=self.cfg.top_p,
|
||||
max_tokens=self.cfg.max_tokens,
|
||||
)
|
||||
content = resp.choices[0].message.content or ""
|
||||
if self.cfg.strip_think:
|
||||
content = _strip_leading_think(content)
|
||||
content = content.strip()
|
||||
if not content:
|
||||
raise ValueError(
|
||||
"teacher 返回空解答(可能:max_tokens 太小把思考截断在半途,"
|
||||
"或模型拒答)。该条不会入缓存。"
|
||||
)
|
||||
return content
|
||||
|
||||
|
||||
def generate_completions(
|
||||
prompts: list[Messages],
|
||||
cache_path: str,
|
||||
teacher: TeacherClient,
|
||||
) -> None:
|
||||
"""批量生成解答并追加写入 JSONL 缓存(每行 {"key", "completion", "preview"})。
|
||||
|
||||
- 已在缓存中的键直接跳过(断点续传);
|
||||
- 并发线程池执行,每完成一条立即写盘并 flush(中断不丢已完成的结果);
|
||||
- 单条失败不中断其余任务(并发中的兄弟请求已经花了钱,先让它们落盘),
|
||||
全部结束后若有失败则汇总显式报错——重跑即续传,绝不静默缺数据。
|
||||
"""
|
||||
path = Path(cache_path)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
done_keys = _cached_keys(path)
|
||||
todo = [(prompt_key(p), p) for p in prompts]
|
||||
todo = [(k, p) for k, p in todo if k not in done_keys]
|
||||
print(
|
||||
f"[teacher] 共 {len(prompts)} 条:缓存命中 {len(prompts) - len(todo)},"
|
||||
f"待生成 {len(todo)},并发 {teacher.cfg.concurrency}",
|
||||
flush=True,
|
||||
)
|
||||
if not todo:
|
||||
return
|
||||
|
||||
failures: list[tuple[str, str]] = []
|
||||
finished = 0
|
||||
start = time.monotonic()
|
||||
# 写盘收口在主线程(as_completed 消费端),工作线程只跑网络请求——
|
||||
# 多线程同写一个文件句柄会交错损坏 JSONL
|
||||
with open(path, "a", encoding="utf-8") as f:
|
||||
with ThreadPoolExecutor(max_workers=teacher.cfg.concurrency) as pool:
|
||||
futures = {pool.submit(teacher.generate, p): (k, p) for k, p in todo}
|
||||
for fut in as_completed(futures):
|
||||
key, p = futures[fut]
|
||||
try:
|
||||
completion = fut.result()
|
||||
except Exception as e: # noqa: BLE001 —— 收集后统一显式报错,非静默吞错
|
||||
failures.append((key, repr(e)))
|
||||
continue
|
||||
finally:
|
||||
finished += 1
|
||||
if finished % 20 == 0 or finished == len(todo):
|
||||
elapsed = time.monotonic() - start
|
||||
rate = finished / elapsed * 60 # 条/分
|
||||
eta = (len(todo) - finished) / rate if rate > 0 else 0
|
||||
print(
|
||||
f"[teacher] {finished}/{len(todo)} 完成 | "
|
||||
f"{rate:.1f} 条/分 | 已用 {elapsed / 60:.1f} 分 | "
|
||||
f"预计剩余 {eta:.0f} 分",
|
||||
flush=True,
|
||||
)
|
||||
record = {
|
||||
"key": key,
|
||||
"completion": completion,
|
||||
# preview 仅供人工抽查缓存文件,消费端(attach)只认 key/completion
|
||||
"preview": p[-1]["content"][:80],
|
||||
}
|
||||
f.write(json.dumps(record, ensure_ascii=False) + "\n")
|
||||
f.flush()
|
||||
|
||||
if failures:
|
||||
examples = "; ".join(f"{k[:12]}…: {err}" for k, err in failures[:3])
|
||||
raise RuntimeError(
|
||||
f"{len(failures)}/{len(todo)} 条生成失败(成功的已入缓存,重跑本命令"
|
||||
f"即断点续传)。前几条错误:{examples}"
|
||||
)
|
||||
|
||||
|
||||
def _cached_keys(path: Path) -> set[str]:
|
||||
"""读取缓存中已有的键集合;文件不存在视为空缓存(首跑)。"""
|
||||
if not path.exists():
|
||||
return set()
|
||||
keys = set()
|
||||
with open(path, encoding="utf-8") as f:
|
||||
for line_no, line in enumerate(f, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
rec = json.loads(line) # 坏行直接炸:缓存损坏必须暴露,不能悄悄重新生成
|
||||
keys.add(rec["key"])
|
||||
return keys
|
||||
@@ -0,0 +1,423 @@
|
||||
"""训练编排(IO 边缘)。层 1 形态:掩码 SFT 损失 + 最小 HF Trainer 子类。
|
||||
|
||||
论文锚点:§3.1 式(1) 的标准交叉熵 SFT(监督目标是 teacher rollout,见 docs/02 §1)。
|
||||
层 5 会在此模块长出式(8) 的 chunk 蒸馏损失与 KL 锚定;届时式(1) 路径保留为基线。
|
||||
|
||||
结构:损失算法收在纯张量函数 `sft_loss`(可用 toy 张量在 CPU 单测),
|
||||
`SFTTrainer` 只做接线——模型前向、调 `sft_loss`、把指标并进日志,
|
||||
其余一切(优化器、调度、DDP、checkpoint 保存)原样继承 HF Trainer。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from transformers import Trainer
|
||||
|
||||
from ars_opd.data import IGNORE_INDEX
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 损失算法(纯张量函数,对拍 distillation_trainer.py:751-759 + 2776-2874)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def compute_prompt_length(attention_mask: torch.Tensor, labels: torch.Tensor) -> int:
|
||||
"""batch 级 prompt 边界 = batch 内最短的"有效长度 - completion 长度"。
|
||||
|
||||
参数(B = batch,T = padding 后长度):
|
||||
- attention_mask: (B, T),左 padding 位置为 0
|
||||
- labels: (B, T),completion 位置为 token id,其余为 -100
|
||||
|
||||
返回标量 pl。非显然约束:取 batch **最小值**是为了不切掉任何行的
|
||||
completion token——左 padding 下所有序列右对齐,第 r 行的 completion 起点
|
||||
索引是 T - comp_r ≥ prompt_r ≥ min,故切片 [pl:] 必然包含全部 completion;
|
||||
代价是长 prompt 行会漏进一些 prompt token,由 sft_loss 用 labels 重掩码兜住。
|
||||
"""
|
||||
full_lengths = attention_mask.sum(dim=1) # (B,) 每行非 padding 的 token 数
|
||||
completion_lengths = (labels != IGNORE_INDEX).sum(dim=1) # (B,)
|
||||
return int((full_lengths - completion_lengths).min().item())
|
||||
|
||||
|
||||
def sft_loss(
|
||||
logits: torch.Tensor,
|
||||
input_ids: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, int]:
|
||||
"""式(1):completion 位置上的移位交叉熵。
|
||||
|
||||
参数(B = batch,T = padding 后长度,V = 词表大小):
|
||||
- logits: (B, T, V) 模型对全序列的输出
|
||||
- input_ids / labels / attention_mask: (B, T),SFTCollator 的产物
|
||||
|
||||
返回 (标量 loss, 本 batch 有效 completion token 数)。
|
||||
|
||||
切片几何(docs/02 §2.4 的四行,此处为权威实现):
|
||||
位置 t-1 的 logit 预测位置 t 的 token,故 logits 取 [pl-1, T-1) 、
|
||||
targets 取 [pl, T),两段长度同为 T-pl,逐位对齐。
|
||||
"""
|
||||
pl = compute_prompt_length(attention_mask, labels)
|
||||
if pl < 1:
|
||||
# 需要位置 pl-1 的 logit 存在;pl=0 意味着某行完全没有 prompt token,
|
||||
# 数据管线出了问题(collator 保证 prompt 至少含模板 token)
|
||||
raise ValueError(f"prompt_length={pl} < 1,存在无 prompt 的数据行")
|
||||
|
||||
shifted_logits = logits[:, pl - 1 : -1, :] # (B, T, V) -> (B, T-pl, V)
|
||||
targets = input_ids[:, pl:].clone() # (B, T-pl);clone: 下面要原地改写
|
||||
|
||||
# 重掩码:切片里漏进的 prompt token(长 prompt 行)与 padding 全部置 -100。
|
||||
# labels 是权威掩码,切片几何只是省算力——正确性完全由这一步保证
|
||||
invalid = labels[:, pl:] == IGNORE_INDEX # (B, T-pl)
|
||||
targets[invalid] = IGNORE_INDEX
|
||||
|
||||
num_valid = int((~invalid).sum().item())
|
||||
if num_valid == 0:
|
||||
# 差异标注:参考实现(trainer:2810-2812)对非有限 loss 返回零梯度标量静默
|
||||
# 继续;我们显式报错——collator 已挡下 prompt-only 行,走到这里仍全被
|
||||
# 掩码只可能是数据坏了(如 teacher 返回空解答),必须暴露而非跳过
|
||||
raise ValueError(
|
||||
"本 batch 没有任何有效 completion token(全被 -100 掩码)。"
|
||||
"检查 teacher 解答是否为空、completion 预算是否被截光。"
|
||||
)
|
||||
|
||||
vocab = shifted_logits.shape[-1]
|
||||
loss = F.cross_entropy(
|
||||
shifted_logits.reshape(-1, vocab), # (B*(T-pl), V)
|
||||
targets.reshape(-1), # (B*(T-pl),)
|
||||
ignore_index=IGNORE_INDEX,
|
||||
)
|
||||
return loss, num_valid
|
||||
|
||||
|
||||
def token_divergence(
|
||||
student_logits: torch.Tensor,
|
||||
teacher_logits: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
beta: float = 1.0,
|
||||
temperature: float = 1.0,
|
||||
) -> tuple[torch.Tensor, int]:
|
||||
"""式(2):completion 位置上的 token 级(广义)KL 散度,全词表精确。
|
||||
|
||||
论文锚点:§3.1 式(2) L = E_{y~π_θ}[Σ_t KL(π_θ(·|y_<t,x) ‖ π_T(·|y_<t,x))]。
|
||||
本函数只管"给定两组**已对齐**的 logits,算散度标量"——on-policy 采样(y~π_θ)
|
||||
与移位对齐([pl-1:-1])由调用方(DistillTrainer, U4)负责,复用 sft_loss 同一套
|
||||
切片几何。故这里不吃 input_ids:散度是分布对分布,不需要目标 token,labels
|
||||
仅用于定位有效位置。
|
||||
|
||||
参数(B=batch,T=已移位对齐长度,V=词表大小):
|
||||
- student_logits: (B, T, V),student 前向输出(带梯度)
|
||||
- teacher_logits: (B, T, V),teacher no_grad 前向输出(无梯度)
|
||||
- labels: (B, T),completion 位为 token id、其余为 IGNORE_INDEX;只做有效位掩码
|
||||
- beta: KL 方向(docs/03 §2.3 三副面孔)。0=前向 KL(π_T‖π_θ)、1=反向 KL(π_θ‖π_T)=
|
||||
式(2)、(0,1)=JSD 插值
|
||||
- temperature: softmax 前除进两侧 logits 的温度(§2.3),软化/锐化分布
|
||||
|
||||
返回 (per-token mean 散度标量, 有效 token 数)。
|
||||
|
||||
差异标注:参考实现(DT:2408-2491)含 top-k 稀疏 + 尾桶快路,我们只保留全词表这
|
||||
一条精确路径(docs/03 §4 删除清单:本地同 tokenizer teacher 放得下)。因全词表
|
||||
log_softmax 对有限 logits 恒有限,也不需要参考在 -inf 支持集上的 nan_to_num 兜底。
|
||||
"""
|
||||
if not 0.0 <= beta <= 1.0:
|
||||
raise ValueError(f"beta 必须在 [0,1](0=前向/1=反向/中间=JSD),收到 {beta}")
|
||||
|
||||
# 温度除进 logits、softmax 之前(§2.3):调分布形状,不是等比缩概率
|
||||
student_logits = student_logits / temperature
|
||||
teacher_logits = teacher_logits / temperature
|
||||
|
||||
# 全程 log 域(§2.3 数值稳定性)。log_probs 两侧都要;probs 按分支只算用得上的
|
||||
# 那一份——(B,T,V) 在真实规模下每份 ~2.5G(§5 显存账),默认 β=1 热路径只需
|
||||
# student_probs,不materialize teacher_probs
|
||||
student_log_probs = F.log_softmax(student_logits, dim=-1) # (B, T, V)
|
||||
teacher_log_probs = F.log_softmax(teacher_logits, dim=-1) # (B, T, V)
|
||||
|
||||
if beta == 1.0:
|
||||
# 反向 KL(π_θ‖π_T) = Σ_v π_θ (log π_θ − log π_T) —— 式(2)
|
||||
# 非显然约束(§4.1 梯度爆炸源):对 student logit 的梯度含 π_θ·log(π_θ/π_T),
|
||||
# on-policy 采到 teacher 眼中烂 token(π_T→0)时 log 比值→∞,单 token 梯度可
|
||||
# 炸掉整个 batch。这正是层 2 要亲眼观察、层 5 用有界乘子 π̂ 替换的病灶
|
||||
per_token = (
|
||||
student_log_probs.exp() * (student_log_probs - teacher_log_probs)
|
||||
).sum(-1) # (B, T, V) -> (B, T)
|
||||
elif beta == 0.0:
|
||||
# 前向 KL(π_T‖π_θ) = Σ_v π_T (log π_T − log π_θ)
|
||||
per_token = (
|
||||
teacher_log_probs.exp() * (teacher_log_probs - student_log_probs)
|
||||
).sum(-1)
|
||||
else:
|
||||
# JSD 插值:m = (1−β)π_θ + β π_T;β·KL(π_T‖m) + (1−β)·KL(π_θ‖m)
|
||||
student_probs = student_log_probs.exp()
|
||||
teacher_probs = teacher_log_probs.exp()
|
||||
mixture = (1.0 - beta) * student_probs + beta * teacher_probs
|
||||
# clamp_min(tiny) 防 log0(§2.3):混合概率理论上恒正,此处是浮点下溢兜底
|
||||
log_mixture = mixture.clamp_min(torch.finfo(mixture.dtype).tiny).log()
|
||||
kl_teacher = (teacher_probs * (teacher_log_probs - log_mixture)).sum(-1)
|
||||
kl_student = (student_probs * (student_log_probs - log_mixture)).sum(-1)
|
||||
per_token = beta * kl_teacher + (1.0 - beta) * kl_student
|
||||
|
||||
# 掩码 + per-token mean(docs/03 §2.3:reduction 实义为 sum/有效token数,
|
||||
# 与层 1 sft_loss 同尺度,两条 loss 曲线才可比)
|
||||
mask = labels != IGNORE_INDEX # (B, T)
|
||||
num_valid = int(mask.sum().item())
|
||||
if num_valid == 0:
|
||||
# 与 sft_loss 同纪律:走到这里全被掩码只可能是数据/生成坏了,显式报错不静默
|
||||
raise ValueError(
|
||||
"本 batch 没有任何有效 completion token(全被 -100 掩码)。"
|
||||
"检查 on-policy 生成是否产出了空 completion。"
|
||||
)
|
||||
loss = per_token[mask].sum() / num_valid
|
||||
return loss, num_valid
|
||||
|
||||
|
||||
def build_generated_batch(
|
||||
prompts: torch.Tensor,
|
||||
prompt_attention_mask: torch.Tensor,
|
||||
gen_output: torch.Tensor,
|
||||
eos_token_id: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""把 model.generate 的输出重建成 (input_ids, attention_mask, labels)。
|
||||
|
||||
对应 docs/03 §2.5"生成结果重建 input_ids/labels 写回"。on-policy 下 completion
|
||||
不来自数据、而是 student 现场生成,故 labels 也在生成后现造:prompt 段全 -100、
|
||||
生成段有效处填 token id,供 token_divergence 掩码。
|
||||
|
||||
参数(B=batch,P=prompt padding 长度,G=本 batch 最大生成长度):
|
||||
- prompts: (B, P) 左 padding 的 prompt(SFTCollator prompt_only 产物)
|
||||
- prompt_attention_mask: (B, P) prompt 左 padding 位为 0
|
||||
- gen_output: (B, P+G) model.generate 输出(前 P 列即 prompts,后 G 列是生成)
|
||||
- eos_token_id: 生成终止符 id
|
||||
|
||||
返回 (input_ids (B,P+G), attention_mask (B,P+G), labels (B,P+G))。
|
||||
|
||||
非显然约束(生成段的右 padding 掩码):generate 对提前结束的序列在右侧补
|
||||
padding 到 batch 最大长度。"首个 eos(含)之前有效"用 `cumsum - self == 0`
|
||||
实现——它精确保留到首个 eos、屏蔽其后一切(无论其后是 eos 还是 pad,也无论
|
||||
pad_token 是否等于 eos),避免"pad==eos 时把补位当解答"或"pad!=eos 时漏掉补位"
|
||||
两种静默错误。无 eos(撞 max_new_tokens)则整段生成全有效。
|
||||
"""
|
||||
b, p = prompts.shape
|
||||
gen_tokens = gen_output[:, p:] # (B, G) 纯生成段
|
||||
is_eos = gen_tokens == eos_token_id # (B, G)
|
||||
# cumsum - self:截至本位、其**之前**出现过的 eos 数;==0 即"首个 eos 及之前"
|
||||
gen_valid = (is_eos.cumsum(dim=1) - is_eos.long()) == 0 # (B, G) bool
|
||||
|
||||
input_ids = gen_output
|
||||
attention_mask = torch.cat(
|
||||
[prompt_attention_mask, gen_valid.long()], dim=1
|
||||
) # (B, P+G)
|
||||
prompt_labels = torch.full(
|
||||
(b, p), IGNORE_INDEX, dtype=torch.long, device=prompts.device
|
||||
)
|
||||
gen_labels = torch.where(
|
||||
gen_valid, gen_tokens, torch.full_like(gen_tokens, IGNORE_INDEX)
|
||||
) # 生成段:有效处填 token id,其余 -100
|
||||
labels = torch.cat([prompt_labels, gen_labels], dim=1) # (B, P+G)
|
||||
return input_ids, attention_mask, labels
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Trainer 接线
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SFTTrainer(Trainer):
|
||||
"""最小 SFT Trainer:只重写 compute_loss 与 log,其余全部继承 HF Trainer。
|
||||
|
||||
用法(见 scripts/ 训练脚本):与 HF Trainer 完全同参构造,
|
||||
data_collator 传 SFTCollator,train_dataset 传 load_sft_dataset 的产物。
|
||||
"""
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
# 关 dropout(参考实现同款,docs/02 §3 保留清单):层 2+ 的蒸馏要求
|
||||
# student/ref 两次前向可比,前向必须确定;Qwen3 默认无 dropout,
|
||||
# 此处是无害的统一前置
|
||||
for module in self.model.modules():
|
||||
if isinstance(module, torch.nn.Dropout):
|
||||
module.p = 0.0
|
||||
# 非显然约束:显式退出 HF 的新式梯度累积契约(v5 trainer.py:1977 文档
|
||||
# 原话:"If you are not using num_items_in_batch ... overwrite
|
||||
# self.model_accepts_loss_kwargs to False")。新式契约要求返回
|
||||
# sum/全局token数,且依赖基类 compute_loss 尾部的 ×world_size 补偿
|
||||
# (trainer.py:2028)——我们整体重写了 compute_loss,那段补偿不会执行,
|
||||
# 曾致 loss 与梯度 ÷4(2026-07-18,第二幕;第一幕是返回裸 mean 被 ×8,
|
||||
# 全程记录见 docs/02 §5)。退出后回到经典契约:返回本微批 mean,
|
||||
# Trainer 负责 ÷累积步数,日志跨卡平均,版本稳定
|
||||
self.model_accepts_loss_kwargs = False
|
||||
self._token_counts: list[int] = []
|
||||
|
||||
def compute_loss(
|
||||
self,
|
||||
model: Any,
|
||||
inputs: dict[str, torch.Tensor],
|
||||
return_outputs: bool = False,
|
||||
num_items_in_batch: int | None = None,
|
||||
) -> torch.Tensor | tuple[torch.Tensor, Any]:
|
||||
# 非显然约束:不把 labels 传给模型前向。HF 模型收到 labels 会自己算
|
||||
# "全序列移位 CE"并放进 outputs.loss,那会绕过 batch-min 切片与重掩码;
|
||||
# 损失的权威实现只能有 sft_loss 一处
|
||||
outputs = model(
|
||||
input_ids=inputs["input_ids"],
|
||||
attention_mask=inputs["attention_mask"],
|
||||
)
|
||||
loss, num_tokens = sft_loss(
|
||||
outputs.logits,
|
||||
inputs["input_ids"],
|
||||
inputs["labels"],
|
||||
inputs["attention_mask"],
|
||||
)
|
||||
# num_items_in_batch 有意忽略:已在 __init__ 退出新式契约(见彼处注释),
|
||||
# 本函数返回微批 mean,÷累积步数由 Trainer.training_step 负责。
|
||||
# 代价是微批按相同权重而非 token 数加权(百分之几的偏差,与参考实现同行为)
|
||||
self._token_counts.append(num_tokens)
|
||||
return (loss, outputs) if return_outputs else loss
|
||||
|
||||
def log(self, logs: dict[str, float], start_time: float | None = None) -> None:
|
||||
"""在 HF 的常规日志里并入每步有效 token 数均值。
|
||||
|
||||
sanity run 时盯这个数:若它远小于预期(≈batch 内解答总长),说明掩码
|
||||
把 completion 也吞了——loss 曲线看不出这种错,token 数看得出。
|
||||
"""
|
||||
if self._token_counts:
|
||||
logs["sft/num_tokens_per_step"] = sum(self._token_counts) / len(
|
||||
self._token_counts
|
||||
)
|
||||
self._token_counts = []
|
||||
super().log(logs, start_time)
|
||||
|
||||
|
||||
class DistillTrainer(Trainer):
|
||||
"""层 2 white-box OPD:on-policy 生成 → teacher no_grad 前向 → token 级反向 KL。
|
||||
|
||||
论文锚点:§3.1 式(2)。每个微批现场编排三步(docs/03 §2.2 主流程的精简版,
|
||||
已按 §4 删除 buffer/稀疏路径/off-policy 抽签):
|
||||
1. student 采样生成轨迹 y~π_θ(on-policy,no_grad——只采样不回传);
|
||||
2. student 带梯度前向 + teacher no_grad 前向,得两份全词表 logits;
|
||||
3. token_divergence 算 KL,梯度只经 student 那一路。
|
||||
|
||||
与 SFTTrainer 的关系:损失几何(移位对齐、batch-min prompt_length、labels 重
|
||||
掩码)完全同款,直接复用 compute_prompt_length;唯一差别是"completion 从哪来"
|
||||
——SFT 读数据缓存,这里 student 现场生成。
|
||||
|
||||
构造(见 scripts/train_whitebox.py):像 HF Trainer 一样传 model(student)/args/
|
||||
train_dataset/data_collator(prompt_only 的 SFTCollator),另用关键字传 teacher_model
|
||||
与 teacher_tokenizer,以及蒸馏超参。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*args: Any,
|
||||
teacher_model: Any,
|
||||
teacher_tokenizer: Any,
|
||||
beta: float = 1.0,
|
||||
kl_temperature: float = 1.0,
|
||||
gen_temperature: float = 1.0,
|
||||
gen_top_p: float = 1.0,
|
||||
max_new_tokens: int = 1024,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
# student tokenizer 从 collator 取(prompt_only collator 必持有它)
|
||||
student_tokenizer = self.data_collator.tokenizer
|
||||
# 白盒前提校验(docs/03 §1、§2.6):KL 逐词表位对齐,同 tokenizer 才有意义。
|
||||
# 构造时就炸——不像参考实现(DT:2876)拖到第一步前向才炸
|
||||
if teacher_tokenizer.get_vocab() != student_tokenizer.get_vocab():
|
||||
raise ValueError(
|
||||
"teacher 与 student 的 tokenizer 词表不一致——白盒 KL 要求逐词表位"
|
||||
"对应(docs/03 §1)。请换用与 student 同 tokenizer 的 teacher。"
|
||||
)
|
||||
self._student_tokenizer = student_tokenizer
|
||||
|
||||
# teacher:eval + 冻结参数 + 每进程一份副本(DDP 每卡一份,docs/03 §2.6)。
|
||||
# 设备迁移推迟到 compute_loss——此刻 student 还没被 Trainer 放到卡上
|
||||
self.teacher = teacher_model.eval()
|
||||
for param in self.teacher.parameters():
|
||||
param.requires_grad_(False)
|
||||
|
||||
self.beta = beta
|
||||
self.kl_temperature = kl_temperature
|
||||
self.gen_temperature = gen_temperature
|
||||
self.gen_top_p = gen_top_p
|
||||
self.max_new_tokens = max_new_tokens
|
||||
|
||||
# 关 dropout:生成、student 前向、teacher 前向三者须确定可比(同 SFTTrainer)
|
||||
for module in self.model.modules():
|
||||
if isinstance(module, torch.nn.Dropout):
|
||||
module.p = 0.0
|
||||
# 同层 1:显式退出 HF 新式梯度累积契约(docs/02 §5,trainer.py:1977)
|
||||
self.model_accepts_loss_kwargs = False
|
||||
self._token_counts: list[int] = []
|
||||
|
||||
def compute_loss(
|
||||
self,
|
||||
model: Any,
|
||||
inputs: dict[str, torch.Tensor],
|
||||
return_outputs: bool = False,
|
||||
num_items_in_batch: int | None = None,
|
||||
) -> torch.Tensor | tuple[torch.Tensor, Any]:
|
||||
prompts = inputs["prompts"]
|
||||
prompt_attention_mask = inputs["prompt_attention_mask"]
|
||||
|
||||
# teacher 迁到 student 所在卡(一次性;.to 幂等,后续步是 no-op)
|
||||
if self.teacher.device != prompts.device:
|
||||
self.teacher = self.teacher.to(prompts.device)
|
||||
|
||||
# 1. on-policy 生成。no_grad:GKD 标准做法是"采样一次、再 teacher-forcing
|
||||
# 前向算分布",梯度经第 3 步的前向回传,不经采样本身。DDP 下须用 unwrap
|
||||
# 后的模型(DDP 包装体不暴露 generate)
|
||||
unwrapped_model = self.accelerator.unwrap_model(model)
|
||||
with torch.no_grad():
|
||||
gen_output = unwrapped_model.generate(
|
||||
input_ids=prompts,
|
||||
attention_mask=prompt_attention_mask,
|
||||
max_new_tokens=self.max_new_tokens,
|
||||
do_sample=True,
|
||||
temperature=self.gen_temperature,
|
||||
top_p=self.gen_top_p,
|
||||
pad_token_id=self._student_tokenizer.pad_token_id,
|
||||
eos_token_id=self._student_tokenizer.eos_token_id,
|
||||
)
|
||||
|
||||
# 2. 重建 input_ids/attention_mask/labels(prompt 段 -100、生成段有效处填 id)
|
||||
input_ids, attention_mask, labels = build_generated_batch(
|
||||
prompts,
|
||||
prompt_attention_mask,
|
||||
gen_output,
|
||||
self._student_tokenizer.eos_token_id,
|
||||
)
|
||||
|
||||
# 3. student 带梯度前向 + teacher no_grad 前向(两份全词表 logits)。
|
||||
# 非显然约束:teacher no_grad 免掉的是它内部几十层激活的反向图(②),
|
||||
# 但输出 logits(①)仍占满 (B,L,V) 显存——两份都要算进 §5 显存账
|
||||
student_outputs = model(input_ids=input_ids, attention_mask=attention_mask)
|
||||
with torch.no_grad():
|
||||
teacher_logits = self.teacher(
|
||||
input_ids=input_ids, attention_mask=attention_mask
|
||||
).logits
|
||||
|
||||
# 4. 移位对齐(与 sft_loss 同款几何,docs/02 §2.4)→ token_divergence。
|
||||
# 切片漏进的 prompt token 与生成段右 padding 由 labels 重掩码兜住
|
||||
pl = compute_prompt_length(attention_mask, labels)
|
||||
if pl < 1:
|
||||
raise ValueError(f"prompt_length={pl} < 1,生成批次存在无 prompt 的行")
|
||||
loss, num_tokens = token_divergence(
|
||||
student_outputs.logits[:, pl - 1 : -1, :],
|
||||
teacher_logits[:, pl - 1 : -1, :],
|
||||
labels[:, pl:],
|
||||
beta=self.beta,
|
||||
temperature=self.kl_temperature,
|
||||
)
|
||||
self._token_counts.append(num_tokens)
|
||||
return (loss, student_outputs) if return_outputs else loss
|
||||
|
||||
def log(self, logs: dict[str, float], start_time: float | None = None) -> None:
|
||||
"""并入每步的生成 token 数均值:远程冒烟盯它,骤降=生成塌成空串。"""
|
||||
if self._token_counts:
|
||||
logs["distill/num_gen_tokens_per_step"] = sum(self._token_counts) / len(
|
||||
self._token_counts
|
||||
)
|
||||
self._token_counts = []
|
||||
super().log(logs, start_time)
|
||||
+29
-7
@@ -1,6 +1,25 @@
|
||||
# 00 · 分层重构路线图
|
||||
|
||||
> 原则:按论文概念的依赖顺序逐层重建,每层完成后代码可运行、可验证。学习路径 = 提交历史。
|
||||
> 后面各层不提前细化——细节在进入该层时随章节文档长出来(依据见 `appendix-claudemd-decisions.md` 的延迟接入哲学)。
|
||||
|
||||
## 当前进度(存档点)
|
||||
|
||||
> 每次断点(层完成/工作暂停)更新此节。恢复上下文时:读 CLAUDE.md → 本节 → 对应章节文档。
|
||||
|
||||
- **日期**: 2026-07-19
|
||||
- **当前层**: 层 3(相似度 φ + MC 估计器),**docs/04 已精读**(用户已理解:logit-free 支点=比文本非比 logits、k_sem 是外部 teacher 信号 π̄ 只是防塌缩地板、detach 命门、chunk-vs-prefix 对应靠"同一前缀现场生成 teacher 续写"非搜索匹配);**下一步开写 E1**。层 3 全本地 CPU 纯逻辑,不碰 GPU/远程。E1 similarity.py(式3 φ+k_sem,两度量统一到词级文本、默认 edit_distance 对齐论文§5.1)、E2 estimator.py(式4 π̄ 几何均值+detach、式5 π̂ 贝叶斯凸组合)、E3 test_estimator_detach.py 接真实现。设计取舍已定见 docs/04 §3。层 2 ✅ 已关账
|
||||
- **层 2 代码构成**: U1 DistillConfig(configs.py,两温度分名/三处刻意缺席/max_grad_norm 显式化);U2 token_divergence + 梯度爆炸单测(trainer.py,全词表 KL/JSD);U3 SFTCollator prompt_only 模式(data.py);U4 DistillTrainer + build_generated_batch(trainer.py,生成→双前向→divergence);U5 train_whitebox.py/.sh(full/sanity/noclip 三模式)。附带:load_sft_dataset 毛刺已磨平(改吃散装参数);hf-mirror 不代理 Xet CAS → HF_HUB_DISABLE_XET=1
|
||||
- **层 2 关账判据全过**(详见 docs/03 §6.1 实证): 66 单测全绿;B=4/T=2048 **实测不 OOM**(§5 估算成立);两次远程跑(sanity 裁到 1.0 / noclip ≈关裁剪+lr5×)均平稳、不 NaN、生成不塌(num_gen ~1900-4096)、loss 0.35→0.21 下降;checkpoint 存下
|
||||
- **层 2 关键发现(勘误"预期见毛刺")**: §4.1 梯度爆炸是真机制(U2 单测坐实单 token π_T→0 暴涨),但真实训练**高度阻尼**——同门 teacher + per-token mean 摊平,batch 级 grad_norm 峰值仅 ~14 且只降不升,关裁剪也不炸。启示:`grad_norm` 日志是**裁剪前**值(14→2 那串即爆炸证据,被 HF 默认 max_grad_norm=1.0 静默压平,现已显式化);层 5 有界乘子 π̂ 真正杀手锏是 **logit-free**(白盒的同 tokenizer 约束把你锁在温和区间)
|
||||
- **层 2 接口回看**(§6.5 每层必做,全部判"深",无需返工的毛刺): DistillConfig/token_divergence/build_generated_batch/load_sft_dataset(已修) 接口均远简于实现。三条**记录不返工**的小注:① DistillTrainer 从 self.data_collator.tokenizer 取 student tokenizer(隐式耦合,但省一个冗余参数,可接受);② SFTCollator 名字略超范(现含 prompt_only 非 SFT 模式),rename 的 churn 不值;③ token_divergence 的 labels 仅作掩码非目标(已在 docstring 标注)
|
||||
- **层 0**: ✅ 已关账(2026-07-18)
|
||||
- **层 1**: ✅ 已关账(2026-07-18)。判据全过:正本缓存 sha `33deb18c…`(1000 条,键唯一,think 残留 0,仅 2 条硬题截断);正式 1 epoch loss 0.94→0.60(56s/16 步);checkpoint 生成通顺(`/data/zym/outputs/sft_qwen3-0.6b_dapo1k`);接口回看完成(全部模块判"深";毛刺记录:load_sft_dataset 吃整个 SFTConfig 迫使诊断脚本填假 output_dir,层 2 第二消费方出现时定夺)
|
||||
- **层 1 疤痕档案**(详见 docs/02 §2.6/§5 勘误): ① 显存大头是 (B,T,V) logits 链与激活(正比 B×T,与参数量无关),B=8 曾爆 80G;② HF 梯度累积契约两幕剧(×8 → ÷4),终解 = model_accepts_loss_kwargs=False 退出新式契约;③ 缓存正本纪律:本地生成一次、单向 scp 分发、sha256 对账,两侧独立生成曾花双份钱且内容漂移
|
||||
- **诊断工具箱**: scripts/diag_collator.py(对齐链逐环)、diag_loss_probe.py(预训练 CE 基准 0.85)、diag_generate.py(生成质量)——层 2+ 数值异常照此三板斧
|
||||
- **已完成学习**: 第一/二/三章全部精讲(第三章含 standard 路径三大反直觉点、两温度、抉择原则、§4.1 梯度爆炸机制+实证)
|
||||
- **未精讲的文档账**: docs/01 的 §3.7(KL 锚三处实现差异)、§3.8(论文外稳定器)、§4(训练步流程走读)
|
||||
- **远程磁盘备忘**: 根分区 100% 的结构性原因是 `/root/zym`(507G 历史工作区)压在根分区,建议择期整体搬迁 `/data`;临时缓解 = 清 `/tmp/pip-unpack-*`、旧 tar.gz、journal。所有新增写盘已改道 `/data/zym`
|
||||
|
||||
## 分层计划
|
||||
|
||||
@@ -9,7 +28,7 @@
|
||||
| 0 | 环境与骨架 | — | 本地/远程 conda 环境、gitea 同步、包骨架 | 两端 `pytest` 空跑通过 |
|
||||
| 1 | SFT 基线 | §3.1 式(1) | 数据管线 + 最小 SFT 训练脚本(Qwen3-0.6B) | 远程 4 卡跑通,loss 正常下降 |
|
||||
| 2 | White-box OPD 基线 | §3.1 式(2) | token 级反向 KL 蒸馏(teacher Qwen3-4B 本地 vLLM) | 远程跑通;理解式(2)梯度爆炸问题(§4.1) |
|
||||
| 3 | 相似度 + MC 估计器 | §3.2.1-3.2.2 式(3)(4)(5) | `similarity.py` + `estimator.py`(纯逻辑) | 本地 CPU 单测,对拍 `validate_mc_estimator.py` |
|
||||
| 3 | 相似度 + MC 估计器 | §3.2.1-3.2.2 式(3)(4)(5) | `similarity.py` + `estimator.py`(纯逻辑) | 本地 CPU 单测,对拍 `validate_mc_estimator.py`;detach 命门测试已预置(`tests/test_estimator_detach.py`),完成后需接入真实实现 |
|
||||
| 4 | Peak-entropy 调度器 | §3.2.3 式(6)(7) | `chunking.py`(纯逻辑) | 本地 CPU 单测:toy 熵序列上验证 chunk 选择与合并 |
|
||||
| 5 | 完整 OmniOPD | §3.2.4 式(8) | `teacher.py`(API 客户端+缓存)+ `trainer.py`(chunk 损失 + KL 锚定) | 远程端到端跑通(DeepSeek/MiniMax teacher) |
|
||||
| 6 | 评测与消融 | §5 | 数学评测脚本;三个消融开关 | MATH-500 子集上 student 有可测提升趋势 |
|
||||
@@ -20,9 +39,9 @@
|
||||
|------|------|------|
|
||||
| `00-roadmap.md` | 本文 | ✅ |
|
||||
| `01-paper-code-map.md` | 论文 §3-§4 精读 + 参考实现全景解剖 + 概念↔代码对照表 | ✅ |
|
||||
| `02-sft-baseline.md` | 层 1:SFT 与数据管线 | ⬜ |
|
||||
| `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性 | ⬜ |
|
||||
| `04-mc-estimator.md` | 层 3:MC 估计 + 贝叶斯平滑 | ⬜ |
|
||||
| `02-sft-baseline.md` | 层 1:SFT 与数据管线 | ✅ 待读 |
|
||||
| `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性(含 §6.1 远程实证) | ✅ 已关账 |
|
||||
| `04-mc-estimator.md` | 层 3:语义相似度 φ + MC 估计 + 贝叶斯平滑(式3/4/5) | ✅ 待读 |
|
||||
| `05-entropy-chunking.md` | 层 4:熵调度 | ⬜ |
|
||||
| `06-omniopd-full.md` | 层 5:完整损失与 teacher 客户端 | ⬜ |
|
||||
| `07-eval-ablation.md` | 层 6:评测与消融 | ⬜ |
|
||||
@@ -35,6 +54,9 @@
|
||||
| rollout 数 N | 10 | 10 | 每个 chunk 的 teacher MC 采样数(§4.2 证明 N=10 是甜点) |
|
||||
| chunk 长度 C | 50 | 50 | token 数 |
|
||||
| 先验强度 α | 1.0 | 1.0 | `chunk_alpha`,Dirichlet 平滑 |
|
||||
| 相似度 φ | ROUGE-1 | ROUGE-1 | 备选 edit_distance |
|
||||
| Student | Qwen3-8B 级 | Qwen3-0.6B | 跑通优先 |
|
||||
| Teacher | Qwen3.5-397B / Claude / Gemini | DeepSeek 或 MiniMax(OpenAI 兼容) | logit-free 主路径 |
|
||||
| 相似度 φ | **edit_distance**(§5.1) | edit_distance | ⚠️ 代码默认 rouge1(config L298)与论文默认背离,须显式指定 |
|
||||
| KL 锚权重 β | **0.1**(§5.1) | 0.1 | ⚠️ 代码默认 `mc_kl_weight=0` 与论文背离,须显式指定 |
|
||||
| 训练数据 | DAPO-Math-17K(prompt-only) | 同(层 1 先抽 ~1k 子集控制 API 成本) | 一份数据服务层 1-6;学生升到 1.7B 后可直接对表论文 Table 1 |
|
||||
| Student | Qwen3-1.7B / 4B | Qwen3-0.6B | 跑通优先;升级 1.7B 即可与论文对比 |
|
||||
| Teacher | Qwen3-32B / Claude-4.5-Haiku / Gemini-2.5-Flash | MiniMax-M3(自建 new-api 网关,OpenAI 兼容;2026-07-18 由 DeepSeek 改定) | logit-free 主路径;M3 是 reasoning 模型,思考段入库前剥离(teacher.py strip_think) |
|
||||
| SFT 基线定义 | teacher rollout 上的离线蒸馏(非人写答案) | 同 | 对应参考实现 `_generate_teacher_completions` + JSONL 缓存路径 |
|
||||
|
||||
@@ -24,6 +24,19 @@ flowchart LR
|
||||
|
||||
$$\mathcal{L}_{\text{OmniOPD}}(\theta) = -\mathbb{E}_{\hat y\sim\pi_\theta}\Big[\sum_{c=1}^{M}\hat\pi^{(c)}_{\text{teacher}}\sum_{t\in c}\log\pi_\theta(y_t\mid x,y_{<t})\Big] + \beta\sum_{t\in\mathcal{U}} D_{KL}\big(\pi_{\text{ref}}\,\|\,\pi_\theta\big)$$
|
||||
|
||||
| 符号 | 含义 | 直观说法 |
|
||||
| --------------------------------- | ------------------------------------ | ------------------------------- |
|
||||
| $\pi_\theta$ | 学生模型($\theta$ 是它的参数,训练改的就是 $\theta$) | 正在被训练的 0.6B |
|
||||
| $\hat y \sim \pi_\theta$ | 轨迹是学生自己生成的 | “on-policy”三个字的全部含义 |
|
||||
| $\mathbb{E}[\cdot]$ | 期望 | 实践中 = 对 batch 里采样出的轨迹求平均,没有更多玄机 |
|
||||
| $c$,共 $M$ 个 | 被熵调度器选中的 chunk(各 $C=50$ 个 token) | 被“抽查”的 $M=10$ 段 |
|
||||
| $\hat\pi^{(c)}_{\text{teacher}}$ | 式(5)算出的贝叶斯估计,$\in [0,1]$ | 老师对这段的认可度打分 |
|
||||
| $\log \pi_\theta(y_t \mid \cdot)$ | 学生给自己当时生成的那个 token 的对数概率 | SFT 里最熟悉的那个量 |
|
||||
| $\mathcal{U}$ | 未被抽查的所有 token | 轨迹的绝大部分 |
|
||||
| $\pi_{\text{ref}}$ | 训练开始前学生的冻结副本 | “初始的自己” |
|
||||
| $\beta$ | 缰绳松紧 | 代码里的 `mc_kl_weight` |
|
||||
|
||||
|
||||
关键设计洞察(§4.1,Theorem 4.1):teacher 估计 π̂ 以**有界乘子** [0,1] 的身份乘在学生 score function 上,而不是像反向 KL 那样出现在分母/log 里——这从结构上消灭了标准 OPD 的梯度爆炸;而贝叶斯先验保证 π̂ ≥ α·π̄/(N+α) > 0,消灭了"teacher 全不匹配 ⇒ 梯度归零"的监督塌缩。
|
||||
|
||||
## 2. 参考实现的真实形态:一个 trainer,三代方法
|
||||
@@ -74,7 +87,7 @@ $$\bar\pi_\theta^{(c)} = \Big(\prod_{t\in c}\pi_\theta(y_t\mid\cdot)\Big)^{1/C}
|
||||
|
||||
代码在 `_compute_chunk_ebopd_loss` 内 L2194-2202:先验 `pi_bar = exp(mean(chunk_lps.detach()))`(几何均值,与式 4 严格一致),`pi_hat = (k + chunk_alpha·pi_bar)/(chunk_mc_samples + chunk_alpha)`,随后 clamp 到 [1e-8, 1] 并 detach。
|
||||
|
||||
> **detach 是命门**:先验和 π̂ 都必须切断梯度,否则学生会通过抬高自己的先验来自我强化(reward hacking 式塌缩)。配置里 `mc_nll_weight`(L246)的注释明确警告开启会塌缩,默认 0。这一点论文只隐含在"π̂ 是目标而非变量"里,代码把它变成了硬约束。
|
||||
> **detach 是命门**:先验和 π̂ 都必须切断梯度(L2196、L2201)。若留梯度通路,损失 π̂·|Σlog π_θ| 中 π̂ 也随 θ 可动,最速下降方向变成**压低**学生对自己 token 的概率、把乘子 π̂ 推向 0(p·ln(1/p)→0,指数快过对数)——在 teacher 全否定(k≈0)的 chunk 上损失可一路逃逸到 0,贝叶斯安全底 α·π̄ 被优化器亲手拆除,Theorem 4.1(a) 的梯度有界性也随之失效(多出的 ∇π̂ 项与惊讶度成正比)。相邻的另一个陷阱:`mc_nll_weight`(config L246-252)给非 MC 位置加自身 NLL 正则,帮助文本明确警告 "non-zero values cause self-reinforcement collapse"(无条件复读自己→熵塌缩),默认 0——两者是方向相反的两种自指失败。论文只隐含在"π̂ 是目标而非变量"里,代码把它变成了硬约束。
|
||||
|
||||
对应理论:Theorem 4.1(b) 下界 π̂ ≥ α·π̄/(N+α) > 0;4.1(c) 偏差-方差分解,α 是噪声-偏移旋钮;Theorem 4.2 证明 N=10 是方差收益的甜点。
|
||||
|
||||
@@ -146,3 +159,5 @@ KL 锚(`mc_kl_weight` = 论文的 β,config L284,**默认 0**)实现与
|
||||
4. 实现的 KL 锚与论文式(8)有哪三处差异?
|
||||
5. API teacher 路径为什么强制 char 级编辑距离?分叉点为什么要对齐词边界?
|
||||
6. `no_bayesian` 消融等价于论文里的哪个估计器?§4.1 预言它会怎么失败?
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
# 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.6B;teacher 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_index,prompt 不产生 loss |
|
||||
| **左** padding | trainer:309-335 | 让整个 batch 能用一个标量 `prompt_length` 切分(见 2.4) |
|
||||
| 边界确定 | 用**未截断**的 prompt 重分词长度切出 completion(trainer:286-289) | 模板渲染后 prompt+completion 的拼接分词 ≠ 分开分词,必须用同一渲染再切 |
|
||||
| `enable_thinking` 透传 + 空 `<think>` 诊断 | trainer:252-266 | Qwen3 模板 no-think 时注入空 `<think>\n\n</think>`;此开关变了,prompt/completion 边界跟着变——错一次全错 |
|
||||
| prompt-only 行 | 全 -100(trainer: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-op(train_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 时 ~20G,cross_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 是),Trainer 默认按"新式契约"对待自定义 compute_loss——第一幕:返回裸 mean 会被 ×累积步数(首跑 loss 7.5 ≈ 0.94×8);第二幕:改成 sum/num_items 后又 ÷world_size(0.244 ≈ 0.85÷4),因为新式契约的 ×num_processes 补偿在**基类** compute_loss 尾部(v5 trainer.py:2028),整体重写 compute_loss 会绕过它。终解 = 按 HF 文档(trainer.py:1977)显式 `self.model_accepts_loss_kwargs = False` 退出新式契约,回到"返回 mean、Trainer ÷累积步数"的经典行为(代价:微批等权而非 token 加权,偏差百分之几,与参考实现同行为)。教训:**整体重写框架方法时,必须检查基类同名方法里除了你替换的逻辑还捎带了什么**。
|
||||
3. 关账判据:1k 子集 1 epoch 跑完,loss 曲线正常,checkpoint 能被 `from_pretrained` 加载并生成通顺文本。
|
||||
@@ -0,0 +1,140 @@
|
||||
# 03 · 层 2:White-box OPD 基线(token 级反向 KL)
|
||||
|
||||
> 本章目标:吃透论文式(2) 的 on-policy 白盒蒸馏及其梯度爆炸脆弱性(§4.1),解剖参考实现 `distillation_mode="standard"` 路径,定出层 2 的重构任务。行号缩写:`DT:` = `references/ars-opd/trl/trl/experimental/distillation/distillation_trainer.py`,`TR:` = `references/ars-opd/train_distillation.py`,`CFG:` = `.../distillation/distillation_config.py`。
|
||||
|
||||
## 1. 论文侧:式(2) 精读
|
||||
|
||||
$$\mathcal{L}_{\text{OPD}} = \mathbb{E}_{y \sim \pi_\theta(\cdot|x)}\Big[\sum_t D_{KL}\big(\pi_\theta(\cdot|y_{<t},x)\,\|\,\pi_T(\cdot|y_{<t},x)\big)\Big]$$
|
||||
|
||||
三个成分逐个看:
|
||||
|
||||
| 成分 | 含义 | 为什么 |
|
||||
|------|------|--------|
|
||||
| $y \sim \pi_\theta$ | **on-policy**:轨迹由 student 自己采样 | 治 SFT 的曝光偏差——式(1) 只在 teacher 轨迹上监督,student 推理时一旦走出熟悉区域就没见过纠正信号;on-policy 让 teacher 在"student 实际会犯错的地方"给监督 |
|
||||
| $D_{KL}(\pi_\theta\|\pi_T)$ | **反向 KL**(student 在前)| mode-seeking:student 容量小,与其平摊质量模仿 teacher 全分布(前向 KL 的 mode-covering),不如集中质量学好 teacher 的主模式 |
|
||||
| 逐 token 求和 | token 级分布对齐 | 这就是"白盒":需要 teacher 每个位置的**完整 logits**,因此 teacher 必须本地可跑、且与 student **同 tokenizer**(词表逐位对应才能算 KL) |
|
||||
|
||||
**§4.1 梯度爆炸(本层的理论主课,OmniOPD 的出发点)**:反向 KL 对 student logit 的梯度含 $\log\frac{\pi_\theta(v)}{\pi_T(v)}$ 项。on-policy 下 student 会采到 teacher 认为极差的 token($\pi_T \to 0$),此时 log 比值 $\to \infty$,单个 token 的梯度可以炸掉整个 batch。第一章已给过一句话版本,本层用代码把它钉死:重构任务里有一个"梯度范数随 $\pi_T$ 衰减而暴涨"的单元测试(对应层 3 detach 测试的姊妹篇——一个证明旧方案为什么坏,一个证明新方案为什么稳)。
|
||||
|
||||
顺带记住对比锚点:层 5 的式(8) 用"有界乘子 π̂"替换这里的 log 比值,这正是两代方法的分水岭。
|
||||
|
||||
## 2. 参考实现解剖(standard 路径)
|
||||
|
||||
### 2.1 一个重要的事实先行
|
||||
|
||||
**训练脚本从不传本地 teacher**(TR:420-424 只传 model/args/dataset):仓库实际跑的 standard 蒸馏全走 **teacher server 路径**(环境变量 `TEACHER_URL` 触发,TR:329-330 → DT:2891-2920);**本地 teacher 路径(DT:2921-2944)只有直接构造 `DistillationTrainer(teacher_model=...)` 才会触发**。我们层 2 恰恰要走后者(4 卡 A800 放得下 4B teacher),所以两条路都要看懂,但以本地路为重构蓝本。
|
||||
|
||||
### 2.2 compute_loss 主流程(DT:2841-2946)
|
||||
|
||||
```
|
||||
守卫: no_teacher and lmbda<1 → 显式报错 (DT:2869)
|
||||
student 前向(带梯度) (DT:2883)
|
||||
prompt_length 切齐 + [pl-1:-1] 移位 (DT:2887-2889) ← 与 SFT 完全同款,T4 已实现
|
||||
teacher logits:
|
||||
server 路 → top-k logprobs 传输 (DT:2898)
|
||||
本地路 → teacher.eval() + no_grad 前向 (DT:2923, 2578-2594)
|
||||
divergence: generalized_jsd_loss / 稀疏快路 (DT:2926-2942)
|
||||
```
|
||||
|
||||
on/off-policy 抽签**不在这里**——在 `_prepare_inputs`→`_fill_buffer`(见 2.5)。
|
||||
|
||||
### 2.3 generalized_jsd_loss(DT:2408-2491)——β 的三副面孔
|
||||
|
||||
记号:$\pi_\theta$ = student,$\pi_T$ = teacher;KL 里"在前"的那个分布是被求期望的一方。
|
||||
|
||||
| β | 数学 | 语义 | 代码 |
|
||||
|---|------|------|------|
|
||||
| 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 |
|
||||
|
||||
> 行号说明:上表指向 **全词表** 路径(`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!(本章最大陷阱)
|
||||
|
||||
默认 `loss_top_k=1, loss_add_tail=True`(CFG:353,368)→ 走 **top-1 稀疏快路**(DT:2507-2549):支持集 = {实际采样 token} ∪ {teacher top-1} + **尾桶**(第 K+1 个桶收纳截断外全部概率质量:$\log(1-\sum e^{\text{top}k})$,DT:108-118,防 top-1 时 loss 平凡为 0)。要严格对齐式(2) 的全词表反向 KL,必须 `loss_top_k=0`。top-k>1 时支持集依 β 取 teacher top-k / student top-k / 两者并集去重(DT:2449-2467)。
|
||||
|
||||
server 路径更粗:teacher 只回传 top-k logprobs 三张表(DT:2676-2679),β>0 时 CFG 强制 top-1(CFG:540-544)。这些截断全是**传输/显存工程妥协**(第一章 §3.7 讨论过),不是论文成分。
|
||||
|
||||
### 2.5 buffer 与 on-policy 生成(DT:846-932, 1071-1263)
|
||||
|
||||
- `_RepeatBatchDataLoader` 把同一 collated batch 重复 `gradient_accumulation_steps` 次(DT:346-371),`_fill_buffer` 按**切片级**伯努利抽签 on/off-policy(`random() <= lmbda`,DT:879,主进程抽签后广播)。
|
||||
- **off-policy 切片 = 在数据集自带 completion 的轨迹上做 KL 蒸馏**(DT:893-894 原样保留数据轨迹;损失仍是 KL,不是交叉熵)——GKD 的 λ 插值本义:λ=0 离线蒸馏、λ=1 纯 on-policy。两个易误解处:轨迹不是 teacher 现场采样的;DAPO prompt-only 下这些切片 labels 全 -100,KL 被掩码归零 = **静默空转的算力浪费**(唯一例外:`lmbda=0`+server 触发 teacher 生成,DT:901-906,即层 1 SFT 的数据来源)。
|
||||
- on-policy 切片:vLLM colocate 生成(按 `vllm_sync_frequency` 同步权重,DT:1094-1106)或 `model.generate`(DT:1113-1168);生成结果重建 input_ids/labels 写回 buffer(DT:1170-1263,labels 只在 completion 段有效)。
|
||||
- loss 前向是对已生成序列的 teacher-forcing(DT:2883)——"采样一次、前向算分布",GKD 标准做法。
|
||||
|
||||
### 2.6 本地 teacher 的基础设施(DT:504-586)
|
||||
|
||||
加载后 `accelerator.prepare_model(teacher, evaluation_mode=True)`(DT:584,随 DDP 每卡一份副本);同 tokenizer 校验比较 `get_vocab()`(DT:734-740),不匹配在 compute_loss 处显式报错(DT:2876);前向 `eval() + no_grad`(DT:2581-2586)。
|
||||
|
||||
### 2.7 顺带发现的坑(解剖副产物)
|
||||
|
||||
| 坑 | 位置 | 说明 |
|
||||
|----|------|------|
|
||||
| `lmbda=1 + no_teacher` 穿过守卫后在深处崩 | DT:2872 vs DT:2593 | 报错文案宣称合法,实际必崩——守卫条件写错 |
|
||||
| off-policy ≠ "teacher 采样的离线数据" | DT:893-894 | 轨迹来自数据集固有 completion(损失仍是 KL);prompt-only 数据下整个切片被掩码归零,静默空转 |
|
||||
| `num_generations>1 且 lmbda<1` 会造重复样本 | CFG:593-596 | 官方注释自己承认 |
|
||||
|
||||
## 3. 与式(2) 的偏差清单(默认配置下)
|
||||
|
||||
抉择原则(下表每一行都由它推出,而非逐条拍板):**式(2) 的算法本质忠于论文(on-policy + 反向 KL + 全词表分布对齐);参考实现的默认近似是为它的处境——API 传输 + 大 teacher——妥协出来的,换了我们的处境(本地同 tokenizer 的 4B teacher + 4×A800)就不继承;不改优化方向的表面差异,选对诊断/教学最有利的;一般性凡免费且未来有用则保留、凡昂贵且当前数据上空转则删除。** 一个反直觉推论:正因处境不同,我们回归论文本质反而比参考默认更贴式(2)(支持集那行)。逐行推导见下,"我们层 2"列即结论。
|
||||
|
||||
| 项 | 参考实现默认 | 严格式(2) | 我们层 2 |
|
||||
|----|--------------|-----------|----------|
|
||||
| 支持集 | top-1 稀疏 + 尾桶 | 全词表 | **全词表**(`top_k=0` 等价;同 tokenizer 本地 teacher 使我们能比参考默认更贴论文) |
|
||||
| on-policy 比例 | lmbda=1.0 ✓ | 纯 on-policy | lmbda=1(不实现混合抽签) |
|
||||
| KL 方向 | beta=1.0 ✓ | 反向 | β 作为参数保留(前向/反向/JSD 同一公式,纯逻辑函数顺手覆盖,also 层 5 KL 锚要用前向) |
|
||||
| 温度 | 1.0 ✓ | 1 | 1.0 |
|
||||
| reduction | per-token mean | 论文 token 求和 | per-token mean(与 SFT loss 同尺度,才能对比曲线;差一个常数因子不改优化方向) |
|
||||
|
||||
## 4. 保留 / 替代 / 删除
|
||||
|
||||
| 决策 | 项目 |
|
||||
|------|------|
|
||||
| **保留** | "生成一次 + teacher-forcing 前向"结构;prompt 边界切齐/重掩码(直接复用 T4 的 `compute_prompt_length`);teacher `eval+no_grad`;同 tokenizer 校验(显式报错);β 语义与温度;per-token mean |
|
||||
| **替代** | vLLM 学生生成 → `model.generate`(0.6B 生成不慢,省掉 colocate+权重同步整套复杂度;卡了吞吐再回来接 vLLM,触发条件记录于此);teacher server → 本地 Qwen3-4B HF 前向(全词表精确);buffer/RepeatBatchDataLoader/切片抽签 → 每个 batch 现场生成现场用(lmbda=1 下 buffer 是纯开销) |
|
||||
| **删除** | top-k/尾桶/top-1 稀疏快路(全词表放得下就不近似;显存账见 §5)、`reverse_kl_top_1_mode`、teacher server 客户端、Liger、off-policy 混训、on/off-policy 指标群 |
|
||||
|
||||
## 5. 重构任务(Claude 写码、你精读提问)
|
||||
|
||||
| # | 任务 | 落点 | 备注 |
|
||||
|---|------|------|------|
|
||||
| U1 | `DistillConfig` dataclass(teacher 模型、β、温度、生成参数) | `ars_opd/configs.py` | 追加,不动 SFTConfig |
|
||||
| U2 | 纯逻辑 divergence:全词表 masked token-KL/JSD(β 参数化)+ 单测 | `ars_opd/trainer.py`(纯张量函数,同 `sft_loss` 地位) | 单测含**梯度爆炸演示**:$\pi_T\to0$ 时梯度范数暴涨的断言(§4.1 的可执行版本) |
|
||||
| U3 | collator 放开 prompt-only(返回 prompts/prompt_attention_mask 供生成) | `ars_opd/data.py` | 兑现 T3 预留的口子;SFT 路径行为不变(回归测试盯住) |
|
||||
| U4 | `DistillTrainer`:生成 → teacher no_grad 前向 → divergence | `ars_opd/trainer.py` | on-policy 生成用 `model.generate`;tokenizer 一致性构造时校验 |
|
||||
| U5 | 自包含脚本 | `scripts/train_whitebox.sh` + `.py` | teacher Qwen/Qwen3-4B;数据复用同一 DAPO 子集(prompt-only,无需 teacher 缓存) |
|
||||
|
||||
**显存账(A800-80G,验证 U 系列前算给自己看)**:全词表 logits (B,T,V) 是大头——B=4、T=2048、V≈151k 的 bf16 logits 单张 ≈2.5G,student+teacher 两份 + log_softmax 中间量 ≈15G 级,加 0.6B 训练全套 ~10G 与 4B teacher 推理副本 ~9G,B=4/T=2048 起步安全;生成长度先压 1024。参数进 `DistillConfig` 显式化。
|
||||
|
||||
**默认参数提案(可否决)**:β=1.0、temperature=1.0、lmbda 固定 1(不做参数)、`max_new_tokens=1024`、per_device batch 4、teacher `Qwen/Qwen3-4B`、数据同 seed 同 1k 子集(复用层 1 的抽取逻辑,不需要 teacher 解答缓存)。
|
||||
|
||||
## 6. 验证方式
|
||||
|
||||
1. **本地 CPU 单测**:divergence 手算对拍(V=5 玩具分布,β∈{0,1,0.5} 三点各一);β=1 与 β=0 的方向性断言(teacher 置信/弥散两种分布下 loss 排序);**梯度爆炸测试**:固定 student,teacher 对采样 token 的概率从 1e-1 衰减到 1e-6,断言梯度范数单调暴涨且超阈值——为层 5 的"有界乘子"对照埋桩。
|
||||
2. **collator 回归**:放开 prompt-only 后,全部既有 SFT 测试必须原样通过。
|
||||
3. **远程冒烟**:盯 KL loss 曲线、`grad_norm`、`distill/num_gen_tokens_per_step`。
|
||||
4. 关账判据:白盒蒸馏跑通不 NaN,生成不塌空,接口回看完成。
|
||||
|
||||
### 6.1 远程实证(2026-07-19,两次跑,勘误当初的"预期见毛刺")
|
||||
|
||||
**当初预测错了**:docs 原写"预期能看到 loss 毛刺(梯度爆炸实况)"。实跑**没有毛刺**,两次都平稳。诚实记录 + 解释:
|
||||
|
||||
| 跑 | 配置 | loss | grad_norm |
|
||||
|----|------|------|-----------|
|
||||
| sanity | 裁到 1.0、lr 1e-6、50 步 | 平滑 0.35→0.22 | 14.42 → ~2(单调降) |
|
||||
| noclip | ≈关裁剪、lr 5e-6、15 步 | 平滑 0.35→0.21 | 14.42 → ~2(无尖峰,且更快收敛) |
|
||||
|
||||
三条实证结论:
|
||||
|
||||
- **爆炸是真机制、但本区间高度阻尼**。§4.1 在单测里坐实(单个 π_T→1e-6 的 token 梯度暴涨),但真实训练里:① **同门 teacher**(Qwen3 0.6B↔4B)使 student 很少采到 teacher 真恨的 token;② **per-token mean 把每步 ~4000 token 的梯度尖峰摊平**(单测看单 token 机制,真实看上千 token 平均后果)。故 batch 级 grad_norm 峰值只 ~14(比健康 ~2 高 5-7 倍,但远非几十上百),且只降不升。
|
||||
- **关键坑:`grad_norm` 日志是裁剪前值**。sanity 的 14→2 那串本身就是爆炸证据,只是 HF 默认 `max_grad_norm=1.0` 把**步长**裁掉了、loss 才平滑——这个静默稳定器现已提进 `DistillConfig`(见其注释)。noclip 关掉它,grad_norm 曲线几乎不变(第 1 步两跑完全相同=14.42,验证确定性),但大步长反而**加速收敛**、仍不炸。
|
||||
- **重构层 5 动机的认知**:白盒的"同 tokenizer"硬约束把你锁在相对温和的区间(换跨家族 teacher 会先被词表校验拦下),所以层 5 有界乘子 π̂ 的真正杀手锏不在"防这个温和爆炸",而在 **logit-free**(teacher 只给文本、拿不到 logits,白盒根本跑不了)。爆炸的干净见证留在 U2 单测。
|
||||
@@ -0,0 +1,154 @@
|
||||
# 04 · 层 3:语义相似度 φ + MC 估计器(式 3/4/5)
|
||||
|
||||
> 本章目标:吃透 OmniOPD 如何把"要 teacher logits"(层 2 白盒的硬约束)换成"比 teacher 文本"(logit-free),并用 Dirichlet 贝叶斯平滑把稀疏的相似度信号变成稳定、非零的监督乘子 π̂。然后建 `ars_opd/similarity.py`(式3)与 `ars_opd/estimator.py`(式4/5)两个**纯逻辑**模块,本地 CPU 对拍参考实现。
|
||||
> 行号缩写:`DT:` = `references/ars-opd/trl/trl/experimental/distillation/distillation_trainer.py`,`CFG:` = `.../distillation_config.py`,`VAL:` = `references/ars-opd/validate_chunk_mc_estimator.py`。
|
||||
|
||||
## 1. 论文侧:从"要 logits"到"比文本"(§3.2.1-3.2.2)
|
||||
|
||||
### 1.1 大局:这一层是 logit-free 的支点
|
||||
|
||||
层 2 白盒 OPD(式2)要 teacher 每个位置的完整分布,还要同 tokenizer——我们实测这把方法锁死在同门 teacher。§3.2.1 用一句话拆掉它:**别在 token 概率上匹配,改在文本语义上匹配**。teacher 只需生成文本 rollout;学生某段 chunk 对不对,由"学生这段文本" vs "teacher 那几段 rollout 文本"的语义相似度判定。
|
||||
|
||||
这一步同时解三个问题(§3.2.1 首段):① teacher 逐 token 查询 O(T) 不可行 → chunk 化降到 O(T/C);② tokenizer 不一致导致学生的精确 token 在 teacher rollout 里根本不出现、产生稀疏零梯度 → 改比语义;③ 由此得到跨架构可用的信号。结构上灵感来自 Speculative Decoding 的验证阶段——把一个 C-token chunk 当作单个"验证单元"。
|
||||
|
||||
### 1.2 式(3):语义相似度聚合 k_sem
|
||||
|
||||
学生生成 on-policy 轨迹 y,从中选 M 个长度 C 的 chunk(选法是层 4 的熵调度)。对某个 chunk c,把它之前的前缀 y_<c 喂给 teacher,teacher 生成 N 个 rollout,逐个与学生 chunk 比相似度、求和:
|
||||
|
||||
$$k_{\text{sem}}^{(c)} = \sum_{i=1}^{N} \phi\big(y_c,\, y_{\text{teacher}}^{(i)}\big),\qquad \phi:(y_c, y_{\text{teacher}})\mapsto[0,1]$$
|
||||
|
||||
| 要点 | 说明 |
|
||||
|------|------|
|
||||
| φ 是**连续**相似度 | ROUGE-1 单词重叠 / 编辑距离,∈[0,1],非 0/1 硬票 |
|
||||
| k_sem ∈ **[0, N]** 是实数 | N 个 [0,1] 相似度之和;`test_estimator_detach.py` 口语叫"票数",实质是连续和 |
|
||||
| 比的是**文本语义**不是 token | 词级重叠对 tokenizer/风格不变——teacher 换词、换词表,只要意思对 φ 就高(§3.2.1 末:不惩罚风格偏差与词表不匹配) |
|
||||
|
||||
### 1.3 式(4):学生先验 π̄(贝叶斯先验)
|
||||
|
||||
$$\bar\pi_\theta^{(c)} = \Big(\prod_{t\in c}\pi_\theta(y_t\mid x,y_{<t})\Big)^{1/C} = \exp\!\Big(\tfrac1C\sum_{t\in c}\log\pi_\theta(y_t\mid\cdot)\Big)$$
|
||||
|
||||
学生对这段 chunk 每 token 概率的**几何均值** = exp(平均 log 概率)。它是学生**自己**"平均每 token 有多自信",∈(0,1],拿来当贝叶斯先验。几何均值(非算术)才是自然的 chunk 级概率:整段联合概率开 C 次方。
|
||||
|
||||
### 1.4 式(5):Dirichlet-Multinomial 贝叶斯平滑 π̂(本层灵魂)
|
||||
|
||||
$$\hat\pi_{\text{teacher}}^{(c)} = \frac{k_{\text{sem}}^{(c)} + \alpha\,\bar\pi_\theta^{(c)}}{N + \alpha}
|
||||
\;\overset{式(10)}{=}\; \underbrace{\tfrac{N}{N+\alpha}}_{}\underbrace{\tfrac{k_{\text{sem}}}{N}}_{\hat\pi_{\text{freq}}\;\text{频率估计}} + \underbrace{\tfrac{\alpha}{N+\alpha}}_{}\underbrace{\bar\pi_\theta}_{\text{学生先验}}$$
|
||||
|
||||
π̂ 是"频率估计 k_sem/N"与"学生先验 π̄"的**凸组合**,权重 N 对 α。α = 先验强度(`chunk_alpha`,默认 1.0)。
|
||||
|
||||
### 1.5 为什么这样设计:定理 4.1 的三性质 + detach 命门
|
||||
|
||||
π̂ 最终进式(8):$\mathcal{L}_{\text{chunk}} = -\hat\pi^{(c)}\sum_{t\in c}\log\pi_\theta(y_t\mid\cdot)$(层 5 的事,此处只需知道 π̂ 是**乘子**)。定理 4.1 证明这套设计同时关掉两个失败模式:
|
||||
|
||||
| 性质 | 内容 | 对照层 2 |
|
||||
|------|------|---------|
|
||||
| (a) 不爆炸 | π̂∈[0,1] 是**乘子**,不在分母、不在 log 里;每 chunk 梯度被学生 score function 卡住有界(式11) | 层 2 反向 KL 的 log(π_θ/π_T) 在 π_T→0 时**无界爆炸** |
|
||||
| (b) 不塌缩 | 先验保证 π̂ ≥ α·π̄/(N+α) **> 0 恒成立**(式12),即便 k_sem=0(teacher 全否定) | 裸频率估计 k_sem/N 在 k_sem=0 时**塌成 0**、梯度死在最该纠正处 |
|
||||
| (c) 方差收缩 | 贝叶斯 MSE 有闭式(式13),方差比频率估计严格缩小 (N/(N+α))²<1 | —— |
|
||||
|
||||
**detach 命门**(`tests/test_estimator_detach.py` 守的,docs/01 §3.4 推导):π̄ 是学生**自己**的概率。若 π̄ 在式(8) 里不 detach,梯度会经它回传——优化器发现**压低学生对自己 token 的概率**能把乘子 π̂ 推向 0、从而逃避 -π̂·log π_θ 的惩罚(p·ln(1/p)→0,指数快过对数),而且**恰好在 k_sem=0 的 chunk 上塌缩**(那里 π̂=α·π̄/(N+α) 纯由 π̄ 驱动)——最需要学的地方最先崩。detach 把 π̄ 变成常量乘子,损失退化为"按 π̂ 权重强化学生 token",梯度方向恒为增大 p。这正是 (a)(b) 得以成立的机制根源,也是层 2 §4.1 的姊妹篇:一个证明旧方案为什么炸,一个建新方案为什么稳。
|
||||
|
||||
## 2. 参考实现解剖(带行号)
|
||||
|
||||
### 2.1 φ 的两个实现——与三处不一致(重构要抹平)
|
||||
|
||||
| 度量 | 位置 | 输入表示 | 算法 |
|
||||
|------|------|---------|------|
|
||||
| rouge1 | DT:1670 `_compute_rouge1` | **词集合**(`.split()` 去重)于 decode 后**文本** | set 重叠的 F1 = 2PR/(P+R) |
|
||||
| edit | DT:1682 `_compute_edit_similarity` | **token id 列表**(顺序敏感) | 1 − 归一化 Levenshtein / max(m,n) |
|
||||
|
||||
⚠️ 三处不一致(都是重构要统一的):
|
||||
1. **rouge1 比文本、edit 比 token id**(DT:1850 vs 1846)。edit 用 token id **重新耦合了 tokenizer**——直接违背 §3.2.1 "跨 tokenizer" 的立身之本;只在 teacher/student 同 tokenizer 的 vLLM 路径侥幸能跑。
|
||||
2. **rouge1 两种算法**:trainer 用**集合**(DT:1672),validate 脚本用**多重集计数** `min(ref_cnt, hyp_cnt)`(VAL:63)——同名不同义。
|
||||
3. edit 的归一化用 max(m,n),标准 ROUGE-1 其实是多重集——参考实现里 rouge/edit/bleu 各行其是(VAL:55-153 有 7 种度量的大杂烩)。
|
||||
|
||||
### 2.2 k_sem 聚合(DT:1841-1852)
|
||||
|
||||
```
|
||||
for teacher_chunk in teacher_chunks (N 个):
|
||||
sim = edit(student_ids, teacher_ids) 或 rouge1(student_text, teacher_text)
|
||||
k_score += sim # 连续求和 = 式(3) 的 k_sem
|
||||
```
|
||||
|
||||
### 2.3 chunk 级 π̄ / π̂ / detach(DT:2195-2205)——正主
|
||||
|
||||
```python
|
||||
# 式(4) 学生先验:几何均值,detach(命门①,DT:2196)
|
||||
log_pi_bar = chunk_lps.detach().sum() / chunk_len
|
||||
pi_bar = log_pi_bar.exp().clamp(min=1e-8, max=1.0)
|
||||
# 式(5) 贝叶斯目标:再 detach(命门②双保险,DT:2201)
|
||||
pi_hat = (k + self.chunk_alpha * pi_bar) / (self.chunk_mc_samples + self.chunk_alpha)
|
||||
pi_hat = pi_hat.clamp(min=1e-8, max=1.0).detach()
|
||||
chunk_loss = -pi_hat * chunk_lps.mean() # 式(8) chunk 项
|
||||
```
|
||||
|
||||
两处 detach(2196 的 `chunk_lps.detach()` 与 2201 的 `pi_hat.detach()`)**任一都足以**切断逃逸路(k 是 python float,π̄ detach 后 π̂ 已无梯度);参考实现两处都留是防御。clamp 的 1e-8 下限防 log0/精确零;上限 1.0 其实自然满足(k≤N、π̄≤1 ⇒ π̂≤1),是防御。**差异标注**:式(8) 论文是 Σlog π_θ,参考用 `.mean()`(除以 chunk 长)——per-token mean,此处是层 5 的事,先记下。
|
||||
|
||||
### 2.4 token 级基线(DT:1465)——论文说"不可行"的朴素版,作对照
|
||||
|
||||
```python
|
||||
pi_hat = ((k_counts.float() + alpha * student_probs_at_token.detach()) / (N + alpha)).detach()
|
||||
mc_loss = -pi_hat * student_log_probs_at_token
|
||||
```
|
||||
|
||||
同一个式(5),但落在**单 token** 上(C=1):teacher 每步查询、数学生精确 token 的经验频率。这就是 §3.2.1 开头说的 O(T) 不可行、且 tokenizer 不一致下 k 恒 0 的朴素方案。`test_estimator_detach.py` 现在的 C=1 简化正对应这个基线。我们重构做 **chunk 级**(2.3)。
|
||||
|
||||
### 2.5 validate 脚本的对拍策略(VAL,本层验证方式的蓝本)
|
||||
|
||||
`validate_chunk_mc_estimator.py` 回答"廉价的文本相似度 π̂ 能否逼近昂贵的真值":
|
||||
|
||||
- **ground truth**(VAL:13,348):teacher 在**学生 chunk 的 token 上**的几何均值概率 π̄_teacher = exp(mean(log P_teacher))——这需要 teacher logprobs(白盒),是 π̂ 想廉价逼近的对象。
|
||||
- **估计**(VAL:417-419):k_continuous = mean(sim)·N;bayes = (k + α·prior)/(N+α)。
|
||||
- **判据**(VAL:433-437):对 7 种度量各算 MSE_freq vs MSE_bayes、Spearman 相关;验证**贝叶斯平滑降 MSE**(定理 4.1c)、哪种 φ 最相关。
|
||||
|
||||
我们重构的纯逻辑单测照此精神,但用 **toy 数据**(不连真 teacher):手构 k_sem/π̄ 断言 π̂ 公式与性质,再用 toy 模拟验证"MSE_bayes < MSE_freq"与"k=0 时 π̂>0"。
|
||||
|
||||
### 2.6 配置默认(CFG)
|
||||
|
||||
| 参数 | 符号 | 默认 | 论文 |
|
||||
|------|------|------|------|
|
||||
| `chunk_mc_samples` | N | 10(CFG:277) | 10(§4.2 甜点) |
|
||||
| `chunk_alpha` | α | 1.0(CFG:281) | 1.0 |
|
||||
| `chunk_length` | C | 50(CFG:273 附近) | 50 |
|
||||
| `chunk_similarity` | φ | **rouge1**(CFG:299) | **edit_distance**(§5.1)——⚠️ 背离,须显式指定 |
|
||||
| `no_bayesian` | — | False(CFG:315) | 消融:直接用 k/N(频率),验证塌缩 |
|
||||
|
||||
## 3. 保留 / 替代 / 删除
|
||||
|
||||
| 决策 | 项目 |
|
||||
|------|------|
|
||||
| **保留** | 式(4) 几何均值先验 + detach(命门);式(5) 贝叶斯凸组合;连续 φ 求和成 k_sem;clamp 下限防零;no_bayesian 消融(留作层 6 消融开关) |
|
||||
| **替代** | φ 统一到**词级文本**(`.split()`)——edit 也比 words,不再比 token id(抹平 2.1 坑①,回归 tokenizer 无关);rouge1 集合/多重集二选一并注明;默认 φ 显式设 edit_distance(对齐论文 §5.1,不用 code 的 rouge1 默认) |
|
||||
| **删除** | token 级基线路径(DT:1465,朴素不可行版);validate 里 bleu/jaccard/exact_match 等多余度量(只留 rouge1 + edit);vLLM/API 采样(那是层 5 teacher.py 的事,层 3 纯逻辑只吃已算好的 k_sem 与 log 概率) |
|
||||
|
||||
## 4. 重构任务(Claude 写码、你精读提问)
|
||||
|
||||
| # | 任务 | 落点 | 备注 |
|
||||
|---|------|------|------|
|
||||
| E1 | `rouge1(hyp, ref)` + `edit_similarity(hyp, ref)` + `phi(hyp, ref, metric)` + `aggregate_similarity(student_chunk, teacher_rollouts, metric)` | `ars_opd/similarity.py`(纯逻辑,只依赖标准库) | 两度量都吃 str、内部 `.split()` 比 words;φ 默认 edit_distance;k_sem = Σφ(式3) |
|
||||
| E2 | `chunk_prior(log_probs)`(式4 几何均值,**detach**)+ `bayesian_target(k_sem, pi_bar, n, alpha)`(式5) | `ars_opd/estimator.py`(纯逻辑,只依赖 torch) | detach 在 `chunk_prior` 内;`bayesian_target` 再 detach 防御;clamp 下限 |
|
||||
| E3 | 把 `test_estimator_detach.py` 接到**真实现** | `tests/test_estimator_detach.py` | 现为独立 C=1 数学测试;追加对 `chunk_prior`/`bayesian_target` 的同名断言,锁死 detach 不被误删 |
|
||||
|
||||
**接口预想**(E1/E2 公共函数,全类型注解):
|
||||
|
||||
```
|
||||
# similarity.py(纯 str/list,可脱离 torch 测)
|
||||
rouge1(hypothesis: str, reference: str) -> float # 词集合 F1 ∈[0,1]
|
||||
edit_similarity(hypothesis: str, reference: str) -> float # 1 − 归一化 Levenshtein ∈[0,1]
|
||||
phi(hypothesis: str, reference: str, metric: str = "edit_distance") -> float
|
||||
aggregate_similarity(student_chunk: str, teacher_rollouts: list[str], metric: str) -> float # 式(3) k_sem
|
||||
|
||||
# estimator.py(torch,toy 张量可测)
|
||||
chunk_prior(log_probs: torch.Tensor) -> torch.Tensor # 式(4) π̄,detach,(C,)->标量
|
||||
bayesian_target(k_sem: float, pi_bar: torch.Tensor, n_rollouts: int, alpha: float) -> torch.Tensor # 式(5) π̂
|
||||
```
|
||||
|
||||
## 5. 验证方式
|
||||
|
||||
1. **similarity.py 单测**:手构字符串断言 φ 值(如全同 chunk→edit=1、rouge1=1;不相交→0;部分重叠手算);k_sem = Σφ 的连续性;空串边界。
|
||||
2. **estimator.py 单测(对拍参考精神)**:
|
||||
- **公式对拍**:手构 log_probs/k_sem,断言 π̄=exp(mean log p)、π̂=(k+απ̄)/(N+α) 与手算一致;凸组合式(10) 恒等。
|
||||
- **性质对拍**:π̂ ∈(0,1] 恒;**k=0 时 π̂ = α·π̄/(N+α) > 0**(定理 4.1b 反塌缩,= detach 测试的正面);
|
||||
- **方差收缩(toy 模拟)**:固定真值 μ、采样 N 个 [0,1] 相似度多次,断言 π̂ 的 MSE < 频率估计 k/N 的 MSE(定理 4.1c)。
|
||||
- **detach 命门**:对真 `chunk_prior`/`bayesian_target` 复刻 `test_estimator_detach.py` 的两个世界断言(detach→k=0 仍增大 p;不 detach→k=0 且 p<1/e 时逃逸)。
|
||||
3. 关账判据:similarity/estimator 全单测本地 CPU 通过;`test_estimator_detach.py` 已接真实现;接口回看完成(两模块判"深")。
|
||||
@@ -0,0 +1,78 @@
|
||||
# 附录 · CLAUDE.md 取舍决策记录
|
||||
|
||||
> 本文回答"为什么本仓库的 CLAUDE.md 这么写"。基准对照物是 Video-Tree-TRM5 项目的 CLAUDE.md(一个积累了五个月的生产级科研工程项目)。当某条规则的存在理由被质疑时,来这里查;当触发点到达时,按 C 节接入。
|
||||
|
||||
## 0. 取舍标准:三问
|
||||
|
||||
CLAUDE.md 是**每轮对话都完整注入模型上下文的提示词,不是文档**。信噪比是第一指标:模型对长指令集中"当前不适用规则"的遵从度会明显下降,而学会忽略 CLAUDE.md 是最糟的结果。每条候选规则过三问:
|
||||
|
||||
1. 从第 0 天起**每轮都生效**吗?
|
||||
2. 是**本项目特有**的信息吗?(通用好实践不用写,那是模型本来就该做的)
|
||||
3. 它指向的**设施真实存在**吗?
|
||||
|
||||
三问全过 → 保留(A 节);不过且无未来场景 → 舍弃(B 节);不过但有明确未来场景 → 延迟接入并写死触发点(C 节)。
|
||||
|
||||
## A. 改造保留
|
||||
|
||||
| 参考章节 | 我们的对应 | 改造点 |
|
||||
|----------|-----------|--------|
|
||||
| URGENT 头(生产级 + 中文) | URGENT 头 | "生产级"改为"学习驱动"——参考项目是 7×24 运行的 Agent 系统,我们是训练实验代码,健壮性需求是局部的(teacher 客户端),不是全局定性 |
|
||||
| §1 项目元数据 | §1 | 换成论文/参考实现/远程机/gitea |
|
||||
| §2.1 Conda + §2.3 ruff | §4 常用命令 | 只留 pytest/ruff |
|
||||
| §2.4 GPU 约定 | §5 远程规则 | 加严:显式选卡之外,加磁盘 12G 红线和"远程不改代码" |
|
||||
| §3 中"前序版本对照" | §6.2 | 参考 SOP 里最值钱的一条(重构前列出旧版全部行为,逐一确认保留/替代/删除),完整移植 |
|
||||
| §4.1-P4 显式优于隐式 | §2 配置显式化 + §7 类型注解 | 具体化:禁硬编码路径,点名参考实现的 `/fsx` 反面教材 |
|
||||
| §4.1-P5 防御性 | §7 硬规则 | 只留两条:禁 `except: pass`、禁默认值兜底 |
|
||||
| §4.1-P6 可测试性 | §2 纯逻辑核心/IO 边缘 | 升级为结构性约束:不是"优先纯函数"的劝导,而是"三个纯逻辑模块禁止 import transformers/vllm/openai"的可执行守则 |
|
||||
| §4.2 中文 docstring | §7 | 保留,另加张量 shape 标注(ML 项目特有痛点) |
|
||||
| §4.7 核心算法保真清单 | §3 模块↔论文映射表 + `01-paper-code-map.md` 差异清单 | 职能相同:防迁移走样的单一事实源;参考的 12 项是五个月长出来的,我们的 5 行随层数增长 |
|
||||
| §7 输出规范 | §6.4 文档规范 | 留核心三条:表格/伪代码优先、代码块 ≤15 行、引用带行号 |
|
||||
| §9 Research Wiki | `docs/` 章节体系 | 同构替代:知识正本在 docs/,CLAUDE.md 只做指针 |
|
||||
|
||||
## B. 舍弃
|
||||
|
||||
| 参考章节 | 舍弃理由 |
|
||||
|----------|----------|
|
||||
| §1.5 PyTorch 类比表 | 参考项目的领域知识;我们的"类比表"就是模块↔论文映射表 |
|
||||
| §3 SOP 全流程 + §8 Skill 门控表 | 引用的 13 个 skill 在本项目 `.claude/` 不存在,写上即死链(三问之③)。那套门控防多人长周期工程走样;我们的防走样机制是对拍参考实现 |
|
||||
| §4.1-P2/P3 可读性、单一职责 | 通用好实践(三问之②),写进提示词边际价值≈0,反而稀释项目特有条目 |
|
||||
| §4.3 feature branch 强制 | 单人学习仓库,主线提交历史 = 学习履历,特意线性;出现并行实验需求再引入 |
|
||||
| §4.6 覆盖率 80% + 三层测试目录 | 训练器/IO 代码需 GPU,全局覆盖率指标会逼出凑数测试;我们的标准更窄更强:纯逻辑三模块必须有对拍测试 |
|
||||
| §4.8 遥测 + §4.9 LLM 治理栈 | 设施不存在;真实需要的部分(teacher API 重试/并发/缓存)在层 5 作为**代码**进 `teacher.py` 而非作为规则;训练可观测性由 W&B 承担 |
|
||||
| §5 硬性目录规则 | "scripts 只放 .sh、根目录无 .py"与 ML 包惯例冲突:我们 `scripts/` 就是放薄 .py 入口 |
|
||||
| §6 迷途指南表 | 仓库目前 4 份文档,README 即地图 |
|
||||
|
||||
## C. 延迟接入(触发点已写死)
|
||||
|
||||
| 参考章节 | 接入触发点 |
|
||||
|----------|-----------|
|
||||
| §2.2 Makefile 收口 | 常用命令超过 3 条时 |
|
||||
| §2.5 自包含实验 sh(写死全参数、零参数复现) | 层 1 第一次远程训练时采纳 |
|
||||
| §4.2.1 非功能性需求覆盖表(持久化/幂等/断点续跑) | 层 5 设计 teacher 缓存与 checkpoint 恢复时 |
|
||||
| §4.5 配置双模式(.env vs 实验 YAML) | 层 6 第一个扫参对比实验时 |
|
||||
| 日志规范(loguru) | 层 1 训练脚本产生第一份需被检查的运行日志时 |
|
||||
|
||||
## D. 新增(参考没有、本项目特有)
|
||||
|
||||
学习优先(每章先讲解、不替用户一次写完,URGENT 级);远程磁盘红线与 `/data/zym` 路径纪律;纯逻辑模块禁 import 清单;教学注释三类型(见 E 节);"每完成一层回填 CLAUDE.md"的增长机制本身。
|
||||
|
||||
## E. 教学注释规范的决策(2026-07-17 补充)
|
||||
|
||||
**问题**:教学项目要不要更重的注释?类型注解是否强制?
|
||||
|
||||
**决策**:分工制——`docs/` 章节讲原理,代码注释做索引,两者不重复。注释只写三类:论文锚点、非显然约束(why + 违反后果)、差异标注。**拒绝逐行解说**:讲解性注释会让代码淹没在散文里、与章节文档重复、且随重构过期。这三类恰好都是"代码自身表达不了的信息"——即 Ousterhout 对注释存在意义的定义。
|
||||
|
||||
**类型注解**:公共接口强制(接口注解本身就是教学信息,成本极低)、私有 helper 从宽(强制到局部就是形式主义)。真正的硬要求是 **shape 标注**:`torch.Tensor` 注解表达不了 shape,而 shape 是 ML 代码可读性的最大杠杆。不引入 jaxtyping 之类的 shape 类型库——多一个依赖、多一层语法噪声,行注释 `# (B,T,V) -> (B,T)` 已够。
|
||||
|
||||
## F. Ousterhout 原则 → 本仓库规则的对照
|
||||
|
||||
**决策**:原则本身不进 CLAUDE.md(书摘是通用内容,三问之②不过),翻译成的可执行规则进。对照关系:
|
||||
|
||||
| 书中原则 | 本仓库的落地 |
|
||||
|----------|--------------|
|
||||
| 深模块(接口简单、实现有料) | 模块按论文概念划分;判据"看公式知文件、开文件知章节"(CLAUDE §2、§3) |
|
||||
| 信息隐藏 | 纯逻辑核心/IO 边缘 + 禁 import 清单(CLAUDE §2) |
|
||||
| 注释写代码表达不了的东西 | 教学注释三类型(CLAUDE §7,本文 E 节) |
|
||||
| 战略式编程(投资设计,不只让代码能跑) | 每层完成后的接口回看:接口复杂度逼近实现复杂度 = 浅模块坏味道,重构后才进下一层(CLAUDE §6.5) |
|
||||
| 适度通用(somewhat general-purpose) | YAGNI + C 节的延迟接入机制:规则和抽象都等真实场景出现才引入 |
|
||||
| Define errors out of existence | 不作为强制规则(与"禁默认值兜底"存在张力),作为设计品味在各章讨论——OmniOPD 本身就是范例:π̂ 的 clamp+先验下界在数学上消灭了零梯度错误态,而不是运行时捕获它 |
|
||||
@@ -0,0 +1,32 @@
|
||||
# 研究方向候选记录
|
||||
|
||||
> 重构过程中冒出的研究想法登记处。每条含:动机、可证伪假设、依托本仓库的最小实验(MVP)、风险。想法不分优先级排序时按登记时间排列。
|
||||
|
||||
## RI-1 · OmniOPD × 零阶黎曼优化(ZO-RGD):全前向"双黑盒"蒸馏
|
||||
|
||||
- **来源**: 合作者论文《Zeroth-Order Riemannian Optimization on Fixed-Rank Update Manifolds for LLM Fine-Tuning》(`references/26_05_subNeruIPS_...pdf`),2026-07 登记。
|
||||
- **背景一句话**: 该文将 LoRA 式增量 ΔW 约束为固定秩流形上的点,用两次前向的有限差分(MeZO 式 ZO)+ 切空间归一化探针 + 截断 SVD retraction 做无反传微调;实验限于 OPT 分类/抽取任务。
|
||||
|
||||
### 1a. 系统组合:teacher 无 logits + 学生无反传
|
||||
|
||||
| 要素 | 说明 |
|
||||
|------|------|
|
||||
| 咬合点 | OmniOPD 的 π̂ 是 **detach 的常数** → ZO 的两次扰动前向 F(ΔW±εZ) 复用同一轨迹与同一组 teacher 打分,不产生额外 API 查询 |
|
||||
| 协同 | ZO 步长极小 → 轨迹+打分可复用 K 个 ZO 步仍近似 on-policy → teacher API 成本摊薄 K 倍(慢优化器 × 贵监督 = 天然互补) |
|
||||
| 系统故事 | 全管线跑在纯推理设施上(teacher=聊天 API,学生=vLLM 前向评分),无训练框架、无激活显存;4×A800 可碰 30B+ 学生 |
|
||||
| **Gate 实验** | ZO-RGD 在本管线上先优化普通 SFT 损失(层 1 复用):长 CoT 生成任务上能否收敛到可用水平。**不过此关全案作废** |
|
||||
| 风险 | ZO 在生成式/推理任务无先例(MeZO 系全是短输出分类);OmniOPD 每轨迹 10 chunk × ZO 每步 1 标量 = 双重稀疏,可能不收敛 |
|
||||
|
||||
### 1b. 秩约束 = 几何 trust region,替代/减弱 β KL 锚
|
||||
|
||||
- **假设**: 式(8) 第二项(行为空间信任域)与固定秩流形约束(参数空间信任域)防的是同一失败模式(未审计区域漂移);流形约束下 β 可调小甚至归零。
|
||||
- **MVP**: 层 6 后消融网格 {全参 / LoRA / 固定秩流形} × {β=0 / β>0},测未审计 token 对 π_ref 的 KL 漂移 + 数学评测分。若"流形+β=0"≈"全参+β>0",得到干净结论。
|
||||
- **依托**: 损失/调度器即本仓库层 3-5 产出,仅换优化端;此实验**不依赖 1a 的 ZO**(一阶梯度 + retraction 即可做),风险远低于 1a,可独立先行。
|
||||
|
||||
### 1c. 方差预算分配(理论附件)
|
||||
|
||||
- OmniOPD Thm 4.2(teacher 采样方差 ∝1/N)× ZO-RGD Prop 1(探针方差 ∝ mn/d_r)串联;固定预算下 N(rollout 数)与 q(探针数)的最优分配。适合作 1a 的理论章节,不独立成文。
|
||||
|
||||
### 诚实评估
|
||||
|
||||
结合发生在**管线/系统层**而非损失数学层——两文公式不冲突也不深融,叙事须立足于:① 双黑盒系统故事(1a);② 信任域替换假设(1b)。建议路径:先做 1b(便宜、独立、可证伪),1a 的 gate 实验穿插进行。
|
||||
@@ -0,0 +1,15 @@
|
||||
# 最小打包配置:只为让 `pip install -e .` 把 ars_opd 注册进环境,
|
||||
# 使脚本/调试器/远程从任意目录都能 import ars_opd(不再依赖 cwd 恰好是仓库根)。
|
||||
# 依赖不在这里声明——统一走 requirements*.txt,避免两处清单漂移。
|
||||
[build-system]
|
||||
requires = ["setuptools>=64"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "ars-opd"
|
||||
version = "0.1.0"
|
||||
description = "OmniOPD (arXiv:2606.01476v2) 分层重构实现"
|
||||
requires-python = ">=3.11"
|
||||
|
||||
[tool.setuptools]
|
||||
packages = ["ars_opd"]
|
||||
@@ -0,0 +1,5 @@
|
||||
# 远程独有依赖(gpu-a800-060):训练与推理重件,本地不装。
|
||||
-r requirements.txt
|
||||
vllm>=0.8 # white-box teacher 服务 + 学生 on-policy 生成
|
||||
accelerate # 多卡训练启动
|
||||
wandb # 训练指标上报
|
||||
@@ -0,0 +1,9 @@
|
||||
# 两端共用核心依赖(本地 + 远程)。版本策略:先宽松安装,两端跑通后按需冻结。
|
||||
torch>=2.6
|
||||
transformers>=4.51 # Qwen3 系列需要 4.51+
|
||||
datasets
|
||||
numpy
|
||||
openai>=1.60 # teacher.py:OpenAI 兼容 API 客户端
|
||||
python-dotenv
|
||||
pytest
|
||||
ruff
|
||||
@@ -0,0 +1,125 @@
|
||||
"""诊断脚本:逐环检验 collator 对齐链(层 1 关账前的疑点排查)。
|
||||
|
||||
背景:远程 sanity 中初始 loss ~7.5(预期 ~2-3),且首样本自检的 completion 段
|
||||
出现 `$k \\50118$`(teacher 原文是 `$k \\leq 2018$`)。本脚本把
|
||||
"缓存文本 → 模板渲染 → 分词 → 边界切片 → 解码"逐环单测,定位腐坏点。
|
||||
|
||||
远程运行(CPU 即可,tokenizer 用已有 HF 缓存):
|
||||
python -u scripts/diag_collator.py
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import sys
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from ars_opd.configs import SFTConfig
|
||||
from ars_opd.data import IGNORE_INDEX, SFTCollator, load_sft_dataset
|
||||
|
||||
CACHE = "data/teacher_completions_dapo1k_minimax-m3.jsonl"
|
||||
DATASET = "data/dapo-math-17k-unique.parquet"
|
||||
MODEL = "Qwen/Qwen3-0.6B"
|
||||
|
||||
|
||||
def check(name: str, ok: bool, detail: str = "") -> bool:
|
||||
print(f"[{'通过' if ok else '失败'}] {name}" + (f" —— {detail}" if detail else ""), flush=True)
|
||||
return ok
|
||||
|
||||
|
||||
def first_diff(a: str, b: str) -> int:
|
||||
n = min(len(a), len(b))
|
||||
for i in range(n):
|
||||
if a[i] != b[i]:
|
||||
return i
|
||||
return -1 if len(a) == len(b) else n
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# 环 0:缓存文件指纹(与本地对比,排除 scp 传坏/版本不一致)
|
||||
digest = hashlib.sha256(open(CACHE, "rb").read()).hexdigest()
|
||||
print(f"缓存文件 sha256: {digest[:16]}… (与本地对比)", flush=True)
|
||||
|
||||
cfg = SFTConfig(
|
||||
dataset_path=DATASET,
|
||||
output_dir="/tmp/diag",
|
||||
subset_size=5,
|
||||
seed=42,
|
||||
teacher_completions_path=CACHE,
|
||||
)
|
||||
ds = load_sft_dataset(
|
||||
cfg.dataset_path,
|
||||
cfg.dataset_split,
|
||||
cfg.subset_size,
|
||||
cfg.seed,
|
||||
cfg.teacher_completions_path,
|
||||
)
|
||||
msgs = ds[0]["messages"]
|
||||
completion_text = msgs[-1]["content"]
|
||||
|
||||
# 环 1:本机数据管线出来的 teacher 文本是否干净
|
||||
check(
|
||||
"环1 缓存→数据集文本干净",
|
||||
"\\leq 2018" in completion_text and "\\50118" not in completion_text,
|
||||
f"开头: {completion_text[:60]!r}",
|
||||
)
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(MODEL)
|
||||
fp = tok.apply_chat_template(
|
||||
msgs[:-1], tokenize=False, add_generation_prompt=True, enable_thinking=False
|
||||
)
|
||||
ff = tok.apply_chat_template(
|
||||
msgs, tokenize=False, add_generation_prompt=False, enable_thinking=False
|
||||
)
|
||||
|
||||
# 环 2:完整渲染必须以 prompt 渲染为前缀(collator 边界法的前提假设!)
|
||||
prefix_ok = ff.startswith(fp)
|
||||
check("环2 完整渲染以 prompt 渲染为前缀", prefix_ok)
|
||||
if not prefix_ok:
|
||||
i = first_diff(ff, fp)
|
||||
print(f" 首个分歧在第 {i} 字符:\n"
|
||||
f" prompt 渲染: …{fp[max(0, i - 60) : i + 60]!r}\n"
|
||||
f" 完整渲染: …{ff[max(0, i - 60) : i + 60]!r}", flush=True)
|
||||
|
||||
# 环 3:完整渲染中 teacher 文本是否原样存在(模板会不会改写 content)
|
||||
check(
|
||||
"环3 完整渲染保留 teacher 原文",
|
||||
"\\leq 2018" in ff and "\\50118" not in ff,
|
||||
"" if "\\leq 2018" in ff else "模板改写了 assistant content!",
|
||||
)
|
||||
|
||||
# 环 4:token 级前缀(坑一:拼接稳定性)
|
||||
full_ids = tok(ff, add_special_tokens=False)["input_ids"]
|
||||
fp_ids = tok(fp, add_special_tokens=False)["input_ids"]
|
||||
tok_prefix_ok = full_ids[: len(fp_ids)] == fp_ids
|
||||
check("环4 token 级前缀一致(无跨界合并)", tok_prefix_ok)
|
||||
if not tok_prefix_ok:
|
||||
i = next(k for k in range(len(fp_ids)) if full_ids[k] != fp_ids[k])
|
||||
lo, hi = max(0, i - 3), i + 4
|
||||
print(f" 首个分歧在 token {i}/{len(fp_ids)}:\n"
|
||||
f" prompt 侧: {[tok.decode([t]) for t in fp_ids[lo:hi]]}\n"
|
||||
f" 完整侧: {[tok.decode([t]) for t in full_ids[lo:hi]]}", flush=True)
|
||||
|
||||
# 环 5:collator 全流程后,completion 解码应等于完整渲染去掉 prompt 的尾段前缀
|
||||
collator = SFTCollator(
|
||||
tok,
|
||||
max_length=cfg.max_length,
|
||||
max_prompt_length=cfg.max_prompt_length,
|
||||
enable_thinking=False,
|
||||
)
|
||||
batch = collator([ds[0]])
|
||||
ids, labels = batch["input_ids"][0], batch["labels"][0]
|
||||
comp_decoded = tok.decode(ids[labels != IGNORE_INDEX], skip_special_tokens=False)
|
||||
expected_tail = ff[len(fp) :] if prefix_ok else "(环2 已失败,无期望值)"
|
||||
tail_ok = prefix_ok and expected_tail.startswith(comp_decoded[:200])
|
||||
check("环5 completion 解码 == 渲染尾段", tail_ok)
|
||||
if prefix_ok and not tail_ok:
|
||||
i = first_diff(comp_decoded, expected_tail)
|
||||
print(f" 首个分歧在第 {i} 字符:\n"
|
||||
f" 解码: …{comp_decoded[max(0, i - 50) : i + 50]!r}\n"
|
||||
f" 期望: …{expected_tail[max(0, i - 50) : i + 50]!r}", flush=True)
|
||||
|
||||
print("\n诊断完成。把全部输出贴回对话。", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,49 @@
|
||||
"""层 1 关账判据 3:训练后 checkpoint 能被 from_pretrained 加载并生成通顺解答。
|
||||
|
||||
远程运行(CPU 即可,0.6B 生成 512 token 约 1-2 分钟):
|
||||
python -u scripts/diag_generate.py
|
||||
"""
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
from ars_opd.data import load_sft_dataset
|
||||
|
||||
MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物
|
||||
|
||||
# 不挂 teacher 解答(teacher_completions_path 缺省):只取题目做推理输入
|
||||
ds = load_sft_dataset(
|
||||
"data/dapo-math-17k-unique.parquet", subset_size=1000, seed=42
|
||||
)
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
|
||||
model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, dtype=torch.float32)
|
||||
model.eval()
|
||||
|
||||
# 取子集第 900+ 行附近的题(训练时见过,此处只验"会不会说话"不验泛化)
|
||||
for i in (900, 950):
|
||||
prompt = tok.apply_chat_template(
|
||||
ds[i]["messages"],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False, # 必须与训练取值一致(docs/02 §2.3 边界契约)
|
||||
)
|
||||
inputs = tok(prompt, return_tensors="pt", add_special_tokens=False)
|
||||
with torch.no_grad():
|
||||
out = model.generate(
|
||||
**inputs, max_new_tokens=512, do_sample=False, temperature=None, top_p=None
|
||||
)
|
||||
completion = tok.decode(
|
||||
out[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True
|
||||
)
|
||||
print(f"===== 样本 {i} 题目 =====")
|
||||
print(ds[i]["messages"][-1]["content"][120:280], "…")
|
||||
print("----- 生成(前 600 字符)-----")
|
||||
print(completion[:600])
|
||||
print()
|
||||
|
||||
print(
|
||||
"判读:应为步骤化数学解答(markdown 风格、以 Answer: 行收尾的倾向);"
|
||||
"乱码/复读/空输出 = 不通过。",
|
||||
flush=True,
|
||||
)
|
||||
@@ -0,0 +1,68 @@
|
||||
"""损失探针:用未训练的预训练模型走完整管线,逐行算 loss(层 1 疑点排查第二步)。
|
||||
|
||||
判读(训练日志初始 loss ≈ 7.5):
|
||||
- 探针也 ≈ 7:管线一致,loss 高是数据/模型现实 → 去查数据(垃圾长文、乱码占比);
|
||||
- 探针 ≈ 2-4:管线(本脚本与训练共用)没问题但训练环节另有妖 → 查训练循环差异。
|
||||
同时打印 HF 模型内建 CE(同一数学的独立实现)交叉验证 sft_loss。
|
||||
|
||||
远程运行(CPU 即可,约 1-2 分钟):
|
||||
python -u scripts/diag_loss_probe.py
|
||||
"""
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
from ars_opd.configs import SFTConfig
|
||||
from ars_opd.data import SFTCollator, load_sft_dataset
|
||||
from ars_opd.trainer import sft_loss
|
||||
|
||||
MODEL = "Qwen/Qwen3-0.6B"
|
||||
|
||||
cfg = SFTConfig(
|
||||
dataset_path="data/dapo-math-17k-unique.parquet",
|
||||
output_dir="/tmp/diag",
|
||||
subset_size=8,
|
||||
seed=42,
|
||||
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
|
||||
)
|
||||
ds = load_sft_dataset(
|
||||
cfg.dataset_path,
|
||||
cfg.dataset_split,
|
||||
cfg.subset_size,
|
||||
cfg.seed,
|
||||
cfg.teacher_completions_path,
|
||||
)
|
||||
tok = AutoTokenizer.from_pretrained(MODEL)
|
||||
collator = SFTCollator(
|
||||
tok,
|
||||
max_length=cfg.max_length,
|
||||
max_prompt_length=cfg.max_prompt_length,
|
||||
enable_thinking=False,
|
||||
)
|
||||
model = AutoModelForCausalLM.from_pretrained(MODEL, dtype=torch.float32)
|
||||
model.eval()
|
||||
|
||||
print(f"{'行':>3} {'sft_loss':>9} {'HF内建CE':>9} {'监督tok':>7} 解答开头")
|
||||
total, total_n = 0.0, 0
|
||||
for i in range(len(ds)):
|
||||
batch = collator([ds[i]])
|
||||
with torch.no_grad():
|
||||
out = model(
|
||||
input_ids=batch["input_ids"], attention_mask=batch["attention_mask"]
|
||||
)
|
||||
ours, n = sft_loss(
|
||||
out.logits, batch["input_ids"], batch["labels"], batch["attention_mask"]
|
||||
)
|
||||
# 交叉验证:HF 内建损失(labels 传入模型,内部自动移位)与 sft_loss
|
||||
# 是同一数学的两个独立实现,单行 batch 下应当几乎相等
|
||||
hf = model(
|
||||
input_ids=batch["input_ids"],
|
||||
attention_mask=batch["attention_mask"],
|
||||
labels=batch["labels"],
|
||||
).loss
|
||||
head = ds[i]["messages"][-1]["content"][:40].replace("\n", " ")
|
||||
print(f"{i:>3} {ours.item():>9.3f} {hf.item():>9.3f} {n:>7} {head}", flush=True)
|
||||
total += ours.item() * n
|
||||
total_n += n
|
||||
|
||||
print(f"\n按 token 加权平均: {total / total_n:.3f}(对照训练日志初始 loss ≈ 7.5)", flush=True)
|
||||
@@ -0,0 +1,36 @@
|
||||
"""层 1:为 DAPO 1k 子集生成 teacher(MiniMax-M3)解答缓存。
|
||||
|
||||
自包含实验脚本:全部参数写死在此,零参数复现。在**本地**运行(纯 API 调用,
|
||||
不需要 GPU;本机可直连自建网关):
|
||||
|
||||
conda activate ars-opd
|
||||
python -u scripts/generate_teacher_completions.py
|
||||
|
||||
前置:
|
||||
1. .env 已填 TEACHER_API_BASE / TEACHER_API_KEY / TEACHER_MODEL;
|
||||
2. DAPO parquet 已下载到 DATASET_PATH(见 docs/02 §4)。
|
||||
|
||||
中断安全:缓存逐条落盘,重跑本脚本自动跳过已完成条目(断点续传)。
|
||||
"""
|
||||
|
||||
from ars_opd.configs import TeacherGenConfig
|
||||
from ars_opd.data import load_sft_dataset
|
||||
from ars_opd.teacher import TeacherClient, generate_completions
|
||||
|
||||
# 非显然约束:这里的 dataset/subset_size/seed 必须与 T5 训练脚本完全一致——
|
||||
# 两侧各自走"加载→归一→抽子集",seed 相同才是同一批题(data.py 有详注)
|
||||
DATASET_PATH = "data/dapo-math-17k-unique.parquet" # DAPO 官方去重版,17917 行
|
||||
CACHE_PATH = "data/teacher_completions_dapo1k_minimax-m3.jsonl"
|
||||
|
||||
# 试跑说明:首次建议把下面 subset_size 临时改成 5,跑通并人工抽查缓存里的解答
|
||||
# 质量(think 是否剥净、格式是否正常)后再改回 1000 重跑。放心改:subset 是对
|
||||
# 同一 seed 的洗牌序列取前缀,前 5 条与前 1000 条的头 5 条完全相同,试跑写入的
|
||||
# 缓存在正式跑时全部命中,一分钱不浪费。
|
||||
|
||||
# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
|
||||
dataset = load_sft_dataset(DATASET_PATH, subset_size=1000, seed=42)
|
||||
prompts = [row["messages"] for row in dataset]
|
||||
|
||||
teacher = TeacherClient(TeacherGenConfig()) # 采样参数全用 configs.py 的显式默认
|
||||
generate_completions(prompts, CACHE_PATH, teacher)
|
||||
print(f"完成。缓存文件:{CACHE_PATH}")
|
||||
Executable
+59
@@ -0,0 +1,59 @@
|
||||
#!/usr/bin/env bash
|
||||
# 远程机 gpu-a800-060 环境搭建(幂等,可重复执行)。
|
||||
# 用法:bash scripts/setup_remote.sh
|
||||
# 硬约束:根分区仅剩 12G —— 环境、缓存、临时目录一律压到 /data/zym 下。
|
||||
set -euo pipefail
|
||||
|
||||
DATA_ROOT=/data/zym
|
||||
ENV_PATH=$DATA_ROOT/envs/ars-opd
|
||||
REPO_DIR=$DATA_ROOT/ars-opd-rebuild
|
||||
|
||||
# ---- 0. 所有会写盘的路径全部改道 /data(防根分区被写爆)----
|
||||
export HF_ENDPOINT=https://hf-mirror.com # huggingface.co 被墙,走镜像
|
||||
export HF_HOME=$DATA_ROOT/hf_cache
|
||||
export CONDA_PKGS_DIRS=$DATA_ROOT/conda_pkgs # conda 包缓存默认在根分区
|
||||
export PIP_CACHE_DIR=$DATA_ROOT/pip_cache # pip 缓存默认在根分区
|
||||
export TMPDIR=$DATA_ROOT/tmp # 大 wheel 解压临时目录
|
||||
mkdir -p "$DATA_ROOT"/{envs,hf_cache,conda_pkgs,pip_cache,tmp}
|
||||
# 清理上次中断可能残留的 pip 临时目录(含曾被 conda run 吞掉 TMPDIR 而落在 /tmp 的)
|
||||
rm -rf "$DATA_ROOT"/tmp/pip-* /tmp/pip-unpack-* 2>/dev/null || true
|
||||
|
||||
# ---- 1. conda 环境(建在 /data,不建在 ~)----
|
||||
if [ ! -d "$ENV_PATH" ]; then
|
||||
conda create -p "$ENV_PATH" python=3.11 -y
|
||||
fi
|
||||
|
||||
# ---- 2. 代码(远程只读:clone 走 HTTPS 匿名,更新只 git pull)----
|
||||
if [ ! -d "$REPO_DIR" ]; then
|
||||
git clone https://gitea.iomgaa.online/iomgaa/ars-opd-rebuild.git "$REPO_DIR"
|
||||
else
|
||||
git -C "$REPO_DIR" pull
|
||||
fi
|
||||
|
||||
# ---- 3. 依赖 ----
|
||||
# 直接调环境内 pip:conda run 会整体缓冲子命令输出(违反"日志实时可查"规矩),弃用
|
||||
"$ENV_PATH/bin/pip" install -r "$REPO_DIR/requirements.txt" -r "$REPO_DIR/requirements-remote.txt"
|
||||
# ars_opd 以 editable 方式注册进环境:脚本从任意目录都能 import,不依赖 cwd
|
||||
"$ENV_PATH/bin/pip" install -e "$REPO_DIR" --no-build-isolation --no-deps
|
||||
|
||||
# ---- 4. 环境变量持久化(写入 ~/.bashrc,幂等)----
|
||||
if ! grep -q "ars-opd-rebuild env" ~/.bashrc; then
|
||||
cat >> ~/.bashrc <<'EOF'
|
||||
|
||||
# --- ars-opd-rebuild env ---
|
||||
export HF_ENDPOINT=https://hf-mirror.com
|
||||
export HF_HOME=/data/zym/hf_cache
|
||||
export CONDA_PKGS_DIRS=/data/zym/conda_pkgs
|
||||
export PIP_CACHE_DIR=/data/zym/pip_cache
|
||||
alias opd='conda activate /data/zym/envs/ars-opd && cd /data/zym/ars-opd-rebuild'
|
||||
EOF
|
||||
fi
|
||||
|
||||
# ---- 5. 验证 ----
|
||||
echo "=== 验证 torch/CUDA ==="
|
||||
"$ENV_PATH/bin/python" -c "import torch; print('torch', torch.__version__, '| cuda可用:', torch.cuda.is_available(), '| 卡数:', torch.cuda.device_count())"
|
||||
echo "=== 验证单元测试 ==="
|
||||
"$ENV_PATH/bin/python" -m pytest "$REPO_DIR/tests" -q
|
||||
echo "=== 磁盘检查(根分区不应有明显增长)==="
|
||||
df -h / /data | tail -2
|
||||
echo "全部完成。日常使用:输入 opd 进入环境与目录。"
|
||||
@@ -0,0 +1,154 @@
|
||||
"""层 1:SFT 基线训练入口(由 train_sft.sh 经 torchrun 启动,勿直接 python 运行)。
|
||||
|
||||
自包含实验脚本:全部参数写死在下方 FULL 配置里,零参数复现;sanity 模式只是
|
||||
对 FULL 的两处显式覆盖(50 步 + 独立输出目录)。
|
||||
"""
|
||||
|
||||
# ---- FSDP 前置块(必须在一切 transformers/accelerate import 之前)----
|
||||
# 非显然约束:FSDP 的激活检查点开关是 accelerate 在 import 时读取的环境变量
|
||||
# (参考实现 train_distillation.py:15-21 的著名坑);写在 import 后会静默无效。
|
||||
# DDP 下本变量是无害 no-op——现在就位是为了未来换 4B 学生/FSDP 时只改此处一行,
|
||||
# 且改完必须 nvidia-smi 实测显存验证生效(docs/02 §2.6)。
|
||||
import os
|
||||
|
||||
os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", "false")
|
||||
|
||||
import dataclasses
|
||||
import sys
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments
|
||||
|
||||
from ars_opd.configs import SFTConfig
|
||||
from ars_opd.data import IGNORE_INDEX, SFTCollator, load_sft_dataset
|
||||
from ars_opd.trainer import SFTTrainer
|
||||
|
||||
STUDENT_MODEL = "Qwen/Qwen3-0.6B"
|
||||
|
||||
FULL = SFTConfig(
|
||||
dataset_path="data/dapo-math-17k-unique.parquet",
|
||||
output_dir="/data/zym/outputs/sft_qwen3-0.6b_dapo1k",
|
||||
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
|
||||
subset_size=1000,
|
||||
seed=42, # 非显然约束:与 generate_teacher_completions.py 一致,否则缓存大面积 miss
|
||||
max_length=4096,
|
||||
max_prompt_length=1024,
|
||||
enable_thinking=False,
|
||||
learning_rate=2e-5,
|
||||
per_device_train_batch_size=2, # B=8 曾爆 80G:大头是 (B,T,V) logits 链与激活,见 SFTConfig 注释
|
||||
gradient_accumulation_steps=8, # 全局 batch = 2 × 4 卡 × 8 = 64
|
||||
num_train_epochs=1,
|
||||
max_steps=-1,
|
||||
lr_scheduler_type="linear",
|
||||
warmup_ratio=0.0,
|
||||
gradient_checkpointing=False,
|
||||
bf16=True,
|
||||
logging_steps=1,
|
||||
save_steps=100,
|
||||
save_total_limit=2,
|
||||
report_to="none", # 层 1 先靠 tmux 实时日志;W&B 触发条件见 appendix C 表
|
||||
)
|
||||
|
||||
|
||||
def build_config() -> SFTConfig:
|
||||
"""按命令行模式产出配置。frozen dataclass 的换参方式:replace 构造新实例。"""
|
||||
mode = sys.argv[1] if len(sys.argv) > 1 else "full"
|
||||
if mode == "full":
|
||||
return FULL
|
||||
if mode == "sanity":
|
||||
return dataclasses.replace(
|
||||
FULL, max_steps=50, output_dir=FULL.output_dir + "-sanity"
|
||||
)
|
||||
raise ValueError(f"未知模式 {mode!r},只接受 full / sanity")
|
||||
|
||||
|
||||
def smoke_check_first_batch(dataset, collator, tokenizer) -> None:
|
||||
"""训练前解码第一个 batch 供肉眼核对(只在 rank0 打印一次)。
|
||||
|
||||
单测用玩具 tokenizer 钉死了预算/边界的算法(tests/test_data.py),但真
|
||||
tokenizer 的模板渲染只能在这里肉眼验证:掩码边界是否落在 assistant 起点、
|
||||
no-think 时空 <think> 块是否在 prompt 侧。这是参考实现"一次性诊断打印"
|
||||
的合理化版本(docs/02 §2.3)。
|
||||
"""
|
||||
batch = collator([dataset[0]])
|
||||
ids, labels = batch["input_ids"][0], batch["labels"][0]
|
||||
masked = labels == IGNORE_INDEX
|
||||
prompt_text = tokenizer.decode(ids[masked], skip_special_tokens=False)
|
||||
completion_text = tokenizer.decode(ids[~masked], skip_special_tokens=False)
|
||||
print(
|
||||
"=" * 30
|
||||
+ " 首样本自检(人工核对掩码边界)"
|
||||
+ "=" * 30
|
||||
+ f"\n[prompt 段 | {int(masked.sum())} tok | 不产生 loss]\n"
|
||||
+ f"…{prompt_text[-300:]}\n"
|
||||
+ f"\n[completion 段 | {int((~masked).sum())} tok | 监督目标]\n"
|
||||
+ f"{completion_text[:300]}…\n"
|
||||
+ "=" * 80,
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
cfg = build_config()
|
||||
rank0 = int(os.environ.get("RANK", "0")) == 0
|
||||
|
||||
# 加载顺序刻意 fail-fast:数据(毫秒级,最易配错)→ tokenizer(几 MB)→
|
||||
# 模型(GB 级下载)。teacher 缓存缺失要在下模型之前炸出来
|
||||
dataset = load_sft_dataset(
|
||||
cfg.dataset_path,
|
||||
cfg.dataset_split,
|
||||
cfg.subset_size,
|
||||
cfg.seed,
|
||||
cfg.teacher_completions_path,
|
||||
)
|
||||
tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
|
||||
collator = SFTCollator(
|
||||
tokenizer,
|
||||
max_length=cfg.max_length,
|
||||
max_prompt_length=cfg.max_prompt_length,
|
||||
enable_thinking=cfg.enable_thinking,
|
||||
)
|
||||
if rank0:
|
||||
smoke_check_first_batch(dataset, collator, tokenizer)
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained(STUDENT_MODEL, dtype=torch.float32)
|
||||
|
||||
args = TrainingArguments(
|
||||
output_dir=cfg.output_dir,
|
||||
# 非显然约束:必须关掉列裁剪。HF Trainer 默认删除模型 forward 签名里
|
||||
# 没有的数据列——"messages" 会被整列删光,collator 收到空字典且不报错
|
||||
remove_unused_columns=False,
|
||||
learning_rate=cfg.learning_rate,
|
||||
per_device_train_batch_size=cfg.per_device_train_batch_size,
|
||||
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
|
||||
num_train_epochs=cfg.num_train_epochs,
|
||||
max_steps=cfg.max_steps,
|
||||
lr_scheduler_type=cfg.lr_scheduler_type,
|
||||
warmup_ratio=cfg.warmup_ratio,
|
||||
gradient_checkpointing=cfg.gradient_checkpointing,
|
||||
bf16=cfg.bf16,
|
||||
seed=cfg.seed,
|
||||
logging_steps=cfg.logging_steps,
|
||||
logging_first_step=True,
|
||||
save_strategy="steps",
|
||||
save_steps=cfg.save_steps,
|
||||
save_total_limit=cfg.save_total_limit,
|
||||
report_to=cfg.report_to,
|
||||
ddp_find_unused_parameters=False, # 全参训练无闲置参数,省一次全模型扫描
|
||||
dataloader_num_workers=2, # collator 逐 batch 分词在 CPU,双 worker 与 GPU 重叠
|
||||
)
|
||||
trainer = SFTTrainer(
|
||||
model=model,
|
||||
args=args,
|
||||
train_dataset=dataset,
|
||||
data_collator=collator,
|
||||
)
|
||||
trainer.train()
|
||||
trainer.save_model() # 终态模型(save_pretrained 格式,含 config)
|
||||
if rank0:
|
||||
tokenizer.save_pretrained(cfg.output_dir)
|
||||
print(f"训练完成,模型已存至 {cfg.output_dir}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,28 @@
|
||||
#!/usr/bin/env bash
|
||||
# 层 1:SFT 基线训练(远程 gpu-a800-060 专用;本地不跑训练)。
|
||||
#
|
||||
# 用法(tmux 内执行,日志实时可查):
|
||||
# bash scripts/train_sft.sh sanity # 50 步冒烟:看首样本自检 + loss 是否从 ~2-3 下降
|
||||
# bash scripts/train_sft.sh # 正式:1k 子集 1 epoch
|
||||
#
|
||||
# 前置检查清单:
|
||||
# 1. nvidia-smi 确认下方 GPUS 四张卡空闲(只许用 8 卡中的 4 张,严禁自动选卡)
|
||||
# 2. data/ 下已有两个文件(gitignore 不随 git 走,本地 scp 上来):
|
||||
# scp data/dapo-math-17k-unique.parquet data/teacher_completions_dapo1k_minimax-m3.jsonl \
|
||||
# <远程>:/data/zym/ars-opd-rebuild/data/
|
||||
# 3. 代码是最新:git -C /data/zym/ars-opd-rebuild pull
|
||||
set -euo pipefail
|
||||
cd "$(dirname "$0")/.." # 锚定仓库根:py 内 data/... 相对路径以此为基准
|
||||
|
||||
GPUS=0,1,2,3 # ⚠️ 改这里前先 nvidia-smi
|
||||
MODE=${1:-full}
|
||||
|
||||
export CUDA_VISIBLE_DEVICES=$GPUS
|
||||
export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5)
|
||||
export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长序列 batch 易碎片化,按需扩段
|
||||
export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com}
|
||||
export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满
|
||||
# hf-mirror 不代理 HF Xet CAS(大权重走 Xet 会 401,见 train_whitebox.sh 详注)
|
||||
export HF_HUB_DISABLE_XET=1
|
||||
|
||||
torchrun --nproc_per_node=4 --master_port=29571 scripts/train_sft.py "$MODE"
|
||||
@@ -0,0 +1,168 @@
|
||||
"""层 2:white-box OPD 训练入口(由 train_whitebox.sh 经 torchrun 启动)。
|
||||
|
||||
自包含实验脚本:全部参数写死在下方 FULL 配置里,零参数复现;sanity 模式只是
|
||||
对 FULL 的两处显式覆盖(50 步 + 独立输出目录)。
|
||||
|
||||
与层 1 train_sft.py 的结构差异:双模型(student + 本地 teacher)、prompt-only
|
||||
数据(无 teacher 缓存,现场 on-policy 生成)、DistillTrainer 编排。
|
||||
"""
|
||||
|
||||
# ---- FSDP 前置块(必须在一切 transformers/accelerate import 之前,同 train_sft.py)----
|
||||
import os
|
||||
|
||||
os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", "false")
|
||||
|
||||
import dataclasses
|
||||
import sys
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments
|
||||
|
||||
from ars_opd.configs import DistillConfig
|
||||
from ars_opd.data import SFTCollator, load_sft_dataset
|
||||
from ars_opd.trainer import DistillTrainer
|
||||
|
||||
STUDENT_MODEL = "Qwen/Qwen3-0.6B" # 被训练的固定基线(同层 1,脚本级常量)
|
||||
|
||||
FULL = DistillConfig(
|
||||
dataset_path="data/dapo-math-17k-unique.parquet",
|
||||
output_dir="/data/zym/outputs/whitebox_qwen3-0.6b_dapo1k",
|
||||
teacher_model="Qwen/Qwen3-4B", # 本地全词表 teacher(须与 student 同 tokenizer)
|
||||
subset_size=1000,
|
||||
seed=42, # 与层 1 一致:同一批题上对比 SFT 与蒸馏
|
||||
max_prompt_length=1024,
|
||||
max_new_tokens=1024, # 与 max_prompt_length 之和 = 序列总长 T≈2048(§5 显存账)
|
||||
enable_thinking=False,
|
||||
beta=1.0, # 反向 KL = 式(2)
|
||||
kl_temperature=1.0,
|
||||
gen_temperature=1.0, # 纯采样自 π_θ(忠实 on-policy)
|
||||
gen_top_p=1.0,
|
||||
learning_rate=1e-6, # 论文 §5.1 蒸馏 lr;小步长也帮训练在梯度爆炸毛刺中存活
|
||||
per_device_train_batch_size=4, # §5 估算,首次远程必须 nvidia-smi 核实不 OOM
|
||||
gradient_accumulation_steps=4, # 全局 batch = 4 × 4 卡 × 4 = 64(同层 1)
|
||||
num_train_epochs=1,
|
||||
max_steps=-1,
|
||||
max_grad_norm=1.0, # 显式写出这个此前静默的稳定器(§4.1 爆炸靠它压平,见 config 注释)
|
||||
bf16=True,
|
||||
logging_steps=1,
|
||||
save_steps=100,
|
||||
save_total_limit=2,
|
||||
report_to="none",
|
||||
)
|
||||
|
||||
|
||||
def build_config() -> DistillConfig:
|
||||
"""按命令行模式产出配置。frozen dataclass 换参方式:replace 构造新实例。"""
|
||||
mode = sys.argv[1] if len(sys.argv) > 1 else "full"
|
||||
if mode == "full":
|
||||
return FULL
|
||||
if mode == "sanity":
|
||||
return dataclasses.replace(
|
||||
FULL, max_steps=50, output_dir=FULL.output_dir + "-sanity"
|
||||
)
|
||||
if mode == "noclip":
|
||||
# §4.1 教学对照:关闭裁剪 + 稍抬 lr,暴露反向 KL 原始爆炸。max_grad_norm
|
||||
# 设远高于实测范数(~14)故永不触发≈无裁剪;lr 5×放大让爆炸在 loss 上可见。
|
||||
# 与 sanity(裁到 1.0、lr 1e-6 的平滑曲线)并排 = 白盒脆弱性活教材,层 5 对照
|
||||
return dataclasses.replace(
|
||||
FULL,
|
||||
max_steps=15,
|
||||
max_grad_norm=1e9,
|
||||
learning_rate=5e-6,
|
||||
output_dir=FULL.output_dir + "-noclip",
|
||||
)
|
||||
raise ValueError(f"未知模式 {mode!r},只接受 full / sanity / noclip")
|
||||
|
||||
|
||||
def smoke_check_first_prompt(dataset, collator, tokenizer) -> None:
|
||||
"""训练前解码第一个 prompt 供肉眼核对(只在 rank0 打印一次)。
|
||||
|
||||
prompt-only 模式的自检重点:prompt 末尾应是生成引导符("...assistant\\n" +
|
||||
no-think 时的空 <think>),student 将从此续写。若末尾不对,生成的分布与
|
||||
训练目标会错位。
|
||||
"""
|
||||
batch = collator([dataset[0]])
|
||||
prompt_ids = batch["prompts"][0]
|
||||
mask = batch["prompt_attention_mask"][0].bool()
|
||||
text = tokenizer.decode(prompt_ids[mask], skip_special_tokens=False)
|
||||
print(
|
||||
"=" * 30
|
||||
+ " 首个 prompt 自检(供 on-policy 生成)"
|
||||
+ "=" * 30
|
||||
+ f"\n[{int(mask.sum())} tok,末尾应为生成引导符]\n…{text[-400:]}\n"
|
||||
+ "=" * 80,
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
cfg = build_config()
|
||||
rank0 = int(os.environ.get("RANK", "0")) == 0
|
||||
|
||||
# 加载顺序 fail-fast(同 train_sft.py):数据(毫秒级)→ tokenizer(几 MB)→
|
||||
# 模型(GB 级)。层 2 无 teacher 缓存,数据是 prompt-only 子集
|
||||
dataset = load_sft_dataset(
|
||||
cfg.dataset_path, cfg.dataset_split, cfg.subset_size, cfg.seed
|
||||
)
|
||||
student_tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
|
||||
teacher_tokenizer = AutoTokenizer.from_pretrained(cfg.teacher_model)
|
||||
collator = SFTCollator(
|
||||
student_tokenizer,
|
||||
max_prompt_length=cfg.max_prompt_length,
|
||||
enable_thinking=cfg.enable_thinking,
|
||||
prompt_only=True, # 层 2:只出 prompt 张量,completion 靠生成
|
||||
)
|
||||
if rank0:
|
||||
smoke_check_first_prompt(dataset, collator, student_tokenizer)
|
||||
|
||||
# student fp32 + bf16 混合精度(同层 1);teacher 直接 bf16(只推理,省显存)
|
||||
student = AutoModelForCausalLM.from_pretrained(STUDENT_MODEL, dtype=torch.float32)
|
||||
teacher = AutoModelForCausalLM.from_pretrained(
|
||||
cfg.teacher_model, dtype=torch.bfloat16
|
||||
)
|
||||
|
||||
args = TrainingArguments(
|
||||
output_dir=cfg.output_dir,
|
||||
remove_unused_columns=False, # 保住 messages 列供 collator(同层 1 注释)
|
||||
learning_rate=cfg.learning_rate,
|
||||
per_device_train_batch_size=cfg.per_device_train_batch_size,
|
||||
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
|
||||
num_train_epochs=cfg.num_train_epochs,
|
||||
max_steps=cfg.max_steps,
|
||||
lr_scheduler_type=cfg.lr_scheduler_type,
|
||||
warmup_ratio=cfg.warmup_ratio,
|
||||
max_grad_norm=cfg.max_grad_norm,
|
||||
gradient_checkpointing=cfg.gradient_checkpointing,
|
||||
bf16=cfg.bf16,
|
||||
seed=cfg.seed,
|
||||
logging_steps=cfg.logging_steps,
|
||||
logging_first_step=True,
|
||||
save_strategy="steps",
|
||||
save_steps=cfg.save_steps,
|
||||
save_total_limit=cfg.save_total_limit,
|
||||
report_to=cfg.report_to,
|
||||
ddp_find_unused_parameters=False,
|
||||
dataloader_num_workers=2,
|
||||
)
|
||||
trainer = DistillTrainer(
|
||||
model=student,
|
||||
args=args,
|
||||
train_dataset=dataset,
|
||||
data_collator=collator,
|
||||
teacher_model=teacher,
|
||||
teacher_tokenizer=teacher_tokenizer, # 构造时校验与 student 同词表
|
||||
beta=cfg.beta,
|
||||
kl_temperature=cfg.kl_temperature,
|
||||
gen_temperature=cfg.gen_temperature,
|
||||
gen_top_p=cfg.gen_top_p,
|
||||
max_new_tokens=cfg.max_new_tokens,
|
||||
)
|
||||
trainer.train()
|
||||
trainer.save_model()
|
||||
if rank0:
|
||||
student_tokenizer.save_pretrained(cfg.output_dir)
|
||||
print(f"训练完成,模型已存至 {cfg.output_dir}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Executable
+36
@@ -0,0 +1,36 @@
|
||||
#!/usr/bin/env bash
|
||||
# 层 2:white-box OPD 训练(远程 gpu-a800-060 专用;本地不跑训练)。
|
||||
#
|
||||
# 用法(tmux 内执行,日志实时可查):
|
||||
# bash scripts/train_whitebox.sh sanity # 50 步冒烟:首 prompt 自检 + KL loss + 生成数
|
||||
# bash scripts/train_whitebox.sh noclip # §4.1 对照:15 步,关裁剪+抬 lr,暴露原始
|
||||
# # 梯度爆炸(loss 毛刺);与 sanity 平滑曲线并排
|
||||
# bash scripts/train_whitebox.sh # 正式:1k 子集 1 epoch
|
||||
#
|
||||
# 前置检查清单:
|
||||
# 1. nvidia-smi 确认下方 GPUS 四张卡空闲(只许用 8 卡中的 4 张,严禁自动选卡)。
|
||||
# ⚠️ 白盒显存比层 1 紧:student 训练全套 + teacher(4B) 推理副本 + **两份**全词表
|
||||
# logits(student/teacher),§5 估算 B=4/T=2048 起步安全,但首跑必须盯 nvidia-smi;
|
||||
# 若 OOM,降 per_device_train_batch_size 到 2,仍不够再开 gradient_checkpointing
|
||||
# (改 DistillConfig,注意 checkpointing 与 generate 的 use_cache 交互)。
|
||||
# 2. data/dapo-math-17k-unique.parquet 已在(层 2 无需 teacher 缓存,纯 prompt-only):
|
||||
# scp data/dapo-math-17k-unique.parquet <远程>:/data/zym/ars-opd-rebuild/data/
|
||||
# 3. 代码最新:git -C /data/zym/ars-opd-rebuild pull
|
||||
# 4. 首跑会下载 teacher Qwen3-4B(GB 级)到 HF_HOME,确保 /data 有空间
|
||||
set -euo pipefail
|
||||
cd "$(dirname "$0")/.." # 锚定仓库根
|
||||
|
||||
GPUS=0,1,2,3 # ⚠️ 改这里前先 nvidia-smi
|
||||
MODE=${1:-full}
|
||||
|
||||
export CUDA_VISIBLE_DEVICES=$GPUS
|
||||
export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5)
|
||||
export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长生成序列易碎片化,按需扩段
|
||||
export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com}
|
||||
export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满
|
||||
# 非显然坑:hf-mirror 不代理 HF 的 Xet CAS——大权重走 Xet 会直连
|
||||
# cas-server.xethub.hf.co 并返 401(2026-07-19 teacher 4B 下载实撞)。禁用 Xet
|
||||
# 退回经典 HTTP/LFS 下载(镜像支持)。若仍不行:pip uninstall hf_xet
|
||||
export HF_HUB_DISABLE_XET=1
|
||||
|
||||
torchrun --nproc_per_node=4 --master_port=29572 scripts/train_whitebox.py "$MODE"
|
||||
@@ -0,0 +1,49 @@
|
||||
"""SFTConfig 的构造校验测试(层 1 / T1)。
|
||||
|
||||
只测"配置错误必须在构造时炸"这一条约定;参数语义本身没有逻辑可测。
|
||||
"""
|
||||
|
||||
import dataclasses
|
||||
|
||||
import pytest
|
||||
|
||||
from ars_opd.configs import SFTConfig
|
||||
|
||||
|
||||
def make(**overrides):
|
||||
"""最小合法配置;单测只关心被覆盖的那个字段。"""
|
||||
base = dict(dataset_path="dummy.parquet", output_dir="/tmp/dummy")
|
||||
base.update(overrides)
|
||||
return SFTConfig(**base)
|
||||
|
||||
|
||||
def test_合法配置可构造():
|
||||
cfg = make()
|
||||
assert cfg.max_length > cfg.max_prompt_length
|
||||
|
||||
|
||||
def test_prompt预算吞掉总预算时报错():
|
||||
# 这是最危险的静默失败:completion 预算为 0 → labels 全 -100 → loss 恒 0
|
||||
with pytest.raises(ValueError, match="max_prompt_length"):
|
||||
make(max_prompt_length=4096, max_length=4096)
|
||||
|
||||
|
||||
def test_非法学习率报错():
|
||||
with pytest.raises(ValueError, match="learning_rate"):
|
||||
make(learning_rate=0.0)
|
||||
|
||||
|
||||
def test_非法子集大小报错():
|
||||
with pytest.raises(ValueError, match="subset_size"):
|
||||
make(subset_size=0)
|
||||
|
||||
|
||||
def test_非法max_steps报错():
|
||||
with pytest.raises(ValueError, match="max_steps"):
|
||||
make(max_steps=0)
|
||||
|
||||
|
||||
def test_配置冻结不可变():
|
||||
cfg = make()
|
||||
with pytest.raises(dataclasses.FrozenInstanceError):
|
||||
cfg.learning_rate = 1e-3
|
||||
@@ -0,0 +1,318 @@
|
||||
"""层 1 / T3:数据管线单测(docs/02 §5.1 规定的验证项)。
|
||||
|
||||
用字符级玩具 tokenizer 在 CPU 上对拍 collator 行为,不依赖网络下载真模型。
|
||||
玩具模板刻意模仿 Qwen3 的关键结构:生成引导符 + no-think 时注入空思考块。
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from datasets import Dataset
|
||||
|
||||
from ars_opd.data import (
|
||||
IGNORE_INDEX,
|
||||
SFTCollator,
|
||||
attach_teacher_completions,
|
||||
prompt_key,
|
||||
to_messages,
|
||||
)
|
||||
|
||||
|
||||
class ToyTokenizer:
|
||||
"""字符级 tokenizer:一个字符一个 token(id = 码点)。
|
||||
|
||||
模板契约与真 chat 模板同构:
|
||||
- 每轮渲染成 "[role]content";
|
||||
- assistant 轮(或生成引导符后)在 no-think 模式下注入 "<T></T>"(模仿
|
||||
Qwen3 的空 <think>\\n\\n</think>);
|
||||
- 完整渲染 == prompt 渲染 + 解答文本,保证边界可精确断言。
|
||||
"""
|
||||
|
||||
def __init__(self, pad_token_id=0, eos_token_id=1):
|
||||
self.pad_token_id = pad_token_id
|
||||
self.eos_token_id = eos_token_id
|
||||
|
||||
def apply_chat_template(
|
||||
self,
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=False,
|
||||
enable_thinking=False,
|
||||
):
|
||||
think = "" if enable_thinking else "<T></T>"
|
||||
parts = []
|
||||
for m in messages:
|
||||
prefix = think if m["role"] == "assistant" else ""
|
||||
parts.append(f"[{m['role']}]{prefix}{m['content']}")
|
||||
text = "".join(parts)
|
||||
if add_generation_prompt:
|
||||
text += f"[assistant]{think}"
|
||||
return text
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
text,
|
||||
truncation=False,
|
||||
max_length=None,
|
||||
padding=False,
|
||||
add_special_tokens=False,
|
||||
):
|
||||
ids = [ord(c) for c in text]
|
||||
if truncation and max_length is not None:
|
||||
ids = ids[:max_length]
|
||||
return {"input_ids": ids}
|
||||
|
||||
|
||||
def ids_of(text):
|
||||
return [ord(c) for c in text]
|
||||
|
||||
|
||||
def row(question, answer=None):
|
||||
msgs = [{"role": "user", "content": question}]
|
||||
if answer is not None:
|
||||
msgs.append({"role": "assistant", "content": answer})
|
||||
return {"messages": msgs}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# to_messages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_dapo_prompt列直接归一():
|
||||
ex = {"prompt": [{"role": "user", "content": "1+1=?"}], "data_source": "dapo"}
|
||||
assert to_messages(ex) == {"messages": [{"role": "user", "content": "1+1=?"}]}
|
||||
|
||||
|
||||
def test_字符串化的列表被还原():
|
||||
ex = {"prompt": "[{'role': 'user', 'content': 'hi'}]"}
|
||||
assert to_messages(ex)["messages"] == [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
def test_坏字符串显式报错而非静默放行():
|
||||
# 参考实现 except:pass 会让这行以字符串形态流进 collator
|
||||
with pytest.raises(ValueError, match="无法解析"):
|
||||
to_messages({"prompt": "[{'role': broken"})
|
||||
|
||||
|
||||
def test_question列包成单user轮():
|
||||
assert to_messages({"question": "2+2=?"}) == {
|
||||
"messages": [{"role": "user", "content": "2+2=?"}]
|
||||
}
|
||||
|
||||
|
||||
def test_无法识别的行报错():
|
||||
with pytest.raises(ValueError, match="无法识别"):
|
||||
to_messages({"foo": 1})
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# prompt_key(与 teacher.py 的缓存契约)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_同题同键_不同题不同键():
|
||||
m1 = [{"role": "user", "content": "q"}]
|
||||
m2 = [{"role": "user", "content": "q'"}]
|
||||
assert prompt_key(m1) == prompt_key(m1)
|
||||
assert prompt_key(m1) != prompt_key(m2)
|
||||
|
||||
|
||||
def test_额外元数据字段不影响键():
|
||||
plain = [{"role": "user", "content": "q"}]
|
||||
noisy = [{"role": "user", "content": "q", "source": "dapo"}]
|
||||
assert prompt_key(plain) == prompt_key(noisy)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# attach_teacher_completions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def write_cache(path, entries):
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
for msgs, completion in entries:
|
||||
f.write(
|
||||
json.dumps({"key": prompt_key(msgs), "completion": completion}) + "\n"
|
||||
)
|
||||
|
||||
|
||||
def test_挂接teacher解答(tmp_path):
|
||||
q = [{"role": "user", "content": "1+1=?"}]
|
||||
cache = tmp_path / "cache.jsonl"
|
||||
write_cache(cache, [(q, "答案是 2")])
|
||||
ds = Dataset.from_list([{"messages": q}])
|
||||
out = attach_teacher_completions(ds, str(cache))
|
||||
assert out[0]["messages"][-1] == {"role": "assistant", "content": "答案是 2"}
|
||||
|
||||
|
||||
def test_缓存缺键一次性报全部缺失(tmp_path):
|
||||
cache = tmp_path / "cache.jsonl"
|
||||
write_cache(cache, [])
|
||||
ds = Dataset.from_list([row("q1"), row("q2")])
|
||||
with pytest.raises(KeyError, match="2/2"):
|
||||
attach_teacher_completions(ds, str(cache))
|
||||
|
||||
|
||||
def test_自带解答的行不被覆盖(tmp_path):
|
||||
cache = tmp_path / "cache.jsonl"
|
||||
write_cache(cache, [])
|
||||
ds = Dataset.from_list([row("q", "人写的答案")])
|
||||
out = attach_teacher_completions(ds, str(cache))
|
||||
assert out[0]["messages"][-1]["content"] == "人写的答案"
|
||||
|
||||
|
||||
def test_缓存文件不存在报错():
|
||||
ds = Dataset.from_list([row("q")])
|
||||
with pytest.raises(FileNotFoundError):
|
||||
attach_teacher_completions(ds, "/不存在/cache.jsonl")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SFTCollator
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def make_collator(**kw):
|
||||
defaults = dict(max_length=1000, max_prompt_length=100, enable_thinking=False)
|
||||
defaults.update(kw)
|
||||
return SFTCollator(ToyTokenizer(), **defaults)
|
||||
|
||||
|
||||
def test_基本形态_掩码与边界():
|
||||
collator = make_collator()
|
||||
batch = collator([row("ab", "cd")])
|
||||
prompt_text = "[user]ab[assistant]<T></T>"
|
||||
|
||||
labels = batch["labels"][0].tolist()
|
||||
# prompt 全 -100,completion 位置是解答的 token
|
||||
assert labels[: len(prompt_text)] == [IGNORE_INDEX] * len(prompt_text)
|
||||
assert labels[len(prompt_text) :] == ids_of("cd")
|
||||
assert batch["input_ids"][0].tolist() == ids_of(prompt_text + "cd")
|
||||
assert batch["attention_mask"][0].tolist() == [1] * (len(prompt_text) + 2)
|
||||
|
||||
|
||||
def test_超长解答不挤占prompt():
|
||||
# 头号正确性卖点:completion 被截,prompt 一个 token 不少
|
||||
prompt_text = "[user]ab[assistant]<T></T>"
|
||||
collator = make_collator(max_length=len(prompt_text) + 3)
|
||||
batch = collator([row("ab", "x" * 50)])
|
||||
|
||||
input_ids = batch["input_ids"][0].tolist()
|
||||
assert input_ids[: len(prompt_text)] == ids_of(prompt_text) # prompt 完整
|
||||
assert len(input_ids) == len(prompt_text) + 3 # completion 只剩预算内 3 个
|
||||
|
||||
|
||||
def test_超长题目截断但边界不错位():
|
||||
# 坑二场景:prompt 超预算被截断,completion 的 token 必须仍然精确
|
||||
# (切分点用未截断长度,而非截断后长度)
|
||||
collator = make_collator(max_prompt_length=10)
|
||||
batch = collator([row("很长的题目" * 20, "答案")])
|
||||
|
||||
labels = batch["labels"][0].tolist()
|
||||
non_masked = [t for t in labels if t != IGNORE_INDEX]
|
||||
assert non_masked == ids_of("答案") # 解答 token 一个不错
|
||||
assert sum(t == IGNORE_INDEX for t in labels) == 10 # prompt 恰被截到预算
|
||||
|
||||
|
||||
def test_enable_thinking两种取值边界都正确():
|
||||
for thinking in (False, True):
|
||||
collator = make_collator(enable_thinking=thinking)
|
||||
batch = collator([row("q", "ans")])
|
||||
non_masked = [t for t in batch["labels"][0].tolist() if t != IGNORE_INDEX]
|
||||
assert non_masked == ids_of("ans"), f"enable_thinking={thinking} 时边界错位"
|
||||
|
||||
|
||||
def test_nothink模板注入空思考块():
|
||||
# 参考实现的一次性诊断打印,在这里变成永久契约
|
||||
text = ToyTokenizer().apply_chat_template(
|
||||
[{"role": "user", "content": "q"}],
|
||||
add_generation_prompt=True,
|
||||
enable_thinking=False,
|
||||
)
|
||||
assert text.endswith("<T></T>")
|
||||
|
||||
|
||||
def test_sft模式仍拒绝prompt_only行():
|
||||
# 回归守卫(docs/03 §5 U3:"SFT 路径行为不变"):SFT 模式下 prompt-only 行
|
||||
# 仍是静默空训练闸门,必须报错——放开只发生在显式 prompt_only=True 模式
|
||||
with pytest.raises(ValueError, match="prompt-only"):
|
||||
make_collator()([row("没有答案的题")])
|
||||
|
||||
|
||||
def test_sft模式缺max_length构造即报错():
|
||||
with pytest.raises(ValueError, match="max_length"):
|
||||
SFTCollator(
|
||||
ToyTokenizer(), max_prompt_length=50
|
||||
) # 非 prompt_only 却无 max_length
|
||||
|
||||
|
||||
def test_左padding对齐():
|
||||
collator = make_collator()
|
||||
batch = collator([row("ab", "cd"), row("a", "c")])
|
||||
t = batch["input_ids"].shape[1]
|
||||
short_mask = batch["attention_mask"][1].tolist()
|
||||
n_pad = t - short_mask.count(1)
|
||||
|
||||
assert n_pad > 0
|
||||
assert short_mask[:n_pad] == [0] * n_pad # padding 在左
|
||||
assert batch["labels"][1].tolist()[:n_pad] == [IGNORE_INDEX] * n_pad
|
||||
assert batch["input_ids"][1].tolist()[:n_pad] == [0] * n_pad # pad_token_id=0
|
||||
|
||||
|
||||
def test_pad回退到eos():
|
||||
collator = SFTCollator(
|
||||
ToyTokenizer(pad_token_id=None, eos_token_id=7),
|
||||
max_length=100,
|
||||
max_prompt_length=50,
|
||||
)
|
||||
assert collator.pad_token_id == 7
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SFTCollator:prompt_only 模式(层 2 / U3)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def make_prompt_collator(**kw):
|
||||
defaults = dict(max_prompt_length=100, enable_thinking=False, prompt_only=True)
|
||||
defaults.update(kw)
|
||||
return SFTCollator(ToyTokenizer(), **defaults)
|
||||
|
||||
|
||||
def test_prompt_only模式返回prompt张量且不报错():
|
||||
# 层 2:prompt-only 行是常态,不再报错;渲染带生成引导符供 generate 续写
|
||||
collator = make_prompt_collator()
|
||||
batch = collator([row("ab")])
|
||||
expected = "[user]ab[assistant]<T></T>"
|
||||
|
||||
assert set(batch.keys()) == {"prompts", "prompt_attention_mask"}
|
||||
assert batch["prompts"][0].tolist() == ids_of(expected)
|
||||
assert batch["prompt_attention_mask"][0].tolist() == [1] * len(expected)
|
||||
|
||||
|
||||
def test_prompt_only模式左padding对齐():
|
||||
# 生成要求左 padding:短 prompt 在左侧补 pad,右边界对齐
|
||||
collator = make_prompt_collator()
|
||||
batch = collator([row("abc"), row("a")])
|
||||
t = batch["prompts"].shape[1]
|
||||
short_mask = batch["prompt_attention_mask"][1].tolist()
|
||||
n_pad = t - short_mask.count(1)
|
||||
|
||||
assert n_pad > 0
|
||||
assert short_mask[:n_pad] == [0] * n_pad # padding 在左
|
||||
assert batch["prompts"][1].tolist()[:n_pad] == [0] * n_pad # pad_token_id=0
|
||||
|
||||
|
||||
def test_prompt_only模式截断到prompt预算():
|
||||
collator = make_prompt_collator(max_prompt_length=8)
|
||||
batch = collator([row("很长的题目" * 20)])
|
||||
assert batch["prompts"].shape[1] == 8 # 恰截到 max_prompt_length
|
||||
|
||||
|
||||
def test_prompt_only模式剥掉末轮assistant():
|
||||
# 若数据碰巧带了 assistant 轮,取生成前上下文(剥掉它再加生成引导符)
|
||||
collator = make_prompt_collator()
|
||||
batch = collator([row("q", "已有答案")])
|
||||
expected = "[user]q[assistant]<T></T>" # 不含"已有答案"
|
||||
assert batch["prompts"][0].tolist() == ids_of(expected)
|
||||
@@ -0,0 +1,91 @@
|
||||
"""层 2 / U4:生成重建纯逻辑单测(docs/03 §2.5)。
|
||||
|
||||
只测 build_generated_batch——把 model.generate 输出重建成 input_ids/attention/labels
|
||||
的张量逻辑。DistillTrainer 的编排(生成→双前向→divergence)需真模型,由远程冒烟
|
||||
验证。重点盯"首个 eos 之后一律屏蔽"这条最易 off-by-one 的规则。
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from ars_opd.data import IGNORE_INDEX
|
||||
from ars_opd.trainer import build_generated_batch
|
||||
|
||||
EOS = 99
|
||||
PAD = 0
|
||||
|
||||
|
||||
def test_eos在中间_其后全屏蔽():
|
||||
# 行内:prompt=[pad,u,u],生成=[a, EOS, pad];a 与 eos 有效,eos 后的 pad 无效
|
||||
prompts = torch.tensor([[PAD, 5, 5]])
|
||||
prompt_mask = torch.tensor([[0, 1, 1]])
|
||||
gen_output = torch.tensor([[PAD, 5, 5, 7, EOS, PAD]]) # (1, P+G)=(1,6)
|
||||
|
||||
ids, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
|
||||
|
||||
assert ids.tolist() == gen_output.tolist() # input_ids 即生成全序列
|
||||
# prompt 段全 -100;生成段 [7, EOS, -100]
|
||||
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX]
|
||||
# attention:prompt 左 pad=0,生成段 eos 及之前=1、其后=0
|
||||
assert attn[0].tolist() == [0, 1, 1, 1, 1, 0]
|
||||
|
||||
|
||||
def test_无eos撞max时整段生成有效():
|
||||
prompts = torch.tensor([[5, 5, 5]])
|
||||
prompt_mask = torch.tensor([[1, 1, 1]])
|
||||
gen_output = torch.tensor([[5, 5, 5, 8, 9, 10]]) # 生成三 token,无 eos
|
||||
|
||||
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
|
||||
|
||||
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [8, 9, 10] # 全监督
|
||||
assert attn[0].tolist() == [1, 1, 1, 1, 1, 1]
|
||||
|
||||
|
||||
def test_batch内不同生成长度_各自正确对齐():
|
||||
# 行0 提前 eos(右侧被补 pad 到 batch 宽度);行1 撞 max。二者共用同一 (B,P+G)
|
||||
prompts = torch.tensor([[PAD, 5, 5], [5, 5, 5]])
|
||||
prompt_mask = torch.tensor([[0, 1, 1], [1, 1, 1]])
|
||||
gen_output = torch.tensor(
|
||||
[
|
||||
[PAD, 5, 5, 7, EOS, PAD], # 行0:生成 [7, EOS],末位 pad 补齐
|
||||
[5, 5, 5, 8, 9, 10], # 行1:生成 [8, 9, 10]
|
||||
]
|
||||
)
|
||||
|
||||
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
|
||||
|
||||
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX]
|
||||
assert labels[1].tolist() == [IGNORE_INDEX] * 3 + [8, 9, 10]
|
||||
assert attn[0].tolist() == [0, 1, 1, 1, 1, 0]
|
||||
assert attn[1].tolist() == [1, 1, 1, 1, 1, 1]
|
||||
|
||||
|
||||
def test_pad等于eos也不误判():
|
||||
# 关键 corner:pad_token == eos_token。首个 eos 有效、其后补位的 eos 全屏蔽
|
||||
prompts = torch.tensor([[5, 5, 5]])
|
||||
prompt_mask = torch.tensor([[1, 1, 1]])
|
||||
gen_output = torch.tensor([[5, 5, 5, 7, EOS, EOS]]) # 末位补的 pad 恰好==eos
|
||||
|
||||
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
|
||||
|
||||
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX]
|
||||
assert attn[0].tolist() == [1, 1, 1, 1, 1, 0] # 第二个 eos 被当补位屏蔽
|
||||
|
||||
|
||||
def test_立即eos_只留一个token():
|
||||
prompts = torch.tensor([[5, 5, 5]])
|
||||
prompt_mask = torch.tensor([[1, 1, 1]])
|
||||
gen_output = torch.tensor([[5, 5, 5, EOS, PAD, PAD]]) # 第一步就 eos
|
||||
|
||||
_, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
|
||||
|
||||
assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [EOS, IGNORE_INDEX, IGNORE_INDEX]
|
||||
assert attn[0].tolist() == [1, 1, 1, 1, 0, 0]
|
||||
|
||||
|
||||
def test_prompt段恒为负100():
|
||||
prompts = torch.tensor([[PAD, PAD, 5, 5]])
|
||||
prompt_mask = torch.tensor([[0, 0, 1, 1]])
|
||||
gen_output = torch.tensor([[PAD, PAD, 5, 5, 8, 9]])
|
||||
|
||||
_, _, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS)
|
||||
assert labels[0, :4].tolist() == [IGNORE_INDEX] * 4 # prompt 段(含左 pad)全 -100
|
||||
@@ -0,0 +1,131 @@
|
||||
"""层 2 / U2:token 级散度单测(docs/03 §2、§4.1、§6.1)。
|
||||
|
||||
只测纯张量函数 token_divergence;DistillTrainer 类由远程冒烟验证。
|
||||
三块:
|
||||
1. 对拍 PyTorch 自带 KL(独立 oracle,非同式自证);
|
||||
2. 反向/前向 KL 的方向性(mode-seeking vs mode-covering);
|
||||
3. §4.1 梯度爆炸演示——teacher 概率趋 0 时 student 梯度暴涨(层 5 有界乘子的对照桩)。
|
||||
|
||||
构造技巧:softmax(log p) = p(p 已归一),故用 `probs.log()` 当 logits 即可精确
|
||||
控制两侧分布,让手算/对拍成为可能。
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.distributions import Categorical, kl_divergence
|
||||
|
||||
from ars_opd.data import IGNORE_INDEX
|
||||
from ars_opd.trainer import token_divergence
|
||||
|
||||
V = 5 # 玩具词表
|
||||
|
||||
|
||||
def logits_of(probs: list[float]) -> torch.Tensor:
|
||||
"""概率向量 -> (1, 1, V) logits,使 log_softmax 后精确还原该分布。"""
|
||||
return torch.tensor(probs).log().reshape(1, 1, V)
|
||||
|
||||
|
||||
ONE_VALID = torch.zeros(1, 1, dtype=torch.long) # 单个有效 token(id 0 ≠ -100)
|
||||
|
||||
|
||||
# ---- 1. 对拍 PyTorch KL ----
|
||||
|
||||
|
||||
@pytest.mark.parametrize("beta", [0.0, 1.0, 0.5])
|
||||
def test_散度对拍pytorch_kl(beta):
|
||||
p_s = [0.10, 0.20, 0.30, 0.25, 0.15]
|
||||
p_t = [0.05, 0.05, 0.40, 0.40, 0.10]
|
||||
loss, num = token_divergence(logits_of(p_s), logits_of(p_t), ONE_VALID, beta=beta)
|
||||
assert num == 1
|
||||
|
||||
cs, ct = Categorical(torch.tensor(p_s)), Categorical(torch.tensor(p_t))
|
||||
if beta == 1.0: # 反向 KL(π_θ‖π_T)
|
||||
oracle = kl_divergence(cs, ct)
|
||||
elif beta == 0.0: # 前向 KL(π_T‖π_θ)
|
||||
oracle = kl_divergence(ct, cs)
|
||||
else: # JSD:对混合分布的两支 KL 加权
|
||||
m = Categorical((1 - beta) * torch.tensor(p_s) + beta * torch.tensor(p_t))
|
||||
oracle = beta * kl_divergence(ct, m) + (1 - beta) * kl_divergence(cs, m)
|
||||
assert torch.allclose(loss, oracle, atol=1e-6)
|
||||
|
||||
|
||||
def test_同分布散度为零():
|
||||
p = [0.1, 0.2, 0.3, 0.25, 0.15]
|
||||
for beta in (0.0, 1.0, 0.5):
|
||||
loss, _ = token_divergence(logits_of(p), logits_of(p), ONE_VALID, beta=beta)
|
||||
assert torch.allclose(loss, torch.zeros(()), atol=1e-6)
|
||||
|
||||
|
||||
# ---- 2. 方向性:反向罚"越界",前向罚"漏覆盖" ----
|
||||
|
||||
|
||||
def test_kl方向性():
|
||||
peaked = [0.90, 0.025, 0.025, 0.025, 0.025]
|
||||
diffuse = [1 / V] * V
|
||||
|
||||
# 情形 A:teacher 尖、student 弥散——student 把质量放到 teacher≈0 处。
|
||||
# 反向 KL(π_θ‖π_T) 因 log(π_θ/π_T) 在越界 token 上爆大而重罚;前向相对轻。
|
||||
rev_A = token_divergence(
|
||||
logits_of(diffuse), logits_of(peaked), ONE_VALID, beta=1.0
|
||||
)[0]
|
||||
fwd_A = token_divergence(
|
||||
logits_of(diffuse), logits_of(peaked), ONE_VALID, beta=0.0
|
||||
)[0]
|
||||
assert rev_A > fwd_A # 反向惩罚 student 越出 teacher 支持集(mode-seeking)
|
||||
|
||||
# 情形 B:student 尖、teacher 弥散——teacher 的质量落在 student≈0 处。
|
||||
# 前向 KL(π_T‖π_θ) 重罚"漏覆盖";反向相对轻。
|
||||
rev_B = token_divergence(
|
||||
logits_of(peaked), logits_of(diffuse), ONE_VALID, beta=1.0
|
||||
)[0]
|
||||
fwd_B = token_divergence(
|
||||
logits_of(peaked), logits_of(diffuse), ONE_VALID, beta=0.0
|
||||
)[0]
|
||||
assert fwd_B > rev_B # 前向惩罚 student 没覆盖 teacher 的质量(mode-covering)
|
||||
|
||||
|
||||
def test_温度升高软化分布降低反向kl():
|
||||
# student 与 teacher 都尖但尖在不同 token;升温软化两侧 → 反向 KL 下降
|
||||
s, t = [0.90, 0.025, 0.025, 0.025, 0.025], [0.025, 0.90, 0.025, 0.025, 0.025]
|
||||
cold = token_divergence(logits_of(s), logits_of(t), ONE_VALID, temperature=1.0)[0]
|
||||
hot = token_divergence(logits_of(s), logits_of(t), ONE_VALID, temperature=4.0)[0]
|
||||
assert hot < cold
|
||||
|
||||
|
||||
# ---- 3. §4.1 梯度爆炸演示 ----
|
||||
|
||||
|
||||
def test_梯度爆炸_teacher概率趋0时student梯度暴涨():
|
||||
# student 固定:对"采样 token"(id 0) 给最高 logit(模拟 on-policy 采到它)
|
||||
base = [1.5, 0.5, 0.3, 0.2, 0.1]
|
||||
epsilons = [1e-1, 1e-2, 1e-3, 1e-4, 1e-5, 1e-6]
|
||||
grad_norms = []
|
||||
for eps in epsilons:
|
||||
student_logits = torch.tensor(base).reshape(1, 1, V).requires_grad_(True)
|
||||
# teacher:token0 概率 = eps(越来越"厌恶"它),其余 (1-eps) 均分
|
||||
t_probs = [(1 - eps) / (V - 1)] * V
|
||||
t_probs[0] = eps
|
||||
teacher_logits = torch.tensor(t_probs).log().reshape(1, 1, V)
|
||||
loss, _ = token_divergence(student_logits, teacher_logits, ONE_VALID, beta=1.0)
|
||||
loss.backward()
|
||||
grad_norms.append(student_logits.grad.norm().item())
|
||||
|
||||
# 单调暴涨:teacher 越否定采样 token,student 梯度范数越大
|
||||
for lo, hi in zip(grad_norms, grad_norms[1:]):
|
||||
assert hi > lo
|
||||
# 末端(π_T=1e-6)远超首端(π_T=1e-1)——§4.1 的可执行证据,
|
||||
# 为层 5"有界乘子 π̂"的稳定性对照埋桩
|
||||
assert grad_norms[-1] > 5 * grad_norms[0]
|
||||
|
||||
|
||||
def test_全掩码batch显式报错():
|
||||
all_masked = torch.full((1, 1), IGNORE_INDEX, dtype=torch.long)
|
||||
with pytest.raises(ValueError, match="有效 completion"):
|
||||
token_divergence(logits_of([0.2] * V), logits_of([0.2] * V), all_masked)
|
||||
|
||||
|
||||
def test_beta越界报错():
|
||||
with pytest.raises(ValueError, match="beta"):
|
||||
token_divergence(
|
||||
logits_of([0.2] * V), logits_of([0.2] * V), ONE_VALID, beta=2.0
|
||||
)
|
||||
@@ -0,0 +1,162 @@
|
||||
"""estimator.py 单测——docs/04 §5.2:公式对拍 + 定理 4.1 性质 + 方差收缩。
|
||||
|
||||
对拍精神源自参考实现 validate_chunk_mc_estimator.py(比 MSE_freq vs
|
||||
MSE_bayes),但全用 toy 数据本地 CPU 跑,不连真 teacher。
|
||||
detach 命门的两个世界断言在 tests/test_estimator_detach.py(E3)。
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from ars_opd.estimator import bayesian_target, chunk_prior
|
||||
|
||||
# ------------------------------------------------------------- chunk_prior
|
||||
|
||||
|
||||
def test_prior_is_geometric_mean():
|
||||
# 式(4) 手算:p = [0.9, 0.1] → π̄ = exp((log .9 + log .1)/2) = √0.09 = 0.3
|
||||
log_probs = torch.log(torch.tensor([0.9, 0.1]))
|
||||
assert math.isclose(chunk_prior(log_probs).item(), 0.3, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_prior_uniform_probs():
|
||||
# 全同概率的几何均值 = 该概率本身
|
||||
log_probs = torch.full((50,), math.log(0.5))
|
||||
assert math.isclose(chunk_prior(log_probs).item(), 0.5, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_prior_shape_and_range():
|
||||
pi_bar = chunk_prior(torch.log(torch.rand(50).clamp(1e-6, 1.0)))
|
||||
assert pi_bar.shape == () # (C,) -> 标量
|
||||
assert 0.0 < pi_bar.item() <= 1.0
|
||||
|
||||
|
||||
def test_prior_log_domain_survives_underflow():
|
||||
# 50 个 p=0.01 直接连乘 = 1e-100(fp32 下溢为 0);log 域算出 0.01
|
||||
log_probs = torch.full((50,), math.log(0.01))
|
||||
assert math.isclose(chunk_prior(log_probs).item(), 0.01, rel_tol=1e-4)
|
||||
|
||||
|
||||
def test_prior_clamp_floor():
|
||||
# 极端负 log 均值 → exp 下溢,clamp 兜到 1e-8 保持严格为正(定理 4.1b 前提)
|
||||
log_probs = torch.full((5,), -1e9)
|
||||
assert chunk_prior(log_probs).item() == pytest.approx(1e-8)
|
||||
|
||||
|
||||
def test_prior_is_detached():
|
||||
# detach 命门:π̄ 不带梯度(逃逸机制的完整断言在 test_estimator_detach.py)
|
||||
log_probs = torch.log(torch.tensor([0.5, 0.5], requires_grad=True))
|
||||
pi_bar = chunk_prior(log_probs)
|
||||
assert not pi_bar.requires_grad
|
||||
|
||||
|
||||
def test_prior_empty_raises():
|
||||
with pytest.raises(ValueError, match="为空"):
|
||||
chunk_prior(torch.tensor([]))
|
||||
|
||||
|
||||
# --------------------------------------------------------- bayesian_target
|
||||
|
||||
|
||||
def test_target_formula_hand_computed():
|
||||
# 式(5) 手算:k=8, π̄=0.5, N=10, α=1 → π̂ = (8 + 0.5)/11 = 0.77272…
|
||||
pi_hat = bayesian_target(8.0, torch.tensor(0.5), n_rollouts=10, alpha=1.0)
|
||||
assert math.isclose(pi_hat.item(), 8.5 / 11, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_target_convex_combination_identity():
|
||||
# 式(10) 恒等:π̂ = N/(N+α)·(k/N) + α/(N+α)·π̄,任取参数逐点核对
|
||||
k, pi_bar, n, alpha = 3.7, torch.tensor(0.42), 10, 1.5
|
||||
direct = bayesian_target(k, pi_bar, n, alpha).item()
|
||||
convex = (n / (n + alpha)) * (k / n) + (alpha / (n + alpha)) * pi_bar.item()
|
||||
assert math.isclose(direct, convex, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_target_anti_collapse_at_k_zero():
|
||||
# 定理 4.1(b):k=0(teacher 全否定)时 π̂ = α·π̄/(N+α) > 0,监督不归零
|
||||
pi_hat = bayesian_target(0.0, torch.tensor(0.3), n_rollouts=10, alpha=1.0)
|
||||
assert math.isclose(pi_hat.item(), 0.3 / 11, rel_tol=1e-6)
|
||||
assert pi_hat.item() > 0
|
||||
|
||||
|
||||
def test_target_full_score_shrinks_below_one():
|
||||
# 贝叶斯收缩:k=N 满分时 π̂ = (N+α·π̄)/(N+α) < 1(只要 π̄<1)——
|
||||
# 先验把估计从两端往中间拉,这正是方差收缩的来源
|
||||
pi_hat = bayesian_target(10.0, torch.tensor(0.5), n_rollouts=10, alpha=1.0)
|
||||
assert math.isclose(pi_hat.item(), 10.5 / 11, rel_tol=1e-6)
|
||||
assert pi_hat.item() < 1.0
|
||||
|
||||
|
||||
def test_target_bounded_in_unit_interval():
|
||||
# 定理 4.1(a):任意合法参数下 π̂ ∈ (0, 1]
|
||||
for k in [0.0, 2.5, 10.0]:
|
||||
for p in [1e-8, 0.5, 1.0]:
|
||||
v = bayesian_target(k, torch.tensor(p), 10, 1.0).item()
|
||||
assert 0.0 < v <= 1.0
|
||||
|
||||
|
||||
def test_target_alpha_zero_is_frequency_estimate():
|
||||
# α=0 退化为 k/N(no_bayesian 消融);k=0 时被 clamp 兜到 1e-8 而非 0
|
||||
assert math.isclose(
|
||||
bayesian_target(7.0, torch.tensor(0.5), 10, 0.0).item(), 0.7, rel_tol=1e-6
|
||||
)
|
||||
assert bayesian_target(0.0, torch.tensor(0.5), 10, 0.0).item() == pytest.approx(
|
||||
1e-8
|
||||
)
|
||||
|
||||
|
||||
def test_target_is_detached_even_with_grad_input():
|
||||
# 第二道防线:pi_bar 带梯度传入,π̂ 仍必须 detach
|
||||
pi_bar = torch.tensor(0.5, requires_grad=True)
|
||||
pi_hat = bayesian_target(5.0, pi_bar, 10, 1.0)
|
||||
assert not pi_hat.requires_grad
|
||||
|
||||
|
||||
def test_target_validation_raises():
|
||||
pi_bar = torch.tensor(0.5)
|
||||
with pytest.raises(ValueError, match="n_rollouts"):
|
||||
bayesian_target(0.0, pi_bar, 0, 1.0)
|
||||
with pytest.raises(ValueError, match="alpha"):
|
||||
bayesian_target(0.0, pi_bar, 10, -0.1)
|
||||
with pytest.raises(ValueError, match="越界"):
|
||||
bayesian_target(11.0, pi_bar, 10, 1.0) # k_sem > N:口径不一致
|
||||
with pytest.raises(ValueError, match="越界"):
|
||||
bayesian_target(-0.5, pi_bar, 10, 1.0)
|
||||
|
||||
|
||||
# --------------------------------------------- 方差收缩(定理 4.1c,toy 模拟)
|
||||
|
||||
|
||||
def test_variance_shrinkage_beats_frequency_estimate():
|
||||
"""toy 模拟对拍 validate_chunk_mc_estimator.py 的 MSE_freq vs MSE_bayes。
|
||||
|
||||
设真值 μ:每次试验采 N=10 个相似度 sim_i(均值 μ 的噪声),
|
||||
频率估计 = mean(sim) = k/N,贝叶斯估计 = (k + α·π̄)/(N+α)。
|
||||
先验 π̄ = μ(理想先验)时,收缩纯降方差、零偏差代价,MSE 必更小。
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
mu, n, alpha = 0.7, 10, 1.0
|
||||
trials = 2000
|
||||
# (trials, N) 的相似度样本:均值 μ、截断到 [0,1]
|
||||
sims = (mu + 0.25 * torch.randn(trials, n)).clamp(0.0, 1.0)
|
||||
k = sims.sum(dim=1) # (trials,) 每次试验的 k_sem
|
||||
freq = k / n
|
||||
bayes = (k + alpha * mu) / (n + alpha)
|
||||
mse_freq = ((freq - mu) ** 2).mean().item()
|
||||
mse_bayes = ((bayes - mu) ** 2).mean().item()
|
||||
assert mse_bayes < mse_freq
|
||||
|
||||
|
||||
def test_variance_shrinkage_robust_to_imperfect_prior():
|
||||
# 先验偏离真值(π̄ = μ±0.1)仍应赢:α=1、N=10 时先验权重仅 1/11,
|
||||
# 引入的偏差平方远小于省下的方差(定理 4.1c 在论文设定下的稳健性)
|
||||
torch.manual_seed(1)
|
||||
mu, n, alpha, trials = 0.6, 10, 1.0, 2000
|
||||
sims = (mu + 0.25 * torch.randn(trials, n)).clamp(0.0, 1.0)
|
||||
k = sims.sum(dim=1)
|
||||
mse_freq = ((k / n - mu) ** 2).mean().item()
|
||||
for prior in [mu - 0.1, mu + 0.1]:
|
||||
bayes = (k + alpha * prior) / (n + alpha)
|
||||
assert ((bayes - mu) ** 2).mean().item() < mse_freq
|
||||
@@ -0,0 +1,72 @@
|
||||
"""π̂ detach 命门约束的守护测试(对应 docs/01 §3.4,论文式(5)(8))。
|
||||
|
||||
背景:chunk 损失 L = -π̂·Σlog π_θ 中,π̂ 的先验 π̄ 由学生自身概率算出。
|
||||
若不切断 π̄ 的梯度通路,最速下降方向会变成压低学生对自己 token 的概率、
|
||||
把乘子 π̂ 推向 0 以逃逸惩罚(p·ln(1/p)→0,指数快过对数),且恰好在
|
||||
teacher 全否定(k≈0)、最需要纠正的 chunk 上塌缩。推导见 docs/01 §3.4。
|
||||
|
||||
现状:独立的数学性质测试,仅依赖 torch(单 token 简化,C=1)。
|
||||
层 3 完成 ars_opd/estimator.py 后,需追加针对真实实现的同名断言,
|
||||
确保重构时 `.detach()` 不被误删(参考实现锚点:trainer:2196、2201)。
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
# 与论文/参考实现默认一致:α=1, N=10
|
||||
ALPHA = 1.0
|
||||
N_ROLLOUTS = 10.0
|
||||
|
||||
|
||||
def chunk_loss_and_grad(p0: float, k: float, detach_prior: bool) -> tuple[float, float]:
|
||||
"""单 token 版式(8) chunk 项,返回 (loss 值, dL/dp)。
|
||||
|
||||
参数:
|
||||
p0: 学生对自己 token 的概率,标量。
|
||||
k: teacher 相似度票数 k_sem,标量(0 = 全否定)。
|
||||
detach_prior: 是否切断先验 π̄ 的梯度通路。
|
||||
返回:
|
||||
(loss.item(), p.grad.item())
|
||||
"""
|
||||
p = torch.tensor(p0, requires_grad=True)
|
||||
log_p = p.log()
|
||||
prior_src = log_p.detach() if detach_prior else log_p
|
||||
pi_bar = prior_src.exp() # 式(4):C=1 时几何均值即 p 本身
|
||||
pi_hat = (k + ALPHA * pi_bar) / (N_ROLLOUTS + ALPHA) # 式(5)
|
||||
loss = -pi_hat * log_p # 式(8) chunk 项
|
||||
loss.backward()
|
||||
return loss.item(), p.grad.item()
|
||||
|
||||
|
||||
def test_detach_reinforces_even_when_teacher_rejects():
|
||||
"""detach 世界:即使 teacher 全否定(k=0),梯度仍为负 → optimizer 增大 p。
|
||||
|
||||
这是贝叶斯兜底的本意:k=0 处仍有非零、方向正确的学习信号。
|
||||
"""
|
||||
_, grad = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=True)
|
||||
assert grad < 0
|
||||
|
||||
|
||||
def test_no_detach_escapes_when_teacher_rejects():
|
||||
"""不 detach 世界:k=0 且 p 低于 1/e 时梯度为正 → optimizer 压低 p(逃逸)。
|
||||
|
||||
此断言若失败(梯度变负),说明有人"修复"了 detach——那恰恰是 bug。
|
||||
"""
|
||||
_, grad = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=False)
|
||||
assert grad > 0
|
||||
|
||||
|
||||
def test_detach_does_not_change_loss_value():
|
||||
"""detach 只剪梯度不改数值:两个世界的前向 loss 必须完全相等。"""
|
||||
loss_detached, _ = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=True)
|
||||
loss_attached, _ = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=False)
|
||||
assert loss_detached == loss_attached
|
||||
|
||||
|
||||
def test_teacher_agreement_blocks_escape_even_without_detach():
|
||||
"""k 大时逃逸被堵死:分子中 k·|log p| 项不受 p 控制,随否认无限增长。
|
||||
|
||||
逃逸条件为 NLL > 1 + k/(α·π̄);k=5、p=0.2 时阈值 ≈ 26,远未达到,
|
||||
故即使不 detach 梯度仍为负。印证"塌缩恰好集中在 k≈0 的 chunk"。
|
||||
"""
|
||||
_, grad = chunk_loss_and_grad(p0=0.2, k=5.0, detach_prior=False)
|
||||
assert grad < 0
|
||||
@@ -0,0 +1,142 @@
|
||||
"""similarity.py 单测——docs/04 §5.1:手构字符串钉死 φ 与 k_sem(式3)。
|
||||
|
||||
全部本地 CPU、纯标准库,不依赖 torch/transformers。
|
||||
关键手算用例在各测试的注释里逐步展开,方便对着验算。
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from ars_opd.similarity import aggregate_similarity, edit_similarity, phi, rouge1
|
||||
|
||||
# ---------------------------------------------------------------- rouge1
|
||||
|
||||
|
||||
def test_rouge1_identical_is_exact_one():
|
||||
# 全同串必须精确 = 1.0(参考实现因分母 +1e-8 只能得 ≈0.99999998)
|
||||
s = "so x = 5 and y = 12"
|
||||
assert rouge1(s, s) == 1.0
|
||||
|
||||
|
||||
def test_rouge1_disjoint_is_zero():
|
||||
assert rouge1("a b c", "x y z") == 0.0
|
||||
|
||||
|
||||
def test_rouge1_partial_overlap_hand_computed():
|
||||
# hyp = {a, b, c}, ref = {a, b, d}:overlap = 2
|
||||
# precision = 2/3, recall = 2/3, F1 = 2·(2/3)(2/3) / (4/3) = 2/3
|
||||
assert math.isclose(rouge1("a b c", "a b d"), 2 / 3)
|
||||
|
||||
|
||||
def test_rouge1_multiset_counts_repeats():
|
||||
# 多重集语义:hyp = [x,x,x,x], ref = [x] → overlap = min(4,1) = 1
|
||||
# precision = 1/4, recall = 1/1, F1 = 2·(1/4)/(5/4) = 0.4
|
||||
# (参考实现的 set 版会给满分 1.0——数学文本重复词多,这是关键失真点)
|
||||
assert math.isclose(rouge1("x x x x", "x"), 0.4)
|
||||
|
||||
|
||||
def test_rouge1_is_bag_of_words_order_blind():
|
||||
# 词袋:只看用了哪些词,不看顺序
|
||||
assert rouge1("a b", "b a") == 1.0
|
||||
|
||||
|
||||
def test_rouge1_empty_sides():
|
||||
assert rouge1("", "a b") == 0.0
|
||||
assert rouge1("a b", "") == 0.0
|
||||
assert rouge1("", "") == 0.0
|
||||
assert rouge1(" ", "a") == 0.0 # 纯空白 split 后无词
|
||||
|
||||
|
||||
# ---------------------------------------------------------- edit_similarity
|
||||
|
||||
|
||||
def test_edit_identical_is_one():
|
||||
s = "so x = 5 and y = 12"
|
||||
assert edit_similarity(s, s) == 1.0
|
||||
|
||||
|
||||
def test_edit_totally_different_is_zero():
|
||||
# ["a","b"] vs ["c","d"]:2 次替换,dist=2, max(m,n)=2 → 1 − 1 = 0
|
||||
assert edit_similarity("a b", "c d") == 0.0
|
||||
|
||||
|
||||
def test_edit_single_substitution_hand_computed():
|
||||
# ["a","b","c"] vs ["a","x","c"]:1 次替换,dist=1, max=3 → 2/3
|
||||
assert math.isclose(edit_similarity("a b c", "a x c"), 2 / 3)
|
||||
|
||||
|
||||
def test_edit_insertion_hand_computed():
|
||||
# ["a","b"] vs ["a","x","b"]:1 次插入,dist=1, max=3 → 2/3
|
||||
assert math.isclose(edit_similarity("a b", "a x b"), 2 / 3)
|
||||
|
||||
|
||||
def test_edit_is_order_sensitive():
|
||||
# ["a","b"] vs ["b","a"]:两次替换 dist=2 → 0.0;与 rouge1 的 1.0 互补
|
||||
assert edit_similarity("a b", "b a") == 0.0
|
||||
assert rouge1("a b", "b a") == 1.0
|
||||
|
||||
|
||||
def test_edit_empty_sides():
|
||||
assert edit_similarity("", "") == 1.0 # 零距离
|
||||
assert edit_similarity("a b", "") == 0.0 # 全删
|
||||
assert edit_similarity("", "a b") == 0.0 # 全插
|
||||
|
||||
|
||||
def test_edit_asymmetric_lengths():
|
||||
# ["a"] vs ["a","b","c","d"]:3 次插入,dist=3, max=4 → 1/4
|
||||
assert math.isclose(edit_similarity("a", "a b c d"), 1 / 4)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- phi
|
||||
|
||||
|
||||
def test_phi_default_is_edit_distance():
|
||||
# 论文 §5.1 默认;"a b" vs "b a" 恰能区分两度量(edit=0, rouge1=1)
|
||||
assert phi("a b", "b a") == edit_similarity("a b", "b a") == 0.0
|
||||
|
||||
|
||||
def test_phi_dispatch():
|
||||
h, r = "a b c", "a b d"
|
||||
assert phi(h, r, metric="rouge1") == rouge1(h, r)
|
||||
assert phi(h, r, metric="edit_distance") == edit_similarity(h, r)
|
||||
|
||||
|
||||
def test_phi_unknown_metric_raises():
|
||||
with pytest.raises(ValueError, match="bleu"):
|
||||
phi("a", "a", metric="bleu")
|
||||
|
||||
|
||||
# ----------------------------------------------------- aggregate_similarity
|
||||
|
||||
|
||||
def test_aggregate_is_sum_of_phi():
|
||||
# 式(3) 手算:rollouts 与 "a b c" 的 edit 相似度分别为 1.0, 2/3, 0.0
|
||||
chunk = "a b c"
|
||||
rollouts = ["a b c", "a x c", "x y z"]
|
||||
expected = 1.0 + 2 / 3 + 0.0
|
||||
assert math.isclose(aggregate_similarity(chunk, rollouts), expected)
|
||||
|
||||
|
||||
def test_aggregate_bounds():
|
||||
# k_sem ∈ [0, N]:全同 → N,全不同 → 0
|
||||
n = 5
|
||||
assert aggregate_similarity("a b", ["a b"] * n) == float(n)
|
||||
assert aggregate_similarity("a b", ["x y"] * n) == 0.0
|
||||
|
||||
|
||||
def test_aggregate_is_continuous_soft_count():
|
||||
# φ 连续 ⇒ k_sem 非整数是常态(区别于 token 精确匹配的硬计数)
|
||||
k = aggregate_similarity("a b c", ["a b c", "a x c"])
|
||||
assert 1.0 < k < 2.0
|
||||
|
||||
|
||||
def test_aggregate_empty_rollouts():
|
||||
assert aggregate_similarity("a b", []) == 0.0
|
||||
|
||||
|
||||
def test_aggregate_metric_passthrough():
|
||||
# "a b" vs "b a":edit 全零,rouge1 全满——验证 metric 真的传下去了
|
||||
chunk, rollouts = "a b", ["b a", "b a"]
|
||||
assert aggregate_similarity(chunk, rollouts, metric="edit_distance") == 0.0
|
||||
assert aggregate_similarity(chunk, rollouts, metric="rouge1") == 2.0
|
||||
@@ -0,0 +1,143 @@
|
||||
"""层 1 / T2:teacher 批量生成与缓存单测。
|
||||
|
||||
用假 OpenAI 客户端注入(TeacherClient 的测试口),验证思考段剥离、缓存契约
|
||||
(与 data.attach_teacher_completions 的端到端闭环)、断点续传、失败汇总。
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from datasets import Dataset
|
||||
|
||||
from ars_opd.configs import TeacherGenConfig
|
||||
from ars_opd.data import attach_teacher_completions, prompt_key
|
||||
from ars_opd.teacher import TeacherClient, _load_teacher_env, generate_completions
|
||||
|
||||
|
||||
class FakeClient:
|
||||
"""最小 OpenAI 客户端替身:chat.completions.create 按 responder 出内容。"""
|
||||
|
||||
def __init__(self, responder):
|
||||
self.calls = []
|
||||
self._responder = responder
|
||||
self.chat = SimpleNamespace(completions=SimpleNamespace(create=self._create))
|
||||
|
||||
def _create(self, model, messages, **kwargs):
|
||||
self.calls.append(messages)
|
||||
content = self._responder(messages)
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=SimpleNamespace(content=content))]
|
||||
)
|
||||
|
||||
|
||||
def make_teacher(responder, **cfg_overrides):
|
||||
cfg = TeacherGenConfig(**cfg_overrides)
|
||||
return TeacherClient(cfg, client=FakeClient(responder), model="fake-m3")
|
||||
|
||||
|
||||
def user(q):
|
||||
return [{"role": "user", "content": q}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TeacherClient.generate
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_剥离开头思考段():
|
||||
teacher = make_teacher(lambda m: "<think>心算一下</think>\n答案是 42")
|
||||
assert teacher.generate(user("q")) == "答案是 42"
|
||||
|
||||
|
||||
def test_正文中的think字样不误删():
|
||||
teacher = make_teacher(lambda m: "<think>x</think>正文提到 <think> 标签本身")
|
||||
assert teacher.generate(user("q")) == "正文提到 <think> 标签本身"
|
||||
|
||||
|
||||
def test_只剩思考段等于空解答_报错():
|
||||
teacher = make_teacher(lambda m: "<think>思考被截断在半途")
|
||||
# 未闭合的 think 段剥不掉,但闭合后为空的要报错
|
||||
teacher_empty = make_teacher(lambda m: "<think>只有思考</think> ")
|
||||
with pytest.raises(ValueError, match="空解答"):
|
||||
teacher_empty.generate(user("q"))
|
||||
# 未闭合时保留原文(宁可保留可疑内容也不静默删成空)
|
||||
assert "<think>" in teacher.generate(user("q"))
|
||||
|
||||
|
||||
def test_关闭strip_think则原样保留():
|
||||
teacher = make_teacher(lambda m: "<think>a</think>b", strip_think=False)
|
||||
assert teacher.generate(user("q")) == "<think>a</think>b"
|
||||
|
||||
|
||||
def test_system_prompt前置():
|
||||
teacher = make_teacher(lambda m: "ok", system_prompt="你是数学助教")
|
||||
teacher.generate(user("q"))
|
||||
sent = teacher.client.calls[0]
|
||||
assert sent[0] == {"role": "system", "content": "你是数学助教"}
|
||||
assert sent[1]["role"] == "user"
|
||||
|
||||
|
||||
def test_注入client但不给model报错():
|
||||
with pytest.raises(ValueError, match="model"):
|
||||
TeacherClient(TeacherGenConfig(), client=FakeClient(lambda m: "x"), model=None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# generate_completions:缓存契约与断点续传
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_端到端契约_生成的缓存能被attach消费(tmp_path):
|
||||
cache = str(tmp_path / "cache.jsonl")
|
||||
prompts = [user("1+1=?"), user("2+2=?")]
|
||||
teacher = make_teacher(lambda m: f"对「{m[-1]['content']}」的解答")
|
||||
|
||||
generate_completions(prompts, cache, teacher)
|
||||
|
||||
ds = Dataset.from_list([{"messages": p} for p in prompts])
|
||||
out = attach_teacher_completions(ds, cache)
|
||||
assert out[0]["messages"][-1]["content"] == "对「1+1=?」的解答"
|
||||
assert out[1]["messages"][-1]["content"] == "对「2+2=?」的解答"
|
||||
|
||||
|
||||
def test_断点续传_已缓存的不重新生成(tmp_path):
|
||||
cache = str(tmp_path / "cache.jsonl")
|
||||
prompts = [user("q1"), user("q2")]
|
||||
teacher = make_teacher(lambda m: "a")
|
||||
|
||||
generate_completions([prompts[0]], cache, teacher)
|
||||
assert len(teacher.client.calls) == 1
|
||||
generate_completions(prompts, cache, teacher) # q1 命中缓存
|
||||
assert len(teacher.client.calls) == 2 # 只多了 q2 一次调用
|
||||
|
||||
|
||||
def test_单条失败_其余落盘_结束时汇总报错(tmp_path):
|
||||
cache = str(tmp_path / "cache.jsonl")
|
||||
prompts = [user("好题"), user("坏题")]
|
||||
|
||||
def responder(m):
|
||||
if m[-1]["content"] == "坏题":
|
||||
raise RuntimeError("网关 500")
|
||||
return "解答"
|
||||
|
||||
teacher = make_teacher(responder)
|
||||
with pytest.raises(RuntimeError, match="1/2"):
|
||||
generate_completions(prompts, cache, teacher)
|
||||
|
||||
# 成功的那条已经在缓存里,重跑只会补坏题
|
||||
from ars_opd.teacher import _cached_keys
|
||||
from pathlib import Path
|
||||
|
||||
assert _cached_keys(Path(cache)) == {prompt_key(prompts[0])}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# .env 读取
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_env缺失显式报错(monkeypatch):
|
||||
for name in ("TEACHER_API_BASE", "TEACHER_API_KEY", "TEACHER_MODEL"):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
with pytest.raises(ValueError, match="TEACHER_API_BASE"):
|
||||
_load_teacher_env(env_file="/不存在的路径/.env")
|
||||
@@ -0,0 +1,118 @@
|
||||
"""层 1 / T4:掩码 SFT 损失单测(docs/02 §2.4 的切片几何与重掩码)。
|
||||
|
||||
只测纯张量函数 sft_loss / compute_prompt_length;SFTTrainer 类是 HF 接线,
|
||||
由远程 sanity run 验证。张量全部手工构造,期望值可手算。
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ars_opd.data import IGNORE_INDEX
|
||||
from ars_opd.trainer import compute_prompt_length, sft_loss
|
||||
|
||||
V = 7 # 玩具词表大小
|
||||
|
||||
|
||||
def onehot_logits(target_ids, scale=10.0):
|
||||
"""构造在 target_ids 处放尖峰的 logits。(T,) -> (T, V)"""
|
||||
t = torch.tensor(target_ids)
|
||||
return F.one_hot(t, V).float() * scale
|
||||
|
||||
|
||||
def batch_of_one(input_ids, labels, attention=None):
|
||||
"""单行 batch 的三件套,logits 另配。"""
|
||||
ids = torch.tensor([input_ids])
|
||||
lab = torch.tensor([labels])
|
||||
att = torch.ones_like(ids) if attention is None else torch.tensor([attention])
|
||||
return ids, lab, att
|
||||
|
||||
|
||||
def test_prompt_length_取batch最小且不数padding():
|
||||
# pad p p c c c p p p c c c
|
||||
labels = torch.tensor(
|
||||
[
|
||||
[IGNORE_INDEX] * 3 + [5, 6, 5],
|
||||
[IGNORE_INDEX] * 3 + [6, 5, 6],
|
||||
]
|
||||
)
|
||||
attention = torch.tensor([[0, 1, 1, 1, 1, 1], [1, 1, 1, 1, 1, 1]])
|
||||
# 行 0 有效长 5、completion 3 → prompt 2;行 1 是 6-3=3;batch 最小 = 2
|
||||
assert compute_prompt_length(attention, labels) == 2
|
||||
|
||||
|
||||
def test_损失与手算交叉熵一致():
|
||||
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX, IGNORE_INDEX, 5, 6])
|
||||
torch.manual_seed(0)
|
||||
logits = torch.randn(1, 4, V)
|
||||
|
||||
loss, num = sft_loss(logits, ids, lab, att)
|
||||
|
||||
# pl=2:位置 1、2 的 logit 分别预测位置 2、3 的 token(5 和 6)
|
||||
expected = F.cross_entropy(logits[0, 1:3], torch.tensor([5, 6]))
|
||||
assert torch.allclose(loss, expected)
|
||||
assert num == 2
|
||||
|
||||
|
||||
def test_移位对齐_预测下一个token而非当前():
|
||||
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX, IGNORE_INDEX, 5, 6])
|
||||
|
||||
# 位置 t 的尖峰指向位置 t+1 的 token(正确的"预测下一个")→ loss ≈ 0
|
||||
next_logits = onehot_logits([4, 5, 6, 0]).unsqueeze(0)
|
||||
loss_next, _ = sft_loss(next_logits, ids, lab, att)
|
||||
# 位置 t 的尖峰指向位置 t 自己的 token(错误的"复读当前")→ loss 大
|
||||
self_logits = onehot_logits([3, 4, 5, 6]).unsqueeze(0)
|
||||
loss_self, _ = sft_loss(self_logits, ids, lab, att)
|
||||
|
||||
assert loss_next.item() < 0.01
|
||||
assert loss_self.item() > 5.0
|
||||
|
||||
|
||||
def test_窗口外的logits不影响损失():
|
||||
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX, IGNORE_INDEX, 5, 6])
|
||||
torch.manual_seed(0)
|
||||
logits = torch.randn(1, 4, V)
|
||||
loss_base, _ = sft_loss(logits, ids, lab, att)
|
||||
|
||||
# pl=2 → 用到的窗口是位置 [1, 3);位置 0(prompt 内部)和 3(末位)不参与
|
||||
perturbed = logits.clone()
|
||||
perturbed[0, 0] += 100.0
|
||||
perturbed[0, 3] -= 100.0
|
||||
loss_pert, _ = sft_loss(perturbed, ids, lab, att)
|
||||
assert torch.allclose(loss_base, loss_pert)
|
||||
|
||||
|
||||
def test_batch_min切片漏进的prompt_token被重掩码():
|
||||
# 行 0:pad1 + prompt2 + comp3;行 1:prompt3 + comp3 → pl = min(2,3) = 2
|
||||
ids = torch.tensor([[0, 3, 4, 5, 6, 5], [3, 4, 3, 6, 5, 6]])
|
||||
lab = torch.tensor(
|
||||
[
|
||||
[IGNORE_INDEX] * 3 + [5, 6, 5],
|
||||
[IGNORE_INDEX] * 3 + [6, 5, 6],
|
||||
]
|
||||
)
|
||||
att = torch.tensor([[0, 1, 1, 1, 1, 1], [1, 1, 1, 1, 1, 1]])
|
||||
torch.manual_seed(1)
|
||||
logits = torch.randn(2, 6, V)
|
||||
|
||||
loss_base, num = sft_loss(logits, ids, lab, att)
|
||||
assert num == 6 # 两行各 3 个 completion token,漏进切片的 prompt 位不计数
|
||||
|
||||
# 位置 1 的 logit 预测位置 2——两行的位置 2 都在切片内但都是 -100(行 0 是
|
||||
# prompt 尾、行 1 是漏进来的 prompt token)。改它不该动 loss
|
||||
perturbed = logits.clone()
|
||||
perturbed[:, 1] += 100.0
|
||||
loss_pert, _ = sft_loss(perturbed, ids, lab, att)
|
||||
assert torch.allclose(loss_base, loss_pert)
|
||||
|
||||
|
||||
def test_全掩码batch显式报错():
|
||||
ids, lab, att = batch_of_one([3, 4, 5, 6], [IGNORE_INDEX] * 4)
|
||||
with pytest.raises(ValueError, match="有效 completion"):
|
||||
sft_loss(torch.randn(1, 4, V), ids, lab, att)
|
||||
|
||||
|
||||
def test_无prompt行显式报错():
|
||||
ids, lab, att = batch_of_one([3, 4], [3, 4]) # labels 全有效 → prompt 长 0
|
||||
with pytest.raises(ValueError, match="prompt_length"):
|
||||
sft_loss(torch.randn(1, 2, V), ids, lab, att)
|
||||
Reference in New Issue
Block a user