Files
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

13 KiB
Raw Permalink Blame History

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.pyCFG: = .../distillation_config.pyVAL: = 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 idDT: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)——正主

# 式(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)——论文说"不可行"的朴素版,作对照

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 truthVAL: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 φ rouge1CFG: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 已接真实现;接口回看完成(两模块判"深")。