tests: 预置 π̂ detach 命门守护测试(对应 docs/01 §3.4,层 3 接入真实 estimator)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -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