Download code/tt_diffusion_planner/tt/smsm_kernel.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 4.8 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/smsm_kernel.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/smsm_kernel.py
-
curl -L -o smsm_kernel.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/smsm_kernel.py
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 | |