File size: 6,620 Bytes
be62f78
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# SPDX-License-Identifier: Apache-2.0
"""Fused fp32 matmul attention as one ``ttnn.generic_op`` (``ATTN_FUSED``, OPT round 3 item 4).

``fused_attention(q, k, v, mask, scale, heads, q_at, k_at, v_at)`` -> the merged-heads output ``[1, 1, Sq, H * 32]``
fp32: ``softmax(scale * Q K^T + mask) V`` per head with head dim 32 (one tile), i.e. the programs
``nlp_create_qkv_heads`` (or ``split_heads``), ``Q K^T`` matmul, the scale + mask + softmax (``ATTN_SMSM``), the
``P V`` matmul and ``nlp_concat_heads`` of ``tt/attention.py`` in one program (``kernels/fattn_*.cpp``). Each unit
is one head x one query tile row: ``Q K^T`` into DEST two key tiles at a time, the smask SFPU sequence, the stock
softmax ``kernel_lib`` calls on the row in L1, then ``P V`` accumulated in DEST in key order, so the scores and the
probabilities never leave L1 and the math is that of the stock programs (meant to be bit-identical: the device
check compares with them).

Q / K / V are read in place from their producers: ``q_at`` / ``k_at`` / ``v_at`` = ``(base, row_stride,
head_stride)`` in tiles, e.g. the self-attention ``qkv`` projection ``[1, 1, S, 768]`` gives Q at ``(0, 24, 1)``,
K at ``(8, 24, 1)``, V at ``(16, 24, 1)``; a ``[1, H, S, 32]`` heads tensor is ``(0, 1, S / 32)``.
"""
from __future__ import annotations

import os
import struct
from typing import Any, Sequence

__all__ = ["fused_attention", "supported", "flat_at", "heads_at"]

_KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels")
TB = 4096
N_RT = 2
NDST = 4


def flat_at(t: Any, col_tile: int):
    """``(base, row_stride, head_stride)`` of heads stored as consecutive 32-column tiles of ``[1, 1, S, W]``,
    the first head at column tile ``col_tile``."""
    return (int(col_tile), int(t.padded_shape[-1]) // 32, 1)


def heads_at(t: Any):
    """``(base, row_stride, head_stride)`` of a ``[1, H, S, 32]`` heads tensor."""
    return (0, 1, int(t.padded_shape[-2]) // 32)


def supported(tensors: Sequence[Any], mask: Any) -> bool:
    import ttnn

    try:
        if not hasattr(ttnn, "generic_op") or mask is None:
            return False
        for t in tensors:
            if t.dtype != ttnn.float32 or t.layout != ttnn.TILE_LAYOUT or t.is_sharded():
                return False
        return (mask.layout == ttnn.TILE_LAYOUT and mask.dtype in (ttnn.bfloat16, ttnn.float32)
                and not mask.is_sharded() and int(mask.padded_shape[1]) == 1)
    except Exception:  # noqa: BLE001 - the host fake ttnn
        return False


def fused_attention(q: Any, k: Any, v: Any, mask: Any, scale: float, heads: int, q_at, k_at, v_at,
                    kcat_ktp: int = 0, memory_config: Any = None):
    """``kcat_ktp`` (``KCAT_EMIT``): write the split operand ``[o_hi | o_hi | o_lo | 1 | 0..]`` of the next
    K-concatenated linear (``[1, 1, Sq, 32 * kcat_ktp]``, the ``kcat_operand`` layout) instead of the output."""
    import ttnn

    dev = q.device()
    Sq, Sk = int(mask.padded_shape[-2]), int(mask.padded_shape[-1])
    Mt, Wt = Sq // 32, Sk // 32
    units = heads * Mt
    Wp = (Wt + 1) // 2
    rb = -(-Wt // NDST) * NDST
    p_pad = rb - Wt
    ow = 32 * kcat_ktp if kcat_ktp else heads * 32
    out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, Sq, ow]), ttnn.float32, ttnn.TILE_LAYOUT, dev,
                                         memory_config or ttnn.DRAM_MEMORY_CONFIG)
    g = dev.compute_with_storage_grid_size()
    n = min(units, g.x * g.y)
    cs = [(i % g.x, i // g.x) for i in range(n)]
    full, rem = divmod(n, g.x)
    rs = []
    if full:
        rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(g.x - 1, full - 1)))
    if rem:
        rs.append(ttnn.CoreRange(ttnn.CoreCoord(0, full), ttnn.CoreCoord(rem - 1, full)))
    crs = ttnn.CoreRangeSet(set(rs))
    base, extra = divmod(units, n)
    rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs()
    u0 = 0
    for i, (cx, cy) in enumerate(cs):
        kk = base + (1 if i < extra else 0)
        rd[cx][cy] = [u0, kk]
        wr[cx][cy] = [u0, kk]
        cp[cx][cy] = [kk, 0]
        u0 += kk
    mbf = mask.dtype == ttnn.bfloat16
    tbm = 2048 if mbf else 4096

    def cb(idx, pages, dt, tb):
        return ttnn.CBDescriptor(total_size=pages * tb, core_ranges=crs, format_descriptors=[
            ttnn.CBFormatDescriptor(buffer_index=idx, data_format=dt, page_size=tb)])

    f32 = ttnn.float32
    cbs = [cb(0, 2, f32, TB), cb(1, 2 * Wp, mask.dtype, tbm), cb(2, 4, f32, TB), cb(3, 1, f32, TB),
           cb(4, 1, f32, TB), cb(5, Wt, f32, TB), cb(16, 2, f32, TB), cb(24, Wt, f32, TB), cb(25, 1, f32, TB),
           cb(26, Wt, f32, TB), cb(27, 1, f32, TB), cb(28, rb, f32, TB)]
    if kcat_ktp:
        cbs += [cb(17, 2, f32, TB), cb(18, 1, f32, TB), cb(29, 1, f32, TB)]
        if kcat_ktp > 3 * heads + 1:
            cbs.append(cb(19, 1, f32, TB))
    um = [ttnn.UnpackToDestMode.Default] * 64
    if kcat_ktp:
        um[29] = ttnn.UnpackToDestMode.UnpackToDestFp32
    if not mbf:
        um[1] = ttnn.UnpackToDestMode.UnpackToDestFp32
    ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True,
                                        math_approx_mode=False)
    ccfg.unpack_to_dest_mode = um
    bits = int.from_bytes(struct.pack("<f", float(scale)), "little")
    SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH

    def acc(t):
        return list(ttnn.TensorAccessorArgs(t).get_compile_time_args())

    def kd(name, ct, rt, common, config):
        return ttnn.KernelDescriptor(kernel_source=os.path.join(_KDIR, name), source_type=SRC, core_ranges=crs,
                                     compile_time_args=ct, defines=[], runtime_args=rt, common_runtime_args=common,
                                     config=config)

    ats = [int(x) for a in (q_at, k_at, v_at) for x in a]
    ks = [kd("fattn_reader.cpp", [Wt, Mt, N_RT, tbm] + ats + acc(q) + acc(k) + acc(v) + acc(mask), rd,
             [q.buffer_address(), k.buffer_address(), v.buffer_address(), mask.buffer_address()],
             ttnn.ReaderConfigDescriptor()),
          kd("fattn_writer.cpp", [Mt, heads, N_RT, int(bool(kcat_ktp)), int(kcat_ktp)] + acc(out), wr,
             [out.buffer_address()],
             ttnn.WriterConfigDescriptor()),
          kd("fattn_compute.cpp", [Wt, N_RT, int(mbf), NDST, p_pad, bits, int(bool(kcat_ktp))], cp, [], ccfg)]
    ins = []
    for t in (q, k, v):
        if all(t is not x for x in ins):
            ins.append(t)
    ins += [mask, out]
    ttnn.generic_op(ins, ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs))
    return out