changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
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
@dataclass
class Lin:
"""``y = x @ w + b``; ``w`` [in, out] float32, ``b`` [out] or None."""
w: np.ndarray
b: Optional[np.ndarray]
@dataclass
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
@dataclass
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
@property
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