tests: 预置 π̂ detach 命门守护测试(对应 docs/01 §3.4,层 3 接入真实 estimator)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
+1
-1
@@ -9,7 +9,7 @@
|
|||||||
| 0 | 环境与骨架 | — | 本地/远程 conda 环境、gitea 同步、包骨架 | 两端 `pytest` 空跑通过 |
|
| 0 | 环境与骨架 | — | 本地/远程 conda 环境、gitea 同步、包骨架 | 两端 `pytest` 空跑通过 |
|
||||||
| 1 | SFT 基线 | §3.1 式(1) | 数据管线 + 最小 SFT 训练脚本(Qwen3-0.6B) | 远程 4 卡跑通,loss 正常下降 |
|
| 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) |
|
| 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 选择与合并 |
|
| 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) |
|
| 5 | 完整 OmniOPD | §3.2.4 式(8) | `teacher.py`(API 客户端+缓存)+ `trainer.py`(chunk 损失 + KL 锚定) | 远程端到端跑通(DeepSeek/MiniMax teacher) |
|
||||||
| 6 | 评测与消融 | §5 | 数学评测脚本;三个消融开关 | MATH-500 子集上 student 有可测提升趋势 |
|
| 6 | 评测与消融 | §5 | 数学评测脚本;三个消融开关 | MATH-500 子集上 student 有可测提升趋势 |
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user