File size: 8,963 Bytes
4d9b003 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 | # 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
|