changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
13 kB
# 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.<cat>.pre": "mixer categories: token_pre_project output transposed back to [E, 64, 128]",
"enc.<cat>.mixer": "mixer categories: output of the 6 MixerBlocks [E, 64, 128]",
"enc.<cat>": "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.<i>": "fusion block i output [564, 256]",
"enc.encoding": "final LayerNorm [564, 256] (the ``encoding`` graph output)",
"dec.<k>.temb": "t_embedder output [321, 256] of evaluation k",
"dec.<k>.x": "preproj + agent embedding [321, 256]",
"dec.<k>.block<i>": "DiT block i output [321, 256]",
"dec.<k>.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])