File size: 11,402 Bytes
4d9b003 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 | # 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
|