Download code/tt_diffusion_planner/reference/model.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 13 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/model.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/reference/model.py
-
curl -L -o model.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/model.py
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) | |
| 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]) | |