File size: 8,963 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# SPDX-License-Identifier: Apache-2.0
"""The CPU reference without ONNX Runtime (no device; needs torch, onnx and the v5.0 weights, else skipped):

- it reproduces the small goldens of the shipped samples (stored by ``code/scripts/ref_golden.py``; the reference
  itself is proven against ONNX Runtime in ``test_reference_cpu.py``) and the stored ``/predict`` reference bodies;
- the exact rewrites the TT port applies give the as-exported results (``reference.rewrites``): per-step adaLN tables
  folded into the LayerNorm affine, hoisted cross K/V, the pad-relative pre-projection island; and the structural
  facts they rest on: uniform-time adaLN rows, bit-identical padded encoder tokens, decoder agent buckets.

    TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_reference_host.py
"""
from __future__ import annotations

import json

import numpy as np
import pytest

torch = pytest.importorskip("torch")
pytest.importorskip("onnx")

from tt_diffusion_planner.api import DiffusionPlanner  # noqa: E402
from tt_diffusion_planner.host import pipeline as hp  # noqa: E402
from tt_diffusion_planner.host.solver import solver_plan  # noqa: E402
from tt_diffusion_planner.reference import config as C  # noqa: E402
from tt_diffusion_planner.reference import rewrites as RW  # noqa: E402
from tt_diffusion_planner.reference.pipeline import ReferencePlanner  # noqa: E402
from tt_diffusion_planner.reference.weights import coverage, find_weights_dir, load_weights  # noqa: E402
from tt_diffusion_planner.tests import _research as R  # noqa: E402
from tt_diffusion_planner.ttaw.golden import TapRegistry  # noqa: E402
from tt_diffusion_planner.ttaw.metrics import ade_fde  # noqa: E402

WEIGHTS = find_weights_dir()
pytestmark = pytest.mark.skipif(WEIGHTS is None, reason="Diffusion Planner v5.0 weights not found "
                                "(set DIFFUSION_PLANNER_WEIGHTS_DIR)")
STEMS = ("kashiwanoha_dense", "straight_road")


@pytest.fixture(scope="module")
def ref():
    return ReferencePlanner(weights=load_weights(WEIGHTS), threads=4)


_RUNS: dict = {}


def _run(ref, stem):
    if stem not in _RUNS:
        taps = TapRegistry(include=["enc.encoding", "enc.tokens", "dec.0.*"])
        _RUNS[stem] = (ref.run(R.sample_raw(stem), taps=taps, keep_eval_io=True), taps.to_dict())
    return _RUNS[stem]


def _output(ref, stem):
    """``ReferencePlanner.__call__`` of the sample, from the cached run (the same two host halves)."""
    res, _ = _run(ref, stem)
    params = {k: spec[3] for k, spec in hp.RUNTIME_PARAMS.items()}
    return hp.make_output(res.final_x0, res.logit, res.prepared, ref.normalization, params,
                          model=DiffusionPlanner.MODEL_NAME)


def test_weight_loader_covers_the_export(ref):
    cov = coverage(ref.weights)
    assert cov["unused"] == [] and len(cov["duplicates"]) == 2


@pytest.mark.parametrize("stem", STEMS)
def test_reproduces_small_goldens(ref, stem):
    path = R.SMALL_GOLDENS / f"{stem}.outputs.npz"
    if not path.is_file():
        pytest.skip("small goldens not generated (code/scripts/ref_golden.py)")
    res, _ = _run(ref, stem)
    with np.load(path) as z:
        g = {k: z[k] for k in z.files}
    rows = g["final_x0.rows"]
    np.testing.assert_array_equal(rows, np.flatnonzero(res.prepared.decoder.agent_valid))
    # another torch build may round differently; the solver keeps such float32 noise small
    assert np.abs(res.final_x0[rows] - g["final_x0"]).max() < 1e-3
    np.testing.assert_allclose(res.logit, g["logit"], rtol=0, atol=1e-3)
    out = _output(ref, stem)
    assert out.turn_indicator["command"] == int(g["turn_command"])
    ade, fde = ade_fde(out.poses[:, :2], g["trajectory"][:, :2])
    assert ade < 1e-3 and fde < 5e-3


@pytest.mark.parametrize("stem", STEMS)
def test_stored_reference_bodies_match(ref, stem):
    """``samples/<stem>.reference.json`` (the smoke test's oracle) is this reference's ``/predict`` body."""
    path = R.SAMPLES / f"{stem}.reference.json"
    if not path.is_file():
        pytest.skip("reference body not generated")
    stored = json.loads(path.read_text())
    body = _output(ref, stem).to_dict()
    assert stored["model"] == body["model"] and stored["columns"] == body["columns"]
    assert stored["turn_indicator"]["command"] == body["turn_indicator"]["command"]
    a, b = np.asarray(stored["trajectory"]), np.asarray(body["trajectory"])
    assert a.shape == b.shape == (80, 7) and np.abs(a[:, :2] - b[:, :2]).max() < 1e-2


# -------------------------------------------------------------------------------------------- rewrites

def test_adaln_is_uniform_over_agents_and_folds_into_layernorm(ref):
    """With one diffusion time for every agent (multi-step mode) the t-embedding rows are identical, so the
    modulation is a per-step constant; the folded decoder equals the exported one at all 11 evaluation times."""
    res, taps = _run(ref, "straight_road")
    temb = taps["dec.0.temb"]
    assert np.array_equal(temb, np.broadcast_to(temb[0], temb.shape))
    plan = solver_plan(C.DPM_SOLVER_STEPS)
    np.testing.assert_allclose(res.eval_times, plan.eval_times, rtol=0, atol=0)
    tables = RW.adaln_tables(ref.weights.params, plan.eval_times)
    np.testing.assert_allclose(tables.temb[0], temb[0], rtol=0, atol=2e-5)
    kv = ref.decoder.cross_kv(torch.from_numpy(res.encoding))
    rows = np.flatnonzero(res.prepared.decoder.agent_valid)
    for k in range(len(plan.eval_times)):
        with torch.no_grad():
            want = ref.decoder.forward(res.eval_inputs[k], plan.eval_times[k], kv, res.prepared.decoder.agent_valid)
            got = RW.decoder_forward_folded(ref.decoder, tables, k, res.eval_inputs[k], kv,
                                            res.prepared.decoder.agent_valid)
        err = float((got - want)[rows].abs().max())
        assert err < 2e-5, f"evaluation {k}: folded decoder differs by {err}"


def test_cross_kv_hoisting_is_exact(ref):
    res, _ = _run(ref, "kashiwanoha_dense")
    hoisted = RW.cross_kv(ref.weights.params, res.encoding)
    for i, kv in enumerate(ref.decoder.cross_kv(torch.from_numpy(res.encoding))):
        np.testing.assert_allclose(hoisted[i], kv.numpy(), rtol=0, atol=1e-5)


@pytest.mark.parametrize("category", ["neighbor", "ego"])
def test_pad_relative_island_is_exact(ref, category):
    res, _ = _run(ref, "kashiwanoha_dense")
    f = res.prepared.features
    x = f.neighbor if category == "neighbor" else f.ego[None]
    cst = RW.island_constants(ref.weights.params, category)
    t1e, t2e = RW.island_forward_exported(ref.weights.params, category, x)
    t1p, t2p = RW.island_forward_pad_relative(ref.weights.params, cst, x)
    assert float((t1e - t1p).abs().max()) < 1e-5 and float((t2e - t2p).abs().max()) < 1e-5
    if category == "neighbor":  # a padded agent IS the pad agent
        pad = int(np.flatnonzero(~f.valid["neighbor"])[0])
        np.testing.assert_allclose(t2e[pad].numpy(), cst.t2_pad, rtol=0, atol=1e-6)
        # the offset dominates the valid agents' fc1 output (SPEC 4.6.7: ~97 % of the mean magnitude on straight,
        # 89 % on this denser scene)
        valid = np.flatnonzero(f.valid["neighbor"])
        share = float(np.abs(cst.t1_pad).mean() / t1e[valid].abs().mean())
        assert share > 0.8


def test_padded_tokens_are_identical(ref):
    """Every invalid entity enters the fusion as the zero vector and is masked as a key, so all padded rows of the
    encoding are bit-identical (SPEC 4.6.3: the basis of exact compaction with a pad multiplicity)."""
    res, taps = _run(ref, "kashiwanoha_dense")
    pad = np.flatnonzero(~res.prepared.features.token_valid)
    assert pad.size > 100
    assert not taps["enc.tokens"][pad].any()
    enc = res.encoding[pad]
    assert np.array_equal(enc, np.broadcast_to(enc[0], enc.shape))
    # turn-head mean pool = (sum of valid rows + m_pad * pad row) / 564
    valid = np.flatnonzero(res.prepared.features.token_valid)
    pooled = ((res.encoding[valid].astype(np.float64).sum(0) + pad.size * enc[0].astype(np.float64))
              / C.ENCODING_TOKEN_NUM)
    np.testing.assert_allclose(pooled, res.encoding.astype(np.float64).mean(0), rtol=0, atol=1e-6)


def test_decoder_agent_buckets_are_exact(ref):
    """The decoder on the first K >= 1 + valid neighbours agent rows gives the same outputs for those agents
    (padded agents are masked keys and nothing else mixes agents; SPEC 4.6.4)."""
    res, _ = _run(ref, "straight_road")
    nv = int(res.prepared.decoder.agent_valid.sum())
    kv = ref.decoder.cross_kv(torch.from_numpy(res.encoding))
    x, t = res.eval_inputs[3], res.eval_times[3]
    with torch.no_grad():
        full = ref.decoder.forward(x, t, kv, res.prepared.decoder.agent_valid).numpy()
        for k in sorted({nv, 16, 32, 64}):
            part = ref.decoder.forward(x[:k], t, kv, res.prepared.decoder.agent_valid[:k]).numpy()
            assert np.abs(part[:nv] - full[:nv]).max() < 2e-5, k