changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
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
@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.<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))