Download code/tt_diffusion_planner/tests/test_tt_params_host.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.66 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_tt_params_host.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_tt_params_host.py
-
curl -L -o test_tt_params_host.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_tt_params_host.py
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] | |
| 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) | |
| 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 | |