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