tests: 预置 π̂ detach 命门守护测试(对应 docs/01 §3.4,层 3 接入真实 estimator)

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-18 02:32:05 -04:00
parent c9eddd5be8
commit 72799a34dc
2 changed files with 73 additions and 1 deletions
+72
View File
@@ -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