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))