File size: 8,750 Bytes
038acee | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 | """Literature-anchored specification tests — step-8 gap closure.
Complements tests/test_method_spec.py with the surfaces the step-8 audit
found untested in this repo: the reverse-posterior/remasking/carry-over
chain, the final greedy cleanup, decode temperature, label smoothing and
the loss weight clip. Every expected value derives from a cited source
or a derivation written in the docstring — never from the current output
of this or the sibling repo. The minihack twin file carries the same
chain/cleanup/temperature/smoothing assertions with the same inputs and
tolerances (weight clip and empty-mask zero already exist there).
References as in tests/test_method_spec.py (MDLM arXiv:2406.07524;
ReMDM arXiv:2503.00307; Holtzman arXiv:1904.09751).
"""
from __future__ import annotations
import math
import jax
import jax.numpy as jnp
import numpy as np
import pytest
from src.diffusion.loss import compute_loss
from src.diffusion.sampling import _decode, sample_plan
from src.diffusion.schedules import SCHEDULE_MAP
ATOL = 1e-6
# ---------------------------------------------------------------------------
# Reverse posterior + remasking + carry-over, jointly, through sample_plan
# ---------------------------------------------------------------------------
def _time_coded_apply(num_actions: int, horizon: int):
"""Stub model whose argmax token encodes the decode time.
token 0 for t > 0.9, token 1 for 0.5 < t <= 0.9, token 2 otherwise.
Deterministic under temperature=0 (argmax decode).
"""
def apply_fn(params, obs, z, t, rng):
idx = jnp.where(t[0] > 0.9, 0, jnp.where(t[0] > 0.5, 1, 2))
logits = jax.nn.one_hot(idx, num_actions) * 30.0
return jnp.broadcast_to(logits, (obs.shape[0], horizon, num_actions))
return apply_fn
@pytest.mark.parametrize("eta", [0.0, 0.5])
def test_posterior_remask_carryover_token_distribution(eta):
"""Final token distribution matches the ReMDM Alg 1 posterior chain.
Source: ReMDM eq (6) masked branch (unmask probability
(alpha_s - (1-sigma) alpha_t) / (1 - alpha_t)), eq (7) sigma_max,
Sec 4.1 rescale sigma = eta * sigma_max, and MDLM Sec 3.2.3 carry-over.
Derivation (linear schedule, K=3, grid t=1,2/3,1/3 with s=2/3,1/3,0;
the stub decodes token 0 at t=1, token 1 at t=2/3, token 2 at
t<=1/3):
step 1 (t=1): alpha_t=0, alpha_s=1/3, p_unmask=1/3 -> token 0.
step 2 (t=2/3): alpha_t=1/3, alpha_s=2/3, sigma_max=1, sigma=eta;
committed remask w.p. eta; masked unmask w.p.
(2/3-(1-eta)/3)/(2/3) = (1+eta)/2 -> token 1.
step 3 (t=1/3): alpha_s=1 -> sigma_max=0 (no remask), p_unmask=1
-> every remaining mask becomes token 2.
Hence P(0) = (1-eta)/3, P(1) = (1+eta)/3, P(2) = 1/3.
eta=0 additionally pins carry-over: a committed token is never
re-decided when sigma=0.
Statistical: N = 512*32 = 16384 independent per-position outcomes;
max sigma = sqrt(0.25/16384) = 0.0039; bound 0.02 = 5.1 sigma.
"""
B, H, V = 512, 32, 3
fn, _ = SCHEDULE_MAP["linear"]
plan = sample_plan(
_time_coded_apply(V, H), None, jax.random.PRNGKey(0),
jnp.zeros((B, 4)), V, H, num_steps=3, schedule_fn=fn,
remask_strategy="rescale", eta=eta, use_loop=False,
temperature=0.0, top_p=None,
)
freq = np.bincount(np.asarray(plan).ravel(), minlength=V) / (B * H)
expected = np.array([(1 - eta) / 3, (1 + eta) / 3, 1 / 3])
assert np.all(np.abs(freq - expected) < 0.02), (freq, expected)
def test_final_cleanup_commits_all_remaining_masks():
"""With zero denoising steps, the final greedy cleanup decodes every
position at t=0 (argmax), leaving no MASK token.
Source: spec-method 4.9 (final-step commit of remaining masks is the
project safety net; craftax cleanup is unconditional). The stub
decodes token 2 at t=0, so the output must be all-2.
"""
B, H, V = 8, 16, 3
fn, _ = SCHEDULE_MAP["linear"]
plan = sample_plan(
_time_coded_apply(V, H), None, jax.random.PRNGKey(0),
jnp.zeros((B, 4)), V, H, num_steps=0, schedule_fn=fn,
remask_strategy="rescale", eta=0.0, use_loop=False,
temperature=0.0, top_p=None,
)
assert (np.asarray(plan) == 2).all()
# ---------------------------------------------------------------------------
# Decode temperature
# ---------------------------------------------------------------------------
def test_decode_temperature_scales_logits_before_sampling():
"""Sampling frequencies follow softmax(logits / temperature).
Source: spec-method 5.2 (softmax temperature before filtering;
standard technique). Derivation: logits [0, 1] at temperature 0.5
give softmax([0, 2]) = [1/(1+e^2), e^2/(1+e^2)] = [0.1192, 0.8808].
Statistical: 8192 draws; sigma = sqrt(0.8808*0.1192/8192) = 0.00358;
bound 0.0143 = 4 sigma.
"""
logits = jnp.broadcast_to(jnp.array([0.0, 1.0]), (8192, 1, 2))
tokens = np.asarray(
_decode(jax.random.PRNGKey(0), logits, temperature=0.5, top_p=None)
).ravel()
p1 = math.exp(2) / (1 + math.exp(2))
assert abs(tokens.mean() - p1) < 0.0143
# ---------------------------------------------------------------------------
# Label smoothing
# ---------------------------------------------------------------------------
def test_label_smoothing_matches_closed_form():
"""Smoothed CE = -[(1-eps)+eps/V] log p_true - (eps/V) sum log p_other.
Source: spec-method 3.7 (smoothing target (1-eps)*onehot + eps/V;
eps=0 is the exact ELBO). Derivation: model probs [0.7,0.1,0.1,0.1],
true class 0, V=4, eps=0.3: coefficient on -log 0.7 is
(1-0.3)+0.3/4 = 0.775; each other class gets 0.3/4 = 0.075.
t pinned to 1 (linear, w=1, everything masked) makes the loss equal
that CE exactly. Same inputs and expectations as the minihack twin.
"""
V, H, B = 4, 8, 4
fn, deriv = SCHEDULE_MAP["linear"]
probs = jnp.array([0.7, 0.1, 0.1, 0.1])
def apply_fn(params, obs, z, t, rng):
return jnp.broadcast_to(jnp.log(probs), (obs.shape[0], H, V))
x0 = jnp.zeros((B, H), dtype=jnp.int32)
obs = jnp.zeros((B, 3))
expected = -0.775 * math.log(0.7) - 0.075 * 3 * math.log(0.1)
loss, _ = compute_loss(
apply_fn, None, jax.random.PRNGKey(0), x0, obs, jnp.ones(B), V,
fn, deriv, label_smoothing=0.3, t_min=1.0, t_max=1.0,
)
assert abs(float(loss) - expected) < 1e-5
loss0, _ = compute_loss(
apply_fn, None, jax.random.PRNGKey(0), x0, obs, jnp.ones(B), V,
fn, deriv, label_smoothing=0.0, t_min=1.0, t_max=1.0,
)
assert abs(float(loss0) - (-math.log(0.7))) < 1e-5
# ---------------------------------------------------------------------------
# Loss weight clip and the empty-mask edge case (minihack twins exist)
# ---------------------------------------------------------------------------
def test_loss_weight_clip_bound():
"""w(t) is clipped at 1000 (project numerics guard, spec-method 3.4;
the same _MAX_WEIGHT the minihack twin pins via an explicit zt).
compute_loss samples zt internally, so this is statistical: at
t=1e-4 (linear) the raw weight is 1/t = 10^4, clipped to 10^3.
E[loss] = 1000 * ln V * P(mask) = 1000 * ln 5 * 1e-4 = 0.1 ln 5.
The unclipped hypothesis gives 1.0 ln 5 (10x larger). Masked count
over N = 131072*8 positions is ~Poisson(104.9); sigma(loss) =
1000 * lnV * sqrt(104.9)/N = 0.0098 lnV; bound 0.04 lnV = 4.1 sigma.
"""
V, H, B = 5, 8, 131072
fn, deriv = SCHEDULE_MAP["linear"]
def apply_fn(params, obs, z, t, rng):
return jnp.zeros((obs.shape[0], H, V))
x0 = jnp.zeros((B, H), dtype=jnp.int32)
loss, _ = compute_loss(
apply_fn, None, jax.random.PRNGKey(0), x0, jnp.zeros((B, 1)),
jnp.ones(B), V, fn, deriv, t_min=1e-4, t_max=1e-4,
)
assert abs(float(loss) - 0.1 * math.log(V)) < 0.04 * math.log(V)
def test_loss_zero_when_nothing_masked():
"""Loss is exactly 0.0 when the batch has no masked positions.
Source: spec-method 3.4 (zero-on-empty project convention; the
minihack twin asserts the same through an explicit unmasked zt).
At t=0, alpha=1 and the forward process keeps every token
(uniform draws in [0,1) are always < 1), so nothing is masked.
"""
V, H, B = 5, 8, 4
fn, deriv = SCHEDULE_MAP["linear"]
def apply_fn(params, obs, z, t, rng):
return jnp.zeros((obs.shape[0], H, V))
x0 = jnp.zeros((B, H), dtype=jnp.int32)
loss, _ = compute_loss(
apply_fn, None, jax.random.PRNGKey(0), x0, jnp.zeros((B, 1)),
jnp.ones(B), V, fn, deriv, t_min=0.0, t_max=0.0,
)
assert float(loss) == 0.0
|