Download code/tt_diffusion_planner/tt/fattn_kernel.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 6.62 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/fattn_kernel.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/fattn_kernel.py
-
curl -L -o fattn_kernel.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/fattn_kernel.py
6.62 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Fused fp32 matmul attention as one ``ttnn.generic_op`` (``ATTN_FUSED``, OPT round 3 item 4). | |
| ``fused_attention(q, k, v, mask, scale, heads, q_at, k_at, v_at)`` -> the merged-heads output ``[1, 1, Sq, H * 32]`` | |
| fp32: ``softmax(scale * Q K^T + mask) V`` per head with head dim 32 (one tile), i.e. the programs | |
| ``nlp_create_qkv_heads`` (or ``split_heads``), ``Q K^T`` matmul, the scale + mask + softmax (``ATTN_SMSM``), the | |
| ``P V`` matmul and ``nlp_concat_heads`` of ``tt/attention.py`` in one program (``kernels/fattn_*.cpp``). Each unit | |
| is one head x one query tile row: ``Q K^T`` into DEST two key tiles at a time, the smask SFPU sequence, the stock | |
| softmax ``kernel_lib`` calls on the row in L1, then ``P V`` accumulated in DEST in key order, so the scores and the | |
| probabilities never leave L1 and the math is that of the stock programs (meant to be bit-identical: the device | |
| check compares with them). | |
| Q / K / V are read in place from their producers: ``q_at`` / ``k_at`` / ``v_at`` = ``(base, row_stride, | |
| head_stride)`` in tiles, e.g. the self-attention ``qkv`` projection ``[1, 1, S, 768]`` gives Q at ``(0, 24, 1)``, | |
| K at ``(8, 24, 1)``, V at ``(16, 24, 1)``; a ``[1, H, S, 32]`` heads tensor is ``(0, 1, S / 32)``. | |
| """ | |
| from __future__ import annotations | |
| import os | |
| import struct | |
| from typing import Any, Sequence | |
| __all__ = ["fused_attention", "supported", "flat_at", "heads_at"] | |
| _KDIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "kernels") | |
| TB = 4096 | |
| N_RT = 2 | |
| NDST = 4 | |
| def flat_at(t: Any, col_tile: int): | |
| """``(base, row_stride, head_stride)`` of heads stored as consecutive 32-column tiles of ``[1, 1, S, W]``, | |
| the first head at column tile ``col_tile``.""" | |
| return (int(col_tile), int(t.padded_shape[-1]) // 32, 1) | |
| def heads_at(t: Any): | |
| """``(base, row_stride, head_stride)`` of a ``[1, H, S, 32]`` heads tensor.""" | |
| return (0, 1, int(t.padded_shape[-2]) // 32) | |
| def supported(tensors: Sequence[Any], mask: Any) -> bool: | |
| import ttnn | |
| try: | |
| if not hasattr(ttnn, "generic_op") or mask is None: | |
| return False | |
| for t in tensors: | |
| if t.dtype != ttnn.float32 or t.layout != ttnn.TILE_LAYOUT or t.is_sharded(): | |
| return False | |
| return (mask.layout == ttnn.TILE_LAYOUT and mask.dtype in (ttnn.bfloat16, ttnn.float32) | |
| and not mask.is_sharded() and int(mask.padded_shape[1]) == 1) | |
| except Exception: # noqa: BLE001 - the host fake ttnn | |
| return False | |
| def fused_attention(q: Any, k: Any, v: Any, mask: Any, scale: float, heads: int, q_at, k_at, v_at, | |
| kcat_ktp: int = 0, memory_config: Any = None): | |
| """``kcat_ktp`` (``KCAT_EMIT``): write the split operand ``[o_hi | o_hi | o_lo | 1 | 0..]`` of the next | |
| K-concatenated linear (``[1, 1, Sq, 32 * kcat_ktp]``, the ``kcat_operand`` layout) instead of the output.""" | |
| import ttnn | |
| dev = q.device() | |
| Sq, Sk = int(mask.padded_shape[-2]), int(mask.padded_shape[-1]) | |
| Mt, Wt = Sq // 32, Sk // 32 | |
| units = heads * Mt | |
| Wp = (Wt + 1) // 2 | |
| rb = -(-Wt // NDST) * NDST | |
| p_pad = rb - Wt | |
| ow = 32 * kcat_ktp if kcat_ktp else heads * 32 | |
| out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, Sq, ow]), ttnn.float32, ttnn.TILE_LAYOUT, dev, | |
| memory_config or ttnn.DRAM_MEMORY_CONFIG) | |
| g = dev.compute_with_storage_grid_size() | |
| n = min(units, 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(units, n) | |
| rd, wr, cp = ttnn.RuntimeArgs(), ttnn.RuntimeArgs(), ttnn.RuntimeArgs() | |
| u0 = 0 | |
| for i, (cx, cy) in enumerate(cs): | |
| kk = base + (1 if i < extra else 0) | |
| rd[cx][cy] = [u0, kk] | |
| wr[cx][cy] = [u0, kk] | |
| cp[cx][cy] = [kk, 0] | |
| u0 += kk | |
| 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, 2, f32, TB), cb(1, 2 * Wp, mask.dtype, tbm), cb(2, 4, f32, TB), cb(3, 1, f32, TB), | |
| cb(4, 1, f32, TB), cb(5, Wt, f32, TB), cb(16, 2, f32, TB), cb(24, Wt, f32, TB), cb(25, 1, f32, TB), | |
| cb(26, Wt, f32, TB), cb(27, 1, f32, TB), cb(28, rb, f32, TB)] | |
| if kcat_ktp: | |
| cbs += [cb(17, 2, f32, TB), cb(18, 1, f32, TB), cb(29, 1, f32, TB)] | |
| if kcat_ktp > 3 * heads + 1: | |
| cbs.append(cb(19, 1, f32, TB)) | |
| um = [ttnn.UnpackToDestMode.Default] * 64 | |
| if kcat_ktp: | |
| um[29] = 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) | |
| ats = [int(x) for a in (q_at, k_at, v_at) for x in a] | |
| ks = [kd("fattn_reader.cpp", [Wt, Mt, N_RT, tbm] + ats + acc(q) + acc(k) + acc(v) + acc(mask), rd, | |
| [q.buffer_address(), k.buffer_address(), v.buffer_address(), mask.buffer_address()], | |
| ttnn.ReaderConfigDescriptor()), | |
| kd("fattn_writer.cpp", [Mt, heads, N_RT, int(bool(kcat_ktp)), int(kcat_ktp)] + acc(out), wr, | |
| [out.buffer_address()], | |
| ttnn.WriterConfigDescriptor()), | |
| kd("fattn_compute.cpp", [Wt, N_RT, int(mbf), NDST, p_pad, bits, int(bool(kcat_ktp))], cp, [], ccfg)] | |
| ins = [] | |
| for t in (q, k, v): | |
| if all(t is not x for x in ins): | |
| ins.append(t) | |
| ins += [mask, out] | |
| ttnn.generic_op(ins, ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs)) | |
| return out | |