142aeb8ab1
论文精读(§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>
155 lines
13 KiB
Markdown
155 lines
13 KiB
Markdown
# 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` 已接真实现;接口回看完成(两模块判"深")。
|