File size: 28,313 Bytes
e0eb79a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
"""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}"
    )