论文精读(§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>
13 KiB
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) |
⚠️ 三处不一致(都是重构要统一的):
- rouge1 比文本、edit 比 token id(DT:1850 vs 1846)。edit 用 token id 重新耦合了 tokenizer——直接违背 §3.2.1 "跨 tokenizer" 的立身之本;只在 teacher/student 同 tokenizer 的 vLLM 路径侥幸能跑。
- rouge1 两种算法:trainer 用集合(DT:1672),validate 脚本用多重集计数
min(ref_cnt, hyp_cnt)(VAL:63)——同名不同义。 - 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)——正主
# 式(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)——论文说"不可行"的朴素版,作对照
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. 验证方式
- similarity.py 单测:手构字符串断言 φ 值(如全同 chunk→edit=1、rouge1=1;不相交→0;部分重叠手算);k_sem = Σφ 的连续性;空串边界。
- 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 时逃逸)。
- 关账判据:similarity/estimator 全单测本地 CPU 通过;
test_estimator_detach.py已接真实现;接口回看完成(两模块判"深")。