Compare commits
10 Commits
0ca60ea93f
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 630a9c4636 | |||
| e06e5ed7a2 | |||
| e7cf27cecc | |||
| 142aeb8ab1 | |||
| fac8e0dcc5 | |||
| 8b362eae09 | |||
| f1b6d1f668 | |||
| 2f4780b69f | |||
| b7f24d635e | |||
| 404abc22bf |
@@ -253,6 +253,14 @@ class DistillConfig:
|
||||
lr_scheduler_type: str = "linear"
|
||||
warmup_ratio: float = 0.0
|
||||
|
||||
max_grad_norm: float = 1.0
|
||||
"""梯度裁剪阈值。此前是 HF Trainer 的静默默认(1.0),现显式化——它是式(2)
|
||||
反向 KL 梯度爆炸(§4.1)的**隐形稳定器**:on-policy 采到 teacher 眼中烂 token
|
||||
时单步梯度范数可炸到十几(2026-07-19 首冒烟实测 grad_norm 14→2),HF 默认
|
||||
裁到 1.0 才让 loss 曲线平稳。把它设得远大于实测范数(≈关闭裁剪)可暴露原始
|
||||
爆炸,供教学对照(train_whitebox.py 的 noclip 模式)。非显然约束:日志里的
|
||||
grad_norm 是**裁剪前**范数,故 14→2 那串本身就是爆炸证据,只是被裁剪掩盖了。"""
|
||||
|
||||
gradient_checkpointing: bool = False
|
||||
"""默认不开(§5 显存账 B=4 富余);OOM 时作为降 batch 之后的第二道降显存手段。
|
||||
注意 FSDP 下此开关是 no-op(docs/02 §2.6),但层 2 坚持 DDP 故此处有效。"""
|
||||
@@ -289,6 +297,9 @@ class DistillConfig:
|
||||
raise ValueError(f"max_new_tokens 必须为正,收到 {self.max_new_tokens}")
|
||||
if self.learning_rate <= 0:
|
||||
raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}")
|
||||
if self.max_grad_norm <= 0:
|
||||
# 用远大于实测范数的值≈关闭裁剪;≤0 无意义(0 会把梯度裁没)
|
||||
raise ValueError(f"max_grad_norm 必须为正,收到 {self.max_grad_norm}")
|
||||
if self.subset_size is not None and self.subset_size <= 0:
|
||||
raise ValueError(
|
||||
f"subset_size 必须为正整数或 None(全量),收到 {self.subset_size}"
|
||||
|
||||
+23
-13
@@ -18,14 +18,11 @@ import ast
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from datasets import Dataset, load_dataset
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ars_opd.configs import SFTConfig
|
||||
|
||||
# F.cross_entropy 的 ignore_index 默认值;标了它的位置不产生 loss
|
||||
IGNORE_INDEX = -100
|
||||
|
||||
@@ -165,23 +162,36 @@ def attach_teacher_completions(dataset: Dataset, jsonl_path: str) -> Dataset:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def load_sft_dataset(cfg: "SFTConfig") -> Dataset:
|
||||
"""层 1 数据管线入口:加载 → 归一 → 抽子集 → 挂 teacher 解答。
|
||||
def load_sft_dataset(
|
||||
dataset_path: str,
|
||||
dataset_split: str = "train",
|
||||
subset_size: int | None = None,
|
||||
seed: int = 42,
|
||||
teacher_completions_path: str | None = None,
|
||||
) -> Dataset:
|
||||
"""数据管线入口:加载 → 归一 → 抽子集 →(可选)挂 teacher 解答。
|
||||
|
||||
返回只含 ``messages`` 一列的 Dataset,每行末轮是 assistant(可直接喂 SFTCollator)。
|
||||
收散装参数而非整个 config(深模块:本函数只用这 5 个字段,不该索要一整个
|
||||
SFTConfig)。这样层 1(SFTConfig)、层 2(DistillConfig,无 teacher 缓存)、
|
||||
诊断脚本都能直接调,无需伪造无关字段。teacher_completions_path=None 时
|
||||
返回 prompt-only 数据集(末轮 user,供 on-policy 生成);给了则挂 teacher
|
||||
解答(末轮 assistant,供 SFT)。
|
||||
|
||||
返回只含 ``messages`` 一列的 Dataset。
|
||||
"""
|
||||
ds = _load_raw(cfg.dataset_path, cfg.dataset_split)
|
||||
ds = _load_raw(dataset_path, dataset_split)
|
||||
ds = ds.map(
|
||||
to_messages,
|
||||
remove_columns=[c for c in ds.column_names if c != "messages"],
|
||||
)
|
||||
if cfg.subset_size is not None and cfg.subset_size < len(ds):
|
||||
if subset_size is not None and subset_size < len(ds):
|
||||
# 非显然约束:抽子集必须在挂接 teacher 解答之前、且由 seed 完全确定——
|
||||
# teacher.py 生成缓存时会走完全相同的"加载→归一→抽子集"路径,两侧 seed
|
||||
# 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。
|
||||
ds = ds.shuffle(seed=cfg.seed).select(range(cfg.subset_size))
|
||||
if cfg.teacher_completions_path is not None:
|
||||
ds = attach_teacher_completions(ds, cfg.teacher_completions_path)
|
||||
# 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。层 2 与层 1
|
||||
# 用同 seed 同 subset_size,才能在同一批题上对比 SFT 与蒸馏。
|
||||
ds = ds.shuffle(seed=seed).select(range(subset_size))
|
||||
if teacher_completions_path is not None:
|
||||
ds = attach_teacher_completions(ds, teacher_completions_path)
|
||||
return ds
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
"""MC 估计与 Dirichlet 贝叶斯平滑——论文 §3.2.2 式(4)(5)。
|
||||
|
||||
把 similarity.py 产出的软计数 k_sem(外部 teacher 信号)与学生自身的
|
||||
chunk 置信度 π̄(内部先验)融合成有界目标 π̂ ∈ (0, 1],供层 5 的 chunk
|
||||
损失当乘子:loss_c = −π̂ · mean(log p)。定理 4.1 三性质由此获得:
|
||||
(a) π̂ 有界 → 无白盒式(2) 的梯度爆炸;(b) π̂ > 0 → k_sem=0 也不塌缩;
|
||||
(c) 先验收缩 → 方差小于频率估计 k/N。
|
||||
|
||||
纯逻辑模块(CLAUDE.md §2):只依赖 torch,toy 张量本地 CPU 可测。
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def chunk_prior(log_probs: torch.Tensor) -> torch.Tensor:
|
||||
"""式(4):π̄ = exp((1/C)·Σ_t log p_t)——学生对整个 chunk 的几何均值置信度。
|
||||
|
||||
C 个 token 概率的几何均值,充当式(5) 的贝叶斯先验:teacher 采样(k_sem)
|
||||
是主信号,π̄ 只是"学生自己觉得这段有多稳"的地板,防 k_sem=0 时目标归零。
|
||||
|
||||
参数:
|
||||
log_probs: 学生对 chunk 内各 token 的对数概率,shape (C,),值 ≤ 0。
|
||||
(调用方从 log_softmax 后 gather 标签位置所得,层 5 负责。)
|
||||
|
||||
返回:
|
||||
π̄,标量张量 shape ()、值 ∈ [1e-8, 1],**不带梯度**。
|
||||
|
||||
实现细节:
|
||||
- log 域先均值再 exp:直接连乘 C=50 个小概率会下溢
|
||||
(50 个 0.01 → 1e-100,超出 fp32 下限 ~1e-38),log 域安全。
|
||||
- detach 命门(参考实现 distillation_trainer.py:2196 同):π̄ 是学生
|
||||
自身概率的函数,若保留梯度,优化器会发现"压低自己的 chunk 概率
|
||||
→ π̄→0 → π̂ 变小 → 损失权重变小"这条逃逸路径——恰在 k_sem=0
|
||||
(teacher 否定)的 chunk 上最有利可图,这些 chunk 最先塌缩。
|
||||
π̄ 只能当常数先验,不能当优化变量。锁死断言见
|
||||
tests/test_estimator_detach.py。
|
||||
- clamp 下限 1e-8:极端负的均值 exp 后可能下溢为 0,而定理 4.1(b)
|
||||
的反塌缩要求 π̄ 严格为正。差异标注:参考实现 clamp(1e-8, 1.0),
|
||||
上限实为冗余——log p ≤ 0 ⇒ mean ≤ 0 ⇒ exp ≤ 1,此处省去。
|
||||
"""
|
||||
if log_probs.numel() == 0:
|
||||
raise ValueError("log_probs 为空:chunk 至少要含 1 个 token")
|
||||
log_pi_bar = log_probs.detach().mean() # (C,) -> ()
|
||||
return log_pi_bar.exp().clamp(min=1e-8)
|
||||
|
||||
|
||||
def bayesian_target(
|
||||
k_sem: float,
|
||||
pi_bar: torch.Tensor,
|
||||
n_rollouts: int,
|
||||
alpha: float,
|
||||
) -> torch.Tensor:
|
||||
"""式(5):π̂ = (k_sem + α·π̄) / (N + α)——chunk 接受概率的贝叶斯估计。
|
||||
|
||||
等价凸组合视角(论文式10):
|
||||
π̂ = N/(N+α) · (k_sem/N) + α/(N+α) · π̄
|
||||
即"teacher 频率估计"与"学生先验"的加权平均;默认 N=10、α=1 时权重
|
||||
约 91% : 9%,teacher 主导,先验只兜底。
|
||||
|
||||
参数:
|
||||
k_sem: 式(3) 的软匹配计数,∈ [0, N](aggregate_similarity 产出)。
|
||||
pi_bar: 式(4) 的先验 π̄,标量张量(chunk_prior 产出)。
|
||||
n_rollouts: teacher rollout 数 N。**必须等于算 k_sem 时的
|
||||
len(teacher_rollouts)**——分子分母口径不一致会系统性偏移 π̂。
|
||||
alpha: 先验强度 α ≥ 0。α=0 退化为频率估计 k/N(层 6 的
|
||||
no_bayesian 消融,参考 config.py:315),失去定理 4.1(b) 保护。
|
||||
|
||||
返回:
|
||||
π̂,标量张量 shape ()、值 ∈ [1e-8, 1],**不带梯度**。
|
||||
|
||||
实现细节:
|
||||
- 差异标注:参考实现(distillation_trainer.py:2205)在此对 π̂ 整体
|
||||
detach;我们的 π̄ 在 chunk_prior 内已 detach,此处的 detach 是
|
||||
第二道防线——防止将来有人把带梯度的张量传进 pi_bar。
|
||||
- clamp(1e-8, 1.0):下限防 α=0 且 k_sem=0 时 π̂=0(乘子归零则该
|
||||
chunk 完全失去监督);上限防 pi_bar 越界传入时 π̂ 溢出概率语义。
|
||||
"""
|
||||
if n_rollouts < 1:
|
||||
raise ValueError(f"n_rollouts 必须 ≥ 1,得到 {n_rollouts}")
|
||||
if alpha < 0:
|
||||
raise ValueError(f"alpha 必须 ≥ 0,得到 {alpha}")
|
||||
if not 0.0 <= k_sem <= n_rollouts:
|
||||
raise ValueError(
|
||||
f"k_sem={k_sem} 越界 [0, {n_rollouts}]:检查是否与"
|
||||
f" len(teacher_rollouts) 口径一致"
|
||||
)
|
||||
# 式(5): π̂ = (k_sem + α·π̄) / (N + α)
|
||||
pi_hat = (k_sem + alpha * pi_bar) / (n_rollouts + alpha) # () -> ()
|
||||
return pi_hat.clamp(1e-8, 1.0).detach()
|
||||
@@ -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_<c(对应关系由"同一前缀现场生成"保证,无需搜索匹配,docs/04 §1)。
|
||||
metric: 传给 phi,默认 "edit_distance"。
|
||||
|
||||
返回:
|
||||
k_sem ∈ [0, N],N = len(teacher_rollouts)。空列表得 0.0(空和)。
|
||||
"""
|
||||
return sum(phi(student_chunk, r, metric) for r in teacher_rollouts)
|
||||
+9
-6
@@ -7,14 +7,17 @@
|
||||
|
||||
> 每次断点(层完成/工作暂停)更新此节。恢复上下文时:读 CLAUDE.md → 本节 → 对应章节文档。
|
||||
|
||||
- **日期**: 2026-07-18
|
||||
- **当前层**: 层 2(white-box OPD),学习阶段——docs/03 已写,用户精读中;U1-U5 待默认参数确认后开写
|
||||
- **层 2 开写前两件待办**: ① docs/03 §5 默认参数提案待用户确认,其中 per-device batch=4 须先用层 1 实测显存账重审(白盒 = student+teacher 两份全词表 logits);② 用户三道自查题待对答案(尾桶丢什么信息 / off-policy 蒸馏何时值得加回+层 1 缓存的角色 / 低频路径 bug 存活率与单测应压向哪里)
|
||||
- **日期**: 2026-07-19
|
||||
- **当前层**: 层 3(相似度 φ + MC 估计器),**docs/04 已精读**(用户已理解:logit-free 支点=比文本非比 logits、k_sem 是外部 teacher 信号 π̄ 只是防塌缩地板、detach 命门、chunk-vs-prefix 对应靠"同一前缀现场生成 teacher 续写"非搜索匹配);**下一步开写 E1**。层 3 全本地 CPU 纯逻辑,不碰 GPU/远程。E1 similarity.py(式3 φ+k_sem,两度量统一到词级文本、默认 edit_distance 对齐论文§5.1)、E2 estimator.py(式4 π̄ 几何均值+detach、式5 π̂ 贝叶斯凸组合)、E3 test_estimator_detach.py 接真实现。设计取舍已定见 docs/04 §3。层 2 ✅ 已关账
|
||||
- **层 2 代码构成**: U1 DistillConfig(configs.py,两温度分名/三处刻意缺席/max_grad_norm 显式化);U2 token_divergence + 梯度爆炸单测(trainer.py,全词表 KL/JSD);U3 SFTCollator prompt_only 模式(data.py);U4 DistillTrainer + build_generated_batch(trainer.py,生成→双前向→divergence);U5 train_whitebox.py/.sh(full/sanity/noclip 三模式)。附带:load_sft_dataset 毛刺已磨平(改吃散装参数);hf-mirror 不代理 Xet CAS → HF_HUB_DISABLE_XET=1
|
||||
- **层 2 关账判据全过**(详见 docs/03 §6.1 实证): 66 单测全绿;B=4/T=2048 **实测不 OOM**(§5 估算成立);两次远程跑(sanity 裁到 1.0 / noclip ≈关裁剪+lr5×)均平稳、不 NaN、生成不塌(num_gen ~1900-4096)、loss 0.35→0.21 下降;checkpoint 存下
|
||||
- **层 2 关键发现(勘误"预期见毛刺")**: §4.1 梯度爆炸是真机制(U2 单测坐实单 token π_T→0 暴涨),但真实训练**高度阻尼**——同门 teacher + per-token mean 摊平,batch 级 grad_norm 峰值仅 ~14 且只降不升,关裁剪也不炸。启示:`grad_norm` 日志是**裁剪前**值(14→2 那串即爆炸证据,被 HF 默认 max_grad_norm=1.0 静默压平,现已显式化);层 5 有界乘子 π̂ 真正杀手锏是 **logit-free**(白盒的同 tokenizer 约束把你锁在温和区间)
|
||||
- **层 2 接口回看**(§6.5 每层必做,全部判"深",无需返工的毛刺): DistillConfig/token_divergence/build_generated_batch/load_sft_dataset(已修) 接口均远简于实现。三条**记录不返工**的小注:① DistillTrainer 从 self.data_collator.tokenizer 取 student tokenizer(隐式耦合,但省一个冗余参数,可接受);② SFTCollator 名字略超范(现含 prompt_only 非 SFT 模式),rename 的 churn 不值;③ token_divergence 的 labels 仅作掩码非目标(已在 docstring 标注)
|
||||
- **层 0**: ✅ 已关账(2026-07-18)
|
||||
- **层 1**: ✅ 已关账(2026-07-18)。判据全过:正本缓存 sha `33deb18c…`(1000 条,键唯一,think 残留 0,仅 2 条硬题截断);正式 1 epoch loss 0.94→0.60(56s/16 步);checkpoint 生成通顺(`/data/zym/outputs/sft_qwen3-0.6b_dapo1k`);接口回看完成(全部模块判"深";毛刺记录:load_sft_dataset 吃整个 SFTConfig 迫使诊断脚本填假 output_dir,层 2 第二消费方出现时定夺)
|
||||
- **层 1 疤痕档案**(详见 docs/02 §2.6/§5 勘误): ① 显存大头是 (B,T,V) logits 链与激活(正比 B×T,与参数量无关),B=8 曾爆 80G;② HF 梯度累积契约两幕剧(×8 → ÷4),终解 = model_accepts_loss_kwargs=False 退出新式契约;③ 缓存正本纪律:本地生成一次、单向 scp 分发、sha256 对账,两侧独立生成曾花双份钱且内容漂移
|
||||
- **诊断工具箱**: scripts/diag_collator.py(对齐链逐环)、diag_loss_probe.py(预训练 CE 基准 0.85)、diag_generate.py(生成质量)——层 2+ 数值异常照此三板斧
|
||||
- **已完成学习**: 第一/二章全部精讲;第三章已写待读(含 standard 路径解剖三大反直觉点精讲)
|
||||
- **已完成学习**: 第一/二/三章全部精讲(第三章含 standard 路径三大反直觉点、两温度、抉择原则、§4.1 梯度爆炸机制+实证)
|
||||
- **未精讲的文档账**: docs/01 的 §3.7(KL 锚三处实现差异)、§3.8(论文外稳定器)、§4(训练步流程走读)
|
||||
- **远程磁盘备忘**: 根分区 100% 的结构性原因是 `/root/zym`(507G 历史工作区)压在根分区,建议择期整体搬迁 `/data`;临时缓解 = 清 `/tmp/pip-unpack-*`、旧 tar.gz、journal。所有新增写盘已改道 `/data/zym`
|
||||
|
||||
@@ -37,8 +40,8 @@
|
||||
| `00-roadmap.md` | 本文 | ✅ |
|
||||
| `01-paper-code-map.md` | 论文 §3-§4 精读 + 参考实现全景解剖 + 概念↔代码对照表 | ✅ |
|
||||
| `02-sft-baseline.md` | 层 1:SFT 与数据管线 | ✅ 待读 |
|
||||
| `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性 | ✅ 待读 |
|
||||
| `04-mc-estimator.md` | 层 3:MC 估计 + 贝叶斯平滑 | ⬜ |
|
||||
| `03-whitebox-opd.md` | 层 2:token 级 KL 蒸馏及其脆弱性(含 §6.1 远程实证) | ✅ 已关账 |
|
||||
| `04-mc-estimator.md` | 层 3:语义相似度 φ + MC 估计 + 贝叶斯平滑(式3/4/5) | ✅ 待读 |
|
||||
| `05-entropy-chunking.md` | 层 4:熵调度 | ⬜ |
|
||||
| `06-omniopd-full.md` | 层 5:完整损失与 teacher 客户端 | ⬜ |
|
||||
| `07-eval-ablation.md` | 层 6:评测与消融 | ⬜ |
|
||||
|
||||
+17
-2
@@ -121,5 +121,20 @@ server 路径更粗:teacher 只回传 top-k logprobs 三张表(DT:2676-2679
|
||||
|
||||
1. **本地 CPU 单测**:divergence 手算对拍(V=5 玩具分布,β∈{0,1,0.5} 三点各一);β=1 与 β=0 的方向性断言(teacher 置信/弥散两种分布下 loss 排序);**梯度爆炸测试**:固定 student,teacher 对采样 token 的概率从 1e-1 衰减到 1e-6,断言梯度范数单调暴涨且超阈值——为层 5 的"有界乘子"对照埋桩。
|
||||
2. **collator 回归**:放开 prompt-only 后,全部既有 SFT 测试必须原样通过。
|
||||
3. **远程冒烟**:50 步,盯 KL loss 曲线与生成样本质量;预期能看到 loss 毛刺(梯度爆炸的实况)——这本身就是教学目标,截图留档给层 5 当对比。
|
||||
4. 关账判据:白盒蒸馏 1k 子集跑通不 NaN(允许毛刺),生成质量肉眼不劣于 SFT 基线;接口回看完成。
|
||||
3. **远程冒烟**:盯 KL loss 曲线、`grad_norm`、`distill/num_gen_tokens_per_step`。
|
||||
4. 关账判据:白盒蒸馏跑通不 NaN,生成不塌空,接口回看完成。
|
||||
|
||||
### 6.1 远程实证(2026-07-19,两次跑,勘误当初的"预期见毛刺")
|
||||
|
||||
**当初预测错了**:docs 原写"预期能看到 loss 毛刺(梯度爆炸实况)"。实跑**没有毛刺**,两次都平稳。诚实记录 + 解释:
|
||||
|
||||
| 跑 | 配置 | loss | grad_norm |
|
||||
|----|------|------|-----------|
|
||||
| sanity | 裁到 1.0、lr 1e-6、50 步 | 平滑 0.35→0.22 | 14.42 → ~2(单调降) |
|
||||
| noclip | ≈关裁剪、lr 5e-6、15 步 | 平滑 0.35→0.21 | 14.42 → ~2(无尖峰,且更快收敛) |
|
||||
|
||||
三条实证结论:
|
||||
|
||||
- **爆炸是真机制、但本区间高度阻尼**。§4.1 在单测里坐实(单个 π_T→1e-6 的 token 梯度暴涨),但真实训练里:① **同门 teacher**(Qwen3 0.6B↔4B)使 student 很少采到 teacher 真恨的 token;② **per-token mean 把每步 ~4000 token 的梯度尖峰摊平**(单测看单 token 机制,真实看上千 token 平均后果)。故 batch 级 grad_norm 峰值只 ~14(比健康 ~2 高 5-7 倍,但远非几十上百),且只降不升。
|
||||
- **关键坑:`grad_norm` 日志是裁剪前值**。sanity 的 14→2 那串本身就是爆炸证据,只是 HF 默认 `max_grad_norm=1.0` 把**步长**裁掉了、loss 才平滑——这个静默稳定器现已提进 `DistillConfig`(见其注释)。noclip 关掉它,grad_norm 曲线几乎不变(第 1 步两跑完全相同=14.42,验证确定性),但大步长反而**加速收敛**、仍不炸。
|
||||
- **重构层 5 动机的认知**:白盒的"同 tokenizer"硬约束把你锁在相对温和的区间(换跨家族 teacher 会先被词表校验拦下),所以层 5 有界乘子 π̂ 的真正杀手锏不在"防这个温和爆炸",而在 **logit-free**(teacher 只给文本、拿不到 logits,白盒根本跑不了)。爆炸的干净见证留在 U2 单测。
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
# 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` 已接真实现;接口回看完成(两模块判"深")。
|
||||
@@ -46,7 +46,13 @@ def main() -> None:
|
||||
seed=42,
|
||||
teacher_completions_path=CACHE,
|
||||
)
|
||||
ds = load_sft_dataset(cfg)
|
||||
ds = load_sft_dataset(
|
||||
cfg.dataset_path,
|
||||
cfg.dataset_split,
|
||||
cfg.subset_size,
|
||||
cfg.seed,
|
||||
cfg.teacher_completions_path,
|
||||
)
|
||||
msgs = ds[0]["messages"]
|
||||
completion_text = msgs[-1]["content"]
|
||||
|
||||
|
||||
@@ -7,19 +7,14 @@
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
from ars_opd.configs import SFTConfig
|
||||
from ars_opd.data import load_sft_dataset
|
||||
|
||||
MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物
|
||||
|
||||
cfg = SFTConfig(
|
||||
dataset_path="data/dapo-math-17k-unique.parquet",
|
||||
output_dir="/tmp/diag",
|
||||
subset_size=1000,
|
||||
seed=42,
|
||||
# 不挂 teacher 解答:只取题目做推理输入
|
||||
# 不挂 teacher 解答(teacher_completions_path 缺省):只取题目做推理输入
|
||||
ds = load_sft_dataset(
|
||||
"data/dapo-math-17k-unique.parquet", subset_size=1000, seed=42
|
||||
)
|
||||
ds = load_sft_dataset(cfg)
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
|
||||
model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, dtype=torch.float32)
|
||||
|
||||
@@ -25,7 +25,13 @@ cfg = SFTConfig(
|
||||
seed=42,
|
||||
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
|
||||
)
|
||||
ds = load_sft_dataset(cfg)
|
||||
ds = load_sft_dataset(
|
||||
cfg.dataset_path,
|
||||
cfg.dataset_split,
|
||||
cfg.subset_size,
|
||||
cfg.seed,
|
||||
cfg.teacher_completions_path,
|
||||
)
|
||||
tok = AutoTokenizer.from_pretrained(MODEL)
|
||||
collator = SFTCollator(
|
||||
tok,
|
||||
|
||||
@@ -13,7 +13,7 @@
|
||||
中断安全:缓存逐条落盘,重跑本脚本自动跳过已完成条目(断点续传)。
|
||||
"""
|
||||
|
||||
from ars_opd.configs import SFTConfig, TeacherGenConfig
|
||||
from ars_opd.configs import TeacherGenConfig
|
||||
from ars_opd.data import load_sft_dataset
|
||||
from ars_opd.teacher import TeacherClient, generate_completions
|
||||
|
||||
@@ -27,15 +27,8 @@ CACHE_PATH = "data/teacher_completions_dapo1k_minimax-m3.jsonl"
|
||||
# 同一 seed 的洗牌序列取前缀,前 5 条与前 1000 条的头 5 条完全相同,试跑写入的
|
||||
# 缓存在正式跑时全部命中,一分钱不浪费。
|
||||
|
||||
sft_cfg = SFTConfig(
|
||||
dataset_path=DATASET_PATH,
|
||||
output_dir="outputs/_unused", # 本脚本不训练,仅复用数据管线配置
|
||||
subset_size=1000,
|
||||
seed=42,
|
||||
# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
|
||||
)
|
||||
|
||||
dataset = load_sft_dataset(sft_cfg)
|
||||
dataset = load_sft_dataset(DATASET_PATH, subset_size=1000, seed=42)
|
||||
prompts = [row["messages"] for row in dataset]
|
||||
|
||||
teacher = TeacherClient(TeacherGenConfig()) # 采样参数全用 configs.py 的显式默认
|
||||
|
||||
@@ -94,7 +94,13 @@ def main() -> None:
|
||||
|
||||
# 加载顺序刻意 fail-fast:数据(毫秒级,最易配错)→ tokenizer(几 MB)→
|
||||
# 模型(GB 级下载)。teacher 缓存缺失要在下模型之前炸出来
|
||||
dataset = load_sft_dataset(cfg)
|
||||
dataset = load_sft_dataset(
|
||||
cfg.dataset_path,
|
||||
cfg.dataset_split,
|
||||
cfg.subset_size,
|
||||
cfg.seed,
|
||||
cfg.teacher_completions_path,
|
||||
)
|
||||
tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
|
||||
collator = SFTCollator(
|
||||
tokenizer,
|
||||
|
||||
@@ -22,5 +22,7 @@ export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5
|
||||
export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长序列 batch 易碎片化,按需扩段
|
||||
export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com}
|
||||
export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满
|
||||
# hf-mirror 不代理 HF Xet CAS(大权重走 Xet 会 401,见 train_whitebox.sh 详注)
|
||||
export HF_HUB_DISABLE_XET=1
|
||||
|
||||
torchrun --nproc_per_node=4 --master_port=29571 scripts/train_sft.py "$MODE"
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
"""层 2:white-box OPD 训练入口(由 train_whitebox.sh 经 torchrun 启动)。
|
||||
|
||||
自包含实验脚本:全部参数写死在下方 FULL 配置里,零参数复现;sanity 模式只是
|
||||
对 FULL 的两处显式覆盖(50 步 + 独立输出目录)。
|
||||
|
||||
与层 1 train_sft.py 的结构差异:双模型(student + 本地 teacher)、prompt-only
|
||||
数据(无 teacher 缓存,现场 on-policy 生成)、DistillTrainer 编排。
|
||||
"""
|
||||
|
||||
# ---- FSDP 前置块(必须在一切 transformers/accelerate import 之前,同 train_sft.py)----
|
||||
import os
|
||||
|
||||
os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", "false")
|
||||
|
||||
import dataclasses
|
||||
import sys
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments
|
||||
|
||||
from ars_opd.configs import DistillConfig
|
||||
from ars_opd.data import SFTCollator, load_sft_dataset
|
||||
from ars_opd.trainer import DistillTrainer
|
||||
|
||||
STUDENT_MODEL = "Qwen/Qwen3-0.6B" # 被训练的固定基线(同层 1,脚本级常量)
|
||||
|
||||
FULL = DistillConfig(
|
||||
dataset_path="data/dapo-math-17k-unique.parquet",
|
||||
output_dir="/data/zym/outputs/whitebox_qwen3-0.6b_dapo1k",
|
||||
teacher_model="Qwen/Qwen3-4B", # 本地全词表 teacher(须与 student 同 tokenizer)
|
||||
subset_size=1000,
|
||||
seed=42, # 与层 1 一致:同一批题上对比 SFT 与蒸馏
|
||||
max_prompt_length=1024,
|
||||
max_new_tokens=1024, # 与 max_prompt_length 之和 = 序列总长 T≈2048(§5 显存账)
|
||||
enable_thinking=False,
|
||||
beta=1.0, # 反向 KL = 式(2)
|
||||
kl_temperature=1.0,
|
||||
gen_temperature=1.0, # 纯采样自 π_θ(忠实 on-policy)
|
||||
gen_top_p=1.0,
|
||||
learning_rate=1e-6, # 论文 §5.1 蒸馏 lr;小步长也帮训练在梯度爆炸毛刺中存活
|
||||
per_device_train_batch_size=4, # §5 估算,首次远程必须 nvidia-smi 核实不 OOM
|
||||
gradient_accumulation_steps=4, # 全局 batch = 4 × 4 卡 × 4 = 64(同层 1)
|
||||
num_train_epochs=1,
|
||||
max_steps=-1,
|
||||
max_grad_norm=1.0, # 显式写出这个此前静默的稳定器(§4.1 爆炸靠它压平,见 config 注释)
|
||||
bf16=True,
|
||||
logging_steps=1,
|
||||
save_steps=100,
|
||||
save_total_limit=2,
|
||||
report_to="none",
|
||||
)
|
||||
|
||||
|
||||
def build_config() -> DistillConfig:
|
||||
"""按命令行模式产出配置。frozen dataclass 换参方式:replace 构造新实例。"""
|
||||
mode = sys.argv[1] if len(sys.argv) > 1 else "full"
|
||||
if mode == "full":
|
||||
return FULL
|
||||
if mode == "sanity":
|
||||
return dataclasses.replace(
|
||||
FULL, max_steps=50, output_dir=FULL.output_dir + "-sanity"
|
||||
)
|
||||
if mode == "noclip":
|
||||
# §4.1 教学对照:关闭裁剪 + 稍抬 lr,暴露反向 KL 原始爆炸。max_grad_norm
|
||||
# 设远高于实测范数(~14)故永不触发≈无裁剪;lr 5×放大让爆炸在 loss 上可见。
|
||||
# 与 sanity(裁到 1.0、lr 1e-6 的平滑曲线)并排 = 白盒脆弱性活教材,层 5 对照
|
||||
return dataclasses.replace(
|
||||
FULL,
|
||||
max_steps=15,
|
||||
max_grad_norm=1e9,
|
||||
learning_rate=5e-6,
|
||||
output_dir=FULL.output_dir + "-noclip",
|
||||
)
|
||||
raise ValueError(f"未知模式 {mode!r},只接受 full / sanity / noclip")
|
||||
|
||||
|
||||
def smoke_check_first_prompt(dataset, collator, tokenizer) -> None:
|
||||
"""训练前解码第一个 prompt 供肉眼核对(只在 rank0 打印一次)。
|
||||
|
||||
prompt-only 模式的自检重点:prompt 末尾应是生成引导符("...assistant\\n" +
|
||||
no-think 时的空 <think>),student 将从此续写。若末尾不对,生成的分布与
|
||||
训练目标会错位。
|
||||
"""
|
||||
batch = collator([dataset[0]])
|
||||
prompt_ids = batch["prompts"][0]
|
||||
mask = batch["prompt_attention_mask"][0].bool()
|
||||
text = tokenizer.decode(prompt_ids[mask], skip_special_tokens=False)
|
||||
print(
|
||||
"=" * 30
|
||||
+ " 首个 prompt 自检(供 on-policy 生成)"
|
||||
+ "=" * 30
|
||||
+ f"\n[{int(mask.sum())} tok,末尾应为生成引导符]\n…{text[-400:]}\n"
|
||||
+ "=" * 80,
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
cfg = build_config()
|
||||
rank0 = int(os.environ.get("RANK", "0")) == 0
|
||||
|
||||
# 加载顺序 fail-fast(同 train_sft.py):数据(毫秒级)→ tokenizer(几 MB)→
|
||||
# 模型(GB 级)。层 2 无 teacher 缓存,数据是 prompt-only 子集
|
||||
dataset = load_sft_dataset(
|
||||
cfg.dataset_path, cfg.dataset_split, cfg.subset_size, cfg.seed
|
||||
)
|
||||
student_tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
|
||||
teacher_tokenizer = AutoTokenizer.from_pretrained(cfg.teacher_model)
|
||||
collator = SFTCollator(
|
||||
student_tokenizer,
|
||||
max_prompt_length=cfg.max_prompt_length,
|
||||
enable_thinking=cfg.enable_thinking,
|
||||
prompt_only=True, # 层 2:只出 prompt 张量,completion 靠生成
|
||||
)
|
||||
if rank0:
|
||||
smoke_check_first_prompt(dataset, collator, student_tokenizer)
|
||||
|
||||
# student fp32 + bf16 混合精度(同层 1);teacher 直接 bf16(只推理,省显存)
|
||||
student = AutoModelForCausalLM.from_pretrained(STUDENT_MODEL, dtype=torch.float32)
|
||||
teacher = AutoModelForCausalLM.from_pretrained(
|
||||
cfg.teacher_model, dtype=torch.bfloat16
|
||||
)
|
||||
|
||||
args = TrainingArguments(
|
||||
output_dir=cfg.output_dir,
|
||||
remove_unused_columns=False, # 保住 messages 列供 collator(同层 1 注释)
|
||||
learning_rate=cfg.learning_rate,
|
||||
per_device_train_batch_size=cfg.per_device_train_batch_size,
|
||||
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
|
||||
num_train_epochs=cfg.num_train_epochs,
|
||||
max_steps=cfg.max_steps,
|
||||
lr_scheduler_type=cfg.lr_scheduler_type,
|
||||
warmup_ratio=cfg.warmup_ratio,
|
||||
max_grad_norm=cfg.max_grad_norm,
|
||||
gradient_checkpointing=cfg.gradient_checkpointing,
|
||||
bf16=cfg.bf16,
|
||||
seed=cfg.seed,
|
||||
logging_steps=cfg.logging_steps,
|
||||
logging_first_step=True,
|
||||
save_strategy="steps",
|
||||
save_steps=cfg.save_steps,
|
||||
save_total_limit=cfg.save_total_limit,
|
||||
report_to=cfg.report_to,
|
||||
ddp_find_unused_parameters=False,
|
||||
dataloader_num_workers=2,
|
||||
)
|
||||
trainer = DistillTrainer(
|
||||
model=student,
|
||||
args=args,
|
||||
train_dataset=dataset,
|
||||
data_collator=collator,
|
||||
teacher_model=teacher,
|
||||
teacher_tokenizer=teacher_tokenizer, # 构造时校验与 student 同词表
|
||||
beta=cfg.beta,
|
||||
kl_temperature=cfg.kl_temperature,
|
||||
gen_temperature=cfg.gen_temperature,
|
||||
gen_top_p=cfg.gen_top_p,
|
||||
max_new_tokens=cfg.max_new_tokens,
|
||||
)
|
||||
trainer.train()
|
||||
trainer.save_model()
|
||||
if rank0:
|
||||
student_tokenizer.save_pretrained(cfg.output_dir)
|
||||
print(f"训练完成,模型已存至 {cfg.output_dir}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Executable
+36
@@ -0,0 +1,36 @@
|
||||
#!/usr/bin/env bash
|
||||
# 层 2:white-box OPD 训练(远程 gpu-a800-060 专用;本地不跑训练)。
|
||||
#
|
||||
# 用法(tmux 内执行,日志实时可查):
|
||||
# bash scripts/train_whitebox.sh sanity # 50 步冒烟:首 prompt 自检 + KL loss + 生成数
|
||||
# bash scripts/train_whitebox.sh noclip # §4.1 对照:15 步,关裁剪+抬 lr,暴露原始
|
||||
# # 梯度爆炸(loss 毛刺);与 sanity 平滑曲线并排
|
||||
# bash scripts/train_whitebox.sh # 正式:1k 子集 1 epoch
|
||||
#
|
||||
# 前置检查清单:
|
||||
# 1. nvidia-smi 确认下方 GPUS 四张卡空闲(只许用 8 卡中的 4 张,严禁自动选卡)。
|
||||
# ⚠️ 白盒显存比层 1 紧:student 训练全套 + teacher(4B) 推理副本 + **两份**全词表
|
||||
# logits(student/teacher),§5 估算 B=4/T=2048 起步安全,但首跑必须盯 nvidia-smi;
|
||||
# 若 OOM,降 per_device_train_batch_size 到 2,仍不够再开 gradient_checkpointing
|
||||
# (改 DistillConfig,注意 checkpointing 与 generate 的 use_cache 交互)。
|
||||
# 2. data/dapo-math-17k-unique.parquet 已在(层 2 无需 teacher 缓存,纯 prompt-only):
|
||||
# scp data/dapo-math-17k-unique.parquet <远程>:/data/zym/ars-opd-rebuild/data/
|
||||
# 3. 代码最新:git -C /data/zym/ars-opd-rebuild pull
|
||||
# 4. 首跑会下载 teacher Qwen3-4B(GB 级)到 HF_HOME,确保 /data 有空间
|
||||
set -euo pipefail
|
||||
cd "$(dirname "$0")/.." # 锚定仓库根
|
||||
|
||||
GPUS=0,1,2,3 # ⚠️ 改这里前先 nvidia-smi
|
||||
MODE=${1:-full}
|
||||
|
||||
export CUDA_VISIBLE_DEVICES=$GPUS
|
||||
export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5)
|
||||
export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长生成序列易碎片化,按需扩段
|
||||
export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com}
|
||||
export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满
|
||||
# 非显然坑:hf-mirror 不代理 HF 的 Xet CAS——大权重走 Xet 会直连
|
||||
# cas-server.xethub.hf.co 并返 401(2026-07-19 teacher 4B 下载实撞)。禁用 Xet
|
||||
# 退回经典 HTTP/LFS 下载(镜像支持)。若仍不行:pip uninstall hf_xet
|
||||
export HF_HUB_DISABLE_XET=1
|
||||
|
||||
torchrun --nproc_per_node=4 --master_port=29572 scripts/train_whitebox.py "$MODE"
|
||||
@@ -0,0 +1,162 @@
|
||||
"""estimator.py 单测——docs/04 §5.2:公式对拍 + 定理 4.1 性质 + 方差收缩。
|
||||
|
||||
对拍精神源自参考实现 validate_chunk_mc_estimator.py(比 MSE_freq vs
|
||||
MSE_bayes),但全用 toy 数据本地 CPU 跑,不连真 teacher。
|
||||
detach 命门的两个世界断言在 tests/test_estimator_detach.py(E3)。
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from ars_opd.estimator import bayesian_target, chunk_prior
|
||||
|
||||
# ------------------------------------------------------------- chunk_prior
|
||||
|
||||
|
||||
def test_prior_is_geometric_mean():
|
||||
# 式(4) 手算:p = [0.9, 0.1] → π̄ = exp((log .9 + log .1)/2) = √0.09 = 0.3
|
||||
log_probs = torch.log(torch.tensor([0.9, 0.1]))
|
||||
assert math.isclose(chunk_prior(log_probs).item(), 0.3, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_prior_uniform_probs():
|
||||
# 全同概率的几何均值 = 该概率本身
|
||||
log_probs = torch.full((50,), math.log(0.5))
|
||||
assert math.isclose(chunk_prior(log_probs).item(), 0.5, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_prior_shape_and_range():
|
||||
pi_bar = chunk_prior(torch.log(torch.rand(50).clamp(1e-6, 1.0)))
|
||||
assert pi_bar.shape == () # (C,) -> 标量
|
||||
assert 0.0 < pi_bar.item() <= 1.0
|
||||
|
||||
|
||||
def test_prior_log_domain_survives_underflow():
|
||||
# 50 个 p=0.01 直接连乘 = 1e-100(fp32 下溢为 0);log 域算出 0.01
|
||||
log_probs = torch.full((50,), math.log(0.01))
|
||||
assert math.isclose(chunk_prior(log_probs).item(), 0.01, rel_tol=1e-4)
|
||||
|
||||
|
||||
def test_prior_clamp_floor():
|
||||
# 极端负 log 均值 → exp 下溢,clamp 兜到 1e-8 保持严格为正(定理 4.1b 前提)
|
||||
log_probs = torch.full((5,), -1e9)
|
||||
assert chunk_prior(log_probs).item() == pytest.approx(1e-8)
|
||||
|
||||
|
||||
def test_prior_is_detached():
|
||||
# detach 命门:π̄ 不带梯度(逃逸机制的完整断言在 test_estimator_detach.py)
|
||||
log_probs = torch.log(torch.tensor([0.5, 0.5], requires_grad=True))
|
||||
pi_bar = chunk_prior(log_probs)
|
||||
assert not pi_bar.requires_grad
|
||||
|
||||
|
||||
def test_prior_empty_raises():
|
||||
with pytest.raises(ValueError, match="为空"):
|
||||
chunk_prior(torch.tensor([]))
|
||||
|
||||
|
||||
# --------------------------------------------------------- bayesian_target
|
||||
|
||||
|
||||
def test_target_formula_hand_computed():
|
||||
# 式(5) 手算:k=8, π̄=0.5, N=10, α=1 → π̂ = (8 + 0.5)/11 = 0.77272…
|
||||
pi_hat = bayesian_target(8.0, torch.tensor(0.5), n_rollouts=10, alpha=1.0)
|
||||
assert math.isclose(pi_hat.item(), 8.5 / 11, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_target_convex_combination_identity():
|
||||
# 式(10) 恒等:π̂ = N/(N+α)·(k/N) + α/(N+α)·π̄,任取参数逐点核对
|
||||
k, pi_bar, n, alpha = 3.7, torch.tensor(0.42), 10, 1.5
|
||||
direct = bayesian_target(k, pi_bar, n, alpha).item()
|
||||
convex = (n / (n + alpha)) * (k / n) + (alpha / (n + alpha)) * pi_bar.item()
|
||||
assert math.isclose(direct, convex, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_target_anti_collapse_at_k_zero():
|
||||
# 定理 4.1(b):k=0(teacher 全否定)时 π̂ = α·π̄/(N+α) > 0,监督不归零
|
||||
pi_hat = bayesian_target(0.0, torch.tensor(0.3), n_rollouts=10, alpha=1.0)
|
||||
assert math.isclose(pi_hat.item(), 0.3 / 11, rel_tol=1e-6)
|
||||
assert pi_hat.item() > 0
|
||||
|
||||
|
||||
def test_target_full_score_shrinks_below_one():
|
||||
# 贝叶斯收缩:k=N 满分时 π̂ = (N+α·π̄)/(N+α) < 1(只要 π̄<1)——
|
||||
# 先验把估计从两端往中间拉,这正是方差收缩的来源
|
||||
pi_hat = bayesian_target(10.0, torch.tensor(0.5), n_rollouts=10, alpha=1.0)
|
||||
assert math.isclose(pi_hat.item(), 10.5 / 11, rel_tol=1e-6)
|
||||
assert pi_hat.item() < 1.0
|
||||
|
||||
|
||||
def test_target_bounded_in_unit_interval():
|
||||
# 定理 4.1(a):任意合法参数下 π̂ ∈ (0, 1]
|
||||
for k in [0.0, 2.5, 10.0]:
|
||||
for p in [1e-8, 0.5, 1.0]:
|
||||
v = bayesian_target(k, torch.tensor(p), 10, 1.0).item()
|
||||
assert 0.0 < v <= 1.0
|
||||
|
||||
|
||||
def test_target_alpha_zero_is_frequency_estimate():
|
||||
# α=0 退化为 k/N(no_bayesian 消融);k=0 时被 clamp 兜到 1e-8 而非 0
|
||||
assert math.isclose(
|
||||
bayesian_target(7.0, torch.tensor(0.5), 10, 0.0).item(), 0.7, rel_tol=1e-6
|
||||
)
|
||||
assert bayesian_target(0.0, torch.tensor(0.5), 10, 0.0).item() == pytest.approx(
|
||||
1e-8
|
||||
)
|
||||
|
||||
|
||||
def test_target_is_detached_even_with_grad_input():
|
||||
# 第二道防线:pi_bar 带梯度传入,π̂ 仍必须 detach
|
||||
pi_bar = torch.tensor(0.5, requires_grad=True)
|
||||
pi_hat = bayesian_target(5.0, pi_bar, 10, 1.0)
|
||||
assert not pi_hat.requires_grad
|
||||
|
||||
|
||||
def test_target_validation_raises():
|
||||
pi_bar = torch.tensor(0.5)
|
||||
with pytest.raises(ValueError, match="n_rollouts"):
|
||||
bayesian_target(0.0, pi_bar, 0, 1.0)
|
||||
with pytest.raises(ValueError, match="alpha"):
|
||||
bayesian_target(0.0, pi_bar, 10, -0.1)
|
||||
with pytest.raises(ValueError, match="越界"):
|
||||
bayesian_target(11.0, pi_bar, 10, 1.0) # k_sem > N:口径不一致
|
||||
with pytest.raises(ValueError, match="越界"):
|
||||
bayesian_target(-0.5, pi_bar, 10, 1.0)
|
||||
|
||||
|
||||
# --------------------------------------------- 方差收缩(定理 4.1c,toy 模拟)
|
||||
|
||||
|
||||
def test_variance_shrinkage_beats_frequency_estimate():
|
||||
"""toy 模拟对拍 validate_chunk_mc_estimator.py 的 MSE_freq vs MSE_bayes。
|
||||
|
||||
设真值 μ:每次试验采 N=10 个相似度 sim_i(均值 μ 的噪声),
|
||||
频率估计 = mean(sim) = k/N,贝叶斯估计 = (k + α·π̄)/(N+α)。
|
||||
先验 π̄ = μ(理想先验)时,收缩纯降方差、零偏差代价,MSE 必更小。
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
mu, n, alpha = 0.7, 10, 1.0
|
||||
trials = 2000
|
||||
# (trials, N) 的相似度样本:均值 μ、截断到 [0,1]
|
||||
sims = (mu + 0.25 * torch.randn(trials, n)).clamp(0.0, 1.0)
|
||||
k = sims.sum(dim=1) # (trials,) 每次试验的 k_sem
|
||||
freq = k / n
|
||||
bayes = (k + alpha * mu) / (n + alpha)
|
||||
mse_freq = ((freq - mu) ** 2).mean().item()
|
||||
mse_bayes = ((bayes - mu) ** 2).mean().item()
|
||||
assert mse_bayes < mse_freq
|
||||
|
||||
|
||||
def test_variance_shrinkage_robust_to_imperfect_prior():
|
||||
# 先验偏离真值(π̄ = μ±0.1)仍应赢:α=1、N=10 时先验权重仅 1/11,
|
||||
# 引入的偏差平方远小于省下的方差(定理 4.1c 在论文设定下的稳健性)
|
||||
torch.manual_seed(1)
|
||||
mu, n, alpha, trials = 0.6, 10, 1.0, 2000
|
||||
sims = (mu + 0.25 * torch.randn(trials, n)).clamp(0.0, 1.0)
|
||||
k = sims.sum(dim=1)
|
||||
mse_freq = ((k / n - mu) ** 2).mean().item()
|
||||
for prior in [mu - 0.1, mu + 0.1]:
|
||||
bayes = (k + alpha * prior) / (n + alpha)
|
||||
assert ((bayes - mu) ** 2).mean().item() < mse_freq
|
||||
@@ -0,0 +1,142 @@
|
||||
"""similarity.py 单测——docs/04 §5.1:手构字符串钉死 φ 与 k_sem(式3)。
|
||||
|
||||
全部本地 CPU、纯标准库,不依赖 torch/transformers。
|
||||
关键手算用例在各测试的注释里逐步展开,方便对着验算。
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from ars_opd.similarity import aggregate_similarity, edit_similarity, phi, rouge1
|
||||
|
||||
# ---------------------------------------------------------------- rouge1
|
||||
|
||||
|
||||
def test_rouge1_identical_is_exact_one():
|
||||
# 全同串必须精确 = 1.0(参考实现因分母 +1e-8 只能得 ≈0.99999998)
|
||||
s = "so x = 5 and y = 12"
|
||||
assert rouge1(s, s) == 1.0
|
||||
|
||||
|
||||
def test_rouge1_disjoint_is_zero():
|
||||
assert rouge1("a b c", "x y z") == 0.0
|
||||
|
||||
|
||||
def test_rouge1_partial_overlap_hand_computed():
|
||||
# hyp = {a, b, c}, ref = {a, b, d}:overlap = 2
|
||||
# precision = 2/3, recall = 2/3, F1 = 2·(2/3)(2/3) / (4/3) = 2/3
|
||||
assert math.isclose(rouge1("a b c", "a b d"), 2 / 3)
|
||||
|
||||
|
||||
def test_rouge1_multiset_counts_repeats():
|
||||
# 多重集语义:hyp = [x,x,x,x], ref = [x] → overlap = min(4,1) = 1
|
||||
# precision = 1/4, recall = 1/1, F1 = 2·(1/4)/(5/4) = 0.4
|
||||
# (参考实现的 set 版会给满分 1.0——数学文本重复词多,这是关键失真点)
|
||||
assert math.isclose(rouge1("x x x x", "x"), 0.4)
|
||||
|
||||
|
||||
def test_rouge1_is_bag_of_words_order_blind():
|
||||
# 词袋:只看用了哪些词,不看顺序
|
||||
assert rouge1("a b", "b a") == 1.0
|
||||
|
||||
|
||||
def test_rouge1_empty_sides():
|
||||
assert rouge1("", "a b") == 0.0
|
||||
assert rouge1("a b", "") == 0.0
|
||||
assert rouge1("", "") == 0.0
|
||||
assert rouge1(" ", "a") == 0.0 # 纯空白 split 后无词
|
||||
|
||||
|
||||
# ---------------------------------------------------------- edit_similarity
|
||||
|
||||
|
||||
def test_edit_identical_is_one():
|
||||
s = "so x = 5 and y = 12"
|
||||
assert edit_similarity(s, s) == 1.0
|
||||
|
||||
|
||||
def test_edit_totally_different_is_zero():
|
||||
# ["a","b"] vs ["c","d"]:2 次替换,dist=2, max(m,n)=2 → 1 − 1 = 0
|
||||
assert edit_similarity("a b", "c d") == 0.0
|
||||
|
||||
|
||||
def test_edit_single_substitution_hand_computed():
|
||||
# ["a","b","c"] vs ["a","x","c"]:1 次替换,dist=1, max=3 → 2/3
|
||||
assert math.isclose(edit_similarity("a b c", "a x c"), 2 / 3)
|
||||
|
||||
|
||||
def test_edit_insertion_hand_computed():
|
||||
# ["a","b"] vs ["a","x","b"]:1 次插入,dist=1, max=3 → 2/3
|
||||
assert math.isclose(edit_similarity("a b", "a x b"), 2 / 3)
|
||||
|
||||
|
||||
def test_edit_is_order_sensitive():
|
||||
# ["a","b"] vs ["b","a"]:两次替换 dist=2 → 0.0;与 rouge1 的 1.0 互补
|
||||
assert edit_similarity("a b", "b a") == 0.0
|
||||
assert rouge1("a b", "b a") == 1.0
|
||||
|
||||
|
||||
def test_edit_empty_sides():
|
||||
assert edit_similarity("", "") == 1.0 # 零距离
|
||||
assert edit_similarity("a b", "") == 0.0 # 全删
|
||||
assert edit_similarity("", "a b") == 0.0 # 全插
|
||||
|
||||
|
||||
def test_edit_asymmetric_lengths():
|
||||
# ["a"] vs ["a","b","c","d"]:3 次插入,dist=3, max=4 → 1/4
|
||||
assert math.isclose(edit_similarity("a", "a b c d"), 1 / 4)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- phi
|
||||
|
||||
|
||||
def test_phi_default_is_edit_distance():
|
||||
# 论文 §5.1 默认;"a b" vs "b a" 恰能区分两度量(edit=0, rouge1=1)
|
||||
assert phi("a b", "b a") == edit_similarity("a b", "b a") == 0.0
|
||||
|
||||
|
||||
def test_phi_dispatch():
|
||||
h, r = "a b c", "a b d"
|
||||
assert phi(h, r, metric="rouge1") == rouge1(h, r)
|
||||
assert phi(h, r, metric="edit_distance") == edit_similarity(h, r)
|
||||
|
||||
|
||||
def test_phi_unknown_metric_raises():
|
||||
with pytest.raises(ValueError, match="bleu"):
|
||||
phi("a", "a", metric="bleu")
|
||||
|
||||
|
||||
# ----------------------------------------------------- aggregate_similarity
|
||||
|
||||
|
||||
def test_aggregate_is_sum_of_phi():
|
||||
# 式(3) 手算:rollouts 与 "a b c" 的 edit 相似度分别为 1.0, 2/3, 0.0
|
||||
chunk = "a b c"
|
||||
rollouts = ["a b c", "a x c", "x y z"]
|
||||
expected = 1.0 + 2 / 3 + 0.0
|
||||
assert math.isclose(aggregate_similarity(chunk, rollouts), expected)
|
||||
|
||||
|
||||
def test_aggregate_bounds():
|
||||
# k_sem ∈ [0, N]:全同 → N,全不同 → 0
|
||||
n = 5
|
||||
assert aggregate_similarity("a b", ["a b"] * n) == float(n)
|
||||
assert aggregate_similarity("a b", ["x y"] * n) == 0.0
|
||||
|
||||
|
||||
def test_aggregate_is_continuous_soft_count():
|
||||
# φ 连续 ⇒ k_sem 非整数是常态(区别于 token 精确匹配的硬计数)
|
||||
k = aggregate_similarity("a b c", ["a b c", "a x c"])
|
||||
assert 1.0 < k < 2.0
|
||||
|
||||
|
||||
def test_aggregate_empty_rollouts():
|
||||
assert aggregate_similarity("a b", []) == 0.0
|
||||
|
||||
|
||||
def test_aggregate_metric_passthrough():
|
||||
# "a b" vs "b a":edit 全零,rouge1 全满——验证 metric 真的传下去了
|
||||
chunk, rollouts = "a b", ["b a", "b a"]
|
||||
assert aggregate_similarity(chunk, rollouts, metric="edit_distance") == 0.0
|
||||
assert aggregate_similarity(chunk, rollouts, metric="rouge1") == 2.0
|
||||
Reference in New Issue
Block a user