# SPDX-License-Identifier: Apache-2.0 """The encoder graph on the device: six MLP-Mixer trunks, the entity heads, the token assembly and the fusion transformer -> ``encoding`` ``[1, 1, 576, 256]`` (564 real tokens + 12 zero pad tokens, ``tt/config.py``). Inputs are the persistent trace inputs written by :mod:`.inputs` (host features of ``host.features``, fp32 TILE): ``ego_x`` ``[1, 1, 6, 4]`` and ``neighbor_x`` ``[1, 320, 6, 9]`` (the 6 time rows the export keeps), the other mixer inputs ``[1, E, T, C]``, the aux columns of the exact embedding rewrites (``tt/params.py``), the small-encoder rows, ``token_valid`` ``[1, 1, 576, 1]``, ``pos_aug`` ``[1, 1, 576, 15]`` and the fusion key-bias row. Mixer trunk (``T4M mixer.py``; SPEC 4.2): ``channel_pre`` (C -> 128 -> 128) and ``token_pre`` over the T axis (T -> 64 -> 64), then 6 blocks ``x += tokens_mlp(LN1(x)^T)^T; x += channels_mlp(LN2(x))`` on ``[1, E, 64, 128]``, token mean -> ``[1, 1, E, 128]``. Token mixing is ``transpose -> 2-D [1, 1, E*128, 64] @ W -> transpose`` (probe P12: 108 us at E = 320, never the broadcast-left batched matmul). Ego and neighbour run ``channel_pre`` + ``token_pre`` as the fp32 **pad-relative island** (``reference.rewrites``, P12): only deviations from the all-zero agent pass the TF32-like matmuls. Fusion block (``encoder.py:327-333``): ``K|V = W_kv x`` from the un-normalised stream, ``Q = W_q LN1(x)``, masked attention over 576 keys (the -inf key row expanded once per plan), ``x += out(.)``; ``x += mlp(LN2(x))``; final LN. The attention is fp32 matmuls + softmax by default (C20 ``attention_matmul``, ``ATTN_MATMUL`` ``enc.fusion.attn``: the bf16 SDPA was the largest encoder error on the sensitive nuScenes instants, PORT_LOG decision 11), else C20 ``sdpa``. """ from __future__ import annotations from typing import Any, Dict, Mapping, Optional import numpy as np from ..reference import config as C from ..ttaw.ops import attention as A from . import config as T from . import params as P from .attention import attention_matmul from .layers import ATTN, Build, Const, LayerNorm, make_linear __all__ = ["MixerTrunk", "TtEncoder", "MIXER_CATS"] MIXER_CATS = ("ego", "neighbor", "lane", "route", "polygon", "line_string") ENTITIES = dict(C.TOKEN_LAYOUT) def _gelu(x): import ttnn return ttnn.gelu(x, fast_and_approximate_mode=False) class MixerTrunk: """One MLP-Mixer category up to the token mean (``[1, 1, E, 128]``).""" def __init__(self, build: Build, p: Mapping[str, np.ndarray], cat: str): self.cat, self.E, self.T = cat, ENTITIES[cat], T.MIXER_T[cat] self.build = build self.island = cat in T.ISLAND_ROWS mods = P.mixer_module(p, cat) mix = f"enc.mixer.{cat}" self.stream = build.stream(mix) if self.island: isl = f"enc.island.{cat}" w = P.island(p, cat) self.isl_dtype = build.stream(isl) def lin(weights, act=None): return make_linear(build, weights, isl, out=self.isl_dtype, activation=act) self.c1 = lin(P.Lin(w["c1_w"], w["c1_b"]), "gelu") self.gelu_b1 = Const(build, w["gelu_b1"], self.isl_dtype) self.c2 = lin(P.Lin(w["c2_w"], None)) self.t1 = lin(P.Lin(w["t1_w"], None)) self.t1_pad = Const(build, w["t1_pad"], self.isl_dtype) self.g_pad = Const(build, w["g_pad"], self.isl_dtype) self.t2 = lin(P.Lin(w["t2_w"], None)) self.t2_pad = Const(build, w["t2_pad"], self.isl_dtype) else: pre = f"enc.pre.{cat}" self.c1 = make_linear(build, mods["c1"], pre, out=build.hidden(pre), activation="gelu") self.c2 = make_linear(build, mods["c2"], pre, out=build.hidden(pre)) self.t1 = make_linear(build, mods["t1"], pre, out=build.hidden(pre), activation="gelu") self.t2 = make_linear(build, mods["t2"], pre, out=self.stream) self.blocks = [] for blk in mods["blocks"]: self.blocks.append({ "n1": LayerNorm(build, blk["n1"], mix), "tk1": make_linear(build, blk["tk1"], mix, out=build.hidden(mix), activation="gelu"), "tk2": make_linear(build, blk["tk2"], mix, out=self.stream), "n2": LayerNorm(build, blk["n2"], mix), "ch1": make_linear(build, blk["ch1"], mix, out=build.hidden(mix), activation="gelu"), "ch2": make_linear(build, blk["ch2"], mix, out=self.stream)}) if build.enc_mem() is not None: # ENC_L1: the mixer blocks' intermediates in L1 (interleaved) for b in self.blocks: b["n1"].mem = b["n2"].mem = build.enc_mem() for key in ("tk1", "tk2", "ch1", "ch2"): b[key].out_mem = build.enc_mem() # ------------------------------------------------------------------------------------------------------------ def _tokens_view(self, x): """``[1, E, 128, 64]`` <-> ``[1, 1, E*128, 64]`` (free: 128 rows are whole tiles).""" import ttnn return ttnn.reshape(x, (1, 1, int(x.shape[1]) * C.MIXER_CHANNELS, x.shape[-1])) def _entities_view(self, x): import ttnn return ttnn.reshape(x, (1, int(x.shape[2]) // C.MIXER_CHANNELS, C.MIXER_CHANNELS, x.shape[-1])) def pre(self, x): """``[1, E, T, C]`` -> ``x0`` ``[1, E, 64, 128]`` (the ``enc..pre`` tap).""" import ttnn if self.island: h = ttnn.subtract(self.c1(x), self.gelu_b1()) # gelu(x W1 + b1) - gelu(b1) dz = self.c2(h) # z - c0 [1, E, 6, 128] dzt = self._tokens_view(ttnn.transpose(dz, -2, -1)) # [1, 1, E*128, 6] t1 = ttnn.add(self._entities_view(self.t1(dzt)), self.t1_pad()) # t1_pad + dz^T W_t1[rows] g = ttnn.subtract(_gelu(t1), self.g_pad()) t2 = ttnn.add(self._entities_view(self.t2(self._tokens_view(g))), self.t2_pad()) x0 = ttnn.transpose(t2, -2, -1) # [1, E, 64, 128] if self.isl_dtype != self.stream: x0 = ttnn.typecast(x0, self.stream_ttnn) return x0 z = self.c2(self.c1(x)) # [1, E, T, 128] zt = self._tokens_view(ttnn.transpose(z, -2, -1)) # [1, 1, E*128, T] t = self._entities_view(self.t2(self.t1(zt))) # [1, E, 128, 64] x0 = ttnn.transpose(t, -2, -1) if x0.dtype != self.stream_ttnn: x0 = ttnn.typecast(x0, self.stream_ttnn) return x0 @property def stream_ttnn(self): from ..ttaw.tensors import ttnn_dtype return ttnn_dtype(self.stream) def mix(self, x): """The 6 MixerBlocks on ``[1, E, 64, 128]`` (the ``enc..mixer`` tap).""" import ttnn if self.build.ln_resid and self.blocks[0]["n1"].can_transpose(x): # LN_TR: the transposes around the token-mixing MLP inside the LN programs (n1 writes LN(h)^T, n2 reads # the token-mixing output transposed): the same values, two programs less per block pending = None for b in self.blocks: if pending is None: n1t = b["n1"].transposed(x) else: x, n1t = b["n1"].transposed(x, res=pending) t = self._entities_view(b["tk2"](b["tk1"](self._tokens_view(n1t)))) # [1, E, 128, 64] x, n2 = b["n2"].residual_t(x, t) pending = b["ch2"](b["ch1"](n2)) return ttnn.add(x, pending, memory_config=ttnn.DRAM_MEMORY_CONFIG) if self.build.ln_resid: # LN_RESID: adds fused into the LNs pending = None for b in self.blocks: if pending is None: n1 = b["n1"](x) else: x, n1 = b["n1"].residual(x, pending) y = self._tokens_view(ttnn.transpose(n1, -2, -1)) y = ttnn.transpose(self._entities_view(b["tk2"](b["tk1"](y))), -2, -1) x, n2 = b["n2"].residual(x, y) pending = b["ch2"](b["ch1"](n2)) return ttnn.add(x, pending) for b in self.blocks: y = self._tokens_view(ttnn.transpose(b["n1"](x), -2, -1)) # LN over the 128 channels, then T y = ttnn.transpose(self._entities_view(b["tk2"](b["tk1"](y))), -2, -1) x = ttnn.add(x, y) x = ttnn.add(x, b["ch2"](b["ch1"](b["n2"](x)))) return x def pool(self, x): """Mean over the 64 tokens -> ``[1, 1, E, 128]``.""" import ttnn m = ttnn.mean(x, dim=2, keepdim=True) # [1, E, 1, 128] return ttnn.reshape(m, (1, 1, int(x.shape[1]), C.MIXER_CHANNELS)) class _Head: """``(m [+ aux @ W_aux]) -> LN(128) -> emb_project (128 -> 256 -> 256)`` (+ the route position embedding).""" def __init__(self, build: Build, p: Mapping[str, np.ndarray], cat: str, out_dtype: str, *, mixer=True): mod = f"enc.head.{cat}" stream = build.stream(mod) m = P.mixer_module(p, cat) if mixer else P.small_module(p, cat) self.aux = None if cat == "neighbor": self.aux = make_linear(build, P.Lin(P.neighbor_aux(p), None), mod, out=stream) elif cat in ("lane", "route"): self.aux = make_linear(build, P.Lin(P.lane_aux(p, cat), None), mod, out=stream) self.norm = LayerNorm(build, m["norm"], mod) self.e1 = make_linear(build, m["e1"], mod, out=build.hidden(mod), activation="gelu") self.e2 = make_linear(build, m["e2"], mod, out=out_dtype) self.route_pos = Const(build, p["encoder.route_position_embedding"], out_dtype) if cat == "route" else None def __call__(self, m, aux=None): import ttnn if self.aux is not None: m = ttnn.add(m, self.aux(aux)) out = self.e2(self.e1(self.norm(m))) if self.route_pos is not None: out = ttnn.add(out, self.route_pos()) return out class _Small: """goal / ego-shape / turn encoders: channel MLP (C -> 128 -> 128) -> LN -> emb_project.""" def __init__(self, build: Build, p: Mapping[str, np.ndarray], cat: str, out_dtype: str): mod = f"enc.head.{cat}" m = P.small_module(p, cat) self.c1 = make_linear(build, m["c1"], mod, out=build.hidden(mod), activation="gelu") self.c2 = make_linear(build, m["c2"], mod, out=build.stream(mod)) self.head = _Head(build, p, cat, out_dtype, mixer=False) def __call__(self, x): return self.head(self.c2(self.c1(x))) class TtEncoder: """Encoder + fusion on the device. ``forward(ctx, taps)`` -> ``encoding`` ``[1, 1, 576, 256]``; ``taps`` (a dict) collects the device tensors of the ``reference.model.TAP_NAMES`` encoder taps when given.""" INPUTS = ("ego_x", "neighbor_x", "neighbor_aux", "static_x", "lane_x", "lane_aux", "route_x", "route_aux", "polygon_x", "line_string_x", "goal_x", "ego_shape_x", "turn_x", "token_valid", "pos_aug", "fusion_key_row") def __init__(self, build: Build, p: Mapping[str, np.ndarray], nb_rows=()): """``nb_rows`` (``COMPACT``): the neighbour entity counts the trunk may run on (the agent buckets): a zero token block for the rest is uploaded per count.""" self.build = build self.fstream = build.stream("enc.fusion") E = ENTITIES["neighbor"] self.nb_zeros = {int(n): Const(build, np.zeros((E - int(n), C.HIDDEN_DIM), np.float32), self.fstream) for n in nb_rows if 0 < int(n) < E} self.trunks = {cat: MixerTrunk(build, p, cat) for cat in MIXER_CATS} self.heads = {cat: _Head(build, p, cat, self.fstream) for cat in MIXER_CATS} st = "enc.head.static" self.static1 = make_linear(build, P.linear(p, "encoder.static_encoder.projection.fc1"), st, out=build.hidden(st), activation="gelu") self.static2 = make_linear(build, P.linear(p, "encoder.static_encoder.projection.fc2"), st, out=self.fstream) self.small = {cat: _Small(build, p, cat, self.fstream) for cat in ("goal", "ego_shape", "turn")} self.pos = make_linear(build, P.Lin(P.pos_aug(p), None), "enc.tokens", out=self.fstream) self.pad_tokens = Const(build, np.zeros((T.TOKENS - T.TOKENS_REAL, C.HIDDEN_DIM), np.float32), self.fstream) fu = "enc.fusion" self.attn_mm = build.attn_matmul("enc.fusion.attn") # fp32 matmul attention instead of bf16 SDPA qkv_out = "float32" if self.attn_mm else ATTN self.blocks = [] for i in range(C.FUSION_DEPTH): B = f"encoder.fusion.blocks.{i}" self.blocks.append({ "kv": make_linear(build, P.linear(p, f"{B}.attn.kv"), fu, out=qkv_out), "n1": LayerNorm(build, P.norm(p, f"{B}.norm1"), fu), "q": make_linear(build, P.linear(p, f"{B}.attn.q"), fu, out=qkv_out), "out": make_linear(build, P.linear(p, f"{B}.attn.out"), fu, out=self.fstream), "n2": LayerNorm(build, P.norm(p, f"{B}.norm2"), fu), "fc1": make_linear(build, P.linear(p, f"{B}.mlp.fc1"), fu, out=build.hidden(fu), activation="gelu"), "fc2": make_linear(build, P.linear(p, f"{B}.mlp.fc2"), fu, out=self.fstream)}) if self.attn_mm and build.attn_mem() is not None and build.attn_l1 == 1: # ATTN_L1=1: Q, K | V to L1 for blk in self.blocks: blk["q"].out_mem = blk["kv"].out_mem = build.attn_mem() self.fus_mem = None if build.fus_l1: # FUS_L1: the fusion blocks' intermediates in L1 (interleaved) import ttnn self.fus_mem = ttnn.L1_MEMORY_CONFIG for blk in self.blocks: blk["n1"].mem = blk["n2"].mem = self.fus_mem for key in ("out", "fc1", "fc2"): blk[key].out_mem = self.fus_mem self.final_norm = LayerNorm(build, P.norm(p, "encoder.fusion.norm"), fu) self.scale = float(C.ATTN_SCALE) self.attn_fp32 = build.attn_fp32_acc("enc.fusion.attn") def _fused(self, qs, kv, mask): """``ATTN_FUSED``: the fusion attention as one program (``tt/fattn_kernel.py``; Q from ``[1, 1, 576, 256]``, K | V from ``[1, 1, 576, 512]`` in place) -> merged heads, or None (the stock chain runs).""" b = self.build if not (b.attn_fused and self.attn_mm and b.attn_smask and b.attn_smsm and b.attn_fast == 1): return None from .fattn_kernel import flat_at, fused_attention, supported if not supported((qs, kv), mask): return None H = C.NUM_HEADS return fused_attention(qs, kv, kv, mask, self.scale, H, flat_at(qs, 0), flat_at(kv, 0), flat_at(kv, H), memory_config=self.fus_mem) def forward(self, ctx: Mapping[str, Any], taps: Optional[Dict[str, Any]] = None, nb: Optional[int] = None): """``nb`` (``COMPACT``): run the neighbour trunk and head on the first ``nb`` entities only and fill the other neighbour tokens with zeros: they are invalid entities, which ``token_valid`` zeroes anyway.""" import ttnn def tap(name, t): if taps is not None: taps[name] = t return t out: Dict[str, Any] = {} for cat in MIXER_CATS: tr = self.trunks[cat] xin = ctx[f"{cat}_x"] aux = ctx.get(f"{cat}_aux") if cat in ("neighbor", "lane", "route") else None cut = cat == "neighbor" and nb in self.nb_zeros if cut: # COMPACT: the first nb neighbours xin = ttnn.slice(xin, [0, 0, 0, 0], [1, nb] + list(xin.shape)[2:]) aux = ttnn.slice(aux, [0, 0, 0, 0], [1, 1, nb, int(aux.shape[-1])]) x0 = tap(f"enc.{cat}.pre", tr.pre(xin)) x = tap(f"enc.{cat}.mixer", tr.mix(x0)) out[cat] = self.heads[cat](tr.pool(x), aux) if cut: out[cat] = ttnn.concat([out[cat], self.nb_zeros[nb]()], dim=2) out["static"] = self.static2(self.static1(ctx["static_x"])) for cat, enc in self.small.items(): out[cat] = enc(ctx[f"{cat}_x"]) for name, _ in C.TOKEN_LAYOUT: tap(f"enc.{name}", out[name]) x = ttnn.concat([out[name] for name, _ in C.TOKEN_LAYOUT] + [self.pad_tokens()], dim=2) # [1,1,576,256] x = ttnn.multiply(x, ctx["token_valid"]) # invalid entities -> 0 x = tap("enc.tokens", ttnn.add(x, self.pos(ctx["pos_aug"]))) # + valid * (pos W + b) mask = A.expand_key_bias(ctx["fusion_key_row"], T.TOKENS, # [1, 1, 576, 576], once per plan memory_config=self.build.attn_mem() if self.attn_mm else None) for i, b in enumerate(self.blocks): kv = b["kv"](x) # K | V from the un-normalised x qs = b["q"](b["n1"](x)) a = self._fused(qs, kv, mask) if a is not None: pass elif self.attn_mm: q, k, v = A.split_q_kv(qs, kv, C.NUM_HEADS) a = A.merge_heads(attention_matmul(q, k, v, scale=self.scale, attn_mask=mask, mode=self.build.attn_fast, smask=self.build.attn_smask, smsm=self.build.attn_smsm)) else: q, k, v = A.split_q_kv(qs, kv, C.NUM_HEADS) a = A.sdpa(q, k, v, scale=self.scale, attn_mask=mask, concat_heads=True, fp32_acc=self.attn_fp32) fkw = {} if self.fus_mem is None else {"memory_config": self.fus_mem} x = ttnn.add(x, b["out"](a), **fkw) x = ttnn.add(x, b["fc2"](b["fc1"](b["n2"](x))), **fkw) tap(f"enc.fusion.{i}", x) return tap("enc.encoding", self.final_norm(x))