Download code/tt_diffusion_planner/tt/kcat_kernel.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 5.14 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kcat_kernel.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kcat_kernel.py
-
curl -L -o kcat_kernel.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kcat_kernel.py
5.14 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Operand build of the K-concatenated split matmul (``SPLIT_KCAT=2``, OPT round 2 item 2) as one ``generic_op``. | |
| ``kcat_operand(x)``: fp32 TILE ``[..., M, K]`` (K a multiple of 32) -> fp32 ``[..., M, 3K + 32]`` = | |
| ``[x_hi | x_hi | x_lo | ones]`` with ``x_hi`` = bf16(x) (the ``ttnn.typecast`` LLK) and ``x_lo = x - x_hi``; the | |
| last tile has columns 0 and 1 = 1 (the exact bias rows of the concatenated weight). It replaces the 5 stock programs | |
| of ``SPLIT_KCAT=1`` (typecast, typecast, subtract, concat over 4 tensors). x tiles are split in contiguous ranges | |
| over the grid (``kernels/kcat_*.cpp``). | |
| """ | |
| from __future__ import annotations | |
| import os | |
| from typing import Any | |
| __all__ = ["kcat_operand", "kcat_tiles", "supported"] | |
| _KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels") | |
| TB = 4096 | |
| N_RT = 2 | |
| def supported(x: Any) -> bool: | |
| """True for a device fp32 TILE interleaved tensor with tile-aligned rows and columns (not the host fake).""" | |
| import ttnn | |
| try: | |
| shp = list(x.padded_shape) | |
| return (x.dtype == ttnn.float32 and x.layout == ttnn.TILE_LAYOUT and not x.is_sharded() | |
| and shp[-1] % 32 == 0 and shp[-2] % 32 == 0 and hasattr(ttnn, "generic_op")) | |
| except Exception: # noqa: BLE001 - the host fake ttnn | |
| return False | |
| def kcat_tiles(k: int, pad: int) -> int: | |
| """Row tiles of X' for an input of K = ``k`` columns: 3 Kt + 1, rounded up to a multiple of ``pad``.""" | |
| return -(-(3 * (k // 32) + 1) // pad) * pad | |
| ACTS = {None: 0, "gelu": 1, "gelu_tanh": 2} | |
| def kcat_operand(x: Any, pad: int = 1, act: Any = None, memory_config: Any = None, act_once: bool = False): | |
| """``pad``: X' row tiles rounded up to a multiple of ``pad`` with zero tiles (the matmul's K block must divide | |
| them; 3 Kt + 1 is odd, e.g. 97 for K = 1024). ``act`` (``KCAT_ACT``): ``"gelu"`` / ``"gelu_tanh"`` applied to | |
| x first (the previous linear's activation, the same LLK as the stock ``ttnn.gelu`` program). ``memory_config``: | |
| of X' (default DRAM interleaved; ``KCAT_L1``: the L1 block-sharded layout the consumer matmul reads in place, the | |
| writer's TensorAccessor resolves the shard of each tile). ``act_once`` (``KCAT_ACT_ONCE``): the activation on one | |
| DEST copy of the tile, copied to the second by ``copy_dest_values`` (instead of on both copies).""" | |
| import ttnn | |
| dev = x.device() | |
| shp = list(x.padded_shape) | |
| K = shp[-1] | |
| assert K % 32 == 0 and shp[-2] % 32 == 0 and x.dtype == ttnn.float32 and x.layout == ttnn.TILE_LAYOUT | |
| Kt = K // 32 | |
| rows = 1 | |
| for d in shp[:-1]: | |
| rows *= d | |
| rows //= 32 | |
| total = rows * Kt | |
| Ktp = kcat_tiles(K, pad) | |
| oshape = list(x.shape)[:-1] + [32 * Ktp] | |
| out = ttnn.allocate_tensor_on_device(ttnn.Shape(oshape), ttnn.float32, ttnn.TILE_LAYOUT, dev, | |
| memory_config or ttnn.DRAM_MEMORY_CONFIG) | |
| g = dev.compute_with_storage_grid_size() | |
| n = min(total, 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(total, n) | |
| rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs() | |
| t0 = 0 | |
| for i, (cx, cy) in enumerate(cs): | |
| k = base + (1 if i < extra else 0) | |
| rd[cx][cy] = [t0, k] | |
| wr[cx][cy] = [t0, k] | |
| cp[cx][cy] = [k, 0] | |
| t0 += k | |
| def cb(idx, pages): | |
| return ttnn.CBDescriptor(total_size=pages * TB, core_ranges=crs, format_descriptors=[ | |
| ttnn.CBFormatDescriptor(buffer_index=idx, data_format=ttnn.float32, page_size=TB)]) | |
| cbs = [cb(0, 2), cb(16, 2), cb(17, 2), cb(18, 1)] + ([cb(19, 1)] if Ktp > 3 * Kt + 1 else []) | |
| um = [ttnn.UnpackToDestMode.Default] * 64 | |
| um[0] = 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 | |
| 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("kcat_reader.cpp", [N_RT] + acc(x), rd, [x.buffer_address()], ttnn.ReaderConfigDescriptor()), | |
| kd("kcat_writer.cpp", [Kt, N_RT, Ktp] + acc(out), wr, [out.buffer_address()], ttnn.WriterConfigDescriptor()), | |
| kd("kcat_compute.cpp", [N_RT, ACTS[act], int(bool(act_once))], cp, [], ccfg)] | |
| ttnn.generic_op([x, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs)) | |
| return out | |