changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
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