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