From 72799a34dc811b65cfbbad9edd2d39bac4f7e650 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 02:32:05 -0400 Subject: [PATCH] =?UTF-8?q?tests:=20=E9=A2=84=E7=BD=AE=20=CF=80=CC=82=20de?= =?UTF-8?q?tach=20=E5=91=BD=E9=97=A8=E5=AE=88=E6=8A=A4=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=EF=BC=88=E5=AF=B9=E5=BA=94=20docs/01=20=C2=A73.4=EF=BC=8C?= =?UTF-8?q?=E5=B1=82=203=20=E6=8E=A5=E5=85=A5=E7=9C=9F=E5=AE=9E=20estimato?= =?UTF-8?q?r=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 --- docs/00-roadmap.md | 2 +- tests/test_estimator_detach.py | 72 ++++++++++++++++++++++++++++++++++ 2 files changed, 73 insertions(+), 1 deletion(-) create mode 100644 tests/test_estimator_detach.py diff --git a/docs/00-roadmap.md b/docs/00-roadmap.md index e35324a..6866541 100644 --- a/docs/00-roadmap.md +++ b/docs/00-roadmap.md @@ -9,7 +9,7 @@ | 0 | 环境与骨架 | — | 本地/远程 conda 环境、gitea 同步、包骨架 | 两端 `pytest` 空跑通过 | | 1 | SFT 基线 | §3.1 式(1) | 数据管线 + 最小 SFT 训练脚本(Qwen3-0.6B) | 远程 4 卡跑通,loss 正常下降 | | 2 | White-box OPD 基线 | §3.1 式(2) | token 级反向 KL 蒸馏(teacher Qwen3-4B 本地 vLLM) | 远程跑通;理解式(2)梯度爆炸问题(§4.1) | -| 3 | 相似度 + MC 估计器 | §3.2.1-3.2.2 式(3)(4)(5) | `similarity.py` + `estimator.py`(纯逻辑) | 本地 CPU 单测,对拍 `validate_mc_estimator.py` | +| 3 | 相似度 + MC 估计器 | §3.2.1-3.2.2 式(3)(4)(5) | `similarity.py` + `estimator.py`(纯逻辑) | 本地 CPU 单测,对拍 `validate_mc_estimator.py`;detach 命门测试已预置(`tests/test_estimator_detach.py`),完成后需接入真实实现 | | 4 | Peak-entropy 调度器 | §3.2.3 式(6)(7) | `chunking.py`(纯逻辑) | 本地 CPU 单测:toy 熵序列上验证 chunk 选择与合并 | | 5 | 完整 OmniOPD | §3.2.4 式(8) | `teacher.py`(API 客户端+缓存)+ `trainer.py`(chunk 损失 + KL 锚定) | 远程端到端跑通(DeepSeek/MiniMax teacher) | | 6 | 评测与消融 | §5 | 数学评测脚本;三个消融开关 | MATH-500 子集上 student 有可测提升趋势 | diff --git a/tests/test_estimator_detach.py b/tests/test_estimator_detach.py new file mode 100644 index 0000000..23d3c9a --- /dev/null +++ b/tests/test_estimator_detach.py @@ -0,0 +1,72 @@ +"""π̂ detach 命门约束的守护测试(对应 docs/01 §3.4,论文式(5)(8))。 + +背景:chunk 损失 L = -π̂·Σlog π_θ 中,π̂ 的先验 π̄ 由学生自身概率算出。 +若不切断 π̄ 的梯度通路,最速下降方向会变成压低学生对自己 token 的概率、 +把乘子 π̂ 推向 0 以逃逸惩罚(p·ln(1/p)→0,指数快过对数),且恰好在 +teacher 全否定(k≈0)、最需要纠正的 chunk 上塌缩。推导见 docs/01 §3.4。 + +现状:独立的数学性质测试,仅依赖 torch(单 token 简化,C=1)。 +层 3 完成 ars_opd/estimator.py 后,需追加针对真实实现的同名断言, +确保重构时 `.detach()` 不被误删(参考实现锚点:trainer:2196、2201)。 +""" + +import torch + +# 与论文/参考实现默认一致:α=1, N=10 +ALPHA = 1.0 +N_ROLLOUTS = 10.0 + + +def chunk_loss_and_grad(p0: float, k: float, detach_prior: bool) -> tuple[float, float]: + """单 token 版式(8) chunk 项,返回 (loss 值, dL/dp)。 + + 参数: + p0: 学生对自己 token 的概率,标量。 + k: teacher 相似度票数 k_sem,标量(0 = 全否定)。 + detach_prior: 是否切断先验 π̄ 的梯度通路。 + 返回: + (loss.item(), p.grad.item()) + """ + p = torch.tensor(p0, requires_grad=True) + log_p = p.log() + prior_src = log_p.detach() if detach_prior else log_p + pi_bar = prior_src.exp() # 式(4):C=1 时几何均值即 p 本身 + pi_hat = (k + ALPHA * pi_bar) / (N_ROLLOUTS + ALPHA) # 式(5) + loss = -pi_hat * log_p # 式(8) chunk 项 + loss.backward() + return loss.item(), p.grad.item() + + +def test_detach_reinforces_even_when_teacher_rejects(): + """detach 世界:即使 teacher 全否定(k=0),梯度仍为负 → optimizer 增大 p。 + + 这是贝叶斯兜底的本意:k=0 处仍有非零、方向正确的学习信号。 + """ + _, grad = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=True) + assert grad < 0 + + +def test_no_detach_escapes_when_teacher_rejects(): + """不 detach 世界:k=0 且 p 低于 1/e 时梯度为正 → optimizer 压低 p(逃逸)。 + + 此断言若失败(梯度变负),说明有人"修复"了 detach——那恰恰是 bug。 + """ + _, grad = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=False) + assert grad > 0 + + +def test_detach_does_not_change_loss_value(): + """detach 只剪梯度不改数值:两个世界的前向 loss 必须完全相等。""" + loss_detached, _ = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=True) + loss_attached, _ = chunk_loss_and_grad(p0=0.2, k=0.0, detach_prior=False) + assert loss_detached == loss_attached + + +def test_teacher_agreement_blocks_escape_even_without_detach(): + """k 大时逃逸被堵死:分子中 k·|log p| 项不受 p 控制,随否认无限增长。 + + 逃逸条件为 NLL > 1 + k/(α·π̄);k=5、p=0.2 时阈值 ≈ 26,远未达到, + 故即使不 detach 梯度仍为负。印证"塌缩恰好集中在 k≈0 的 chunk"。 + """ + _, grad = chunk_loss_and_grad(p0=0.2, k=5.0, detach_prior=False) + assert grad < 0