Download code/tt_diffusion_planner/tt/ln_kernel.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 13.5 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/ln_kernel.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/ln_kernel.py
-
curl -L -o ln_kernel.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/ln_kernel.py
13.5 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Fused fp32 LayerNorm as one ``ttnn.generic_op`` program (``LN_KERNEL``, OPT round 2 item 3). | |
| :func:`layer_norm_fp32_fused` computes exactly what :func:`tt.layers.layer_norm_fp32` computes with 7-9 stock | |
| programs (mean, subtract, square, mean, add eps, rsqrt, multiply, gamma, beta), in one program: the compute kernel | |
| (``kernels/ln32_compute.cpp``) issues the same SFPU LLK calls in the same order as the stock kernels, and every | |
| intermediate the stock graph writes to an fp32 DRAM tensor stays fp32 in L1 / DST, so the output is meant to be | |
| bit-identical (checked on the device: ``code/scripts/ln_kernel_check.py``). | |
| Layout: x is an fp32 TILE DRAM-interleaved tensor ``[..., R, W]`` (R a multiple of 32, W = 32 * Wt); the tile rows | |
| are split in contiguous blocks over ``min(rows, grid)`` cores (one core per tile row up to the grid). gamma / beta | |
| are ``[1, 1, 1, W]`` fp32 TILE rows (row 0 valid) or None. | |
| """ | |
| from __future__ import annotations | |
| import os | |
| import struct | |
| from typing import Any, Optional | |
| __all__ = ["layer_norm_fp32_fused", "supported"] | |
| _KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels") | |
| TB = 4096 # fp32 tile bytes | |
| N_RT = 2 # per-core RT args of every kernel ([row0, n_rows] / [n_rows] + pad): a CT arg (probe P3) | |
| def _bits(v: float) -> int: | |
| return int.from_bytes(struct.pack("<f", float(v)), "little") | |
| def supported(x: Any) -> bool: | |
| """True when ``x`` is a device fp32 TILE interleaved tensor with tile-aligned rows and columns.""" | |
| 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 int(x.shape[-1]) == shp[-1] | |
| and hasattr(ttnn, "generic_op")) | |
| except Exception: # noqa: BLE001 - the host fake ttnn | |
| return False | |
| def _cores(n: int, g): | |
| """The first ``n`` cores of the grid in row-major order: a CoreRangeSet and the coordinate list.""" | |
| import ttnn | |
| 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))) | |
| return ttnn.CoreRangeSet(set(rs)), cs | |
| def layer_norm_fp32_fused(x, gamma=None, beta=None, *, eps: float, residual=None, rgate=None, write_h: bool = True, | |
| lean: bool = True, sfpu_bcast: bool = False, res_t: bool = False, out_t: bool = False, | |
| memory_config=None): | |
| """LayerNorm over the last dim of fp32 ``x`` -> fp32 tensor of x's shape (see the module docstring). | |
| With ``residual`` (``LN_RESID``): ``h = x + residual (* rgate)`` first (the stock ``ttnn.add(x, | |
| ttnn.multiply(residual, rgate))``, fp32 SFPU ops), then ``LN(h)``; returns ``(h, LN(h))`` (``h`` is None when | |
| ``write_h`` is False). ``residual``: fp32 like ``x``; ``rgate``: a ``[1, 1, 1, W]`` fp32 row. | |
| ``lean``: one unpacker / SFPU-binary init per phase instead of one per tile (same LLK math calls). | |
| ``sfpu_bcast`` (``LN_SFPU_BCAST``): the SFPU row reduce writes the row statistics to every column | |
| (``kernels/ln32_sfpu.h``) instead of the writer's RISC-V column fill (same values). | |
| ``res_t`` / ``out_t`` (``LN_TR``, x ``[.., E, T, W]``): the residual comes as ``[.., E, W, T]`` (its per-entity | |
| transpose) / the output is written as ``[.., E, W, T]``, the tiles transposed in the kernel with the stock | |
| ``ttnn.transpose`` LLK (exact): the two transposes around the mixer's token-mixing MLP. ``memory_config``: of | |
| the outputs (default DRAM; interleaved L1 for ``ENC_L1``).""" | |
| import ttnn | |
| dev = x.device() | |
| shp = list(x.padded_shape) | |
| W = shp[-1] | |
| Wt = W // 32 | |
| rows = 1 | |
| for d in shp[:-1]: | |
| rows *= d | |
| rows //= 32 | |
| hr = int(residual is not None) | |
| hrg = int(hr and rgate is not None) | |
| wh = int(hr and write_h) | |
| Tt = shp[-2] // 32 | |
| assert not (res_t and hrg), "LN_TR: no residual gate" | |
| oshape = x.shape if not out_t else ttnn.Shape(list(x.shape)[:-2] + [shp[-1], shp[-2]]) | |
| omem = memory_config or ttnn.DRAM_MEMORY_CONFIG | |
| out = ttnn.allocate_tensor_on_device(oshape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem) | |
| h = (ttnn.allocate_tensor_on_device(x.shape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem) | |
| if wh else None) | |
| g = dev.compute_with_storage_grid_size() | |
| n = min(rows, g.x * g.y) | |
| crs, cs = _cores(n, g) | |
| 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 | |
| hg, hb = int(gamma is not None), int(beta is not None) | |
| def acc(t): | |
| return list(ttnn.TensorAccessorArgs(t).get_compile_time_args()) | |
| 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 * Wt), cb(3, 1), cb(4, 2), cb(5, 2), cb(6, Wt), cb(16, 2 * Wt if Wt <= 4 else Wt)] | |
| if hg: | |
| cbs.append(cb(1, Wt)) | |
| if hb: | |
| cbs.append(cb(2, Wt)) | |
| if hr: | |
| cbs += [cb(7, Wt), cb(9, Wt)] | |
| if hrg: | |
| cbs.append(cb(8, Wt)) | |
| if wh: | |
| cbs.append(cb(17, 2)) | |
| if out_t: | |
| cbs.append(cb(18, 1)) | |
| um = [ttnn.UnpackToDestMode.Default] * 64 | |
| for i in (0, 1, 2, 3, 5, 6, 7, 8, 9, 18): | |
| um[i] = 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 | |
| g_t = gamma if hg else x | |
| b_t = beta if hb else x | |
| r_t = residual if hr else x | |
| rg_t = rgate if hrg else x | |
| h_t = h if wh else out | |
| reader_ct = [Wt, hg, hb, _bits(eps), N_RT, hr, hrg, int(bool(res_t)), Tt] + acc(x) + acc(g_t) + acc(b_t) + acc( | |
| r_t) + acc(rg_t) | |
| writer_ct = [Wt, N_RT, wh, int(bool(sfpu_bcast)), int(bool(out_t)), Tt] + acc(out) + acc(h_t) | |
| compute_ct = [Wt, hg, hb, _bits(1.0 / W), N_RT, hr, hrg, wh, int(bool(lean)), int(bool(sfpu_bcast)), | |
| int(bool(res_t)), int(bool(out_t))] | |
| SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH | |
| 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, compiler_include_paths=[_KDIR]) | |
| ks = [kd("ln32_reader.cpp", reader_ct, rd, [x.buffer_address(), g_t.buffer_address(), b_t.buffer_address(), | |
| r_t.buffer_address(), rg_t.buffer_address()], | |
| ttnn.ReaderConfigDescriptor()), | |
| kd("ln32_writer.cpp", writer_ct, wr, [out.buffer_address(), h_t.buffer_address()], | |
| ttnn.WriterConfigDescriptor()), | |
| kd("ln32_compute.cpp", compute_ct, cp, [], ccfg)] | |
| ins = [x] + ([gamma] if hg else []) + ([beta] if hb else []) + ([residual] if hr else []) + ( | |
| [rgate] if hrg else []) + ([h] if wh else []) | |
| ttnn.generic_op(ins + [out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs)) | |
| if hr: | |
| return h, out | |
| return out | |
| def reference_decomposition(x, gamma: Optional[Any] = None, beta: Optional[Any] = None, *, eps: float): | |
| """The stock decomposition (``tt.layers.layer_norm_fp32``), for the device check.""" | |
| from .layers import layer_norm_fp32 | |
| return layer_norm_fp32(x, gamma, beta, eps=eps) | |
| def split_supported(x: Any) -> bool: | |
| """The split-row form needs every tile row's Wt column tiles on distinct cores (rows * Wt <= grid cores).""" | |
| if not supported(x): | |
| return False | |
| shp = list(x.padded_shape) | |
| rows = 1 | |
| for d in shp[:-1]: | |
| rows *= d | |
| rows //= 32 | |
| g = x.device().compute_with_storage_grid_size() | |
| wt = shp[-1] // 32 | |
| return 2 <= wt <= g.x * g.y and rows * wt <= g.x * g.y | |
| def layer_norm_fp32_split(x, gamma=None, beta=None, *, eps: float, residual=None, rgate=None, write_h: bool = True, | |
| sfpu_bcast: bool = False, kcat_ktp: int = 0, memory_config=None): | |
| """:func:`layer_norm_fp32_fused` with each tile row spread over Wt cores (``LN_SPLIT``, ``kernels/ln32s_*.cpp``): | |
| member j owns column tile j; the root (member 0) folds the gathered tiles in order and broadcasts the mean and | |
| rstd, so the result is the same bit for bit. For few rows (the decoder's 11 tile rows: 88 cores instead of 11). | |
| ``kcat_ktp`` (``KCAT_EMIT``): instead of y, write the split operand ``[y_hi | y_hi | y_lo | 1 | 0..]`` of the next | |
| K-concatenated linear (``[..., R, 32 * kcat_ktp]``, the ``kcat_operand`` layout and LLK calls). | |
| ``memory_config``: of the outputs (default DRAM; interleaved L1 for ``DEC_L1``).""" | |
| import ttnn | |
| dev = x.device() | |
| shp = list(x.padded_shape) | |
| W = shp[-1] | |
| Wt = W // 32 | |
| omem = memory_config or ttnn.DRAM_MEMORY_CONFIG | |
| rows = 1 | |
| for d in shp[:-1]: | |
| rows *= d | |
| rows //= 32 | |
| hg, hb = int(gamma is not None), int(beta is not None) | |
| hr = int(residual is not None) | |
| hrg = int(hr and rgate is not None) | |
| wh = int(hr and write_h) | |
| oshape = x.shape if not kcat_ktp else ttnn.Shape(list(x.shape)[:-1] + [32 * kcat_ktp]) | |
| out = ttnn.allocate_tensor_on_device(oshape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem) | |
| h = (ttnn.allocate_tensor_on_device(x.shape, ttnn.float32, ttnn.TILE_LAYOUT, dev, omem) | |
| if wh else None) | |
| g = dev.compute_with_storage_grid_size() | |
| G = min(rows, (g.x * g.y) // Wt) | |
| n = G * Wt | |
| crs, cs = _cores(n, g) | |
| phys = [dev.worker_core_from_logical_core(ttnn.CoreCoord(cx, cy)) for cx, cy in cs] | |
| rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs() | |
| n_rt = 6 + 2 * Wt | |
| for gi in range(G): | |
| n_rows = len(range(gi, rows, G)) | |
| root = phys[gi * Wt] | |
| members = [] | |
| for m in range(Wt): | |
| members += [phys[gi * Wt + m].x, phys[gi * Wt + m].y] | |
| for j in range(Wt): | |
| cx, cy = cs[gi * Wt + j] | |
| args = [gi, n_rows, G, j, root.x, root.y] + members | |
| rd[cx][cy] = args | |
| wr[cx][cy] = args | |
| cp[cx][cy] = [n_rows, int(j == 0)] | |
| def acc(t): | |
| return list(ttnn.TensorAccessorArgs(t).get_compile_time_args()) | |
| 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(3, 1), cb(4, 2), cb(5, 2), cb(6, 2), cb(10, Wt), cb(11, 2), cb(12, 1), cb(16, 2)] | |
| if hg: | |
| cbs.append(cb(1, 1)) | |
| if hb: | |
| cbs.append(cb(2, 1)) | |
| if hr: | |
| cbs += [cb(7, 2), cb(9, 2)] | |
| if hrg: | |
| cbs.append(cb(8, 1)) | |
| if wh: | |
| cbs.append(cb(17, 2)) | |
| if kcat_ktp: | |
| cbs += [cb(18, 1), cb(19, 2), cb(20, 1)] + ([cb(21, 1)] if kcat_ktp > 3 * Wt + 1 else []) | |
| um = [ttnn.UnpackToDestMode.Default] * 64 | |
| for i in (0, 1, 2, 3, 5, 6, 7, 8, 9, 10, 18): | |
| um[i] = 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 | |
| g_t = gamma if hg else x | |
| b_t = beta if hb else x | |
| r_t = residual if hr else x | |
| rg_t = rgate if hrg else x | |
| h_t = h if wh else out | |
| reader_ct = [Wt, hg, hb, _bits(eps), n_rt, hr, hrg] + acc(x) + acc(g_t) + acc(b_t) + acc(r_t) + acc(rg_t) | |
| writer_ct = [Wt, n_rt, wh, int(bool(sfpu_bcast)), int(bool(kcat_ktp)), int(kcat_ktp)] + acc(out) + acc(h_t) | |
| compute_ct = [Wt, hg, hb, _bits(1.0 / W), 2, hr, hrg, wh, int(bool(sfpu_bcast)), int(bool(kcat_ktp))] | |
| SRC = ttnn.KernelDescriptor.SourceType.FILE_PATH | |
| 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, compiler_include_paths=[_KDIR]) | |
| ks = [kd("ln32s_reader.cpp", reader_ct, rd, [x.buffer_address(), g_t.buffer_address(), b_t.buffer_address(), | |
| r_t.buffer_address(), rg_t.buffer_address()], | |
| ttnn.ReaderConfigDescriptor()), | |
| kd("ln32s_writer.cpp", writer_ct, wr, [out.buffer_address(), h_t.buffer_address()], | |
| ttnn.WriterConfigDescriptor()), | |
| kd("ln32s_compute.cpp", compute_ct, cp, [], ccfg)] | |
| sems = [ttnn.SemaphoreDescriptor(id=i, core_ranges=crs, initial_value=0) for i in range(2)] | |
| ins = [x] + ([gamma] if hg else []) + ([beta] if hb else []) + ([residual] if hr else []) + ( | |
| [rgate] if hrg else []) + ([h] if wh else []) | |
| ttnn.generic_op(ins + [out], ttnn.ProgramDescriptor(kernels=ks, semaphores=sems, cbs=cbs)) | |
| if hr: | |
| return h, out | |
| return out | |