# 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/.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