File size: 16,199 Bytes
4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 be62f78 4d9b003 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 | # 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))
|