# SPDX-License-Identifier: Apache-2.0 """Attention score scale + additive mask + softmax as one ``ttnn.generic_op`` (``ATTN_SMSM``, OPT round 3 item 1). ``scale_mask_softmax(s, scale, mask)``: fp32 TILE scores ``[1, H, Sq, Sk]`` and a ``[1, 1, Sq, Sk]`` mask (bf16 or fp32, broadcast over the heads) -> ``softmax(s * scale + mask, dim=-1)`` fp32: the smask program (``tt/smask_kernel.py``) and the stock ``ttnn.softmax(numeric_stable=True)`` program in one pass (``kernels/smsm_*.cpp``). Phase A is the smask kernel's SFPU sequence, phase B the stock softmax compute's ``kernel_lib`` calls in the same order, so the result is meant to be bit-identical to the two programs; the scaled scores stay in L1 instead of a 4-8 MB DRAM round trip. Work split: one tile row (head, 32 query rows) at a time, the rows in contiguous ranges over the grid (the stock softmax's split). """ from __future__ import annotations import os import struct from typing import Any __all__ = ["scale_mask_softmax", "supported"] _KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels") TB = 4096 N_RT = 2 NDST = 4 # fp32 DEST half-sync: 4 tiles (the stock softmax block size) def supported(s: Any, mask: Any) -> bool: from .smask_kernel import supported as sm_supported try: return sm_supported(s, mask) and int(s.padded_shape[-1]) // 32 <= 40 except Exception: # noqa: BLE001 - the host fake ttnn return False def scale_mask_softmax(s: Any, scale: float, mask: Any, scale_mode: int = 0): """``scale_mode`` 0: the scale multiply as the smask kernel does it (``mul_binary_tile`` against a tile filled with the scale); 1: ``mul_unary_tile`` with the scale's bits (the same fp32 SFPU multiply without the tile copy).""" import ttnn dev = s.device() _, H, Sq, Sk = list(s.padded_shape) Mt, Wt = Sq // 32, Sk // 32 rows = H * Mt Wp = (Wt + 1) // 2 rb = -(-Wt // NDST) * NDST out_pad = rb - Wt out = ttnn.allocate_tensor_on_device(s.shape, ttnn.float32, ttnn.TILE_LAYOUT, dev, ttnn.DRAM_MEMORY_CONFIG) g = dev.compute_with_storage_grid_size() n = min(rows, 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(rows, n) rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs() r0 = 0 for i, (cx, cy) in enumerate(cs): k = base + (1 if i < extra else 0) rd[cx][cy] = [r0, k] wr[cx][cy] = [r0, k] cp[cx][cy] = [k, 0] r0 += k 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, 8, f32, TB), cb(1, 2 * Wp, mask.dtype, tbm), cb(2, 1, f32, TB), cb(3, 1, f32, TB), cb(4, 1, f32, TB), cb(16, 2 * rb, f32, TB), cb(24, Wt, f32, TB), cb(25, 1, f32, TB), cb(26, Wt, f32, TB), cb(27, 1, f32, TB)] um = [ttnn.UnpackToDestMode.Default] * 64 um[0] = um[2] = 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("