Files
ars-opd-rebuild/docs/04-mc-estimator.md
T
iomgaa 142aeb8ab1 docs: 层 3 章节文档 docs/04(语义相似度 φ + MC 估计器,式3/4/5)
论文精读(§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) <noreply@anthropic.com>
2026-07-19 09:40:58 -04:00

155 lines
13 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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 喂给 teacherteacher 生成 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=0teacher 全否定) | 裸频率估计 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 级 π̄ / π̂ / detachDT: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 项
```
两处 detach2196 的 `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)·Nbayes = (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 | 10CFG:277 | 10(§4.2 甜点) |
| `chunk_alpha` | α | 1.0CFG:281 | 1.0 |
| `chunk_length` | C | 50CFG:273 附近) | 50 |
| `chunk_similarity` | φ | **rouge1**CFG:299 | **edit_distance**(§5.1)——⚠️ 背离,须显式指定 |
| `no_bayesian` | — | FalseCFG: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_distancek_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.pytorchtoy 张量可测)
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` 已接真实现;接口回看完成(两模块判"深")。