# 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