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