# SPDX-License-Identifier: Apache-2.0 """Pure-PyTorch fp32 re-implementation of the deployed v5.0 graphs (encoder, DiT decoder, turn-indicator head). Semantics follow the exported ONNX (SPEC 4.2-4.4; tier4/Diffusion-Planner ``encoder.py`` / ``mixer.py`` / ``dit.py`` / ``decoder.py`` read for reference, never imported), operation by operation where the order matters for float32: attention scales Q before ``Q.K^T`` and adds the ``-inf`` key bias before the softmax; the fusion takes Q from ``LN1(x)`` but K and V from the un-normalised ``x``; SiLU is ``x * sigmoid(x)``; GELU is exact (erf) everywhere in the encoder and in the decoder's ``preproj`` / ``t_embedder``, tanh-approximate in the six DiT MLPs and the final projection; every LayerNorm uses epsilon 1e-5. Inputs are the host features of :mod:`tt_diffusion_planner.host.features` (computed from normalised inputs), so this module holds only the network. Weights are the canonical dict of :mod:`.weights` (``w`` is ``[in, out]``). Every module output can be recorded with a ``ttaw.golden.TapRegistry`` (names in :data:`TAP_NAMES`); valid-row compaction of the taps is left to the caller. """ from __future__ import annotations from typing import Dict, Mapping, Optional, Tuple import numpy as np import torch import torch.nn.functional as F from . import config as C from ..ttaw.golden import NULL_TAPS, TapRegistry __all__ = ["Encoder", "Decoder", "TurnHead", "torch_params", "TAP_NAMES"] Tensor = torch.Tensor TAP_NAMES = { "enc..pre": "mixer categories: token_pre_project output transposed back to [E, 64, 128]", "enc..mixer": "mixer categories: output of the 6 MixerBlocks [E, 64, 128]", "enc.": "category output [E, 256] (masked; route with its position embedding), the fusion input rows", "enc.tokens": "fusion input [564, 256] (category outputs + masked positional embedding)", "enc.fusion.": "fusion block i output [564, 256]", "enc.encoding": "final LayerNorm [564, 256] (the ``encoding`` graph output)", "dec..temb": "t_embedder output [321, 256] of evaluation k", "dec..x": "preproj + agent embedding [321, 256]", "dec..block": "DiT block i output [321, 256]", "dec..out": "model_output [321, 81, 4]", "turn.pool": "mean of the encoding over the 564 tokens [256]", "turn.logit": "turn_indicator_logit [5]", } def torch_params(params: Mapping[str, np.ndarray], dtype: torch.dtype = torch.float32) -> Dict[str, Tensor]: """``{name: tensor}`` copies (the ONNX-backed arrays are read-only).""" return {k: torch.from_numpy(np.array(v, copy=True)).to(dtype) for k, v in params.items()} def _t(a, dtype=torch.float32) -> Tensor: if isinstance(a, torch.Tensor): return a.to(dtype) return torch.from_numpy(np.ascontiguousarray(a)).to(dtype) class _Ops: """Shared building blocks over the canonical parameter dict.""" def __init__(self, p: Mapping[str, Tensor]): self.p = p def linear(self, x: Tensor, name: str) -> Tensor: return torch.matmul(x, self.p[f"{name}.w"]) + self.p[f"{name}.b"] def ln(self, x: Tensor, name: str) -> Tensor: return F.layer_norm(x, x.shape[-1:], self.p[f"{name}.gamma"], self.p[f"{name}.beta"], eps=C.LN_EPS) def mlp(self, x: Tensor, name: str, approximate: str = "none") -> Tensor: return self.linear(F.gelu(self.linear(x, f"{name}.fc1"), approximate=approximate), f"{name}.fc2") def attention(self, q_in: Tensor, kv: Tensor, name_q: str, key_bias: Optional[Tensor]) -> Tensor: """Multi-head attention with 8 heads of 32: ``q_in`` [Sq, 256] -> Q (``name_q`` linear); ``kv`` [Sk, 512] is the already projected K|V; ``key_bias`` [Sk] is 0 / -inf. Returns the heads concatenated [Sq, 256] (before the output projection).""" q = self.linear(q_in, name_q) return self.attend(q, kv[:, :C.HIDDEN_DIM], kv[:, C.HIDDEN_DIM:], key_bias) @staticmethod def attend(q: Tensor, k: Tensor, v: Tensor, key_bias: Optional[Tensor]) -> Tensor: sq, sk = q.shape[0], k.shape[0] qh = q.reshape(sq, C.NUM_HEADS, C.HEAD_DIM).transpose(0, 1) * float(C.ATTN_SCALE) # [H, Sq, 32] kh = k.reshape(sk, C.NUM_HEADS, C.HEAD_DIM).permute(1, 2, 0) # [H, 32, Sk] vh = v.reshape(sk, C.NUM_HEADS, C.HEAD_DIM).transpose(0, 1) # [H, Sk, 32] scores = torch.matmul(qh, kh) if key_bias is not None: scores = scores + key_bias.reshape(1, 1, sk) o = torch.matmul(torch.softmax(scores, dim=-1), vh) # [H, Sq, 32] return o.transpose(0, 1).reshape(sq, C.HIDDEN_DIM) def _key_bias(valid: np.ndarray) -> Tensor: b = np.where(np.asarray(valid, bool), np.float32(0.0), np.float32(-np.inf)).astype(np.float32) return torch.from_numpy(b) class Encoder(_Ops): """``diffusion_planner_encoder.onnx`` on host features -> ``encoding`` [564, 256].""" def mixer_trunk(self, x: Tensor, mod: str, taps: TapRegistry, cat: str) -> Tensor: """channel_pre (C_in -> 128 -> 128), token_pre over the time / point axis (T -> 64 -> 64), 6 MixerBlocks, mean over the 64 tokens -> [E, 128].""" N = f"encoder.{mod}" x = self.mlp(x, f"{N}.channel_pre_project") # [E, T, 128] x = self.mlp(x.transpose(1, 2), f"{N}.token_pre_project").transpose(1, 2) # [E, 64, 128] taps.tap(f"enc.{cat}.pre", x) for i in range(C.MIXER_DEPTH): B = f"{N}.blocks.{i}" y = self.mlp(self.ln(x, f"{B}.norm1").transpose(1, 2), f"{B}.tokens_mlp").transpose(1, 2) x = x + y x = x + self.mlp(self.ln(x, f"{B}.norm2"), f"{B}.channels_mlp") taps.tap(f"enc.{cat}.mixer", x) return x.mean(dim=1) def head(self, x: Tensor, mod: str) -> Tensor: """LayerNorm(128) + emb_project (128 -> 256 -> 256).""" return self.mlp(self.ln(x, f"encoder.{mod}.norm"), f"encoder.{mod}.emb_project") def lanes(self, x: Tensor, attr: Tensor, speed: Tensor, has_speed: Tensor, valid: Tensor, mod: str, taps: TapRegistry, cat: str) -> Tensor: N = f"encoder.{mod}" h = self.mixer_trunk(x, mod, taps, cat) speed_emb = torch.where(has_speed, self.linear(speed, f"{N}.speed_limit_emb"), self.p[f"{N}.unknown_speed_emb"].reshape(1, -1)) h = h + speed_emb + self.linear(attr, f"{N}.attribute_emb") return self.head(h, mod) * valid def small(self, x: Tensor, mod: str) -> Tensor: """goal / ego-shape / turn-indicator encoders: channel MLP -> LayerNorm -> projection.""" return self.head(self.mlp(x, f"encoder.{mod}.channel_pre_project"), mod) def forward(self, f, taps: TapRegistry = NULL_TAPS) -> Tensor: valid = {k: _t(v.astype(np.float32)).reshape(-1, 1) for k, v in f.valid.items()} out: Dict[str, Tensor] = {} e = self.mixer_trunk(_t(f.ego)[None], "ego_encoder", taps, "ego") out["ego"] = self.head(e, "ego_encoder") n = self.mixer_trunk(_t(f.neighbor), "neighbor_encoder", taps, "neighbor") n = n + self.linear(_t(f.neighbor_type), "encoder.neighbor_encoder.type_emb") out["neighbor"] = self.head(n, "neighbor_encoder") * valid["neighbor"] out["static"] = self.mlp(_t(f.static), "encoder.static_encoder.projection") * valid["static"] out["lane"] = self.lanes(_t(f.lane), _t(f.lane_attr), _t(f.lane_speed), torch.from_numpy(f.lane_has_speed), valid["lane"], "lane_encoder", taps, "lane") route = self.lanes(_t(f.route), _t(f.route_attr), _t(f.route_speed), torch.from_numpy(f.route_has_speed), valid["route"], "route_encoder", taps, "route") out["route"] = route + self.p["encoder.route_position_embedding"] * valid["route"] poly = self.mixer_trunk(_t(f.polygon), "polygon_encoder", taps, "polygon") out["polygon"] = self.head(poly, "polygon_encoder") * valid["polygon"] ls = self.mixer_trunk(_t(f.line_string), "line_string_encoder", taps, "line_string") out["line_string"] = self.head(ls, "line_string_encoder") * valid["line_string"] out["goal"] = self.small(_t(f.goal)[None], "goal_pose_encoder") out["ego_shape"] = self.small(_t(f.ego_shape)[None], "ego_shape_encoder") out["turn"] = self.small(_t(f.turn)[None], "turn_indicator_encoder") for name, _ in C.TOKEN_LAYOUT: taps.tap(f"enc.{name}", out[name]) x = torch.cat([out[name] for name, _ in C.TOKEN_LAYOUT], dim=0) # [564, 256] pos = self.linear(_t(f.pos), "encoder.pos_emb") * _t(f.token_valid.astype(np.float32)).reshape(-1, 1) x = taps.tap("enc.tokens", x + pos) key_bias = _key_bias(f.key_valid) for i in range(C.FUSION_DEPTH): B = f"encoder.fusion.blocks.{i}" kv = self.linear(x, f"{B}.attn.kv") # K and V from the un-normalised x heads = self.attention(self.ln(x, f"{B}.norm1"), kv, f"{B}.attn.q", key_bias) x = x + self.linear(heads, f"{B}.attn.out") x = x + self.mlp(self.ln(x, f"{B}.norm2"), f"{B}.mlp") taps.tap(f"enc.fusion.{i}", x) return taps.tap("enc.encoding", self.ln(x, "encoder.fusion.norm")) class Decoder(_Ops): """``diffusion_planner_decoder.onnx``: one DiT evaluation (x0 prediction) per call. ``cross_kv(encoding)`` computes the K|V of the three cross-attention blocks once per plan (the decoder graph recomputes the same product at every call; hoisting it is exact, SPEC 4.6.1).""" def cross_kv(self, encoding: Tensor) -> Tuple[Tensor, ...]: return tuple(self.linear(encoding, f"decoder.dit.blocks.{i}.cross_attn.kv") for i in range(C.DIT_DEPTH)) def time_embedding(self, t, agents: int = C.MAX_NUM_AGENTS) -> Tensor: """t_embedder of the 81 per-point diffusion times [P, 81] (uniform in the multi-step mode) -> [P, 256].""" tt = torch.full((agents, C.DIT_TIME_DIM), float(np.float32(t))) if np.ndim(t) == 0 else _t(t) return self.mlp(tt, "decoder.dit.t_embedder") def forward(self, x_t: np.ndarray, t, kv: Tuple[Tensor, ...], agent_valid: np.ndarray, taps: TapRegistry = NULL_TAPS, prefix: str = "dec") -> Tensor: """``x_t`` [P, 81, 4] with P = 321 as exported, or any agent bucket ``P >= 1 + valid neighbours`` (ego first; exact for the agents present, SPEC 4.6.4); ``agent_valid`` [P]; ``t`` a scalar time (or [P, 81]).""" P = int(np.shape(x_t)[0]) x = self.mlp(_t(x_t).reshape(P, C.DIT_INPUT_DIM), "decoder.dit.preproj") c = taps.tap(f"{prefix}.temb", self.time_embedding(t, P)) emb = self.p["decoder.dit.agent_embedding"] x = x + torch.cat([emb[0:1], emb[1:2].expand(P - 1, -1)], dim=0) x = taps.tap(f"{prefix}.x", x) silu_c = c * torch.sigmoid(c) key_bias = _key_bias(agent_valid) for i in range(C.DIT_DEPTH): B = f"decoder.dit.blocks.{i}" mod = self.linear(silu_c, f"{B}.adaLN_modulation") shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = torch.split(mod, C.HIDDEN_DIM, dim=-1) h = self.ln(x, f"{B}.norm1") * (scale_msa + 1.0) + shift_msa qkv = self.linear(h, f"{B}.attn.qkv") q, k, v = torch.split(qkv, C.HIDDEN_DIM, dim=-1) x = x + gate_msa * self.linear(self.attend(q, k, v, key_bias), f"{B}.attn.out") h = self.ln(x, f"{B}.norm2") * (scale_mlp + 1.0) + shift_mlp x = x + gate_mlp * self.mlp(h, f"{B}.mlp1", "tanh") heads = self.attention(self.ln(x, f"{B}.norm3"), kv[i], f"{B}.cross_attn.q", None) # no key mask x = x + self.linear(heads, f"{B}.cross_attn.out") x = x + self.mlp(self.ln(x, f"{B}.norm4"), f"{B}.mlp2", "tanh") taps.tap(f"{prefix}.block{i}", x) Fn = "decoder.dit.final_layer" shift, scale = torch.split(self.linear(silu_c, f"{Fn}.adaLN_modulation"), C.HIDDEN_DIM, dim=-1) h = self.ln(x, f"{Fn}.norm_final") * (scale + 1.0) + shift h = F.gelu(self.linear(self.ln(h, f"{Fn}.proj.0"), f"{Fn}.proj.1"), approximate="tanh") out = self.linear(self.ln(h, f"{Fn}.proj.3"), f"{Fn}.proj.4").reshape(P, C.OUTPUT_T + 1, C.POSE_DIM) return taps.tap(f"{prefix}.out", out) class TurnHead(_Ops): """``diffusion_planner_turn_indicator.onnx``: ``W [272 -> 5]`` over ``final_x0[0, 1::10, :2]`` (16 values) and the token mean of the encoding (256).""" def forward(self, encoding: Tensor, final_x0: np.ndarray, taps: TapRegistry = NULL_TAPS) -> Tensor: pool = taps.tap("turn.pool", encoding.mean(dim=0)) ego = _t(final_x0)[0, 1::10, :2].reshape(-1) feat = torch.cat([ego, pool], dim=0)[None] return taps.tap("turn.logit", self.linear(feat, "decoder.turn_indicator_predictor")[0])