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