# SPDX-License-Identifier: Apache-2.0 """The DiT decoder, the on-device DPM-Solver++(2M) loop and the turn head (SPEC 4.3-4.4, 5.1; ``dpm_solver.cpp``). One evaluation (``T4M dit.py``; rows = the 352 decoder agents of ``tt/config.py``): - ``preproj``: ``x [352, 324] (fp32 state) -> 512 (exact GELU) -> 256`` + agent embedding rows (folded bias rows); - 3 DiT blocks with the per-step adaLN tables **folded into the LayerNorm affine** (``reference.rewrites``): ``x += gate_msa * out(attn_masked(qkv(LN1_k(x))))``, ``x += gate_mlp * mlp1(LN2_k(x))`` (tanh GELU), ``x += out(attn(q(LN3(x)), K_i, V_i))`` on the hoisted cross K / V, ``x += mlp2(LN4(x))``; - final layer: ``LN_k`` (folded) -> ``LN(256)`` -> ``1024`` (tanh GELU) -> ``LN(1024)`` -> ``324`` with the t = 0 output columns zeroed (``m * mask0``, fp32). Solver (``tt/config.py``): state ``y = x * mask0``; evaluation ``k`` reads ``x = y + cs``; updates ``y' = A_k y - B_k m_k + Cm_k m_(k-1)`` (``tt.params.solver_coefficients``); after the 11th evaluation (t = 1e-3, denoise to zero) ``final_x0 = m_10 + cs``. All state and model outputs are fp32; the scalars are constants of the unrolled loop (one trace per plan). Attention (``ATTN_MATMUL`` knob, default ``dec.*``): fp32 Q / K / V, two fp32 matmuls + softmax (C20 ``attention_matmul``); the cross-attention then reads the 576 encoder rows with the 12 pad tokens masked. With the knob off: the bf16 C20 ``sdpa`` (self-attention masked on 352 keys; cross-attention on the 564 real tokens, no mask). Every linear is a split matmul by default (``SPLIT_MATMUL`` ``dec.*``: fp32-accurate, PORT_LOG decision 11). Turn head: ``logit = final_x0[0] @ W_sel + (ones_564 @ encoding) @ (W_pool / 564) + b`` (fp32). """ from __future__ import annotations from typing import Any, Dict, List, Mapping, Sequence, Tuple 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 KcatOperand, operand_ktp, ATTN, Build, Const, LayerNorm, chain2, make_linear __all__ = ["TtDecoder", "TtTurnHead", "STEP_KEYS"] STEP_KEYS = ("n1_g", "n1_b", "gate_msa", "n2_g", "n2_b", "gate_mlp") class TtDecoder: """DiT weights, the per-step tables as device rows, and the evaluation / solver graph.""" def __init__(self, build: Build, p: Mapping[str, np.ndarray], tables: P.StepTables, rows: Sequence[int] = (T.AGENTS,)): """``rows``: the decoder row counts the graph runs on (``T.AGENTS`` and, with ``COMPACT``, the agent buckets): the agent embedding rows and the cross-attention mask are uploaded once per count.""" self.build = build self.tables = tables D = "dec.block" self.stream = build.stream(D) # the pre-projection reads the fp32 solver state: with split matmuls (SPLIT_MATMUL "dec.preproj.*") the # positions are not truncated to the TF32-like operand of a device fp32 matmul (probe P12) self.pre1 = make_linear(build, P.linear(p, "decoder.dit.preproj.fc1"), "dec.preproj.fc1", out=build.hidden("dec.preproj.fc1"), activation="gelu") self.pre2 = make_linear(build, P.linear(p, "decoder.dit.preproj.fc2", bias=False), "dec.preproj.fc2", out=self.stream) self.rows = tuple(sorted({int(r) for r in rows} | {T.AGENTS})) self.agent_rows = {r: Const(build, P.agent_rows(p, r), self.stream) for r in self.rows} self.blocks = [] for i in range(C.DIT_DEPTH): B = f"decoder.dit.blocks.{i}" self.blocks.append({ "ln": LayerNorm(build, None, D), # affine rows per step (folded adaLN) "qkv": make_linear(build, P.linear(p, f"{B}.attn.qkv"), D, out="float32" if build.attn_matmul("dec.self_attn") else ATTN), "out": make_linear(build, P.linear(p, f"{B}.attn.out"), D, out=self.stream), "m1a": make_linear(build, P.linear(p, f"{B}.mlp1.fc1"), D, out=build.hidden(D), activation="gelu_tanh"), "m1b": make_linear(build, P.linear(p, f"{B}.mlp1.fc2"), D, out=self.stream), "n3": LayerNorm(build, P.norm(p, f"{B}.norm3"), D), "cq": make_linear(build, P.linear(p, f"{B}.cross_attn.q"), D, out="float32" if build.attn_matmul("dec.cross_attn") else ATTN), "cout": make_linear(build, P.linear(p, f"{B}.cross_attn.out"), D, out=self.stream), "n4": LayerNorm(build, P.norm(p, f"{B}.norm4"), D), "m2a": make_linear(build, P.linear(p, f"{B}.mlp2.fc1"), D, out=build.hidden(D), activation="gelu_tanh"), "m2b": make_linear(build, P.linear(p, f"{B}.mlp2.fc2"), D, out=self.stream)}) kvw = [P.linear(p, f"decoder.dit.blocks.{i}.cross_attn.kv") for i in range(C.DIT_DEPTH)] H = C.HIDDEN_DIM kv_out = "float32" if build.attn_matmul("dec.cross_attn") else ATTN self.k_lin = [make_linear(build, P.Lin(lin.w[:, :H], lin.b[:H]), "dec.kv", out=kv_out) for lin in kvw] self.v_lin = [make_linear(build, P.Lin(lin.w[:, H:], lin.b[H:]), "dec.kv", out=kv_out) for lin in kvw] F = "dec.final" self.fin_ln = LayerNorm(build, None, F) self.p0 = LayerNorm(build, P.norm(p, "decoder.dit.final_layer.proj.0"), F) self.p1 = make_linear(build, P.linear(p, "decoder.dit.final_layer.proj.1"), F, out=build.hidden(F), activation="gelu_tanh") self.p3 = LayerNorm(build, P.norm(p, "decoder.dit.final_layer.proj.3"), F) self.p4 = make_linear(build, P.final_projection_masked(p), F, out="float32") # per-step rows (RT-dev constants of the unrolled loop): [k][i][key] / [k]["g" | "b"] self.step_rows = [self.upload_step(tables, k) for k in range(tables.nfe)] self.scale = float(C.ATTN_SCALE) self.attn_fp32 = {"self": build.attn_fp32_acc("dec.self_attn"), "cross": build.attn_fp32_acc("dec.cross_attn")} self.attn_mm = {"self": build.attn_matmul("dec.self_attn"), "cross": build.attn_matmul("dec.cross_attn")} if self.attn_mm["cross"]: # fp32 cross-attention over the 576 encoder rows, the 12 pad tokens masked row = np.where(np.arange(T.TOKENS) < T.TOKENS_REAL, 0.0, -np.inf).astype(np.float32) self.cross_mask = {r: Const(build, np.broadcast_to(row, (r, T.TOKENS)), "bfloat16", memory_config=build.attn_mem()) for r in self.rows} if build.dec_mem() is not None: # DEC_L1: the blocks' intermediates in L1 (interleaved) for blk in self.blocks: for key in ("out", "m1a", "m1b", "cout", "m2a", "m2b"): blk[key].out_mem = build.dec_mem() for key in ("ln", "n3", "n4"): blk[key].mem = build.dec_mem() self.edge_mem = build.dec_mem() if build.dec_l1 >= 2 else None if self.edge_mem is not None: # DEC_L1=2: the pre-projection and the final layer too self.pre1.out_mem = self.pre2.out_mem = self.p1.out_mem = self.edge_mem self.fin_ln.mem = self.p0.mem = self.p3.mem = self.edge_mem if build.attn_mem() is not None: # ATTN_L1: the fused attention reads Q / K / V from L1 for blk in self.blocks: blk["qkv"].out_mem = blk["cq"].out_mem = build.attn_mem() def upload_step(self, tables: P.StepTables, k: int) -> Dict[str, Any]: b = self.build return {"blocks": [{key: b.upload(np.asarray(blk[key]).reshape(1, -1), "float32") for key in STEP_KEYS} for blk in tables.blocks[k]], "g": b.upload(np.asarray(tables.final[k]["g"]).reshape(1, -1), "float32"), "b": b.upload(np.asarray(tables.final[k]["b"]).reshape(1, -1), "float32")} # ------------------------------------------------------------------------------------------------------------ def cross_kv(self, enc) -> List[Tuple[Any, Any]]: """Hoisted cross-attention K / V heads of the three blocks (once per plan) from the encoding ``[1, 1, 576, 256]``: ``[1, 8, 564, 32]`` (the 564 real tokens; SDPA masks its own tile padding) or, with the fp32 matmul attention, ``[1, 8, 576, 32]`` fp32 (the pad tokens are masked by ``cross_mask``).""" import ttnn if not self.attn_mm["cross"] and int(enc.shape[-2]) != T.TOKENS_REAL: enc = ttnn.slice(enc, [0, 0, 0, 0], [1, 1, T.TOKENS_REAL, C.HIDDEN_DIM]) kv = [(A.split_heads(k(enc), C.NUM_HEADS), A.split_heads(v(enc), C.NUM_HEADS)) for k, v in zip(self.k_lin, self.v_lin)] mem = self.build.attn_mem() if mem is not None and self.attn_mm["cross"]: # ATTN_L1: the hoisted heads L1-resident for 11 evaluations kv = [(ttnn.to_memory_config(k, mem), ttnn.to_memory_config(v, mem)) for k, v in kv] return kv def _attend(self, kind: str, q, k, v, mask): """Attention of ``kind`` ("self" / "cross") -> concatenated heads ``[1, 1, 352, 256]``.""" if self.attn_mm[kind]: return 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)) return A.sdpa(q, k, v, scale=self.scale, attn_mask=mask, concat_heads=True, fp32_acc=self.attn_fp32[kind]) def _fused(self, kind: str, q, k, v, mask, flat_heads, kcat: int = 0): """``ATTN_FUSED``: the attention of ``kind`` as one program (``tt/fattn_kernel.py``) reading Q / K / V in place, returning the merged heads ``[1, 1, 352, 256]``; ``flat_heads`` = the head-block column offsets (in heads) of Q, K, V inside one projection output (self), None for the cross-attention (Q flat, K / V the hoisted heads). None when the fused program does not apply (then the stock chain runs).""" if not (self.build.attn_fused and self.attn_mm[kind] and mask is not None and self.build.attn_smask and self.build.attn_smsm and self.build.attn_fast == 1): return None from .fattn_kernel import flat_at, fused_attention, heads_at, supported if not supported((q, k, v), mask): return None H = C.NUM_HEADS if flat_heads is not None: ats = [flat_at(t, o * H) for t, o in zip((q, k, v), flat_heads)] else: ats = [flat_at(q, 0), heads_at(k), heads_at(v)] out = fused_attention(q, k, v, mask, self.scale, H, *ats, kcat_ktp=kcat, memory_config=self.build.dec_mem()) return KcatOperand(out, kcat) if kcat else out def evaluate(self, x, rows: Mapping[str, Any], kv: Sequence[Tuple[Any, Any]], self_mask): """One decoder evaluation: ``x`` ``[1, 1, R, 324]`` fp32 (prefix-constrained; R = 352 or an agent bucket) -> ``m * mask0`` fp32.""" import ttnn R = int(tuple(x.shape)[-2]) h = ttnn.add(chain2(self.pre1, self.pre2, x, self.build.kcat_act), self.agent_rows[R](), # [1, 1, R, 256] **({} if self.edge_mem is None else {"memory_config": self.edge_mem})) if self.build.ln_resid: # LN_RESID: the stream adds fused into the LNs pending = None for b, r, (kc, vc) in zip(self.blocks, rows["blocks"], kv): ek = {key: operand_ktp(self.build, b[key]) for key in ("qkv", "out", "m1a", "cq", "cout", "m2a")} if pending is None: n1 = b["ln"](h, r["n1_g"], r["n1_b"], kcat=ek["qkv"]) else: h, n1 = b["ln"].residual(h, pending, None, r["n1_g"], r["n1_b"], kcat=ek["qkv"]) qkv = b["qkv"](n1) a = self._fused("self", qkv, qkv, qkv, self_mask, (0, 1, 2), ek["out"]) if a is None: q, k, v = A.split_qkv(qkv, C.NUM_HEADS) a = self._attend("self", q, k, v, self_mask) a = b["out"](a) h, n2 = b["ln"].residual(h, a, r["gate_msa"], r["n2_g"], r["n2_b"], kcat=ek["m1a"]) m = chain2(b["m1a"], b["m1b"], n2, self.build.kcat_act) h, n3 = b["n3"].residual(h, m, r["gate_mlp"], kcat=ek["cq"]) cq = b["cq"](n3) cmask = self.cross_mask[R]() if self.attn_mm["cross"] else None c = self._fused("cross", cq, kc, vc, cmask, None, ek["cout"]) if c is None: c = self._attend("cross", A.split_heads(cq, C.NUM_HEADS), kc, vc, cmask) h, n4 = b["n4"].residual(h, b["cout"](c), kcat=ek["m2a"]) pending = chain2(b["m2a"], b["m2b"], n4, self.build.kcat_act) _, f = self.fin_ln.residual(h, pending, None, rows["g"], rows["b"], write_h=False) return self.p4(self.p3(self.p1(self.p0(f)))) for b, r, (kc, vc) in zip(self.blocks, rows["blocks"], kv): q, k, v = A.split_qkv(b["qkv"](b["ln"](h, r["n1_g"], r["n1_b"])), C.NUM_HEADS) a = b["out"](self._attend("self", q, k, v, self_mask)) h = ttnn.add(h, ttnn.multiply(a, r["gate_msa"])) m = b["m1b"](b["m1a"](b["ln"](h, r["n2_g"], r["n2_b"]))) h = ttnn.add(h, ttnn.multiply(m, r["gate_mlp"])) qc = A.split_heads(b["cq"](b["n3"](h)), C.NUM_HEADS) cmask = self.cross_mask[R]() if self.attn_mm["cross"] else None h = ttnn.add(h, b["cout"](self._attend("cross", qc, kc, vc, cmask))) h = ttnn.add(h, b["m2b"](b["m2a"](b["n4"](h)))) f = self.fin_ln(h, rows["g"], rows["b"]) return self.p4(self.p3(self.p1(self.p0(f)))) def solve(self, y0, cs, kv, self_mask, *, ego_steps: bool = True): """The 11 evaluations with the DPM-Solver++(2M) updates. Returns ``(final_x0, ego_rows)``: ``final_x0`` ``[1, 1, 352, 324]`` fp32 and, with ``ego_steps``, the ego row of the 11 published iterates ``[1, 1, 11, 324]`` (``~/debug/denoising_steps``; ``None`` otherwise).""" import ttnn y, m_prev, ego = y0, None, [] for k in range(self.tables.nfe): x = ttnn.add(y, cs) # prefix constraint if ego_steps and k > 0: ego.append(ttnn.slice(x, [0, 0, 0, 0], [1, 1, 1, T.STATE_COLS])) m = self.evaluate(x, self.step_rows[k], kv, self_mask) if k == self.tables.nfe - 1: break a_k, b_k, c_k = self.tables.solver[k] y = ttnn.subtract(ttnn.multiply(y, a_k), ttnn.multiply(m, b_k)) if m_prev is not None: y = ttnn.add(y, ttnn.multiply(m_prev, c_k)) m_prev = m final = ttnn.add(m, cs) if not ego_steps: return final, None ego.append(ttnn.slice(final, [0, 0, 0, 0], [1, 1, 1, T.STATE_COLS])) return final, ttnn.concat(ego, dim=2) class TtTurnHead: """``logit [1, 1, 1, 5]`` from ``final_x0`` (row 0) and the token sum of the 564 real encoding rows.""" def __init__(self, build: Build, p: Mapping[str, np.ndarray]): w = P.turn_weights(p) self.sel = make_linear(build, P.Lin(w["w_sel"], w["b"]), "turn", out="float32") self.pool = make_linear(build, P.Lin(w["w_pool"], None), "turn", out="float32") ones = np.zeros((1, T.TOKENS), np.float32) ones[0, :T.TOKENS_REAL] = 1.0 self.ones = Const(build, ones, "bfloat16") self.cfg = build.cfg("turn") def __call__(self, final_x0, encoding): import ttnn row0 = ttnn.slice(final_x0, [0, 0, 0, 0], [1, 1, 1, T.STATE_COLS]) s = ttnn.matmul(self.ones(), encoding, dtype=ttnn.float32, compute_kernel_config=self.cfg) # [1,1,1,256] return ttnn.add(self.sel(row0), self.pool(s))