Download code/tt_diffusion_planner/tests/test_reference_cpu.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 15.3 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_reference_cpu.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_reference_cpu.py
-
curl -L -o test_reference_cpu.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_reference_cpu.py
15.3 kB
| # 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"} | |
| def weights(): | |
| return load_weights(WEIGHTS) | |
| def ref(weights): | |
| return ReferencePlanner(weights=weights, threads=THREADS) | |
| 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"]) | |
| 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 | |
| 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 | |
| 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 | |
| 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) | |
| 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 | |
| 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 | |
| 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 | |