Download code/tt_diffusion_planner/tt/params.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/params.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/params.py
-
curl -L -o params.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/params.py
11.4 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Host-side (numpy) preparation of every tensor the ttnn graph holds: the canonical weights of | |
| ``reference.weights`` plus the exact rewrites of the port, folded in float64 and rounded to float32 once. | |
| No ttnn here (host-testable); the modules of :mod:`.encoder` / :mod:`.decoder` upload what they need with their | |
| precision. Rewrites (each exact in real arithmetic; ``tests/test_tt_params_host.py`` proves them against the CPU | |
| reference): | |
| - **pad-relative island** (ego / neighbour ``channel_pre`` + ``token_pre``): ``reference.rewrites.island_constants``; | |
| - **neighbour type embedding** ``n + type @ W_t + b_t`` = ``n + [type, 1] @ [W_t; b_t]`` (:func:`neighbor_aux`); | |
| - **lane / route speed + attribute embeddings**: ``where(has, speed * w_s + b_s, unk) + attr @ W_a + b_a`` with | |
| ``has`` in {0, 1} = ``[speed * has, has, 1 - has, attr, 1] @ [w_s; b_s; unk; W_a; b_a]`` (:func:`lane_aux`); | |
| - **masked positional embedding** ``valid * (pos @ W + b)`` with invalid pos rows already zero = | |
| ``[pos, valid] @ [W; b]`` (:func:`pos_aug`); | |
| - **agent embedding** folded into the pre-projection bias rows: ``preproj.fc2(h) + emb[ego / neighbour]`` = | |
| ``h @ W2 + rows`` with ``rows[0] = b2 + emb[0]``, ``rows[i > 0] = b2 + emb[1]`` (:func:`agent_rows`); | |
| - **masked last projection**: the t = 0 output columns (0..3) of ``final_layer.proj.4`` zeroed, so the model output | |
| is ``m * mask0`` (the prefix constraint overwrites that slot anyway; ``tt/config.py`` solver state); | |
| - **turn head** ``W [272 -> 5]`` over ``(final_x0[0, 1::10, :2], mean(encoding))`` = ``x0_row0 @ W_sel + | |
| sum_tokens(encoding) @ (W_pool / 564) + b`` with ``W_sel`` the 16 used rows scattered to ``[324, 5]`` | |
| (:func:`turn_weights`); | |
| - **per-step adaLN tables** folded into the LayerNorm affine (``reference.rewrites.adaln_tables``) and the solver | |
| update ``x' = a x - b m0 - c (m0 - m1) / r0`` as ``y' = A y - B m0 + Cm m1`` with ``A = a``, ``B = b + c / r0``, | |
| ``Cm = c / r0`` (:func:`solver_coefficients`; float64, rounded once). | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Dict, List, Mapping, Optional, Tuple | |
| import numpy as np | |
| from ..host.solver import SolverPlan, solver_plan | |
| from ..reference import config as C | |
| from ..reference import rewrites as R | |
| from . import config as T | |
| __all__ = ["Lin", "Norm", "linear", "norm", "mixer_module", "neighbor_aux", "lane_aux", "lane_aux_features", | |
| "pos_aug", "agent_rows", "final_projection_masked", "turn_weights", "solver_coefficients", "StepTables", | |
| "step_tables", "island", "small_module", "pad_rows", "MIXER_MODULES", "SMALL_MODULES"] | |
| MIXER_MODULES = {"ego": "encoder.ego_encoder", "neighbor": "encoder.neighbor_encoder", | |
| "lane": "encoder.lane_encoder", "route": "encoder.route_encoder", | |
| "polygon": "encoder.polygon_encoder", "line_string": "encoder.line_string_encoder"} | |
| SMALL_MODULES = {"goal": "encoder.goal_pose_encoder", "ego_shape": "encoder.ego_shape_encoder", | |
| "turn": "encoder.turn_indicator_encoder"} | |
| f32 = np.float32 | |
| class Lin: | |
| """``y = x @ w + b``; ``w`` [in, out] float32, ``b`` [out] or None.""" | |
| w: np.ndarray | |
| b: Optional[np.ndarray] | |
| class Norm: | |
| gamma: np.ndarray | |
| beta: np.ndarray | |
| def _a(p: Mapping[str, np.ndarray], name: str) -> np.ndarray: | |
| return np.ascontiguousarray(np.asarray(p[name], f32)) | |
| def linear(p: Mapping[str, np.ndarray], name: str, *, bias: bool = True) -> Lin: | |
| return Lin(_a(p, f"{name}.w"), _a(p, f"{name}.b") if bias else None) | |
| def norm(p: Mapping[str, np.ndarray], name: str) -> Norm: | |
| return Norm(_a(p, f"{name}.gamma"), _a(p, f"{name}.beta")) | |
| def mixer_module(p: Mapping[str, np.ndarray], cat: str) -> Dict[str, object]: | |
| """Weights of one MLP-Mixer category: ``channel_pre`` / ``token_pre`` (fc1, fc2), 6 blocks (norm1, tokens fc1 / | |
| fc2, norm2, channels fc1 / fc2), ``norm``, ``emb_project`` (fc1, fc2).""" | |
| N = MIXER_MODULES[cat] | |
| out: Dict[str, object] = { | |
| "c1": linear(p, f"{N}.channel_pre_project.fc1"), "c2": linear(p, f"{N}.channel_pre_project.fc2"), | |
| "t1": linear(p, f"{N}.token_pre_project.fc1"), "t2": linear(p, f"{N}.token_pre_project.fc2"), | |
| "norm": norm(p, f"{N}.norm"), "e1": linear(p, f"{N}.emb_project.fc1"), "e2": linear(p, f"{N}.emb_project.fc2"), | |
| "blocks": [{"n1": norm(p, f"{N}.blocks.{i}.norm1"), "tk1": linear(p, f"{N}.blocks.{i}.tokens_mlp.fc1"), | |
| "tk2": linear(p, f"{N}.blocks.{i}.tokens_mlp.fc2"), "n2": norm(p, f"{N}.blocks.{i}.norm2"), | |
| "ch1": linear(p, f"{N}.blocks.{i}.channels_mlp.fc1"), | |
| "ch2": linear(p, f"{N}.blocks.{i}.channels_mlp.fc2")} for i in range(C.MIXER_DEPTH)], | |
| } | |
| return out | |
| def island(p: Mapping[str, np.ndarray], cat: str) -> Dict[str, np.ndarray]: | |
| """Pad-relative island of ``cat`` in {"ego", "neighbor"} (``reference.rewrites``, probe P12): ``c1`` (fc1 of | |
| channel_pre, with bias), ``gelu_b1`` [1, 128], ``c2_w`` [128, 128] (no bias: deviation), ``t1_w`` = W_t1 of the | |
| 6 kept rows [6, 64], ``t1_pad`` / ``g_pad`` / ``t2_pad`` [128, 64], ``t2_w`` [64, 64].""" | |
| N = MIXER_MODULES[cat] | |
| cst = R.island_constants(p, cat) | |
| if tuple(cst.valid_rows) != T.ISLAND_ROWS[cat]: | |
| raise AssertionError(f"island rows of {cat} differ from tt.config.ISLAND_ROWS") | |
| return {"c1_w": _a(p, f"{N}.channel_pre_project.fc1.w"), "c1_b": _a(p, f"{N}.channel_pre_project.fc1.b"), | |
| "gelu_b1": cst.gelu_b1.reshape(1, -1), "c2_w": _a(p, f"{N}.channel_pre_project.fc2.w"), | |
| "t1_w": np.ascontiguousarray(cst.w_t1_valid), "t1_pad": cst.t1_pad, "g_pad": cst.g_pad, | |
| "t2_pad": cst.t2_pad, "t2_w": _a(p, f"{N}.token_pre_project.fc2.w")} | |
| def neighbor_aux(p: Mapping[str, np.ndarray]) -> np.ndarray: | |
| """``[W_type; b_type]`` [4, 128] for the host's ``[type one-hot, 1]`` columns.""" | |
| N = MIXER_MODULES["neighbor"] | |
| return np.concatenate([_a(p, f"{N}.type_emb.w"), _a(p, f"{N}.type_emb.b")[None]], 0).astype(f32) | |
| def lane_aux(p: Mapping[str, np.ndarray], cat: str) -> np.ndarray: | |
| """``[w_s; b_s; unk; W_a; b_a]`` [29, 128] for ``[speed * has, has, 1 - has, attributes (25), 1]``.""" | |
| N = MIXER_MODULES[cat] | |
| rows = [_a(p, f"{N}.speed_limit_emb.w").reshape(1, -1), _a(p, f"{N}.speed_limit_emb.b")[None], | |
| _a(p, f"{N}.unknown_speed_emb").reshape(1, -1), _a(p, f"{N}.attribute_emb.w"), | |
| _a(p, f"{N}.attribute_emb.b")[None]] | |
| out = np.concatenate(rows, 0).astype(f32) | |
| assert out.shape == (T.LANE_AUX_DIM, C.MIXER_CHANNELS), out.shape | |
| return out | |
| def lane_aux_features(speed: np.ndarray, has_speed: np.ndarray, attr: np.ndarray) -> np.ndarray: | |
| """Host columns of :func:`lane_aux`: ``[E, 29]`` float32 (``has`` is the node's ``> FLT_EPSILON`` mask).""" | |
| s = np.asarray(speed, f32).reshape(-1, 1) | |
| h = np.asarray(has_speed, bool).reshape(-1, 1).astype(f32) | |
| a = np.asarray(attr, f32).reshape(s.shape[0], -1) | |
| return np.concatenate([s * h, h, f32(1.0) - h, a, np.ones_like(h)], axis=1).astype(f32) | |
| def pos_aug(p: Mapping[str, np.ndarray]) -> np.ndarray: | |
| """``[W_pos; b_pos]`` [15, 256] for the host's ``[pos (14), token valid]`` columns.""" | |
| return np.concatenate([_a(p, "encoder.pos_emb.w"), _a(p, "encoder.pos_emb.b")[None]], 0).astype(f32) | |
| def agent_rows(p: Mapping[str, np.ndarray], agents: int = T.AGENTS) -> np.ndarray: | |
| """``preproj.fc2`` bias + agent embedding per decoder row [agents, 256] (row 0 ego, others neighbour), folded in | |
| float64 and rounded once.""" | |
| b2 = np.asarray(p["decoder.dit.preproj.fc2.b"], np.float64) | |
| emb = np.asarray(p["decoder.dit.agent_embedding"], np.float64) | |
| rows = np.repeat((b2 + emb[1])[None], agents, axis=0) | |
| rows[0] = b2 + emb[0] | |
| return rows.astype(f32) | |
| def final_projection_masked(p: Mapping[str, np.ndarray]) -> Lin: | |
| """``final_layer.proj.4`` (1024 -> 324) with the t = 0 output columns zeroed.""" | |
| lin = linear(p, "decoder.dit.final_layer.proj.4") | |
| w, b = lin.w.copy(), lin.b.copy() | |
| w[:, :T.STATE_COLS_T0] = 0.0 | |
| b[:T.STATE_COLS_T0] = 0.0 | |
| return Lin(w, b) | |
| def turn_weights(p: Mapping[str, np.ndarray]) -> Dict[str, np.ndarray]: | |
| """``w_sel`` [324, 5] (the 16 used rows of W scattered to the ``(t, d)`` columns of a state row), ``w_pool`` | |
| [256, 5] = ``W[16:] / 564`` (float64, rounded once), ``b`` [5].""" | |
| w = np.asarray(p["decoder.turn_indicator_predictor.w"], np.float64) # [272, 5] | |
| w_sel = np.zeros((T.STATE_COLS, w.shape[1]), np.float64) | |
| for i, t in enumerate(C.TURN_HEAD_STEPS): | |
| for d in range(2): | |
| w_sel[t * C.POSE_DIM + d] = w[2 * i + d] | |
| n_sel = 2 * len(C.TURN_HEAD_STEPS) | |
| return {"w_sel": w_sel.astype(f32), "w_pool": (w[n_sel:] / T.TOKENS_REAL).astype(f32), | |
| "b": np.asarray(p["decoder.turn_indicator_predictor.b"], f32)} | |
| def solver_coefficients(plan: SolverPlan) -> List[Tuple[float, float, float]]: | |
| """``(A, B, Cm)`` per update (float32 values) of ``y' = A y - B m0 + Cm m1`` (``Cm = 0`` for the first-order | |
| update).""" | |
| out = [] | |
| for u in plan.updates: | |
| a, b, c, r0 = (float(v) for v in (u.a, u.b, u.c, u.r0)) | |
| cm = c / r0 if u.order > 1 else 0.0 | |
| out.append((float(f32(a)), float(f32(b + cm)), float(f32(cm)))) | |
| return out | |
| class StepTables: | |
| """Per evaluation ``k``: per DiT block the folded ``norm1`` / ``norm2`` affine rows and the gates, and the folded | |
| ``norm_final`` affine; ``solver`` = :func:`solver_coefficients` (one fewer than evaluations).""" | |
| eval_times: Tuple[float, ...] | |
| blocks: List[List[Dict[str, np.ndarray]]] # [k][i] -> {n1_g, n1_b, gate_msa, n2_g, n2_b, gate_mlp} | |
| final: List[Dict[str, np.ndarray]] # [k] -> {g, b} | |
| solver: List[Tuple[float, float, float]] | |
| plan: SolverPlan | |
| def nfe(self) -> int: | |
| return len(self.eval_times) | |
| def step_tables(p: Mapping[str, np.ndarray], steps: int = C.DPM_SOLVER_STEPS) -> StepTables: | |
| plan = solver_plan(steps) | |
| tab = R.adaln_tables(p, plan.eval_times) | |
| blocks = [[{"n1_g": blk["norm1_gamma"][k], "n1_b": blk["norm1_beta"][k], "gate_msa": blk["gate_msa"][k], | |
| "n2_g": blk["norm2_gamma"][k], "n2_b": blk["norm2_beta"][k], "gate_mlp": blk["gate_mlp"][k]} | |
| for blk in tab.blocks] for k in range(len(plan.eval_times))] | |
| final = [{"g": tab.final_gamma[k], "b": tab.final_beta[k]} for k in range(len(plan.eval_times))] | |
| return StepTables(tuple(plan.eval_times), blocks, final, solver_coefficients(plan), plan) | |
| def small_module(p: Mapping[str, np.ndarray], cat: str) -> Dict[str, object]: | |
| """goal / ego-shape / turn encoders: channel MLP (fc1, fc2), ``norm``, ``emb_project`` (fc1, fc2).""" | |
| N = SMALL_MODULES[cat] | |
| return {"c1": linear(p, f"{N}.channel_pre_project.fc1"), "c2": linear(p, f"{N}.channel_pre_project.fc2"), | |
| "norm": norm(p, f"{N}.norm"), "e1": linear(p, f"{N}.emb_project.fc1"), | |
| "e2": linear(p, f"{N}.emb_project.fc2")} | |
| def pad_rows(a: np.ndarray, rows: int) -> np.ndarray: | |
| """Zero-pad the first axis to ``rows``.""" | |
| a = np.asarray(a) | |
| if a.shape[0] > rows: | |
| raise ValueError(f"{a.shape[0]} rows exceed {rows}") | |
| out = np.zeros((rows,) + a.shape[1:], a.dtype) | |
| out[:a.shape[0]] = a | |
| return out | |