仓库骨架:README、CLAUDE.md、分层重构路线图(docs/00-roadmap.md)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
+15
@@ -0,0 +1,15 @@
|
|||||||
|
references/
|
||||||
|
__pycache__/
|
||||||
|
*.pyc
|
||||||
|
.venv/
|
||||||
|
checkpoints/
|
||||||
|
wandb/
|
||||||
|
outputs/
|
||||||
|
slurm/
|
||||||
|
.env
|
||||||
|
*.egg-info/
|
||||||
|
data/
|
||||||
|
logs/
|
||||||
|
*.log
|
||||||
|
.pytest_cache/
|
||||||
|
.ruff_cache/
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
# CLAUDE.md
|
||||||
|
|
||||||
|
> [!URGENT]
|
||||||
|
> 1. 本项目是**学习驱动的科研重构项目**:通过分层重建 OmniOPD 参考实现来学习论文方法,同时产出可作为后续研究基座的干净代码库。学习优先——用户要逐章理解每个模块,**不要一次性替用户写完所有代码**;每章先讲清楚,重构任务由用户主导、你配合。
|
||||||
|
> 2. 所有思考过程和回复必须使用**简体中文**。
|
||||||
|
|
||||||
|
## 1. 项目元数据
|
||||||
|
|
||||||
|
- **论文**: OmniOPD (arXiv:2606.01476v2),logit-free 的 chunk 级 on-policy 蒸馏。PDF 在 `references/2606.01476v2.pdf`
|
||||||
|
- **参考实现**: `references/ars-opd/`(**只读、不入 git**),核心在 `trl/trl/experimental/distillation/`
|
||||||
|
- **Conda 环境**: `ars-opd`(Python 3.11);远程环境建在 `/data` 下(根分区已满)
|
||||||
|
- **版本管理**: git,远程 `git@gitea.iomgaa.online:iomgaa/ars-opd-rebuild.git`,本地 push → 远程 pull 同步
|
||||||
|
|
||||||
|
## 2. 架构哲学(Ousterhout 深模块,非 Clean Architecture)
|
||||||
|
|
||||||
|
- **模块按论文概念划分,不按技术角色划分**。判据:看到论文某公式能说出它在哪个文件;打开一个文件能说出它对应论文哪一节。
|
||||||
|
- **纯逻辑核心 / IO 边缘**:`similarity.py`、`chunking.py`、`estimator.py` 是纯逻辑模块——只依赖 torch/numpy 标准运算,**禁止** import transformers / vllm / openai / requests;所有外部 IO 收口在 `teacher.py`(API/vLLM 客户端)与 `trainer.py`(模型编排)。纯逻辑模块必须能在本地 CPU 上用 toy 数据测试。
|
||||||
|
- **不做防御性膨胀**:不写接口层/工厂/注册表;一个模块从上读到下能看懂全部流程。但纯逻辑模块必须有单元测试(对拍参考实现)。
|
||||||
|
- **配置显式化**:所有实验参数走 `configs.py` 的 dataclass;**严禁**硬编码路径、模型名、URL(参考实现里的 `/fsx` 硬编码是反面教材)。密钥走 `.env`。
|
||||||
|
|
||||||
|
## 3. 模块 ↔ 论文映射(单一事实源)
|
||||||
|
|
||||||
|
| 模块 | 论文 | 内容 |
|
||||||
|
|------|------|------|
|
||||||
|
| `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/trainer.py` | §3.2.4 式(8) | chunk 蒸馏损失 + trust-region KL 锚定,编排全流程 |
|
||||||
|
|
||||||
|
## 4. 常用命令
|
||||||
|
|
||||||
|
```bash
|
||||||
|
conda activate ars-opd && pytest tests/ -x -q # 单元测试(本地 CPU)
|
||||||
|
conda activate ars-opd && ruff check ars_opd/ --fix && ruff format ars_opd/
|
||||||
|
```
|
||||||
|
|
||||||
|
## 5. 本地-远程规则
|
||||||
|
|
||||||
|
> [!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`),确保可实时检查。
|
||||||
|
> - 远程机器上不改代码,只 `git pull` + 跑脚本;本地不跑训练。
|
||||||
|
|
||||||
|
## 6. 学习工作流
|
||||||
|
|
||||||
|
1. 章节文档在 `docs/`,按 `docs/00-roadmap.md` 的分层顺序推进;每章结构:论文精读 → 参考实现解剖(带行号)→ 重构任务 → 验证方式。
|
||||||
|
2. 重构一个模块前,先对照参考实现列出其全部行为(含 trick 和 workaround),逐一确认保留/替代/删除。
|
||||||
|
3. 每个纯逻辑模块完成后,用 toy 数据对拍参考实现的对应逻辑(参考其 `validate_mc_estimator.py` / `validate_chunk_mc_estimator.py`)。
|
||||||
|
4. 文档规范:优先表格与公式,代码块 ≤15 行(展示思路用伪代码,完整代码引用文件路径),引用参考实现必须带 `文件:行号`。
|
||||||
|
|
||||||
|
## 7. 代码规范
|
||||||
|
|
||||||
|
- 公共函数完整类型注解;模块/类/函数写中文 docstring(功能、参数、返回、关键实现细节)。
|
||||||
|
- **严禁** `except Exception: pass`;出错直接报错,不用默认值兜底。
|
||||||
|
- 张量函数在 docstring 中注明各参数的 shape。
|
||||||
|
- 提交信息用中文,说明"这一步对应哪一章/哪个模块"。
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
# ars-opd-rebuild
|
||||||
|
|
||||||
|
对论文 **OmniOPD: Logit-Free On-Policy Distillation via Speculative Verification**(arXiv:2606.01476)参考实现的分层重构。
|
||||||
|
|
||||||
|
## 项目双重目标
|
||||||
|
|
||||||
|
1. **学习**:通过按论文概念逐层重建代码,吃透 OmniOPD 的方法(chunk 级 MC 估计、peak-entropy 调度、贝叶斯平滑、trust-region 锚定)。
|
||||||
|
2. **研究基座**:产出一个结构清晰、模块化、可测试的代码库,作为后续研究的起点——替代原版 3200 行自包含 trainer 的组织方式。
|
||||||
|
|
||||||
|
原始参考实现位于 `references/ars-opd/`(只读、不入 git),核心是其 trl fork 中的 `trl/experimental/distillation/`。
|
||||||
|
|
||||||
|
## 方法一句话
|
||||||
|
|
||||||
|
学生模型在线生成推理轨迹;在轨迹的高熵"推理分叉点"选取 M 个 C-token 的 chunk,向黑盒 teacher 按前缀采样 N 条 rollout,用语义相似度(ROUGE-1)做 Monte Carlo 估计并经 Dirichlet 先验平滑,得到有界的 chunk 级监督权重;未被审计的 token 用对冻结初始模型的 KL 锚定防漂移(论文式 8)。全程不需要 teacher 的 logits。
|
||||||
|
|
||||||
|
## 目录结构
|
||||||
|
|
||||||
|
```
|
||||||
|
ars_opd/ 核心包(模块 ↔ 论文概念一一对应)
|
||||||
|
configs.py dataclass 配置
|
||||||
|
similarity.py 语义相似度 φ:rouge1 / edit_distance(式 3) [纯逻辑]
|
||||||
|
chunking.py peak-entropy chunk 调度器(式 6-7) [纯逻辑]
|
||||||
|
estimator.py MC 估计 + Dirichlet 贝叶斯平滑(式 4-5) [纯逻辑]
|
||||||
|
teacher.py teacher 客户端(OpenAI 兼容 API / vLLM) [IO 边缘]
|
||||||
|
trainer.py 薄编排层:生成 → 审计 → 损失(式 8)
|
||||||
|
scripts/ 训练/评测入口脚本
|
||||||
|
tests/ 纯逻辑模块的单元测试(本地 CPU 可跑)
|
||||||
|
docs/ 循序渐进的学习章节(论文↔代码映射)
|
||||||
|
references/ 原论文 PDF + 参考实现(gitignore,只读对照)
|
||||||
|
```
|
||||||
|
|
||||||
|
## 规模选型(跑通优先,不追论文数值)
|
||||||
|
|
||||||
|
| 角色 | 模型 | 部署 |
|
||||||
|
|------|------|------|
|
||||||
|
| Student | Qwen3-0.6B | 远程 4×A800 训练 |
|
||||||
|
| API teacher(OmniOPD 主路径) | DeepSeek / MiniMax | OpenAI 兼容 API |
|
||||||
|
| White-box teacher(baseline 对照) | Qwen3-4B | 远程本地 vLLM |
|
||||||
|
|
||||||
|
## 工作流
|
||||||
|
|
||||||
|
- **本地**(这台机器):读论文、写代码、跑 `tests/` 单元测试(CPU toy 数据对拍原实现)。
|
||||||
|
- **远程**(gpu-a800-060,4×A800-80G):只跑训练与评测。代码经 gitea 同步:本地 push → 远程 pull → 启动脚本。
|
||||||
|
- 学习按 `docs/` 章节推进:每章 = 论文一节精读 + 原实现解剖 + 重构任务 + 验证方式。路线图见 `docs/00-roadmap.md`。
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
# 00 · 分层重构路线图
|
||||||
|
|
||||||
|
> 原则:按论文概念的依赖顺序逐层重建,每层完成后代码可运行、可验证。学习路径 = 提交历史。
|
||||||
|
|
||||||
|
## 分层计划
|
||||||
|
|
||||||
|
| 层 | 主题 | 论文对应 | 产出 | 验证方式 |
|
||||||
|
|----|------|----------|------|----------|
|
||||||
|
| 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` |
|
||||||
|
| 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 有可测提升趋势 |
|
||||||
|
|
||||||
|
## 章节文档索引
|
||||||
|
|
||||||
|
| 文档 | 内容 | 状态 |
|
||||||
|
|------|------|------|
|
||||||
|
| `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 估计 + 贝叶斯平滑 | ⬜ |
|
||||||
|
| `05-entropy-chunking.md` | 层 4:熵调度 | ⬜ |
|
||||||
|
| `06-omniopd-full.md` | 层 5:完整损失与 teacher 客户端 | ⬜ |
|
||||||
|
| `07-eval-ablation.md` | 层 6:评测与消融 | ⬜ |
|
||||||
|
|
||||||
|
## 关键设定(与论文默认对齐,规模缩小)
|
||||||
|
|
||||||
|
| 参数 | 论文默认 | 本项目 | 说明 |
|
||||||
|
|------|----------|--------|------|
|
||||||
|
| chunk 数 M | 10 | 10 | 每条轨迹审计的 chunk 数 |
|
||||||
|
| 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 主路径 |
|
||||||
Reference in New Issue
Block a user