# 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