changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
6.62 kB
# 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