remdm-planner-minihack / tests /test_method_spec.py
AnonMLuser's picture
Anonymous artefact release
e0eb79a verified
Raw
History Blame Contribute Delete
28.3 kB
"""Literature-anchored specification tests.
Every expected value here is derived from a primary source (cited
per test; full references below) or from a derivation written out in
the docstring - never from the current output of this or the sibling repo. The craftax repo
carries the same assertions with the same inputs and tolerances wherever
the mathematics is parameter-free and shared.
Tolerances: closed-form checks use atol=1e-6 (an order of magnitude above
float32 round-off, which the parity probes measured at <4e-8 on these
functions). Statistical checks state their sampling distribution and use a
4-sigma bound with the derivation in the docstring.
References:
- MDLM: Sahoo et al., "Simple and Effective Masked Diffusion Language
Models", NeurIPS 2024. arXiv:2406.07524.
- Shi: Shi et al., "Simplified and Generalized Masked Diffusion for
Discrete Data", NeurIPS 2024. arXiv:2406.04329.
- ReMDM: Wang et al., "Remasking Discrete Diffusion Models with
Inference-Time Scaling", NeurIPS 2025. arXiv:2503.00307.
- MaskGIT: Chang et al., "MaskGIT: Masked Generative Image
Transformer", CVPR 2022. arXiv:2202.04200.
- Nichol & Dhariwal, "Improved Denoising Diffusion Probabilistic
Models", ICML 2021. arXiv:2102.09672.
- Holtzman et al., "The Curious Case of Neural Text Degeneration",
ICLR 2020. arXiv:1904.09751.
"""
from __future__ import annotations
import math
import numpy as np
from types import SimpleNamespace
import pytest
import torch
from src.diffusion.forward import q_sample
from src.diffusion.loss import mdlm_loss
from src.diffusion.sampling import (
_compute_remask_prob,
greedy_sample,
remdm_sample,
top_p_filter,
)
from src.diffusion.schedules import (
cosine_schedule,
cosine_schedule_deriv,
cosine_sq_schedule,
cosine_sq_schedule_deriv,
get_schedule,
linear_schedule,
linear_schedule_deriv,
)
ATOL = 1e-6
T_GRID = torch.tensor([0.0, 1.0 / 3.0, 0.5, 2.0 / 3.0, 1.0], dtype=torch.float64)
# ---------------------------------------------------------------------------
# Noise schedules: closed forms and derivatives
# ---------------------------------------------------------------------------
def test_linear_schedule_closed_form():
"""alpha(t) = 1 - t; alpha'(t) = -1.
Source: MDLM (Sahoo et al.) App E.1 eq (90) family (log-linear alpha_t
= 1 - t); ReMDM (Wang et al.) Sec 3 uses the same convention
alpha(0)=1, alpha(1)=0.
"""
expected = torch.tensor([1.0, 2.0 / 3.0, 0.5, 1.0 / 3.0, 0.0], dtype=torch.float64)
assert torch.allclose(linear_schedule(T_GRID), expected, atol=ATOL)
assert torch.allclose(
linear_schedule_deriv(T_GRID), torch.full_like(T_GRID, -1.0), atol=ATOL
)
def test_cosine_schedule_closed_form():
"""alpha(t) = cos(pi t / 2); alpha'(t) = -(pi/2) sin(pi t / 2).
Source: MDLM App E.1 eq (92) ("Cosine"): sigma(t) = -log cos(pi/2 (1-t))
i.e. alpha = cos(pi/2 (1-t)) on MDLM's reversed time axis, equal to
cos(pi t / 2) under this repo's alpha(0)=1 orientation. Values:
cos(pi/6) = sqrt(3)/2, cos(pi/4) = sqrt(2)/2, cos(pi/3) = 1/2.
"""
expected = torch.tensor(
[1.0, math.sqrt(3) / 2, math.sqrt(2) / 2, 0.5, 0.0], dtype=torch.float64
)
assert torch.allclose(cosine_schedule(T_GRID), expected, atol=ATOL)
expected_d = torch.tensor(
[0.0, -math.pi / 4, -(math.pi / 2) * math.sqrt(2) / 2,
-(math.pi / 2) * math.sqrt(3) / 2, -math.pi / 2],
dtype=torch.float64,
)
assert torch.allclose(cosine_schedule_deriv(T_GRID), expected_d, atol=ATOL)
def test_cosine_sq_schedule_closed_form():
"""alpha(t) = cos^2(pi t / 2); alpha'(t) = -(pi/2) sin(pi t).
Source: MDLM App E.1 eq (91) ("Cosine Squared"), attributed to Nichol &
Dhariwal (their eq for alpha-bar with s=0). cos^2 at the grid:
[1, 3/4, 1/2, 1/4, 0]; derivative via 2 cos(x)(-sin(x))(pi/2) =
-(pi/2) sin(pi t).
"""
expected = torch.tensor([1.0, 0.75, 0.5, 0.25, 0.0], dtype=torch.float64)
assert torch.allclose(cosine_sq_schedule(T_GRID), expected, atol=ATOL)
expected_d = torch.tensor(
[0.0, -(math.pi / 2) * math.sqrt(3) / 2, -math.pi / 2,
-(math.pi / 2) * math.sqrt(3) / 2, 0.0],
dtype=torch.float64,
)
assert torch.allclose(cosine_sq_schedule_deriv(T_GRID), expected_d, atol=ATOL)
def test_schedule_registry_names_follow_mdlm_e1():
"""The label "cosine" must denote MDLM eq (92), not eq (91).
Source: MDLM App E.1. Guards against the two repos' "cosine"
diverging again.
"""
t = torch.tensor([0.5], dtype=torch.float64)
assert torch.allclose(get_schedule("cosine")(t),
torch.tensor([math.sqrt(2) / 2], dtype=torch.float64),
atol=ATOL)
assert torch.allclose(get_schedule("cosine_sq")(t),
torch.tensor([0.5], dtype=torch.float64), atol=ATOL)
# ---------------------------------------------------------------------------
# Forward corruption q(z_t | x_0)
# ---------------------------------------------------------------------------
def test_forward_marginal_endpoints_and_pad():
"""q(z_t|x) = Cat(alpha_t x + (1-alpha_t) m): t=0 identity, t=1 all-MASK;
PAD positions are never corrupted.
Source: MDLM Sec 3.2.1 (forward masking marginal); PAD exclusion is the
benchmark-forced extension. Endpoints
are deterministic: at t=1 the mask draw u<1 always holds (u in [0,1));
at t=0 u<0 never holds.
"""
torch.manual_seed(0)
x0 = torch.randint(0, 12, (4, 64))
x0[:, -8:] = 13 # PAD tail
t0 = torch.zeros(4)
t1 = torch.ones(4)
z0 = q_sample(x0, t0, mask_token=12, pad_token=13, schedule_fn=linear_schedule)
assert torch.equal(z0, x0), "t=0 must leave the sequence unchanged"
z1 = q_sample(x0, t1, mask_token=12, pad_token=13, schedule_fn=linear_schedule)
real = x0 != 13
assert torch.all(z1[real] == 12), "t=1 must mask every non-PAD position"
assert torch.all(z1[~real] == 13), "PAD positions must never be masked"
def test_forward_marginal_rate_matches_one_minus_alpha():
"""Empirical mask rate at t=0.5 (linear) is 1-alpha = 0.5 within 4 sigma.
Source: MDLM Sec 3.2.1. N = 200*64 = 12800 independent Bernoulli(0.5)
draws; sigma = sqrt(0.25/12800) = 0.00442; bound = 4 sigma = 0.0177.
"""
torch.manual_seed(0)
x0 = torch.randint(0, 12, (200, 64))
t = torch.full((200,), 0.5)
zt = q_sample(x0, t, mask_token=12, pad_token=13, schedule_fn=linear_schedule)
rate = (zt == 12).float().mean().item()
assert abs(rate - 0.5) < 0.0177
# ---------------------------------------------------------------------------
# Loss: NELBO estimator
# ---------------------------------------------------------------------------
def _uniform_logits_case(B=2, L=8, V=12):
logits = torch.zeros(B, L, V)
x0 = torch.randint(0, V, (B, L), generator=torch.Generator().manual_seed(0))
return logits, x0
def test_loss_all_masked_uniform_logits_t1():
"""Loss = w(1) * log V with everything masked and uniform logits.
Source: MDLM eq (10) integrand alpha'_t/(1-alpha_t) * CE summed over
masked positions, per-token normalised. Derivation: uniform logits give
CE = log V at every position; all L positions masked so sum/L = log V;
linear schedule w(1) = -alpha'(1)/(1-alpha(1)) = 1/1 = 1.
"""
logits, x0 = _uniform_logits_case()
zt = torch.full_like(x0, 12)
t = torch.ones(2)
loss = mdlm_loss(logits, x0, zt, t, 12, 13, linear_schedule)
assert abs(loss.item() - math.log(12)) < ATOL
def test_loss_denominator_is_per_token_not_per_masked():
"""Half-masked at t=0.5 (linear): loss = w(0.5)*logV*(4/8) = log V.
Source: MDLM eq (8)/(10); Shi eq (4). The equations contain no division
by the realised masked count; dividing by it would
return 2*log V here, off by the factor L/n_masked = 2. This is the
regression test for the loss denominator.
"""
logits, x0 = _uniform_logits_case()
zt = x0.clone()
zt[:, :4] = 12 # exactly half masked
t = torch.full((2,), 0.5)
loss = mdlm_loss(logits, x0, zt, t, 12, 13, linear_schedule)
assert abs(loss.item() - math.log(12)) < ATOL
def test_loss_weight_clip_bound():
"""w(t) is clipped at weight_clip as t -> 0.
Source: the divergence of w(t)=1/t as t->0 is a property of MDLM
eq (10) under the linear schedule; the finite bound (1000, with the
denominator floored at 1e-5) is this codebase's documented stability
policy (_MAX_WEIGHT), shared with the craftax repo. Derivation: at
t=1e-6, 1-alpha=1e-6 floors to 1e-5 giving w=1e5, clipped to 1000;
half-masked uniform-logits loss = 1000 * log V * 0.5.
"""
logits, x0 = _uniform_logits_case()
zt = x0.clone()
zt[:, :4] = 12
t = torch.full((2,), 1e-6)
loss = mdlm_loss(logits, x0, zt, t, 12, 13, linear_schedule)
assert abs(loss.item() - 1000 * math.log(12) * 0.5) < 1e-3
def test_loss_excludes_pad_and_empty_mask_is_zero():
"""PAD positions contribute nothing; no masked positions -> loss 0.
Source: MDLM Sec 3.2.3 (loss over masked positions only); PAD handling
is the benchmark-forced extension. Derivation: 6 of 8 positions are
real and masked at t=1 -> loss = 1 * log V * 6/8. The all-unmasked case
has an empty diffusion term.
"""
logits, x0 = _uniform_logits_case()
x0[:, -2:] = 13
zt = torch.full_like(x0, 12)
t = torch.ones(2)
loss = mdlm_loss(logits, x0, zt, t, 12, 13, linear_schedule)
assert abs(loss.item() - math.log(12) * 6 / 8) < ATOL
clean = mdlm_loss(logits, x0, x0.clone(), t, 12, 13, linear_schedule)
assert clean.item() == 0.0
@pytest.mark.parametrize("reduction", ["none", "mean"])
def test_the_empty_mask_loss_is_zero_and_still_differentiable(reduction):
"""An all-unmasked draw contributes zero *and* stays in the graph.
An empty mask is a legitimate draw, not an error: at a t where alpha(t)
is near 1 nothing gets masked. The value has to be zero, and the result
has to remain differentiable in ``logits``, because the caller adds this
term to others and back-propagates the sum. Returning a freshly
allocated zero gave the right number with no ``grad_fn``, and
``backward()`` then raised "element 0 of tensors does not require grad
and does not have a grad_fn" for every ablation whose other loss terms
could not carry the graph either.
The gradient must also be exactly zero: an empty mask means no
supervision, so the iteration is a no-op, not an arbitrary update.
"""
logits, x0 = _uniform_logits_case()
logits = logits.clone().requires_grad_(True)
t = torch.full((x0.shape[0],), 0.5)
loss = mdlm_loss(
logits, x0, x0.clone(), t, 12, 13, linear_schedule, reduction=reduction
)
assert float(loss.sum()) == 0.0
assert loss.grad_fn is not None, "empty-mask loss detached from the graph"
loss.sum().backward()
assert logits.grad is not None
assert bool((logits.grad == 0).all())
@pytest.mark.parametrize("reduction", ["none", "mean"])
def test_an_empty_batch_is_zero_rather_than_nan(reduction):
"""A zero-length batch reduces to zero, not to a NaN mean.
``mean`` over an empty tensor is NaN; the removed early return happened
to cover this, so the reduction is ``sum / max(B, 1)``, which is
identical to ``mean`` for every non-empty batch (asserted below).
"""
logits = torch.zeros(0, 8, 12, requires_grad=True)
empty = torch.zeros(0, 8, dtype=torch.long)
loss = mdlm_loss(
logits, empty, empty, torch.zeros(0), 12, 13, linear_schedule,
reduction=reduction,
)
assert bool(torch.isfinite(loss).all())
assert float(loss.sum()) == 0.0
def test_the_batch_reduction_is_exactly_the_per_sample_mean():
"""`sum / max(B, 1)` must be bit-identical to `mean` when B > 0."""
logits, x0 = _uniform_logits_case()
zt = x0.clone()
zt[:, ::2] = 12
t = torch.full((x0.shape[0],), 0.5)
args = (logits, x0, zt, t, 12, 13, linear_schedule)
per_sample = mdlm_loss(*args, reduction="none")
scalar = mdlm_loss(*args, reduction="mean")
assert torch.equal(scalar, per_sample.mean())
def test_the_supervised_auxiliary_goal_loss_is_the_mse_over_visible_rows():
"""The non-degenerate branch: MSE over the visible rows only.
Derivation, hand-computed. `find_staircase_from_glyphs` normalises to
(row/(H-1), col/(W-1)) on a 21x79 map, so a staircase at (0, 0) is the
target (0.0, 0.0) and one at (20, 78) is (1.0, 1.0); a row with no
staircase glyph is (-1, -1) and is excluded by `valid`. With predictions
(0.0, 0.5), (0.5, 1.0) and an arbitrary third row, the squared errors on
the two supervised rows are 0.25 and 0.25, and `diff.mean()` averages over
all 2 x 2 entries: (0.25 + 0.25) / 4 = 0.125.
This branch had no test of its own -- the only coverage of this function
was its degenerate branch below -- so its value was pinned by nothing
while the branch beside it was twice rewritten. The excluded row's
gradient is asserted too: exclusion is what the degenerate branch was
finally made consistent with.
"""
from src.diffusion.loss import auxiliary_goal_loss
global_obs = torch.zeros(3, 21, 79, dtype=torch.long)
global_obs[0, 0, 0] = 62 # '>' at the top-left -> (0.0, 0.0)
global_obs[1, 20, 78] = 62 # '>' at the bottom-right -> (1.0, 1.0)
# row 2 carries no staircase
goal_pred = torch.tensor(
[[0.0, 0.5], [0.5, 1.0], [9.0, 9.0]], requires_grad=True
)
loss = auxiliary_goal_loss(goal_pred, global_obs)
assert float(loss) == 0.125
loss.backward()
assert bool((goal_pred.grad[2] == 0).all()), "an unsupervised row was scored"
assert not bool((goal_pred.grad[:2] == 0).all()), "supervised rows got no gradient"
@pytest.mark.parametrize(
"poison", [None, float("nan"), float("inf"), float("-inf")]
)
def test_the_auxiliary_goal_loss_is_a_differentiable_zero_when_unsupervised(poison):
"""No visible staircase contributes zero, in the graph, for any prediction.
Three properties, all of them needed. The value must be exactly 0.0,
because this term is summed with the ELBO term before `backward()` and
anything else moves the whole loss. It must keep `grad_fn`: a freshly
allocated zero has the right number and no graph, and `backward()` then
raises for every arm whose other terms cannot carry the graph either --
the defect `0cfc632` fixed. And the gradient must be zero, because no
supervision means a no-op iteration, not an arbitrary update.
The `poison` cases are the regression for the NaN this branch returned
between `0cfc632` and its repair. It computed its zero as
`goal_pred * valid.unsqueeze(1)`, and `nan * False` is `nan`, so a
non-finite prediction gave NaN rather than zero -- while the supervised
branch, on the same input, excluded exactly those rows. Both branches now
select the same way.
"""
from src.diffusion.loss import auxiliary_goal_loss
goal_pred = torch.randn(4, 2)
if poison is not None:
goal_pred[0, 0] = poison
goal_pred = goal_pred.clone().requires_grad_(True)
no_staircase = torch.zeros(4, 21, 79, dtype=torch.long)
loss = auxiliary_goal_loss(goal_pred, no_staircase)
assert float(loss) == 0.0, f"non-finite goal_pred leaked: {float(loss)}"
assert loss.grad_fn is not None, "empty-supervision aux loss detached"
loss.backward()
assert bool((goal_pred.grad == 0).all())
# ---------------------------------------------------------------------------
# Reverse step: remasking schedules and the sigma bound
# ---------------------------------------------------------------------------
def test_sigma_strategies_closed_form_and_bound():
"""sigma_max = min(1, (1-alpha_s)/alpha_t); rescale = eta*sigma_max;
cap = min(eta, sigma_max); every sigma <= sigma_max.
Source: ReMDM eq (7) and Sec 4.1 (Max-Capped and Rescaled schedules).
Grid: linear schedule, K=10 reverse steps, eta=0.5.
"""
eta = 0.5
for k in range(1, 10):
alpha_t = 1 - k / 10
alpha_s = 1 - (k + 1) / 10
sigma_max = min(1.0, (1 - alpha_s) / alpha_t)
rescale = _compute_remask_prob("rescale", eta, sigma_max, None)
cap = _compute_remask_prob("cap", eta, sigma_max, None)
assert abs(rescale - eta * sigma_max) < ATOL
assert abs(cap - min(eta, sigma_max)) < ATOL
assert rescale <= sigma_max + ATOL and cap <= sigma_max + ATOL
def test_conf_strategy_softmax_of_stored_psi():
"""sigma_conf(l) = softmax(-psi)_l * eta * sigma_max over committed
positions, zero at masked ones; lower psi => higher remask probability.
Source: ReMDM Sec 4.1 (Confidence-Based Schedule): eta_conf =
exp(-psi_l)/sum exp(-psi_l'), with psi the decoding probability stored
when the token was last unmasked. Sum over committed positions =
eta * sigma_max.
"""
eta, sigma_max = 0.5, 0.8
psi = torch.tensor([[0.9, 0.2, float("inf"), 0.5]])
committed = torch.tensor([[True, True, False, True]])
sigma = _compute_remask_prob("conf", eta, sigma_max, psi, committed)
assert sigma[0, 2].item() == 0.0
assert sigma[0, 1] > sigma[0, 3] > sigma[0, 0], "lower psi must remask more"
assert abs(sigma[0, committed[0]].sum().item() - eta * sigma_max) < 1e-5
assert torch.all(sigma <= sigma_max + ATOL)
# ---------------------------------------------------------------------------
# Reverse chain behaviour (ReMDM Algorithm 1) via a deterministic stub
# ---------------------------------------------------------------------------
class _StubModel(torch.nn.Module):
"""Position-dependent peaked logits: argmax token = position % V."""
def __init__(self, seq_len: int, v: int):
super().__init__()
self.seq_len, self.v = seq_len, v
def forward(self, local_obs, global_obs, action_seq, t_discrete):
B = action_seq.shape[0]
logits = torch.full((B, self.seq_len, self.v), -5.0)
for pos in range(self.seq_len):
logits[:, pos, pos % self.v] = 5.0
return {"actions": logits}
def _stub_cfg(**over):
cfg = dict(
seq_len=64, mask_token=12, action_dim=12, num_diffusion_steps=100,
diffusion_steps_eval=4, temperature=1.0, top_p=1.0, eta=0.0,
remask_strategy="rescale", noise_schedule="linear", crop_size=9,
map_h=21, map_w=79,
)
cfg.update(over)
return SimpleNamespace(**cfg)
def test_carryover_committed_tokens_persist_when_sigma_zero():
"""With sigma = 0 a committed token is never changed or remasked.
Source: ReMDM Algorithm 1, z_t != m branch: Cat(z_s; (1-sigma) x_theta
+ sigma m) with x_theta carrying over unmasked inputs (MDLM Sec 3.2.3,
Carry-Over Unmasking). With sigma=0 (eta=0 rescale) the branch is the
identity. Uses the sampler's analytics trace across a 4-step chain.
"""
torch.manual_seed(0)
cfg = _stub_cfg()
model = _StubModel(cfg.seq_len, cfg.action_dim)
local = torch.zeros(2, 9, 9, dtype=torch.long)
glob = torch.zeros(2, 21, 79, dtype=torch.long)
_, path, _, _ = remdm_sample(
model, local, glob, cfg, "cpu", physics_aware=False,
return_analytics=True,
)
for earlier, later in zip(path, path[1:]):
committed = earlier != cfg.mask_token
assert (later[committed] == earlier[committed]).all(), (
"a committed token changed or was remasked despite sigma=0"
)
def test_locked_prefix_survives_the_full_chain():
"""Conditioned positions are fixed for the whole of denoising.
Source: planning-as-inpainting (Diffuser Sec. 3.3: conditioned
values are fixed throughout denoising) on top of ReMDM Alg 1;
spec-method §6.1/§6.2 are SHARED, and the craftax twin's
sample_plan_inpainting is pinned by the same assertion. Author
decision 2026-08-16 brought this repo into line.
A conf-strategy chain with eta > 0 is the hard case: remasking is
live, so a prefix that is merely written once would be eroded.
"""
torch.manual_seed(0)
cfg = _stub_cfg(seq_len=8, action_dim=5, mask_token=5,
remask_strategy="conf", eta=0.5, diffusion_steps_eval=6)
model = _StubModel(cfg.seq_len, cfg.action_dim)
B = 3
local = torch.zeros(B, 9, 9, dtype=torch.long)
glob = torch.zeros(B, 21, 79, dtype=torch.long)
history = torch.arange(cfg.seq_len).remainder(cfg.action_dim).repeat(B, 1)
hist_len = torch.tensor([0, 3, cfg.seq_len])
seq = remdm_sample(
model, local, glob, cfg, "cpu", physics_aware=False,
history=history, hist_len=hist_len,
)
assert (seq != cfg.mask_token).all(), "output contains MASK tokens"
assert (seq[1, :3] == history[1, :3]).all(), "prefix was overwritten"
assert (seq[2] == history[2]).all(), "fully locked plan changed"
def test_greedy_sampler_locks_the_prefix_too():
"""The DAgger collection sampler honours the same lock, so the data
the model trains on comes from history-conditioned plans."""
torch.manual_seed(0)
cfg = _stub_cfg(seq_len=8, action_dim=5, mask_token=5, diffusion_steps_eval=4)
model = _StubModel(cfg.seq_len, cfg.action_dim)
local = torch.zeros(2, 9, 9, dtype=torch.long)
glob = torch.zeros(2, 21, 79, dtype=torch.long)
history = torch.arange(cfg.seq_len).remainder(cfg.action_dim).repeat(2, 1)
hist_len = torch.tensor([2, 5])
seq = greedy_sample(
model, local, glob, cfg, "cpu", history=history, hist_len=hist_len
)
assert (seq != cfg.mask_token).all()
assert (seq[0, :2] == history[0, :2]).all()
assert (seq[1, :5] == history[1, :5]).all()
def test_locked_prefix_bookkeeping_rolls_the_window():
"""LockedPrefix records executed actions and opens a fresh window
once the plan is used up - the receding-horizon half of the
contract."""
from src.diffusion.sampling import LockedPrefix
prefix = LockedPrefix(n=2, seq_len=4, mask_token=9)
assert prefix.hist_len.tolist() == [0, 0]
for step, action in enumerate((1, 2, 3, 4)):
assert not prefix.is_full(0), f"window full early at step {step}"
prefix.record(0, action)
assert prefix.is_full(0)
assert prefix.history[0].tolist() == [1, 2, 3, 4]
prefix.start_window(np.array([0, 1]))
assert prefix.hist_len.tolist() == [0, 0]
assert prefix.history[0].tolist() == [9, 9, 9, 9]
# Row 1 never filled, so its (empty) window is untouched.
prefix.record(1, 7)
prefix.start_window(np.array([1]))
assert prefix.hist_len[1] == 1, "a partial window must not be reset"
def test_posterior_unmask_rate_first_step():
"""First reverse step (t=1 -> s=1/2, sigma=0, linear) unmasks each
masked token independently with p = (alpha_s - alpha_t)/(1 - alpha_t)
= 0.5.
Source: ReMDM Algorithm 1 approximate posterior, z_t = m branch.
Statistical: 128*64 = 8192 Bernoulli(0.5) draws; sigma = sqrt(0.25/8192)
= 0.0055; bound = 4 sigma = 0.0221. A MaskGIT count-based rule
would deterministically unmask exactly ceil(L/2) per row and, at the
old first step (k=1/K), only L/K tokens - both outside this bound.
"""
import numpy as np
torch.manual_seed(0)
cfg = _stub_cfg(seq_len=64, diffusion_steps_eval=2)
model = _StubModel(cfg.seq_len, cfg.action_dim)
local = torch.zeros(1, 9, 9, dtype=torch.long)
glob = torch.zeros(1, 21, 79, dtype=torch.long)
# The analytics trace records row 0 only, so run 128 single-row chains
# (independent draws from the shared global RNG): 128*64 = 8192 tokens.
seqs = []
for _ in range(128):
_, path, _, _ = remdm_sample(
model, local, glob, cfg, "cpu", physics_aware=False,
return_analytics=True,
)
seqs.append(path[0]) # state after the first reverse step
rate = float(np.mean([(s != cfg.mask_token).mean() for s in seqs]))
assert abs(rate - 0.5) < 0.0221
# ---------------------------------------------------------------------------
# Nucleus filtering
# ---------------------------------------------------------------------------
def test_top_p_filter_known_distribution():
"""Nucleus keeps the smallest prefix with cumulative mass >= p.
Source: ReMDM Sec 5 adopts nucleus sampling (Holtzman et al.): the
candidate set is the smallest prefix of the descending-sorted
distribution whose cumulative probability reaches p. For probs
[0.5, 0.3, 0.15, 0.05]: p=0.9 keeps {0,1,2} (cumulative 0.95 >= 0.9
reached at the third token); p=0.5 keeps {0} alone.
"""
logits = torch.log(torch.tensor([[0.5, 0.3, 0.15, 0.05]]))
kept_09 = top_p_filter(logits, 0.9).isfinite()
assert kept_09.tolist() == [[True, True, True, False]]
kept_05 = top_p_filter(logits, 0.5).isfinite()
assert kept_05.tolist() == [[True, False, False, False]]
# p >= 1 disables filtering
assert top_p_filter(logits, 1.0).isfinite().all()
def test_psi_is_the_raw_posterior_not_the_filtered_one():
"""psi is the model's probability for the token it commits, taken from
the raw posterior, before temperature and nucleus filtering.
Source: ReMDM Sec 4.1 defines psi as the decoding probability of the
token at the step it was unmasked — a property of the model, not of the
decoding settings. Read after `top_p_filter`, it is renormalised over the
nucleus, so whenever the nucleus collapses to one token psi is exactly
1.0 however uncertain the model is, and `sigma_conf` then never remasks
that position. The craftax twin takes the raw quantity
(`src/diffusion/sampling.py`, `probs = jax.nn.softmax(logits)`); this is
the shared canon, decided 2026-08-18.
Derivation: logits [4.0, 2.0, -6.0] at temperature 0.5 give [8, 4, -12];
softmax of that puts 0.9820 on index 0, so the top_p=0.9 exclusive-cumsum
nucleus keeps index 0 alone and its filtered probability is exactly 1.0.
The raw softmax of [4, 2, -6] is e^4/(e^4 + e^2 + e^-6) =
54.5982/61.9895 = 0.8808. psi must be 0.8808, not 1.0.
"""
from src.diffusion.sampling import top_p_filter
logits = torch.tensor([[[4.0, 2.0, -6.0]]])
temperature, top_p = 0.5, 0.9
filtered = torch.softmax(top_p_filter(logits / temperature, top_p), dim=-1)
raw = torch.softmax(logits, dim=-1)
chosen = torch.zeros(1, 1, dtype=torch.long)
assert filtered[0, 0, 0].item() == 1.0, "the nucleus did not collapse"
expected = math.exp(4.0) / (math.exp(4.0) + math.exp(2.0) + math.exp(-6.0))
assert abs(raw[0, 0, 0].item() - expected) < ATOL
assert abs(expected - 0.8808) < 1e-4
psi = raw.gather(-1, chosen.unsqueeze(-1)).squeeze(-1)
assert abs(psi.item() - expected) < ATOL
assert psi.item() < 1.0
# And the production sampler agrees: run one step and read the psi it
# stored, against a stub whose logits are the ones above.
from types import SimpleNamespace
from src.diffusion.sampling import remdm_sample
class _Peaked(torch.nn.Module):
def forward(self, local_obs, global_obs, seq, t_discrete):
b, length = seq.shape
row = torch.tensor([4.0, 2.0, -6.0])
return {
"actions": row.view(1, 1, 3).expand(b, length, 3).clone(),
"goal_pred": torch.zeros(b, 2),
}
cfg = SimpleNamespace(
seq_len=4,
mask_token=3,
action_dim=3,
diffusion_steps_eval=2,
temperature=temperature,
top_p=top_p,
eta=0.15,
remask_strategy="conf",
noise_schedule="linear",
num_diffusion_steps=10,
)
torch.manual_seed(0)
seq, _, confidences, _ = remdm_sample(
_Peaked(),
torch.zeros(1, 9, 9, dtype=torch.long),
torch.zeros(1, 21, 79, dtype=torch.long),
cfg,
"cpu",
physics_aware=False,
return_analytics=True,
)
assert confidences, "the sampler recorded no confidences"
assert all(c < 1.0 for c in confidences), (
f"psi saturated at 1.0 on a collapsed nucleus: {confidences}"
)