changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
8.66 kB
# SPDX-License-Identifier: Apache-2.0
"""The exact rewrites of ``tt/params.py`` and the input packing of ``tt/inputs.py`` against the CPU reference
(float64, no device): every constant the ttnn graph holds is the reference's math rearranged.
TT_VISIBLE_DEVICES=none python -m pytest -q -p fake_ttnn_plugin code/tt_diffusion_planner/tests/test_tt_params_host.py
"""
from __future__ import annotations
import numpy as np
import pytest
from tt_diffusion_planner.host import pipeline as hp
from tt_diffusion_planner.host.solver import apply_prefix_constraint, dpm_solver_sample
from tt_diffusion_planner.reference import config as C
from tt_diffusion_planner.tt import config as T
from tt_diffusion_planner.tt import inputs as I
from tt_diffusion_planner.tt import params as P
torch = pytest.importorskip("torch")
F = torch.nn.functional
PKG = __import__("pathlib").Path(__file__).resolve().parents[1]
@pytest.fixture(scope="module")
def weights():
from tt_diffusion_planner.reference.weights import find_weights_dir, load_weights
wd = find_weights_dir()
if wd is None:
pytest.skip("Diffusion Planner v5.0 weights not found")
return load_weights(wd)
@pytest.fixture(scope="module")
def prepared(weights):
from tt_diffusion_planner.ttaw.io import load_named_arrays
raw = load_named_arrays(str(PKG / "samples" / "kashiwanoha_dense.npz"), C.INPUT_SCHEMA)
return hp.prepare(raw, weights.normalization.observation)
def d(a):
return np.asarray(a, np.float64)
def test_lane_aux_is_the_speed_and_attribute_embeddings(weights):
p = weights.params
rng = np.random.default_rng(0)
for cat, mod in (("lane", "encoder.lane_encoder"), ("route", "encoder.route_encoder")):
e = 40
speed = rng.random(e).astype(np.float32) * 2
has = rng.random(e) < 0.6
speed[~has] = 0.0
attr = (rng.random((e, C.LANE_ATTRIBUTE_DIM)) < 0.2).astype(np.float32)
got = d(P.lane_aux_features(speed, has, attr)) @ d(P.lane_aux(p, cat))
lin = d(speed)[:, None] @ d(p[f"{mod}.speed_limit_emb.w"]) + d(p[f"{mod}.speed_limit_emb.b"])
want = np.where(has[:, None], lin, d(p[f"{mod}.unknown_speed_emb"])[None]) \
+ d(attr) @ d(p[f"{mod}.attribute_emb.w"]) + d(p[f"{mod}.attribute_emb.b"])
np.testing.assert_allclose(got, want, rtol=1e-12, atol=1e-12)
def test_neighbor_aux_and_pos_aug(weights, prepared):
p = weights.params
f = prepared.features
t = d(f.neighbor_type)
got = np.concatenate([t, np.ones((t.shape[0], 1))], 1) @ d(P.neighbor_aux(p))
want = t @ d(p["encoder.neighbor_encoder.type_emb.w"]) + d(p["encoder.neighbor_encoder.type_emb.b"])
np.testing.assert_allclose(got, want, rtol=1e-12, atol=1e-12)
inp = I.plan_inputs(prepared)
pos = d(inp["pos_aug"][0, 0]) @ d(P.pos_aug(p))
tv = np.asarray(f.token_valid, bool)
ref = (d(f.pos) @ d(p["encoder.pos_emb.w"]) + d(p["encoder.pos_emb.b"])) * tv[:, None]
np.testing.assert_allclose(pos[:T.TOKENS_REAL], ref, rtol=1e-12, atol=1e-12)
assert not pos[T.TOKENS_REAL:].any()
def test_agent_rows_and_masked_projection(weights):
p = weights.params
rows = P.agent_rows(p)
emb = d(p["decoder.dit.agent_embedding"])
b2 = d(p["decoder.dit.preproj.fc2.b"])
np.testing.assert_allclose(d(rows[0]), b2 + emb[0], rtol=1e-6)
np.testing.assert_allclose(d(rows[1:]), np.repeat((b2 + emb[1])[None], T.AGENTS - 1, 0), rtol=1e-6)
lin = P.final_projection_masked(p)
assert not lin.w[:, :4].any() and not lin.b[:4].any()
np.testing.assert_array_equal(lin.w[:, 4:], np.asarray(p["decoder.dit.final_layer.proj.4.w"])[:, 4:])
def test_turn_weights_reproduce_the_head(weights):
from tt_diffusion_planner.reference.model import TurnHead, torch_params
p = weights.params
rng = np.random.default_rng(1)
enc = rng.standard_normal((T.TOKENS_REAL, C.HIDDEN_DIM))
x0 = rng.standard_normal((C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM))
w = P.turn_weights(p)
got = x0[0].reshape(-1) @ d(w["w_sel"]) + enc.sum(0) @ d(w["w_pool"]) + d(w["b"])
head = TurnHead(torch_params(p, torch.float64))
want = head.forward(torch.from_numpy(enc), x0).numpy()
np.testing.assert_allclose(got, want, rtol=1e-6, atol=1e-6)
def test_step_tables_and_solver_recursion(weights):
"""The y-recursion with (A, B, Cm) and the masked model output equals the node's DPM-Solver++(2M) with the
prefix constraint (dpm_solver_sample) on a deterministic stand-in model."""
tab = P.step_tables(weights.params)
assert tab.nfe == C.DPM_SOLVER_STEPS + 1 and len(tab.solver) == C.DPM_SOLVER_STEPS
assert len(tab.blocks) == tab.nfe and len(tab.blocks[0]) == C.DIT_DEPTH
rng = np.random.default_rng(2)
agents = 5
W = rng.standard_normal((324, 324)).astype(np.float32) * 0.05
cs = rng.standard_normal((agents, 4)).astype(np.float32)
x_T = rng.standard_normal((agents, 81, 4)).astype(np.float32)
def model(x, t):
return np.tanh(x.reshape(agents, -1) @ W + float(t)).reshape(agents, 81, 4).astype(np.float32)
ref = dpm_solver_sample(x_T, model, lambda x: apply_prefix_constraint(x, cs))
mask = np.ones((81, 4), np.float32)
mask[0] = 0
y = (x_T * mask).astype(np.float64)
m_prev = None
iters = []
for k in range(tab.nfe):
x = y + np.concatenate([cs[:, None], np.zeros((agents, 80, 4))], 1)
if k > 0:
iters.append(x)
m = model(x.astype(np.float32), tab.eval_times[k]).astype(np.float64) * mask
if k == tab.nfe - 1:
break
a, b, c = tab.solver[k]
y = a * y - b * m + (c * m_prev if m_prev is not None else 0.0)
m_prev = m
final = m + np.concatenate([cs[:, None], np.zeros((agents, 80, 4))], 1)
iters.append(final)
np.testing.assert_allclose(final, ref.final_x, rtol=1e-5, atol=1e-5)
for got, want in zip(iters, ref.denoising_steps):
np.testing.assert_allclose(got, want, rtol=1e-5, atol=1e-5)
assert np.allclose(tab.eval_times, ref.eval_times)
def test_plan_inputs_layout(prepared):
inp = I.plan_inputs(prepared)
assert set(inp) == set(I.INPUT_SPECS)
f = prepared.features
np.testing.assert_array_equal(inp["neighbor_x"][0], f.neighbor[:, 25:31])
assert not f.neighbor[:, :25].any() and not f.ego[6:].any() # the island's zero rows
np.testing.assert_array_equal(inp["ego_x"][0, 0], f.ego[:6])
row = inp["fusion_key_row"][0, 0, 0]
assert np.isneginf(row[T.TOKENS_REAL:]).all() and (row[:T.TOKENS_REAL][f.key_valid] == 0).all()
arow = inp["agent_key_row"][0, 0, 0]
assert np.isneginf(arow[C.MAX_NUM_AGENTS:]).all() and arow[0] == 0
cs, y0 = inp["cs"][0, 0], inp["y0"][0, 0]
np.testing.assert_array_equal(cs[:C.MAX_NUM_AGENTS, :4], prepared.decoder.current_states)
assert not cs[:, 4:].any() and not y0[:, :4].any() and not cs[C.MAX_NUM_AGENTS:].any()
np.testing.assert_array_equal(y0[:C.MAX_NUM_AGENTS, 4:], prepared.x_T.reshape(C.MAX_NUM_AGENTS, -1)[:, 4:])
warm = I.warmup_inputs()
assert all(warm[k].shape == v for k, v in I.INPUT_SPECS.items())
def test_compact_bucket_covers_every_needed_row():
"""``COMPACT``: the bucket holds the ego, every valid self-attention key and every emitted neighbour row (the
rows past it are masked keys, never read back), and falls back to the full 352 rows."""
from types import SimpleNamespace
from tt_diffusion_planner.tt.model import bucket_rows, needed_rows
def prep(valid_idx, emitted):
v = np.zeros(C.MAX_NUM_AGENTS, bool)
v[0] = True
v[list(valid_idx)] = True
nb = np.zeros(C.MAX_NUM_NEIGHBORS, bool)
nb[[i - 1 for i in valid_idx]] = True
return SimpleNamespace(decoder=SimpleNamespace(agent_valid=v), neighbor_rows=np.asarray(emitted, int),
features=SimpleNamespace(valid={"neighbor": nb}))
b = (32, 64, 96, 128, 192)
assert needed_rows(prep([], [])) == 1 and bucket_rows(prep([], []), b) == 32
assert needed_rows(prep([1, 2, 31], [0, 1, 30])) == 32 and bucket_rows(prep([31], []), b) == 32
assert bucket_rows(prep([32], []), b) == 64 # the 33rd row is a valid key
assert bucket_rows(prep([5], [70]), b) == 96 # emitted neighbour 70 = decoder row 71
assert bucket_rows(prep([200], []), b) == T.AGENTS
assert bucket_rows(prep([5], []), ()) == T.AGENTS # COMPACT off
p = prep([3], [])
p.features.valid["neighbor"][40] = True # a valid neighbour token (row 41)
assert needed_rows(p) == 42 and bucket_rows(p, b) == 64