File size: 7,534 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
"""Pins for the GPU-side performance optimisations.

Every change in that set is meant to be arithmetic-preserving: it removes
device syncs, kernel launches or PCIe traffic without touching the values
the training step computes. These tests hold that line, so a future edit
that quietly changes the maths fails here rather than in a training curve.

CPU-only, like the rest of the suite.
"""

from __future__ import annotations

import numpy as np
import pytest
import torch

from src.buffer import ReplayBuffer
from src.diffusion.loss import find_staircase_from_glyphs
from src.models.denoiser import ModelEMA, make_model
from src.planners.online import Trainer

STAIR_GLYPHS = (62, 2310, 2368, 2383)


def _reference_staircase(global_obs: torch.Tensor) -> torch.Tensor:
    """The unvectorised reference: a Python loop over the batch.

    Kept here as the specification of the vectorised version.
    """
    if global_obs.ndim == 2:
        global_obs = global_obs.unsqueeze(0)
    B, H, W = global_obs.shape
    is_stair = torch.zeros_like(global_obs, dtype=torch.bool)
    for g in STAIR_GLYPHS:
        is_stair |= global_obs == g
    coords = torch.full((B, 2), -1.0, dtype=torch.float32)
    for b in range(B):
        positions = is_stair[b].nonzero(as_tuple=False)
        if positions.shape[0] > 0:
            coords[b, 0] = positions[0, 0].float() / max(1, H - 1)
            coords[b, 1] = positions[0, 1].float() / max(1, W - 1)
    return coords


# ── Vectorised staircase lookup ─────────────────────────────


@pytest.mark.parametrize("glyph", STAIR_GLYPHS)
def test_every_staircase_glyph_variant_is_found(glyph):
    maps = torch.ones(1, 21, 79, dtype=torch.long)
    maps[0, 7, 13] = glyph

    coords = find_staircase_from_glyphs(maps)

    assert coords[0, 0].item() == pytest.approx(7 / 20)
    assert coords[0, 1].item() == pytest.approx(13 / 78)


def test_absent_staircase_yields_the_pad_sentinel():
    maps = torch.ones(4, 21, 79, dtype=torch.long)

    coords = find_staircase_from_glyphs(maps)

    assert torch.equal(coords, torch.full((4, 2), -1.0))


def test_first_staircase_in_row_major_order_wins():
    """Several staircases: the lowest flat index is the one reported."""
    maps = torch.ones(1, 21, 79, dtype=torch.long)
    maps[0, 9, 40] = 62
    maps[0, 3, 70] = 62  # earlier row, so this is the one that must win
    maps[0, 15, 2] = 62

    coords = find_staircase_from_glyphs(maps)

    assert coords[0, 0].item() == pytest.approx(3 / 20)
    assert coords[0, 1].item() == pytest.approx(70 / 78)


def test_matches_the_per_sample_loop_on_random_maps():
    torch.manual_seed(0)
    maps = torch.randint(0, 3000, (32, 21, 79))
    # Guarantee a mix of present and absent staircases across the batch.
    for b in range(0, 32, 3):
        maps[b, b % 21, (b * 7) % 79] = STAIR_GLYPHS[b % len(STAIR_GLYPHS)]

    assert torch.equal(find_staircase_from_glyphs(maps), _reference_staircase(maps))


def test_unbatched_input_is_accepted():
    maps = torch.ones(21, 79, dtype=torch.long)
    maps[4, 4] = 62

    coords = find_staircase_from_glyphs(maps)

    assert coords.shape == (1, 2)
    assert torch.equal(coords, _reference_staircase(maps))


def test_result_is_float32_regardless_of_input_dtype():
    """The buffer holds glyphs as int16; the goal loss expects float32."""
    for dtype in (torch.int16, torch.int32, torch.int64):
        maps = torch.ones(2, 21, 79, dtype=dtype)
        maps[0, 1, 1] = 62
        assert find_staircase_from_glyphs(maps).dtype == torch.float32


# ── Fused EMA update ────────────────────────────────────────


def test_fused_ema_matches_the_per_parameter_loop(tiny_cfg):
    torch.manual_seed(0)
    model = make_model(tiny_cfg)
    ema = ModelEMA(model, decay=tiny_cfg.ema_decay)
    reference = {n: p.data.clone() for n, p in model.named_parameters()}
    decay = tiny_cfg.ema_decay

    for _ in range(5):
        with torch.no_grad():
            for p in model.parameters():
                p.add_(torch.randn_like(p) * 0.01)
        ema.update(model)
        with torch.no_grad():
            for name, p in model.named_parameters():
                reference[name].mul_(decay).add_(p.data, alpha=1.0 - decay)

    for name, expected in reference.items():
        assert torch.equal(ema._shadow[name], expected), name


def test_ema_rebuilds_its_cache_for_a_different_model(tiny_cfg):
    """The cached operand lists must not leak between source models."""
    torch.manual_seed(0)
    model_a = make_model(tiny_cfg)
    model_b = make_model(tiny_cfg)
    ema = ModelEMA(model_a, decay=0.5)

    # Same update sequence applied to a plain dict, as the reference.
    shadow = {n: p.data.clone() for n, p in model_a.named_parameters()}
    for source in (model_a, model_b):
        ema.update(source)
        for name, p in source.named_parameters():
            shadow[name] = shadow[name] * 0.5 + p.data * 0.5

    for name, expected in shadow.items():
        assert torch.allclose(ema._shadow[name], expected), name


# ── Device-side step metrics ────────────────────────────────


def _trainer(cfg, trajectory=None, device="cpu"):
    model = make_model(cfg).to(device)
    buffer = ReplayBuffer(cfg.buffer_capacity, cfg.seq_len, cfg.pad_token)
    if trajectory is not None:
        buffer.add(trajectory)
    return Trainer(
        model,
        ModelEMA(model, decay=cfg.ema_decay),
        torch.optim.AdamW(model.parameters(), lr=cfg.dagger_lr),
        None,
        buffer,
        collector=None,
        evaluator=None,
        log=None,
        cfg=cfg,
        device=device,
        raw_model=model,
    )


def test_device_step_returns_tensors_and_the_wrapper_returns_floats(
    tiny_cfg, tiny_trajectory
):
    trainer = _trainer(tiny_cfg, tiny_trajectory)
    trainer.model.train()

    device_metrics = trainer._train_step_device()
    float_metrics = trainer._train_step()

    assert set(device_metrics) == set(float_metrics)
    for key, value in device_metrics.items():
        assert isinstance(value, torch.Tensor), key
        assert value.ndim == 0, key
        assert not value.requires_grad, key
    for key, value in float_metrics.items():
        assert isinstance(value, float), key


def test_empty_buffer_step_is_a_no_op_in_both_forms(tiny_cfg):
    trainer = _trainer(tiny_cfg)

    assert trainer._train_step() == {
        "loss": 0.0,
        "loss_diff": 0.0,
        "loss_aux": 0.0,
        "grad_norm": 0.0,
    }
    assert all(
        float(v) == 0.0 for v in trainer._train_step_device().values()
    )


# ── The int16 transfer widens exactly ───────────────────────


def test_buffer_glyphs_widen_from_int16_without_changing_values(tiny_cfg):
    """C1 sends glyph maps as int16 and widens on the device instead.

    Both are exact for the glyph range, so the widened values must equal
    the CPU-side cast the code used to do.
    """
    glyphs = np.array(
        [[0, 1, 62, 2310, 2368, 2383, 5999, np.iinfo(np.int16).max]],
        dtype=np.int16,
    )

    widened_on_device = torch.from_numpy(glyphs).to("cpu").long()
    widened_on_host = torch.from_numpy(glyphs).long().to("cpu")

    assert torch.equal(widened_on_device, widened_on_host)
    assert widened_on_device.tolist() == glyphs.astype(np.int64).tolist()