File size: 7,193 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
"""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 through remdm_sample, the final greedy cleanup, decode temperature
and label smoothing. 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 craftax twin file carries the same
assertions with the same inputs and tolerances.

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
from types import SimpleNamespace

import numpy as np
import pytest
import torch
from torch import nn

from src.diffusion.loss import mdlm_loss
from src.diffusion.sampling import remdm_sample
from src.diffusion.schedules import linear_schedule, linear_schedule_deriv


class _TimeCodedModel(nn.Module):
    """Stub whose argmax token encodes the decode time bin.

    token 0 for t_discrete > 90, token 1 for 50 < t_discrete <= 90,
    token 2 otherwise (num_diffusion_steps = 100 in the cfg below).
    Logit margin 30 makes categorical sampling deterministic to ~1e-13.
    """

    def __init__(self, seq_len: int, vocab: int):
        super().__init__()
        self.seq_len = seq_len
        self.vocab = vocab

    def forward(self, local_obs, global_obs, seq, t_discrete):
        td = int(t_discrete[0])
        idx = 0 if td > 90 else (1 if td > 50 else 2)
        logits = torch.full((seq.shape[0], self.seq_len, self.vocab), -30.0)
        logits[:, :, idx] = 30.0
        return {"actions": logits, "goal_pred": torch.zeros(seq.shape[0], 2)}


def _chain_cfg(eta: float) -> SimpleNamespace:
    return SimpleNamespace(
        seq_len=32, mask_token=3, action_dim=3, diffusion_steps_eval=3,
        temperature=1.0, top_p=1.0, eta=eta, remask_strategy="rescale",
        noise_schedule="linear", num_diffusion_steps=100,
    )


@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, eq (7) sigma_max, Sec 4.1
    rescale sigma = eta * sigma_max, 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 (sigma=0: committed tokens are
    never re-decided). Same derivation and bound as the craftax twin:
    N = 512*32 = 16384; max sigma = 0.0039; bound 0.02 = 5.1 sigma.
    """
    B, L = 512, 32
    cfg = _chain_cfg(eta)
    model = _TimeCodedModel(L, 4)
    seq = remdm_sample(
        model, torch.zeros(B, 9, 9), torch.zeros(B, 21, 79), cfg, "cpu",
        physics_aware=False,
    )
    freq = np.bincount(seq.numpy().ravel(), minlength=4) / (B * L)
    expected = np.array([(1 - eta) / 3, (1 + eta) / 3, 1 / 3, 0.0])
    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; the
    craftax twin asserts the same). The stub decodes token 2 at t=0,
    so the output must be all-2.
    """
    B, L = 8, 16
    cfg = _chain_cfg(0.0)
    cfg.seq_len = L
    seq = remdm_sample(
        _TimeCodedModel(L, 4), torch.zeros(B, 9, 9), torch.zeros(B, 21, 79),
        cfg, "cpu", physics_aware=False, num_steps=0,
    )
    assert (seq.numpy() == 2).all()


class _FixedLogitsModel(nn.Module):
    """Stub returning fixed logits [0, 1] over two real actions."""

    def __init__(self, seq_len: int):
        super().__init__()
        self.seq_len = seq_len

    def forward(self, local_obs, global_obs, seq, t_discrete):
        logits = torch.zeros(seq.shape[0], self.seq_len, 3)
        logits[:, :, 1] = 1.0
        logits[:, :, 2] = 0.0  # masked out by the sampler (>= action_dim)
        return {"actions": logits, "goal_pred": torch.zeros(seq.shape[0], 2)}


def test_decode_temperature_scales_logits_before_sampling():
    """Sampling frequencies follow softmax(logits / temperature).

    Source: spec-method 5.2. Derivation: logits [0, 1] at temperature
    0.5 give softmax([0, 2]) = [0.1192, 0.8808] over the two real
    actions; a single denoising step (K=1: t=1, s=0, p_unmask=1)
    commits every position in one draw. Statistical: 8192 draws;
    sigma = sqrt(0.8808*0.1192/8192) = 0.00358; bound 0.0143 = 4 sigma.
    Same numbers as the craftax twin.
    """
    B, L = 256, 32  # 8192 positions
    cfg = SimpleNamespace(
        seq_len=L, mask_token=2, action_dim=2, diffusion_steps_eval=1,
        temperature=0.5, top_p=1.0, eta=0.0, remask_strategy="rescale",
        noise_schedule="linear", num_diffusion_steps=100,
    )
    seq = remdm_sample(
        _FixedLogitsModel(L), torch.zeros(B, 9, 9), torch.zeros(B, 21, 79),
        cfg, "cpu", physics_aware=False,
    )
    p1 = math.exp(2) / (1 + math.exp(2))
    assert abs(seq.float().mean().item() - p1) < 0.0143


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,
    which is exactly torch.nn.functional.cross_entropy's label_smoothing
    semantics; 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.075.
    t pinned to 1 (linear, w=1, everything masked) makes the loss equal
    that CE exactly. Same inputs and expectations as the craftax twin.
    """
    B, L, V = 4, 8, 4
    logits = torch.log(torch.tensor([0.7, 0.1, 0.1, 0.1])).expand(B, L, V).clone()
    x0 = torch.zeros(B, L, dtype=torch.long)
    zt = torch.full((B, L), 5, dtype=torch.long)  # mask_token=5, pad=6
    t = torch.ones(B)
    expected = -0.775 * math.log(0.7) - 0.075 * 3 * math.log(0.1)
    got = mdlm_loss(
        logits, x0, zt, t, mask_token=5, pad_token=6,
        schedule_fn=linear_schedule, schedule_deriv_fn=linear_schedule_deriv,
        label_smoothing=0.3,
    )
    assert abs(float(got) - expected) < 1e-5
    got0 = mdlm_loss(
        logits, x0, zt, t, mask_token=5, pad_token=6,
        schedule_fn=linear_schedule, schedule_deriv_fn=linear_schedule_deriv,
        label_smoothing=0.0,
    )
    assert abs(float(got0) - (-math.log(0.7))) < 1e-5