changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
4.8 kB
# 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("<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)
ks = [kd("smsm_reader.cpp", [Wt, Mt, bits, N_RT, tbm] + acc(s) + acc(mask), rd,
[s.buffer_address(), mask.buffer_address()], ttnn.ReaderConfigDescriptor()),
kd("smsm_writer.cpp", [Wt, NDST, N_RT, out_pad] + acc(out), wr, [out.buffer_address()],
ttnn.WriterConfigDescriptor()),
kd("smsm_compute.cpp", [Wt, N_RT, int(mbf), NDST, out_pad, int(scale_mode), bits], cp, [], ccfg)]
ttnn.generic_op([s, mask, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs))
return out