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
|