changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
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"}
@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