# 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