File size: 15,263 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
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
# SPDX-License-Identifier: Apache-2.0
"""The fp32 CPU reference against ONNX Runtime on the deployed v5.0 ONNX files (PLAN.md 0.3 item 1).

    # research venv (torch, onnx, onnxruntime, pytest), from the bundle root:
    PYTHONPATH=code OMP_NUM_THREADS=4 tools/research-venv/bin/python -m pytest -q -p no:cacheprovider \
        code/tt_diffusion_planner/tests/test_reference_cpu.py

Skipped where onnxruntime or the weights are absent. The model has no data-dependent selection (no top-k, no NMS),
so every comparison is a PCC / max-error check; PCC thresholds are >= 0.9999 on the valid rows of every tap:

- the weight loader re-labels every initializer of the three files exactly once (by consuming node);
- the host features (``host.features``) equal the encoder graph's own pre-processing tensors (masks exactly);
- every encoder module tap (mixer trunks, category outputs, fusion input, 6 fusion blocks, encoding);
- every decoder evaluation, teacher-forced on ORT's own solver inputs and encoding (t-embedding, pre-projection,
  3 DiT blocks, model output), and the turn head on ORT's encoding / final x0;
- end to end (free-running solver on both sides): final x0, logits, turn command, the post-processed ego
  trajectory and the predicted neighbour paths; and the stored research goldens of ``dp_reference.py``.
"""
from __future__ import annotations

from pathlib import Path

import numpy as np
import pytest

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

from tt_diffusion_planner.host import pipeline as hp  # noqa: E402
from tt_diffusion_planner.reference import config as C  # noqa: E402
from tt_diffusion_planner.reference.ort import OrtPlanner  # 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.ttaw.golden import TapRegistry  # noqa: E402
from tt_diffusion_planner.ttaw.metrics import ade_fde, pcc  # 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)")

PKG = Path(__file__).resolve().parents[1]
RESEARCH = PKG.parents[3] / "research" / "diffusion-planner"
PCC_MIN = 0.9999
THREADS = 4


def _scenes():
    """``{scene: raw inputs}``: the shipped samples, plus the ORT golden scenes of the research directory (when present,
    deduplicated by content)."""
    out, seen = {}, set()
    for p in sorted((PKG / "samples").glob("*.npz")):
        with np.load(p) as z:
            out[p.stem] = {k: z[k] for k in C.INPUT_NAMES}
    for p in sorted((RESEARCH / "ort").glob("golden_*.npz")):
        with np.load(p) as z:
            raw = {k: z["raw/" + k] for k in C.INPUT_NAMES}
        key = b"".join(raw[k].tobytes() for k in C.INPUT_NAMES)
        if any(key == b"".join(v[k].tobytes() for k in C.INPUT_NAMES) for v in out.values()):
            continue
        if key not in seen:
            seen.add(key)
            out[p.stem[len("golden_"):]] = raw
    return out


SCENES = _scenes()
MIXERS = {"ego": "ego_encoder", "neighbor": "neighbor_encoder", "lane": "lane_encoder", "route": "route_encoder",
          "polygon": "polygon_encoder", "line_string": "line_string_encoder"}
HOST_TAPS = {  # host feature -> encoder graph tensor (batch dim dropped where the graph keeps it)
    "ego": "/encoder/Concat_output_0", "neighbor": "/encoder/neighbor_encoder/Where_output_0",
    "neighbor_type": "/encoder/neighbor_encoder/Reshape_3_output_0",
    "static": "/encoder/static_encoder/Where_output_0",
    "lane": "/encoder/lane_encoder/Where_2_output_0", "lane_attr": "/encoder/lane_encoder/Reshape_5_output_0",
    "lane_speed": "/encoder/lane_encoder/Reshape_3_output_0",
    "route": "/encoder/route_encoder/Where_2_output_0", "route_attr": "/encoder/route_encoder/Reshape_5_output_0",
    "route_speed": "/encoder/route_encoder/Reshape_3_output_0",
    "polygon": "/encoder/polygon_encoder/Where_2_output_0",
    "line_string": "/encoder/line_string_encoder/Where_2_output_0", "turn": "/encoder/Slice_4_output_0",
}
MASK_TAPS = {"token_invalid": "/encoder/Concat_4_output_0", "key_invalid": "/encoder/fusion/Concat_output_0",
             "pos": "/encoder/Concat_5_output_0"}
ENC_TAPS = {**{f"enc.{c}.pre": f"/encoder/{m}/Transpose_1_output_0" for c, m in MIXERS.items()},
            **{f"enc.{c}.mixer": f"/encoder/{m}/blocks.{C.MIXER_DEPTH - 1}/Add_1_output_0" for c, m in MIXERS.items()},
            "enc.categories": "/encoder/Concat_3_output_0", "enc.tokens": "/encoder/Add_1_output_0",
            **{f"enc.fusion.{i}": f"/encoder/fusion/blocks.{i}/Add_1_output_0" for i in range(C.FUSION_DEPTH)}}
DEC_TAPS = {"temb": "/dit/t_embedder/fc2/Add_output_0", "x": "/dit/Add_output_0",
            **{f"block{i}": f"/dit/blocks.{i}/Add_8_output_0" for i in range(C.DIT_DEPTH)},
            "agent_invalid": "/dit/Concat_6_output_0"}


@pytest.fixture(scope="module")
def weights():
    return load_weights(WEIGHTS)


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


@pytest.fixture(scope="module")
def ort_taps():
    return OrtPlanner(WEIGHTS, threads=THREADS, encoder_taps=list(HOST_TAPS.values()) + list(MASK_TAPS.values())
                      + list(ENC_TAPS.values()), decoder_taps=list(DEC_TAPS.values()),
                      turn_taps=["/ReduceMean_output_0"])


@pytest.fixture(scope="module")
def ort_plain():
    return OrtPlanner(WEIGHTS, threads=THREADS)


_CACHE: dict = {}


def _ref_run(ref, scene):
    key = ("ref", scene)
    if key not in _CACHE:
        taps = TapRegistry()
        _CACHE[key] = (ref.run(SCENES[scene], taps=taps, keep_eval_io=True), taps.to_dict())
    return _CACHE[key]


def _ort_run(ort_taps, scene):
    key = ("ort", scene)
    if key not in _CACHE:
        _CACHE[key] = ort_taps.run(SCENES[scene], keep_eval_io=True)
    return _CACHE[key]


def _check_pcc(name, test, ref_arr, rows=None, pcc_min=PCC_MIN):
    t, r = np.asarray(test, np.float64), np.asarray(ref_arr, np.float64)
    if rows is not None:
        t, r = t[rows], r[rows]
    if t.size == 0:
        return None
    v = pcc(t, r)
    err = float(np.abs(t - r).max())
    assert v >= pcc_min, f"{name}: PCC {v:.7f} < {pcc_min} (max |err| {err:.3g})"
    return v, err


# ------------------------------------------------------------------------------------------------- weights

def test_weights_cover_every_initializer(weights):
    cov = coverage(weights)
    assert cov["unused"] == [], cov["unused"][:10]
    # the only shared tensors are the two biases the export itself deduplicated (SPEC 6.2)
    assert cov["duplicates"] == [("encoder.route_encoder.attribute_emb.b", "encoder.route_encoder.speed_limit_emb.b"),
                                 ("encoder.static_encoder.projection.fc1.b", "encoder.static_encoder.projection.fc2.b")]
    assert weights.num_parameters() == 14_545_305  # 14,544,921 float initializers + the two deduplicated biases
    assert weights.facts["encoder"]["nodes"] == 1530 and weights.facts["decoder"]["nodes"] == 402


# ------------------------------------------------------------------------------------------ host features

@pytest.mark.parametrize("scene", sorted(SCENES))
def test_host_features_match_graph(ref, ort_taps, scene):
    res, _ = _ref_run(ref, scene)
    o = _ort_run(ort_taps, scene)
    f = res.prepared.features
    for name, tensor in HOST_TAPS.items():
        got = np.asarray(getattr(f, name), np.float32)
        want = o.taps[f"enc:{tensor}"].reshape(got.shape)
        np.testing.assert_array_equal(got, want, err_msg=f"host feature {name} != {tensor}")
    np.testing.assert_array_equal(~f.token_valid, o.taps["enc:/encoder/Concat_4_output_0"].reshape(-1))
    np.testing.assert_array_equal(~f.key_valid, o.taps["enc:/encoder/fusion/Concat_output_0"].reshape(-1))
    pos = o.taps["enc:/encoder/Concat_5_output_0"].reshape(f.pos.shape)
    np.testing.assert_allclose(f.pos[f.token_valid], pos[f.token_valid], rtol=0, atol=2e-6)
    np.testing.assert_array_equal(~res.prepared.decoder.agent_valid, o.taps["dec0:/dit/Concat_6_output_0"].reshape(-1))
    for k in C.INPUT_NAMES:  # normalization: bit-exact with the ORT path's numpy port of preprocessing_utils.cpp
        np.testing.assert_array_equal(res.prepared.norm[k], o.norm[k], err_msg=k)


# --------------------------------------------------------------------------------------------- encoder

@pytest.mark.parametrize("scene", sorted(SCENES))
def test_encoder_taps_match_ort(ref, ort_taps, scene):
    res, taps = _ref_run(ref, scene)
    o = _ort_run(ort_taps, scene)
    f = res.prepared.features
    report = {}
    for c in MIXERS:
        rows = np.flatnonzero(f.valid[c])
        for part in ("pre", "mixer"):
            name = f"enc.{c}.{part}"
            want = o.taps[f"enc:{ENC_TAPS[name]}"].reshape(taps[name].shape)
            report[name] = _check_pcc(name, taps[name], want, rows)
    cats = o.taps["enc:/encoder/Concat_3_output_0"][0]
    for name, sl in C.TOKEN_SLICES.items():
        rows = np.flatnonzero(f.valid[name])
        report[f"enc.{name}"] = _check_pcc(f"enc.{name}", taps[f"enc.{name}"], cats[sl], rows)
        # invalid entities are exactly zero on both sides
        inv = np.flatnonzero(~f.valid[name])
        assert np.all(taps[f"enc.{name}"][inv] == 0) and np.all(cats[sl][inv] == 0), name
    rows = np.flatnonzero(f.token_valid)
    report["enc.tokens"] = _check_pcc("enc.tokens", taps["enc.tokens"], o.taps["enc:/encoder/Add_1_output_0"][0], rows)
    for i in range(C.FUSION_DEPTH):
        name = f"enc.fusion.{i}"
        report[name] = _check_pcc(name, taps[name], o.taps[f"enc:{ENC_TAPS[name]}"][0], rows)
    report["enc.encoding"] = _check_pcc("enc.encoding", res.encoding, o.encoding[0], rows)
    # padded tokens leave the encoder bit-identical to each other (SPEC 4.6.3): one pad row stands for all
    pad = np.flatnonzero(~f.token_valid)
    if pad.size > 1:
        assert np.array_equal(res.encoding[pad], np.broadcast_to(res.encoding[pad[0]], res.encoding[pad].shape))
    print(scene, {k: (round(v[0], 8), f"{v[1]:.2e}") for k, v in report.items() if v})


# --------------------------------------------------------------------------------------------- decoder

@pytest.mark.parametrize("scene", sorted(SCENES))
def test_decoder_evaluations_teacher_forced(ref, ort_taps, scene):
    """Each of the 11 decoder calls of ORT's own solver run, replayed through the reference decoder with ORT's
    inputs (x, t) and ORT's encoding: the decoder is checked in isolation, at every diffusion time."""
    res, _ = _ref_run(ref, scene)
    o = _ort_run(ort_taps, scene)
    rows = np.flatnonzero(res.prepared.decoder.agent_valid)
    enc = torch.from_numpy(o.encoding[0])
    kv = ref.decoder.cross_kv(enc)
    assert len(o.eval_inputs) == C.DPM_SOLVER_STEPS + 1 == len(o.eval_times)
    worst = 1.0
    for k, (x, t) in enumerate(zip(o.eval_inputs, o.eval_times)):
        taps = TapRegistry()
        with torch.no_grad():
            out = ref.decoder.forward(x, t, kv, res.prepared.decoder.agent_valid, taps, prefix="d").numpy()
        for name in ("temb", "x", "block0", "block1", "block2"):
            v = _check_pcc(f"dec.{k}.{name}", taps[f"d.{name}"], o.taps[f"dec{k}:{DEC_TAPS[name]}"][0], rows)
            worst = min(worst, v[0])
        v = _check_pcc(f"dec.{k}.out", out, o.eval_outputs[k], rows)
        worst = min(worst, v[0])
        assert float(np.abs(out[rows] - o.eval_outputs[k][rows]).max()) < 1e-3, f"eval {k}"
    print(scene, "worst decoder PCC", worst)


@pytest.mark.parametrize("scene", sorted(SCENES))
def test_turn_head_matches_ort(ref, ort_taps, scene):
    o = _ort_run(ort_taps, scene)
    taps = TapRegistry()
    with torch.no_grad():
        logit = ref.turn.forward(torch.from_numpy(o.encoding[0]), o.final_x0[0], taps).numpy()
    # float32 mean of 564 rows: the summation order differs (torch vs ORT ReduceMean), ~4e-6 on values ~0.5
    _check_pcc("turn.pool", taps["turn.pool"], o.taps["turn:/ReduceMean_output_0"][0])
    np.testing.assert_allclose(taps["turn.pool"], o.taps["turn:/ReduceMean_output_0"][0], rtol=0, atol=2e-5)
    np.testing.assert_allclose(logit, o.logit[0], rtol=0, atol=2e-4)


# ------------------------------------------------------------------------------------------ end to end

@pytest.mark.parametrize("scene", sorted(SCENES))
def test_end_to_end_matches_ort(ref, ort_plain, scene):
    """Free-running reference vs free-running ORT (``ORT_ENABLE_ALL``, as the node): the DPM loop amplifies nothing
    beyond float32 noise."""
    res, _ = _ref_run(ref, scene)
    o = ort_plain.run(SCENES[scene])
    rows = np.flatnonzero(res.prepared.decoder.agent_valid)
    _check_pcc("final_x0", res.final_x0, o.final_x0[0], rows)
    assert float(np.abs(res.final_x0[rows] - o.final_x0[0][rows]).max()) < 1e-3
    np.testing.assert_allclose(res.logit, o.logit[0], rtol=0, atol=1e-3)
    np.testing.assert_array_equal(np.asarray(res.denoising_timesteps), np.asarray(o.denoising_timesteps))
    params = {k: spec[3] for k, spec in hp.RUNTIME_PARAMS.items()}
    a = hp.make_output(res.final_x0, res.logit, res.prepared, ref.normalization, params)
    b = hp.make_output(o.final_x0[0], o.logit[0], res.prepared, ref.normalization, params)
    assert a.turn_indicator["command"] == b.turn_indicator["command"]
    ade, fde = ade_fde(a.poses[:, :2], b.poses[:, :2])
    assert ade < 1e-3 and fde < 5e-3, (ade, fde)
    assert np.abs(a.poses - b.poses).max() < 5e-2  # velocity / acceleration are finite differences of positions
    if a.predicted_agents.size:
        assert np.abs(a.predicted_agents[..., :2] - b.predicted_agents[..., :2]).max() < 1e-2


@pytest.mark.skipif(not (RESEARCH / "ort").is_dir(), reason="research goldens not present")
def test_matches_research_goldens(ref):
    """The stored ``dp_reference.py`` goldens (ORT on the shipped ONNX, normalization by ``dp_common.py``)."""
    n = 0
    by_content = {b"".join(v[k].tobytes() for k in C.INPUT_NAMES): name for name, v in SCENES.items()}
    for p in sorted((RESEARCH / "ort").glob("golden_*.npz")):
        z = np.load(p)
        raw = {k: z["raw/" + k] for k in C.INPUT_NAMES}
        scene = by_content.get(b"".join(raw[k].tobytes() for k in C.INPUT_NAMES))
        res = _ref_run(ref, scene)[0] if scene else ref.run(raw)
        for k in C.INPUT_NAMES:
            if k in ("delay",):
                continue
            np.testing.assert_array_equal(res.prepared.norm[k], z["in/" + k], err_msg=f"{p.name}: in/{k}")
        rows = np.flatnonzero(res.prepared.decoder.agent_valid)
        _check_pcc(f"{p.name}: encoding", res.encoding, z["encoding"][0],
                   np.flatnonzero(res.prepared.features.token_valid))
        assert float(np.abs(res.final_x0[rows] - z["final_x_normalized"][0][rows]).max()) < 1e-3, p.name
        np.testing.assert_allclose(res.logit, z["logit_multi"][0], rtol=0, atol=1e-3, err_msg=p.name)
        np.testing.assert_array_equal(np.asarray(res.denoising_timesteps, np.float32), z["denoising_timesteps"])
        n += 1
    assert n >= 1