Download code/tt_diffusion_planner/tests/test_reference_host.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.96 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_reference_host.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_reference_host.py
-
curl -L -o test_reference_host.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_reference_host.py
8.96 kB
| # 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") | |
| 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 | |
| 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 | |
| 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) | |
| 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 | |