diff --git a/ars_opd/similarity.py b/ars_opd/similarity.py new file mode 100644 index 0000000..d43efef --- /dev/null +++ b/ars_opd/similarity.py @@ -0,0 +1,131 @@ +"""语义相似度 φ 与 chunk 级聚合 k_sem——论文 §3.2.1 式(3)。 + +logit-free 的支点:学生 chunk 对不对,不再比 token 概率(层 2 白盒式(2)), +改比"学生 chunk 文本" vs "teacher rollout 文本"的语义相似度。φ 只依赖 +文本本身,与两侧 tokenizer 无关,teacher 只需能吐文本(任何 API 均可)。 + +纯逻辑模块(CLAUDE.md §2):只依赖标准库,可脱离 torch 在本地 CPU 测试。 +teacher rollout 怎么采出来是层 5 teacher.py 的事,本模块只吃现成字符串。 +""" + +from collections import Counter + + +def rouge1(hypothesis: str, reference: str) -> float: + """ROUGE-1 F1(unigram 重叠率),φ 的候选度量之一(论文 §3.2.1)。 + + 以词为单位(空白切分)统计两串的 unigram 重叠,算 F1。 + 词袋语义:只看"用了哪些词",不看词序——"a b" vs "b a" 得 1.0。 + + 参数: + hypothesis: 学生 chunk 文本。 + reference: teacher rollout 文本。 + + 返回: + F1 ∈ [0, 1];任一侧无词(空串/纯空白)时为 0.0。 + + 实现细节: + - 差异标注:参考实现(distillation_trainer.py:1670)用 set 去重后求交, + 会把 "x x x x" vs "x" 判成满分 1.0;此处用 Counter 多重集 + (ROUGE-1 标准定义),重复词按 min 计数配对,同例只得 0.4。 + 数学推理文本里重复 token(数字、"="、变量名)极常见,去重会失真。 + - 差异标注:参考实现分母加 1e-8 防零除,代价是全同串 F1≈0.99999998 + 而非精确 1;此处 overlap==0 时提前返回,分母恒正,无需平滑。 + """ + hyp_counts = Counter(hypothesis.split()) + ref_counts = Counter(reference.split()) + if not hyp_counts or not ref_counts: + return 0.0 + # 多重集交:每个词按两侧出现次数的 min 配对 + overlap = sum((hyp_counts & ref_counts).values()) + if overlap == 0: + return 0.0 + precision = overlap / sum(hyp_counts.values()) + recall = overlap / sum(ref_counts.values()) + return 2 * precision * recall / (precision + recall) + + +def edit_similarity(hypothesis: str, reference: str) -> float: + """归一化编辑相似度 1 − Levenshtein/max(m,n),论文 §5.1 的默认 φ。 + + 以词为单位(空白切分)算 Levenshtein 距离(插入/删除/替换各计 1), + 再归一化到 [0, 1] 取反。顺序敏感:"a b" vs "b a" 距离 2,相似度 0—— + 与 rouge1 的词袋语义形成互补。 + + 参数: + hypothesis: 学生 chunk 文本。 + reference: teacher rollout 文本。 + + 返回: + 相似度 ∈ [0, 1];两侧均空为 1.0(零距离),仅一侧空为 0.0(全删/全插)。 + + 实现细节: + - 差异标注:参考实现(distillation_trainer.py:1682)吃 token id 列表, + 相似度随 tokenizer 切法漂移,违背本层"文本是公共语言"的初衷 + (docs/04 §2.1 坑①);此处吃 str、内部按词切,与 rouge1 统一口径。 + - 两行滚动 DP(同参考实现):空间 O(n) 而非 O(m·n)。 + """ + hyp_words = hypothesis.split() + ref_words = reference.split() + m, n = len(hyp_words), len(ref_words) + if m == 0 and n == 0: + return 1.0 + if m == 0 or n == 0: + return 0.0 + # prev[j] = 前一行的 dist(hyp[:i-1], ref[:j]);curr 原地滚动复用 + prev = list(range(n + 1)) + curr = [0] * (n + 1) + for i in range(1, m + 1): + curr[0] = i + for j in range(1, n + 1): + cost = 0 if hyp_words[i - 1] == ref_words[j - 1] else 1 + curr[j] = min( + prev[j] + 1, # 删除 hyp[i-1] + curr[j - 1] + 1, # 插入 ref[j-1] + prev[j - 1] + cost, # 替换(相同则免费) + ) + prev, curr = curr, prev + return 1.0 - prev[n] / max(m, n) + + +def phi(hypothesis: str, reference: str, metric: str = "edit_distance") -> float: + """语义相似度 φ(y_c, ŷ_c) ∈ [0, 1],论文 §3.2.1 式(3) 的原子度量。 + + 参数: + hypothesis: 学生 chunk 文本。 + reference: teacher rollout 文本。 + metric: "edit_distance"(默认)或 "rouge1"。 + 差异标注:参考实现配置默认 rouge1(config.py:299),与论文 §5.1 + 的 edit_distance 背离;此处从论文。 + + 返回: + 相似度 ∈ [0, 1]。 + """ + if metric == "edit_distance": + return edit_similarity(hypothesis, reference) + if metric == "rouge1": + return rouge1(hypothesis, reference) + raise ValueError(f"未知相似度度量: {metric!r}(可选 'edit_distance' / 'rouge1')") + + +def aggregate_similarity( + student_chunk: str, + teacher_rollouts: list[str], + metric: str = "edit_distance", +) -> float: + """式(3):k_sem = Σ_{i=1}^{N} φ(y_c, ŷ_c^{(i)}),chunk 的语义匹配计数。 + + 学生 chunk 与 N 个 teacher rollout 逐一算 φ 后求和。φ 连续,故 k_sem 是 + [0, N] 上的实数——"软计数":k_sem≈N 意为学生这段与 teacher 高度一致, + k_sem≈0 意为 teacher 从不这么写。它是 π̂(式5)里唯一的外部 teacher 信号。 + + 参数: + student_chunk: 学生 chunk 文本(C 个 token 解码所得)。 + teacher_rollouts: N 段 teacher 续写文本,与学生 chunk 共享同一前缀 + y_