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