From 142aeb8ab11c500b775385b753be6057f2695ff6 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sun, 19 Jul 2026 09:40:58 -0400 Subject: [PATCH] =?UTF-8?q?docs:=20=E5=B1=82=203=20=E7=AB=A0=E8=8A=82?= =?UTF-8?q?=E6=96=87=E6=A1=A3=20docs/04=EF=BC=88=E8=AF=AD=E4=B9=89?= =?UTF-8?q?=E7=9B=B8=E4=BC=BC=E5=BA=A6=20=CF=86=20+=20MC=20=E4=BC=B0?= =?UTF-8?q?=E8=AE=A1=E5=99=A8=EF=BC=8C=E5=BC=8F3/4/5=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 论文精读(§3.2.1-3.2.2 式3/4/5)+ 参考实现解剖(带行号)+ 保留/替代/删除 + E1-E3 重构任务 + 验证方式。核心:logit-free 支点(比文本非比 logits)、Dirichlet 贝叶斯 π̂ 反塌缩、detach 命门(层 2 §4.1 姊妹篇)。解剖抓出参考实现三处不一致(φ 比文本 vs 比 token id、rouge1 集合 vs 多重集)作重构清理依据。roadmap 存档点/索引同步。 Co-Authored-By: Claude Opus 4.8 (1M context) --- docs/00-roadmap.md | 4 +- docs/04-mc-estimator.md | 154 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 156 insertions(+), 2 deletions(-) create mode 100644 docs/04-mc-estimator.md diff --git a/docs/00-roadmap.md b/docs/00-roadmap.md index cdb6e33..b042391 100644 --- a/docs/00-roadmap.md +++ b/docs/00-roadmap.md @@ -8,7 +8,7 @@ > 每次断点(层完成/工作暂停)更新此节。恢复上下文时:读 CLAUDE.md → 本节 → 对应章节文档。 - **日期**: 2026-07-19 -- **当前层**: 层 2(white-box OPD)**✅ 已关账**;下一步进层 3(相似度 + MC 估计器) +- **当前层**: 层 3(相似度 φ + MC 估计器),**学习阶段**——docs/04 已写待读;E1-E3 待精读后动代码。层 3 全本地 CPU 纯逻辑,不碰 GPU/远程。E1 similarity.py(式3 φ+k_sem)、E2 estimator.py(式4 π̄ detach + 式5 π̂)、E3 test_estimator_detach.py 接真实现。层 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 约束把你锁在温和区间) @@ -41,7 +41,7 @@ | `01-paper-code-map.md` | 论文 §3-§4 精读 + 参考实现全景解剖 + 概念↔代码对照表 | ✅ | | `02-sft-baseline.md` | 层 1:SFT 与数据管线 | ✅ 待读 | | `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性(含 §6.1 远程实证) | ✅ 已关账 | -| `04-mc-estimator.md` | 层 3:MC 估计 + 贝叶斯平滑 | ⬜ | +| `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:评测与消融 | ⬜ | diff --git a/docs/04-mc-estimator.md b/docs/04-mc-estimator.md new file mode 100644 index 0000000..9f9ade0 --- /dev/null +++ b/docs/04-mc-estimator.md @@ -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_ 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` 已接真实现;接口回看完成(两模块判"深")。