Download code/tt_diffusion_planner/tt/attention.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 5.4 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/attention.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/attention.py
-
curl -L -o attention.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/attention.py
5.4 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """The fp32 matmul attention of the fusion transformer and the decoder (C20 ``attention_matmul`` math) with stock-op | |
| program choices tuned for the planner's shapes (``ATTN_FAST``, OPT round 1 item 4a). | |
| ``softmax(scale * Q K^T + mask) V`` as in ``ttaw.ops.attention.attention_matmul``, with: | |
| - (``mode >= 1``) ``P V`` on a ``MatmulMultiCoreReuseProgramConfig`` with one output tile per core (8 heads x | |
| Sq / 32 = 88 cores at Sq = 352) instead of the auto config's 11-18 cores; the K-sum in one block | |
| (``in0_block_w`` = Sk / 32): bit-identical to the auto config in the device sweep (73 -> 21 us self, 71 -> 32 us | |
| cross, 81 -> 42 us fusion; ``logs/diffusion-planner/opt_r1/attn_sweep.log``); | |
| - (``mode == 2``) the scale applied to Q (``[1, 8, Sq, 32]``, 11x smaller than the scores) before ``Q K^T`` | |
| instead of to the scores: one small multiply instead of the large one; a precision change (rounding order); | |
| - (``mode == 3``) the scale as an SFPU pre-activation of the mask add (``input_tensor_a_activations``): one | |
| program instead of two over the scores (75 -> 42 us), but NOT bit-identical (final_x0 rel 4e-3; e2e ego mean | |
| 0.143 -> 0.173 m over the 99 scenes): measured and rejected (OPT_REPORT round 1). | |
| - ``Q K^T`` keeps the auto config (the reuse configs of the sweep were slower). | |
| """ | |
| from __future__ import annotations | |
| from typing import Any, Optional | |
| __all__ = ["attention_matmul", "pv_config"] | |
| TILE = 32 | |
| # per-core L1 budget for the P block of one core (fp32, double-buffered): 18 tiles x 4 KiB x 2 = 144 KiB | |
| MAX_KB_TILES = 18 | |
| def pv_config(device: Any, heads: int, sq: int, sk: int): | |
| """``MatmulMultiCoreReuseProgramConfig`` for ``P [1, H, Sq, Sk] @ V [1, H, Sk, 32]``: one ``[32, 32]`` output | |
| tile per core (``per_core_M = 1`` when ``H * Sq / 32`` fits the grid, else 2).""" | |
| import ttnn | |
| grid = device.compute_with_storage_grid_size() | |
| mt, kt = sq // TILE, sk // TILE | |
| pcm = 1 if heads * mt <= grid.x * grid.y else 2 | |
| kb = max(d for d in range(1, min(kt, MAX_KB_TILES) + 1) if kt % d == 0) | |
| return ttnn.MatmulMultiCoreReuseProgramConfig(compute_with_storage_grid_size=grid, in0_block_w=kb, | |
| out_subblock_h=1, out_subblock_w=1, per_core_M=pcm, per_core_N=1) | |
| def attention_matmul(q: Any, k: Any, v: Any, *, scale: float, attn_mask: Any = None, | |
| compute_kernel_config: Any = None, mode: int = 1, pv_pc: Optional[Any] = None, | |
| smask: bool = False, smsm: int = 0): | |
| """``q`` ``[1, H, Sq, D]``, ``k`` / ``v`` ``[1, H, Sk, D]`` (tile-aligned ``Sk``) -> ``[1, H, Sq, D]``. | |
| ``mode=0``: exactly ``ttaw.ops.attention.attention_matmul``. ``smask`` (``ATTN_SMASK``, mode 1 with a mask): | |
| the scale multiply and the mask add as one bit-identical generic_op (``tt/smask_kernel.py``); ``smsm`` | |
| (``ATTN_SMSM`` 1 / 2, with ``smask``): the scale, the mask and the softmax as one generic_op | |
| (``tt/smsm_kernel.py``; 2 = the scale multiply by an immediate).""" | |
| import ttnn | |
| from ..ttaw.ops import attention as A | |
| from ..ttaw.precision import compute_kernel_config as ckc | |
| cfg = compute_kernel_config if compute_kernel_config is not None else ckc("HiFi4", fp32_acc=True) | |
| dev = q.device() if callable(getattr(q, "device", None)) else None | |
| if mode and (dev is None or not hasattr(dev, "compute_with_storage_grid_size") | |
| or not hasattr(getattr(ttnn, "UnaryOpType", None), "MUL_UNARY_SFPU")): | |
| mode = 0 # the host fake ttnn: the stock path (same math) | |
| if not mode: | |
| return A.attention_matmul(q, k, v, scale=scale, attn_mask=attn_mask, compute_kernel_config=cfg) | |
| sk = int(k.shape[-2]) | |
| if sk % TILE: | |
| raise ValueError(f"attention_matmul needs a tile-aligned Sk (got {sk})") | |
| if mode == 2: | |
| q = ttnn.multiply(q, float(scale)) | |
| scores = ttnn.matmul(q, k, transpose_b=True, compute_kernel_config=cfg) | |
| fused = False | |
| if smsm and smask and mode == 1 and attn_mask is not None: | |
| from .smsm_kernel import scale_mask_softmax | |
| from .smsm_kernel import supported as smsm_supported | |
| if smsm_supported(scores, attn_mask): | |
| probs = scale_mask_softmax(scores, scale, attn_mask, scale_mode=int(smsm) - 1) | |
| if pv_pc is None: | |
| pv_pc = pv_config(dev, int(q.shape[1]), int(q.shape[-2]), sk) | |
| return ttnn.matmul(probs, v, compute_kernel_config=cfg, program_config=pv_pc) | |
| if smask and mode == 1 and attn_mask is not None: | |
| from .smask_kernel import scale_mask, supported | |
| if supported(scores, attn_mask): | |
| scores, fused = scale_mask(scores, scale, attn_mask), True | |
| if fused: | |
| pass | |
| elif mode == 3 and attn_mask is not None: | |
| act = [ttnn.UnaryWithParam(ttnn.UnaryOpType.MUL_UNARY_SFPU, float(scale))] | |
| scores = ttnn.add(scores, attn_mask, input_tensor_a_activations=act) | |
| else: | |
| if mode != 2: | |
| scores = ttnn.multiply(scores, float(scale)) | |
| if attn_mask is not None: | |
| scores = ttnn.add(scores, attn_mask) | |
| probs = ttnn.softmax(scores, dim=-1, numeric_stable=True, compute_kernel_config=cfg) | |
| if pv_pc is None: | |
| pv_pc = pv_config(dev, int(q.shape[1]), int(q.shape[-2]), sk) | |
| return ttnn.matmul(probs, v, compute_kernel_config=cfg, program_config=pv_pc) | |