# 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("