Compare commits

...
Author SHA1 Message Date
Will Lin a24b8fc02a [perf] Add fine-grained FA4 VSA kernels 2026-08-17 17:39:20 +00:00
64 changed files with 52248 additions and 56 deletions
+56 -16
View File
@@ -64,28 +64,42 @@ cd fastvideo-kernel
./build.sh --rocm
```
### Optional: FA4 CuTe block-sparse backend (VSA-256 fastpath)
### Optional: FA4 CuTe block-sparse backend
The VSA-256 fastpath (tile volume 256, on NVIDIA Blackwell / sm_100) routes to the
FlashAttention-4 CuTe-DSL block-sparse kernel exposed as `flash_attn.cute`. This is
an **optional** dependency: it is imported lazily, and `video_sparse_attn`
transparently falls back to the Triton backend when it is absent (so the package is
fully usable without it).
The NVIDIA Blackwell / sm_100 VSA fastpath routes to FastVideo's in-tree
FlashAttention-4 CuTe-DSL source fork, exposed as `flash_attn.cute`. This is an
**optional** dependency: it is imported lazily, and `video_sparse_attn`
transparently falls back to the Triton backend when it is absent (so the package
is fully usable without it).
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`,
`flash_attn.cute.interface._flash_attn_fwd`) are provided upstream by
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention). Pin to
commit `940cd9680f3315f2f06b43ab5bea2c2cf2d96806`, the revision FastVideo pins as
the `flash-attn-4` source in the repo-root `pyproject.toml`; other revisions may
have an incompatible `_flash_attn_fwd` signature.
The fork lives in `fastvideo-kernel/fa4`, retains the `flash-attn-4` distribution
name, and records its exact upstream revision and refresh procedure in
`fa4/UPSTREAM.md`. Install it editable while developing the kernels:
```bash
pip install "nvidia-cutlass-dsl>=4.5.0" torchvision
pip install "git+https://github.com/Dao-AILab/flash-attention.git@940cd9680f3315f2f06b43ab5bea2c2cf2d96806#subdirectory=flash_attn/cute"
uv pip install --no-deps --editable fastvideo-kernel/fa4
```
The production VSA-256 entrypoint keeps its existing logical Q256/KV256 mask.
When the CuTe backend is enabled, an optional shape override expands that mask
to the finer physical schedules without changing selected token pairs:
```bash
FASTVIDEO_VSA_CUTEDSL=1 FASTVIDEO_VSA_FA4_BLOCK_SHAPE=128x64 ...
# At the original Q256 API boundary, logical Q64 children are safely paired
# onto the optimized physical Q128/KV64 schedule:
FASTVIDEO_VSA_CUTEDSL=1 FASTVIDEO_VSA_FA4_BLOCK_SHAPE=64x64 ...
```
The default `FASTVIDEO_VSA_FA4_BLOCK_SHAPE=256x256` preserves the historical
route (its logical KV256 edges are implemented with physical KV128 tiles).
Direct callers with arbitrary, distinct Q64 masks still use the native
Q64/KV64 kernel. The Q128/KV64 path enables score/probability double buffering
by default; set `FASTVIDEO_FA4_VSA_SP_DOUBLE_BUFFER=0` only to run its retained
legacy schedule for comparison or rollback.
The CuTe kernel JIT-compiles on first use. Verified on Blackwell (sm_100) against
`tests/test_vsa256_forward*.py`.
`tests/test_fa4_vsa_block_shapes.py` and `tests/test_vsa256_forward*.py`.
## Usage
@@ -137,7 +151,33 @@ statistics and bitwise-compatible `dV`.
### VSA (block-sparse) TFLOPs
After building/installing `fastvideo-kernel`, run:
For the GB200 FA4 comparison, install the in-tree source fork and run the
issue-#4554-derived harness. Its default `exact256` mask mode preserves the
same selected token pairs across Q×KV block shapes, reports sparse-aware
TFLOP/s, MFU against the 2.5-PFLOP/s dense-BF16 GB200 peak, and efficiency
relative to the raw FA4 256×256 path:
```bash
uv pip install --no-deps --editable fastvideo-kernel/fa4
CUDA_VISIBLE_DEVICES=0 CUTE_DSL_ENABLE_TVM_FFI=1 \
uv run --no-sync python fastvideo-kernel/benchmarks/bench_vsa_blackwell.py \
--seq_lens 32768 --sparsities dense 90 \
--block_shapes 256x256 128x64 64x64
```
To exercise every GPU in a four-GB200 tray, launch one independent replica per
device. Each rank writes its own suffixed JSON (for example,
`/tmp/fa4_tray.rank0.json`):
```bash
CUTE_DSL_ENABLE_TVM_FFI=1 uv run --no-sync torchrun --standalone --nproc-per-node=4 \
fastvideo-kernel/benchmarks/bench_vsa_blackwell.py \
--seq_lens 32768 --sparsities dense 90 \
--block_shapes 256x256 128x64 64x64 --out /tmp/fa4_tray.json
```
The older generic benchmark below measures FastVideo's 64-token block-sparse
wrapper; it does not exercise the FA4 fine-grained source fork:
```bash
cd fastvideo-kernel
@@ -0,0 +1,984 @@
#!/usr/bin/env python3
"""Block-sparse attention micro-benchmark: flashinfer ``vsa_blackwell`` and
FA4 CuTe-DSL BSA, on identical block masks where the backends can represent
them (GB200 / SM100).
This harness originated in FlashInfer issue #4554 and its benchmark gist:
https://github.com/flashinfer-ai/flashinfer/issues/4554
https://gist.github.com/SolitaryThinker/90a1d1447929fc38dc509c1852e76532
The imported baseline is pinned to immutable gist revision
``e15ac9066f23ef3690e33e1cc1fdac45b4b9099f``.
Measures achieved sparse-aware algorithmic TFLOP/s over a sequence-length x
sparsity x (Q block, KV block) grid. Every successful arm is checked against
the same fp32 masked-SDPA reference for that cell. FlashInfer's plan()/run()
split is timed separately (VSA masks are data-dependent and change every
layer x step, so mask-per-call deployments pay both).
The default ``--mask_mode exact256`` samples one Q256/KV256 mask and expands
it into each finer block shape. Thus all shapes select exactly the same token
pairs and their output and relative-efficiency comparisons are apples-to-
apples. ``--mask_mode native`` samples independently at each requested shape.
Arms:
flashinfer — BlockSparseAttentionWrapper(backend="vsa_blackwell"),
R=C=128. It runs only when the requested logical mask is
exactly representable by 128-token tiles.
cutedsl256 — original FastVideo 256x256 wrapper arm from the pinned gist;
used as a measured relative-efficiency reference.
fa4_wrapper — FastVideo's direct fine-grained BSHD wrapper, including
block-map conversion and variable-block-size mask plumbing.
fa4 — direct ``flash_attn.cute.interface._flash_attn_fwd`` call.
Sparse Q and KV block sizes are configured independently.
For the stock 256x256 FA4 path, a logical KV256 edge expands into two physical
KV128 edges, matching ``block_sparse_attn_256``. The 128x64 and 64x64 paths
call FA4 with physical tiles (128, 64) and (64, 64), respectively. Mask-index
construction and output allocation are outside the raw FA4 timing region.
The direct ``fa4_wrapper`` 64x64 row is therefore the native Q64/KV64 kernel;
it is not the public VSA-256 ``FASTVIDEO_VSA_FA4_BLOCK_SHAPE=64x64`` adapter,
which safely pairs identical Q64 children onto physical Q128/KV64 and is
measured by selecting that environment variable on the ``cutedsl256`` arm.
The vendored FA4 source is a flat ``flash_attn.cute`` package. It can be used
without installation by putting that directory on ``PYTHONPATH``; this script
detects the flat layout and mounts it under the expected import namespace:
PYTHONPATH="$PWD/fastvideo-kernel/fa4${PYTHONPATH:+:$PYTHONPATH}" \\
uv run --no-sync python \\
fastvideo-kernel/benchmarks/bench_vsa_blackwell.py --quick
MFU is reported separately from TFLOP/s. Its default denominator is the
official 2,500 TFLOP/s dense BF16 tensor-core peak for one GB200 GPU and is
explicitly overrideable with ``--peak_bf16_tflops``. Relative efficiency
against measured 256x256 baselines does not depend on that nominal peak.
Examples:
python bench_vsa_blackwell.py
python bench_vsa_blackwell.py --quick
python bench_vsa_blackwell.py --seq_lens 32768 --sparsities dense 90 --block_shapes 256x256 128x64 64x64
python bench_vsa_blackwell.py --mask_mode native --block_shapes 128x64 64x64
torchrun --standalone --nproc-per-node=4 bench_vsa_blackwell.py --quick --out /tmp/fa4_tray.json
"""
from __future__ import annotations
import argparse
import importlib
import importlib.util
import inspect
import json
import math
import os
from pathlib import Path
import random
import subprocess
import sys
import time
import traceback
import types
os.environ.setdefault("FASTVIDEO_VSA_CUTEDSL", "1")
import numpy as np
import torch
FI_BLOCK = 128 # flashinfer vsa_blackwell R=C
DEFAULT_BLOCK_SHAPES = ((256, 256), (128, 64), (64, 64))
DEFAULT_GB200_BF16_DENSE_TFLOPS = 2500.0
GB200_DATASHEET_URL = (
"https://dam-cdn.nvd.orangelogic.com/AssetLink/"
"y441155802qub41q118b2852i557jem5.pdf"
)
ISSUE_URL = "https://github.com/flashinfer-ai/flashinfer/issues/4554"
GIST_URL = "https://gist.github.com/SolitaryThinker/90a1d1447929fc38dc509c1852e76532"
GIST_REVISION = "e15ac9066f23ef3690e33e1cc1fdac45b4b9099f"
KEEP_FRAC = {
"dense": 1.0,
"50": 0.5,
"60": 0.4,
"70": 0.3,
"75": 0.25,
"80": 0.2,
"87.5": 0.125,
"90": 0.1,
}
class UnsupportedConfig(ValueError):
"""The backend cannot exactly represent the requested logical mask."""
def parse_block_shape(value: str) -> tuple[int, int]:
parts = value.lower().replace(",", "x").split("x")
if len(parts) != 2:
raise argparse.ArgumentTypeError(f"expected QxKV, got {value!r}")
try:
q_block, kv_block = (int(part) for part in parts)
except ValueError as exc:
raise argparse.ArgumentTypeError(f"expected integer QxKV, got {value!r}") from exc
if q_block <= 0 or kv_block <= 0:
raise argparse.ArgumentTypeError("Q and KV block sizes must be positive")
return q_block, kv_block
def flops_sparse_attention(
bs: int,
d: int,
selected_edges: int,
q_block: int,
kv_block: int,
) -> float:
"""QK^T + PV FLOPs over selected token pairs only.
``selected_edges`` already includes the head dimension. Each selected
logical edge covers ``q_block * kv_block`` token pairs, and each of QK^T
and PV costs two FLOPs per head-dimension element.
"""
return 4.0 * bs * d * selected_edges * q_block * kv_block
def set_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def make_logical_mask(h: int, nq: int, nkv: int, keep: int, seed: int) -> torch.Tensor:
"""[H, NQ, NKV] bool with one diagonal anchor and random other edges.
CPU generator so the pattern is device-independent and fully seeded.
"""
if not 1 <= keep <= nkv:
raise ValueError(f"keep must be in [1, {nkv}], got {keep}")
g = torch.Generator(device="cpu").manual_seed(seed)
scores = torch.rand(h, nq, nkv, generator=g)
q_rows = torch.arange(nq)
# For equal sequence lengths this is the KV block containing the first
# token of the Q block. It keeps every softmax row nonempty even when
# Q and KV block sizes differ.
diagonal_anchor = torch.div(q_rows * nkv, nq, rounding_mode="floor").clamp_max(nkv - 1)
scores[:, q_rows, diagonal_anchor] = 2.0 # rand() < 1, so this always wins topk
idx = torch.topk(scores, keep, dim=-1).indices
m = torch.zeros(h, nq, nkv, dtype=torch.bool)
m.scatter_(-1, idx, True)
assert int(m.sum(-1).min()) == keep and int(m.sum(-1).max()) == keep
return m.cuda()
def expand_exact256_mask(base_mask: torch.Tensor, q_block: int, kv_block: int) -> torch.Tensor:
"""Expand a Q256/KV256 mask without changing its selected token pairs."""
if 256 % q_block or 256 % kv_block:
raise ValueError(f"QxKV block shape {q_block}x{kv_block} must divide 256")
return base_mask.repeat_interleave(256 // q_block, dim=1).repeat_interleave(
256 // kv_block, dim=2
)
def ref_masked_sdpa_fp32(q, k, v, logical_mask, q_block, kv_block, scale, qchunk=2048):
"""fp32 SDPA with the expanded boolean token mask, chunked over q.
q,k,v: [1,S,H,D] bf16. Returns [S,H,D] fp32.
TF32 is disabled in main().
"""
_, seq_len, heads, head_dim = q.shape
if seq_len % q_block or k.shape[1] % kv_block:
raise ValueError("reference requires sequence lengths divisible by their logical block sizes")
qchunk = max(q_block, qchunk // q_block * q_block)
q32 = q[0].permute(1, 0, 2).float() # [H,S,D]
k32 = k[0].permute(1, 0, 2).float()
v32 = v[0].permute(1, 0, 2).float()
out = torch.empty(heads, seq_len, head_dim, dtype=torch.float32, device=q.device)
for i in range(0, seq_len, qchunk):
j = min(i + qchunk, seq_len)
scores = torch.matmul(q32[:, i:j], k32.transpose(-1, -2)).mul_(scale)
mask_rows = logical_mask[:, i // q_block:j // q_block]
mask_tokens = mask_rows.repeat_interleave(q_block, dim=1)
mask_tokens = mask_tokens.repeat_interleave(kv_block, dim=2) # [H,c,S]
scores.masked_fill_(~mask_tokens, float("-inf"))
out[:, i:j] = torch.matmul(torch.softmax(scores, dim=-1), v32)
del scores, mask_tokens
return out.permute(1, 0, 2).contiguous() # [S,H,D]
def arm_cutedsl(q, k, v, logical_mask, q_block, kv_block):
"""FA4-lineage CuTe-DSL 256 path. Returns (bench_fn, out [S,H,D], extra).
Timed as the full wrapper call (mask expansion + map->index included) —
that is its per-call hot path in a mask-per-call deployment.
"""
if (q_block, kv_block) != (256, 256):
raise UnsupportedConfig("cutedsl256 is defined only for logical Q256/KV256")
_mount_flat_fa4_from_pythonpath()
from fastvideo_kernel import block_sparse_attn_256 as wrapper_module
block_sparse_attn_256_bshd = wrapper_module.block_sparse_attn_256_bshd
nkv = logical_mask.shape[-1]
vbs = torch.full((nkv, ), kv_block, dtype=torch.int32, device=q.device)
mask = logical_mask.unsqueeze(0) # [1,H,NQ,NKV]
def fn():
return block_sparse_attn_256_bshd(q, k, v, mask, vbs)
out, _lse = fn()
interface = importlib.import_module("flash_attn.cute.interface")
return fn, out[0], {
"fa4_interface": str(Path(interface.__file__).resolve()),
"fastvideo_wrapper": str(Path(wrapper_module.__file__).resolve()),
"fa4_import_mode": _FA4_IMPORT_MODE,
"route_semantics": "public_vsa256",
"requested_vsa256_block_shape": os.environ.get(
"FASTVIDEO_VSA_FA4_BLOCK_SHAPE", "256x256"
),
}
def arm_fa4_wrapper(q, k, v, logical_mask, q_block, kv_block):
"""FastVideo's direct fine-grained BSHD wrapper.
This arm deliberately includes boolean-map conversion and the VBS
``mask_mod``/aux-tensor plumbing in the timed region. Fixed-length cells
use fully valid KV blocks. In particular, its 64x64 row stays on native
Q64/KV64 rather than entering the public VSA-256 coalescing adapter.
"""
if (q_block, kv_block) not in ((128, 64), (64, 64)):
raise UnsupportedConfig("fa4_wrapper supports Q128/KV64 and Q64/KV64")
_mount_flat_fa4_from_pythonpath()
from fastvideo_kernel import block_sparse_attn_cute_fwd as wrapper_module
block_sparse_attn_cute_fwd_bshd = wrapper_module.block_sparse_attn_cute_fwd_bshd
nkv = logical_mask.shape[-1]
vbs = torch.full((nkv,), kv_block, dtype=torch.int32, device=q.device)
mask = logical_mask.unsqueeze(0)
def fn():
return block_sparse_attn_cute_fwd_bshd(q, k, v, mask, vbs)
out, _lse = fn()
interface = importlib.import_module("flash_attn.cute.interface")
return fn, out[0], {
"fa4_interface": str(Path(interface.__file__).resolve()),
"fastvideo_wrapper": str(Path(wrapper_module.__file__).resolve()),
"fa4_import_mode": _FA4_IMPORT_MODE,
"route_semantics": "direct_native_fine_grained",
"physical_tile": [q_block, kv_block],
}
def arm_flashinfer(q, k, v, logical_mask, q_block, kv_block):
"""flashinfer vsa_blackwell (blk128). Returns (bench_fn, out, extra).
run() timed alone (flashinfer's documented plan-once/run-many model);
steady-state plan() wall time recorded separately per cell.
"""
if q_block % FI_BLOCK or kv_block % FI_BLOCK:
raise UnsupportedConfig(
"flashinfer vsa_blackwell R=C=128 cannot exactly represent this logical mask"
)
from flashinfer.sparse import BlockSparseAttentionWrapper
if "block_mask" not in inspect.signature(BlockSparseAttentionWrapper.plan).parameters:
raise UnsupportedConfig(
"installed FlashInfer lacks the per-head block_mask planning API used by "
"issue #4554 (the issue used flashinfer 0.6.16.post2)"
)
_, seq_len, heads, head_dim = q.shape
mask128 = logical_mask.repeat_interleave(q_block // FI_BLOCK, dim=1).repeat_interleave(
kv_block // FI_BLOCK, dim=2
)
workspace = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=q.device)
wrapper = BlockSparseAttentionWrapper(workspace, backend="vsa_blackwell")
def do_plan():
wrapper.plan(
None,
None,
M=seq_len,
N=seq_len,
R=FI_BLOCK,
C=FI_BLOCK,
num_qo_heads=heads,
num_kv_heads=heads,
head_dim=head_dim,
block_mask=mask128,
q_data_type=torch.bfloat16,
o_data_type=torch.bfloat16,
)
torch.cuda.synchronize()
t0 = time.perf_counter()
do_plan()
torch.cuda.synchronize()
plan_first = (time.perf_counter() - t0) * 1e3 # includes JIT on first cell
plan_steady = []
for _ in range(3):
torch.cuda.synchronize()
t0 = time.perf_counter()
do_plan()
torch.cuda.synchronize()
plan_steady.append((time.perf_counter() - t0) * 1e3)
qn, kn, vn = q[0], k[0], v[0] # NHD views of the same storage
def fn():
return wrapper.run(qn, kn, vn)
out = fn()
return fn, out, {
"plan_ms_first": round(plan_first, 2),
"plan_ms": round(min(plan_steady), 2),
"keep_alive": wrapper,
}
_FA4_IMPORT_MODE = "normal"
def _flat_fa4_pythonpath() -> Path | None:
"""Find FastVideo's flat ``flash_attn.cute`` source on PYTHONPATH."""
for entry in os.environ.get("PYTHONPATH", "").split(os.pathsep):
if not entry:
continue
candidate = Path(entry).resolve()
if all((candidate / name).is_file() for name in ("UPSTREAM.md", "interface.py", "block_sparsity.py")):
return candidate
return None
def _mount_flat_fa4_from_pythonpath() -> None:
"""Mount the flat vendored tree as ``flash_attn.cute`` without installing.
A regular upstream checkout already has ``flash_attn/cute`` and needs no
special handling. FastVideo intentionally vendors only that subpackage,
so its source root itself is the package directory.
"""
global _FA4_IMPORT_MODE
source = _flat_fa4_pythonpath()
if source is None:
return
try:
current = importlib.util.find_spec("flash_attn.cute.interface")
except (ImportError, AttributeError, ValueError):
current = None
if current is not None and current.origin is not None and Path(current.origin).resolve().parent == source:
_FA4_IMPORT_MODE = "local-pythonpath"
return
# Load the package under its real namespace so the vendored source's
# absolute ``flash_attn.cute.*`` imports resolve back to the same tree.
try:
import flash_attn
except ModuleNotFoundError as exc:
if exc.name != "flash_attn":
raise
# The fork is a standalone ``flash_attn.cute`` distribution. Create
# its namespace parent when FA2's top-level package is not installed.
flash_attn = types.ModuleType("flash_attn")
flash_attn.__path__ = []
sys.modules["flash_attn"] = flash_attn
for module_name in tuple(sys.modules):
if module_name == "flash_attn.cute" or module_name.startswith("flash_attn.cute."):
del sys.modules[module_name]
spec = importlib.util.spec_from_file_location(
"flash_attn.cute",
source / "__init__.py",
submodule_search_locations=[str(source)],
)
if spec is None or spec.loader is None:
raise ImportError(f"could not create an import spec for local FA4 source {source}")
package = importlib.util.module_from_spec(spec)
sys.modules["flash_attn.cute"] = package
setattr(flash_attn, "cute", package)
spec.loader.exec_module(package)
_FA4_IMPORT_MODE = "flat-pythonpath"
def _load_fa4_symbols():
_mount_flat_fa4_from_pythonpath()
block_sparsity = importlib.import_module("flash_attn.cute.block_sparsity")
interface = importlib.import_module("flash_attn.cute.interface")
return block_sparsity.BlockSparseTensorsTorch, interface._flash_attn_fwd, interface
def _fa4_physical_tiles(q_block: int, kv_block: int) -> tuple[int, int]:
"""Map logical sparse blocks to FA4's physical MMA tiles.
KV256 is the historical logical VSA block and expands to KV128. Smaller
requested KV blocks remain physical. Q256 is two Q128 stages; Q128 and
Q64 select the single-stage paths being optimized.
"""
tile_m = min(q_block, 128)
tile_n = min(kv_block, 128)
if q_block % tile_m or kv_block % tile_n:
raise UnsupportedConfig(
f"logical QxKV={q_block}x{kv_block} is not divisible by physical "
f"FA4 tile {tile_m}x{tile_n}"
)
return tile_m, tile_n
def arm_fa4(q, k, v, logical_mask, q_block, kv_block):
"""Direct FA4 block-sparse forward; mask preparation is not timed."""
BlockSparseTensorsTorch, _flash_attn_fwd, interface = _load_fa4_symbols()
tile_m, tile_n = _fa4_physical_tiles(q_block, kv_block)
kv_expansion = kv_block // tile_n
physical_mask = logical_mask.repeat_interleave(kv_expansion, dim=-1).contiguous()
heads, nq, nkv_physical = physical_mask.shape
selected_per_row = int(physical_mask[0, 0].sum().item())
if selected_per_row <= 0 or not bool((physical_mask.sum(dim=-1) == selected_per_row).all()):
raise ValueError("direct FA4 arm requires the same positive selected-block count per row")
physical_indices = torch.arange(nkv_physical, dtype=torch.int32, device=q.device)
physical_indices = physical_indices.view(1, 1, -1).expand_as(physical_mask)
full_idx = physical_indices.masked_select(physical_mask).view(1, heads, nq, selected_per_row)
full_cnt = torch.full((1, heads, nq), selected_per_row, dtype=torch.int32, device=q.device)
# Every selected logical block is fully valid in this fixed-length benchmark.
# Keep a compact empty partial-block representation to exercise FA4's full path.
mask_cnt = torch.zeros((1, heads, nq), dtype=torch.int32, device=q.device)
mask_idx = torch.zeros((1, heads, nq, 1), dtype=torch.int32, device=q.device)
sparse_tensors = BlockSparseTensorsTorch(
mask_block_cnt=mask_cnt,
mask_block_idx=mask_idx,
full_block_cnt=full_cnt,
full_block_idx=full_idx.contiguous(),
block_size=(q_block, tile_n),
)
batch, seq_len, _, _ = q.shape
out_buffer = torch.empty_like(q)
lse_buffer = torch.empty((batch, heads, seq_len), dtype=torch.float32, device=q.device)
scale = 1.0 / math.sqrt(q.shape[-1])
def fn():
return _flash_attn_fwd(
q,
k,
v,
out=out_buffer,
lse=lse_buffer,
softmax_scale=scale,
tile_mn=(tile_m, tile_n),
num_splits=1,
pack_gqa=False,
block_sparse_tensors=sparse_tensors,
causal=False,
return_lse=True,
)[:2]
out, _lse = fn()
return fn, out[0], {
"fa4_interface": str(Path(interface.__file__).resolve()),
"fa4_import_mode": _FA4_IMPORT_MODE,
"fa4_tile_m": tile_m,
"fa4_tile_n": tile_n,
"fa4_sparse_q_block": q_block,
"fa4_sparse_kv_block": tile_n,
"fa4_kv_expansion": kv_expansion,
"keep_alive": (sparse_tensors, out_buffer, lse_buffer),
}
ARMS = {
"flashinfer": arm_flashinfer,
"cutedsl256": arm_cutedsl,
"fa4_wrapper": arm_fa4_wrapper,
"fa4": arm_fa4,
}
def detect_arms(requested):
if requested != ["auto"]:
return requested
arms = []
try:
from flashinfer.sparse import BlockSparseAttentionWrapper # noqa: F401
if "block_mask" not in inspect.signature(BlockSparseAttentionWrapper.plan).parameters:
raise RuntimeError(
"per-head block_mask planning API unavailable; issue #4554 used "
"flashinfer 0.6.16.post2"
)
arms.append("flashinfer")
except Exception as exc: # noqa: BLE001
print(f"# arm flashinfer unavailable: {exc!r}", flush=True)
fa4_available = False
try:
_load_fa4_symbols()
fa4_available = True
except Exception as exc: # noqa: BLE001
print(f"# arm fa4 unavailable: {exc!r}", flush=True)
if fa4_available:
try:
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd # noqa: F401
arms.append("cutedsl256")
except Exception as exc: # noqa: BLE001
print(f"# arm cutedsl256 unavailable (optional): {exc!r}", flush=True)
try:
from fastvideo_kernel.block_sparse_attn_cute_fwd import ( # noqa: F401
block_sparse_attn_cute_fwd_bshd,
)
arms.append("fa4_wrapper")
except Exception as exc: # noqa: BLE001
print(f"# arm fa4_wrapper unavailable (optional): {exc!r}", flush=True)
arms.append("fa4")
if not arms:
raise RuntimeError(
"no benchmark arm importable — need flashinfer-python, fastvideo_kernel, and/or FA4"
)
return arms
def print_table(rows, arms):
header = "| seq | sparsity | Q blk | KV blk | keep/NKV |"
separator = "|---|---|---|---|---|"
for arm in arms:
header += f" {arm} ms | {arm} TFLOP/s | {arm} MFU % |"
separator += "---|---|---|"
if arm == "flashinfer":
header += " plan ms |"
separator += "---|"
if arm == "fa4":
header += " vs raw 256 % | vs cutedsl256 % |"
separator += "---|---|"
if arm == "fa4_wrapper":
header += " vs cutedsl256 % |"
separator += "---|"
print("\n" + header + "\n" + separator)
for row in rows:
line = (
f"| {row['seq_len']} | {row['sparsity']} | {row['q_block']} | {row['kv_block']} | "
f"{row['keep_kv_blocks']}/{row['nkv_blocks']} |"
)
for arm in arms:
cell = row.get(arm, {})
if cell.get("status") == "ok":
line += f" {cell['fwd_ms']:.3f} | {cell['tflops']:.1f} | {cell['mfu_pct']:.1f} |"
if arm == "flashinfer":
line += f" {cell.get('plan_ms', float('nan')):.2f} |"
if arm == "fa4":
raw_relative = cell.get("vs_fa4_256_pct", float("nan"))
wrapper_relative = cell.get("vs_cutedsl256_pct", float("nan"))
line += f" {raw_relative:.1f} | {wrapper_relative:.1f} |"
if arm == "fa4_wrapper":
wrapper_relative = cell.get("vs_cutedsl256_pct", float("nan"))
line += f" {wrapper_relative:.1f} |"
else:
status = cell.get("status", "FAILED")
line += f" {status} | — | — |"
if arm == "flashinfer":
line += " — |"
if arm == "fa4":
line += " — | — |"
if arm == "fa4_wrapper":
line += " — |"
print(line)
print()
def add_relative_efficiency(cell_rows):
"""Attach denominator-free efficiency relative to the 256x256 baselines."""
baseline = next(
(row for row in cell_rows if (row["q_block"], row["kv_block"]) == (256, 256)),
None,
)
if baseline is None:
return
raw_tflops = baseline.get("fa4", {}).get("tflops")
wrapper_tflops = baseline.get("cutedsl256", {}).get("tflops")
for row in cell_rows:
fa4_result = row.get("fa4", {})
if fa4_result.get("status") == "ok":
if raw_tflops:
fa4_result["vs_fa4_256_pct"] = round(
100.0 * fa4_result["tflops"] / raw_tflops,
2,
)
if wrapper_tflops:
fa4_result["vs_cutedsl256_pct"] = round(
100.0 * fa4_result["tflops"] / wrapper_tflops,
2,
)
fine_wrapper_result = row.get("fa4_wrapper", {})
if fine_wrapper_result.get("status") == "ok" and wrapper_tflops:
fine_wrapper_result["vs_cutedsl256_pct"] = round(
100.0 * fine_wrapper_result["tflops"] / wrapper_tflops,
2,
)
def main():
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
rank = int(os.environ.get("RANK", "0"))
world_size = int(os.environ.get("WORLD_SIZE", "1"))
if world_size > 1:
if local_rank >= torch.cuda.device_count():
raise RuntimeError(
f"LOCAL_RANK={local_rank} exceeds visible CUDA device count={torch.cuda.device_count()}"
)
torch.cuda.set_device(local_rank)
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--seq_lens", type=int, nargs="+", default=[4096, 8192, 16384, 32768, 49152, 65536])
parser.add_argument("--sparsities", type=str, nargs="+", default=list(KEEP_FRAC), choices=list(KEEP_FRAC))
parser.add_argument("--arms", type=str, nargs="+", default=["auto"], choices=["auto", *ARMS])
parser.add_argument(
"--block_shapes",
type=parse_block_shape,
nargs="+",
default=list(DEFAULT_BLOCK_SHAPES),
metavar="QxKV",
help="independent logical QxKV sparse block pairs (default: 256x256 128x64 64x64)",
)
parser.add_argument(
"--mask_mode",
choices=("exact256", "native"),
default="exact256",
help=(
"exact256 expands one shared Q256/KV256 mask into every shape; "
"native samples each shape independently (default: exact256)"
),
)
parser.add_argument(
"--peak_bf16_tflops",
type=float,
default=DEFAULT_GB200_BF16_DENSE_TFLOPS,
help=(
"per-GPU dense BF16 tensor-core peak used only for MFU %% "
f"(default: official GB200 {DEFAULT_GB200_BF16_DENSE_TFLOPS:.0f} TFLOP/s)"
),
)
parser.add_argument(
"--quick",
action="store_true",
help="sanity grid: {8192,32768} x {dense,90} x requested block shapes",
)
parser.add_argument("--out", default="bsa_bench_results.json")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--warmup", type=int, default=5)
parser.add_argument("--rep", type=int, default=20)
parser.add_argument("--num_heads", type=int, default=12)
parser.add_argument("--head_dim", type=int, default=128)
args = parser.parse_args()
if args.quick:
args.seq_lens, args.sparsities = [8192, 32768], ["dense", "90"]
if "auto" in args.arms and args.arms != ["auto"]:
parser.error("--arms auto cannot be combined with explicit arms")
if args.peak_bf16_tflops <= 0:
parser.error("--peak_bf16_tflops must be positive")
args.block_shapes = list(dict.fromkeys(args.block_shapes))
for seq_len in args.seq_lens:
if args.mask_mode == "exact256" and seq_len % 256:
parser.error(f"exact256 mask mode requires sequence length divisible by 256; got {seq_len}")
for q_block, kv_block in args.block_shapes:
if seq_len % q_block or seq_len % kv_block:
parser.error(
f"sequence length {seq_len} must be divisible by QxKV block shape "
f"{q_block}x{kv_block}"
)
if args.mask_mode == "exact256" and (256 % q_block or 256 % kv_block):
parser.error(
"exact256 mask mode requires both block sizes to divide 256; "
f"got {q_block}x{kv_block}"
)
from triton.testing import do_bench
# The fp32 reference must be true fp32, not TF32.
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
arms = detect_arms(args.arms)
batch, heads, head_dim = 1, args.num_heads, args.head_dim
scale = 1.0 / math.sqrt(head_dim)
meta = {
"device": torch.cuda.get_device_name(),
"cuda_device_index": torch.cuda.current_device(),
"rank": rank,
"local_rank": local_rank,
"world_size": world_size,
"capability": list(torch.cuda.get_device_capability()),
"torch": torch.__version__,
"cuda": torch.version.cuda,
"seed": args.seed,
"warmup": args.warmup,
"rep": args.rep,
"batch": batch,
"heads": heads,
"head_dim": head_dim,
"dtype": "bfloat16",
"arms": arms,
"block_shapes": [list(shape) for shape in args.block_shapes],
"mask_mode": args.mask_mode,
"mask_base_block": 256 if args.mask_mode == "exact256" else None,
"peak_bf16_dense_tflops": args.peak_bf16_tflops,
"peak_note": (
"official single-GB200 dense BF16 tensor-core peak; override for clock-derived diagnostics"
),
"peak_source": GB200_DATASHEET_URL,
"flop_formula": "4*bs*d*selected_logical_edges*q_block*kv_block",
"mfu_formula": "100*sparse_aware_algorithmic_tflops/peak_bf16_dense_tflops",
"correctness_gate": "finite and max_abs_vs_fp32 <= max(1e-2, 8*bf16_rounding_floor)",
"timing_scope": {
"flashinfer": "run only; plan reported separately",
"cutedsl256": "full FastVideo wrapper including mask expansion and map-to-index",
"fa4_wrapper": (
"FastVideo direct fine-grained BSHD wrapper including map-to-index and VBS mask "
"plumbing; 64x64 is native Q64/KV64, not the public VSA-256 coalescing adapter"
),
"fa4": "raw _flash_attn_fwd dispatch; mask/index preparation and output allocation excluded",
},
"provenance": {
"issue": ISSUE_URL,
"gist": GIST_URL,
"gist_revision": GIST_REVISION,
},
"fa4_env": {
name: os.environ.get(name)
for name in (
"FASTVIDEO_FA4_VSA_DUAL_STREAM",
"FASTVIDEO_FA4_VSA_SP_DOUBLE_BUFFER",
"FASTVIDEO_VSA_FA4_BLOCK_SHAPE",
)
},
}
try:
meta["driver"] = subprocess.check_output(
["nvidia-smi", "--query-gpu=driver_version", "--format=csv,noheader"], text=True
).split()[0]
except Exception: # noqa: BLE001
meta["driver"] = None
for module in ("flashinfer", "triton", "fastvideo_kernel"):
try:
meta[module] = __import__(module).__version__
except Exception: # noqa: BLE001
pass
if any(arm in arms for arm in ("cutedsl256", "fa4_wrapper", "fa4")):
try:
_, _, fa4_interface = _load_fa4_symbols()
meta["fa4_interface"] = str(Path(fa4_interface.__file__).resolve())
meta["fa4_import_mode"] = _FA4_IMPORT_MODE
except Exception as exc: # noqa: BLE001
meta["fa4_import_error"] = f"{type(exc).__name__}: {exc}"
print("META " + json.dumps(meta), flush=True)
rows = []
for seq_len in args.seq_lens:
# Same q,k,v for every sparsity cell of this seq len (isolates sparsity).
set_seed(args.seed)
q = torch.randn(batch, seq_len, heads, head_dim, dtype=torch.bfloat16, device="cuda")
k = torch.randn(batch, seq_len, heads, head_dim, dtype=torch.bfloat16, device="cuda")
v = torch.randn(batch, seq_len, heads, head_dim, dtype=torch.bfloat16, device="cuda")
for label in args.sparsities:
fraction = KEEP_FRAC[label]
base_mask_seed = args.seed + seq_len + int(fraction * 1000)
if args.mask_mode == "exact256":
base_nblocks_256 = seq_len // 256
base_keep_256 = (
base_nblocks_256
if label == "dense"
else max(1, round(fraction * base_nblocks_256))
)
base_mask_256 = make_logical_mask(
heads,
base_nblocks_256,
base_nblocks_256,
base_keep_256,
base_mask_seed,
)
else:
base_nblocks_256 = None
base_keep_256 = None
base_mask_256 = None
exact_ref = (
ref_masked_sdpa_fp32(
q,
k,
v,
base_mask_256,
256,
256,
scale,
)
if base_mask_256 is not None
else None
)
cell_rows = []
fa4_shape_outputs = {}
for q_block, kv_block in args.block_shapes:
nq_blocks, nkv_blocks = seq_len // q_block, seq_len // kv_block
if args.mask_mode == "exact256":
if base_mask_256 is None: # pragma: no cover - guarded above
raise AssertionError("exact256 mode requires a base mask")
keep = base_keep_256 * (256 // kv_block)
mask_seed = base_mask_seed
logical_mask = expand_exact256_mask(base_mask_256, q_block, kv_block)
else:
keep = nkv_blocks if label == "dense" else max(1, round(fraction * nkv_blocks))
mask_seed = (
base_mask_seed
+ q_block * 1_000_003
+ kv_block * 9_176
)
logical_mask = make_logical_mask(
heads,
nq_blocks,
nkv_blocks,
keep,
mask_seed,
)
selected_edges = heads * nq_blocks * keep
flops = flops_sparse_attention(
batch,
head_dim,
selected_edges,
q_block,
kv_block,
)
row = {
"seq_len": seq_len,
"sparsity": label,
"mask_mode": args.mask_mode,
"base_nblocks_256": base_nblocks_256,
"base_keep_256": base_keep_256,
"base_mask_seed": base_mask_seed,
"q_block": q_block,
"kv_block": kv_block,
"nq_blocks": nq_blocks,
"nkv_blocks": nkv_blocks,
"keep_kv_blocks": keep,
"selected_logical_edges": selected_edges,
"actual_sparsity": round(1.0 - keep / nkv_blocks, 4),
"mask_seed": mask_seed,
"flops": flops,
}
ref = (
exact_ref
if exact_ref is not None
else ref_masked_sdpa_fp32(
q,
k,
v,
logical_mask,
q_block,
kv_block,
scale,
)
)
# Irreducible bf16 output-rounding floor: even a bit-exact kernel
# emitting bf16 cannot beat this vs the fp32 reference.
row["bf16_floor_max_abs"] = float(
(ref - ref.to(torch.bfloat16).float()).abs().max()
)
row["correctness_max_abs_limit"] = max(
1e-2,
8.0 * row["bf16_floor_max_abs"],
)
outputs = {}
for name in arms:
result = {}
try:
fn, out, extra = ARMS[name](
q,
k,
v,
logical_mask,
q_block,
kv_block,
)
torch.cuda.synchronize()
result["max_abs_vs_fp32"] = float((out.float() - ref).abs().max())
if not bool(torch.isfinite(out).all()):
raise RuntimeError(f"{name} produced NaN or Inf")
if result["max_abs_vs_fp32"] > row["correctness_max_abs_limit"]:
raise RuntimeError(
f"{name} max_abs_vs_fp32={result['max_abs_vs_fp32']:.6g} "
f"exceeds limit={row['correctness_max_abs_limit']:.6g}"
)
outputs[name] = out
elapsed_ms = float(
do_bench(fn, warmup=args.warmup, rep=args.rep, quantiles=None)
)
achieved_tflops = flops / elapsed_ms * 1e-9
result["fwd_ms"] = round(elapsed_ms, 4)
result["tflops"] = round(achieved_tflops, 2)
result["mfu_pct"] = round(
100.0 * achieved_tflops / args.peak_bf16_tflops,
2,
)
for key, value in extra.items():
if key != "keep_alive":
result[key] = value
result["status"] = "ok"
del fn, extra
except UnsupportedConfig as exc:
result["status"] = "SKIPPED"
result["error"] = str(exc)
except Exception as exc: # noqa: BLE001
result["status"] = "FAILED"
result["error"] = f"{type(exc).__name__}: {exc}"[:500]
traceback.print_exc()
row[name] = result
torch.cuda.empty_cache()
output_names = list(outputs)
if len(output_names) > 1:
row["cross_max_abs"] = {}
for lhs_idx, lhs in enumerate(output_names):
for rhs in output_names[lhs_idx + 1:]:
row["cross_max_abs"][f"{lhs}_vs_{rhs}"] = float(
(outputs[lhs].float() - outputs[rhs].float()).abs().max()
)
if args.mask_mode == "exact256" and "fa4" in outputs:
fa4_shape_outputs[(q_block, kv_block)] = outputs["fa4"]
outputs.clear()
del ref, logical_mask
torch.cuda.empty_cache()
cell_rows.append(row)
if args.mask_mode == "exact256":
baseline_output = fa4_shape_outputs.get((256, 256))
if baseline_output is not None:
for row in cell_rows:
shape = (row["q_block"], row["kv_block"])
target_output = fa4_shape_outputs.get(shape)
if target_output is not None:
row["fa4"]["max_abs_vs_fa4_256"] = float(
(target_output.float() - baseline_output.float()).abs().max()
)
fa4_shape_outputs.clear()
add_relative_efficiency(cell_rows)
for row in cell_rows:
print("ROW " + json.dumps(row), flush=True)
rows.extend(cell_rows)
del base_mask_256, exact_ref
del q, k, v
torch.cuda.empty_cache()
output_path = Path(args.out)
if world_size > 1:
output_path = output_path.with_name(
f"{output_path.stem}.rank{rank}{output_path.suffix}"
)
with output_path.open("w") as output_file:
json.dump({"meta": meta, "rows": rows}, output_file, indent=2)
print_table(rows, arms)
print(f"wrote {output_path}")
return 0
if __name__ == "__main__":
if not torch.cuda.is_available():
raise RuntimeError("CUDA required")
sys.exit(main())
+4
View File
@@ -0,0 +1,4 @@
[flake8]
max-line-length = 100
# W503: line break before binary operator
ignore = E731, E741, F841, W503
+8
View File
@@ -0,0 +1,8 @@
Tri Dao
Jay Shah
Ted Zadouri
Markus Hoehnerbach
Vijay Thakkar
Timmy Liu
Driss Guessous
Reuben Stern
+29
View File
@@ -0,0 +1,29 @@
BSD 3-Clause License
Copyright (c) 2022, the respective contributors, as shown by the AUTHORS file.
All rights reserved.
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met:
* Redistributions of source code must retain the above copyright notice, this
list of conditions and the following disclaimer.
* Redistributions in binary form must reproduce the above copyright notice,
this list of conditions and the following disclaimer in the documentation
and/or other materials provided with the distribution.
* Neither the name of the copyright holder nor the names of its
contributors may be used to endorse or promote products derived from
this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+6
View File
@@ -0,0 +1,6 @@
include UPSTREAM.md
global-exclude *.egg-info/*
prune flash_attn_4.egg-info
prune flash_attn.egg-info
prune build
prune dist
+33
View File
@@ -0,0 +1,33 @@
# FlashAttention-4 (CuTeDSL)
FlashAttention-4 is a CuTeDSL-based implementation of FlashAttention for Hopper and Blackwell GPUs.
## Installation
```sh
pip install flash-attn-4
```
If you're on CUDA 13, install with the `cu13` extra for best performance:
```sh
pip install "flash-attn-4[cu13]"
```
## Usage
```python
from flash_attn.cute import flash_attn_func, flash_attn_varlen_func
out = flash_attn_func(q, k, v, causal=True)
```
## Development
```sh
git clone https://github.com/Dao-AILab/flash-attention.git
cd flash-attention
pip install -e "flash_attn/cute[dev]" # CUDA 12.x
pip install -e "flash_attn/cute[dev,cu13]" # CUDA 13.x (e.g. B200)
pytest tests/cute/
```
+37
View File
@@ -0,0 +1,37 @@
# FastVideo FA4 source fork
This directory vendors the `flash_attn.cute` Python package from
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention) at
commit [`82d6441eec5d4dfec120153db2c0145ae855a083`](https://github.com/Dao-AILab/flash-attention/commit/82d6441eec5d4dfec120153db2c0145ae855a083).
The upstream source path is `flash_attn/cute/`, and its BSD-3-Clause license is
preserved in `LICENSE`.
FastVideo keeps this as a first-class source fork because its video sparse
attention kernels need FA4 scheduling and tile-shape changes that are developed
and benchmarked together with `fastvideo-kernel`. The distribution intentionally
retains the upstream name `flash-attn-4` and import namespace `flash_attn.cute`.
The forward-kernel delta currently contains FastVideo's single-Q-block,
two-KV-stream Q128/KV64 schedule, four-slot score/probability double buffering,
and FP32 online-softmax merge. The Q64/KV64 primitives were adapted from
upstream's guarded smaller-tile commit
[`526c18d25bcbc7fc7d6740ab3c7c84ed2d42cb0b`](https://github.com/Dao-AILab/flash-attention/commit/526c18d25bcbc7fc7d6740ab3c7c84ed2d42cb0b),
then extended with the same two-stream sparse traversal and merge. These are
forward-only specializations; the existing upstream paths remain the fallback
outside their explicit SM100/SM110 dtype, shape, and metadata gates.
Install the working tree without resolving dependencies:
```bash
uv pip install --no-deps --editable fastvideo-kernel/fa4
```
The initial performance comparison came from FlashInfer issue
[#4554](https://github.com/flashinfer-ai/flashinfer/issues/4554) and its
[benchmark harness](https://gist.github.com/SolitaryThinker/90a1d1447929fc38dc509c1852e76532)
at gist revision `e15ac9066f23ef3690e33e1cc1fdac45b4b9099f`.
When refreshing from upstream, copy only `flash_attn/cute/` from an explicitly
checked-out commit, retain this file and FastVideo-specific changes, update the
commit and package version above, and rerun the FA4 correctness and GB200
performance gates before accepting the refresh.
+18
View File
@@ -0,0 +1,18 @@
"""Flash Attention CUTE (CUDA Template Engine) implementation."""
from importlib.metadata import PackageNotFoundError, version
try:
__version__ = version("flash-attn-4")
except PackageNotFoundError:
__version__ = "0.0.0"
from .interface import (
flash_attn_func,
flash_attn_varlen_func,
)
__all__ = [
"flash_attn_func",
"flash_attn_varlen_func",
]
+103
View File
@@ -0,0 +1,103 @@
# Copyright (c) 2025, Tri Dao.
from typing import Type, Callable, Optional
import cutlass
import cutlass.cute as cute
def get_smem_layout_atom(dtype: Type[cutlass.Numeric], k_dim: int) -> cute.ComposedLayout:
dtype_byte = cutlass.const_expr(dtype.width // 8)
bytes_per_row = cutlass.const_expr(k_dim * dtype_byte)
smem_k_block_size = (
cutlass.const_expr(
128
if bytes_per_row % 128 == 0
else (64 if bytes_per_row % 64 == 0 else (32 if bytes_per_row % 32 == 0 else 16))
)
// dtype_byte
)
swizzle_bits = (
4
if smem_k_block_size == 128
else (3 if smem_k_block_size == 64 else (2 if smem_k_block_size == 32 else 1))
)
swizzle_base = 2 if dtype_byte == 4 else (3 if dtype_byte == 2 else 4)
return cute.make_composed_layout(
cute.make_swizzle(swizzle_bits, swizzle_base, swizzle_base),
0,
cute.make_ordered_layout(
(8 if cutlass.const_expr(k_dim % 32 == 0) else 16, smem_k_block_size), order=(1, 0)
),
)
@cute.jit
def gemm(
tiled_mma: cute.TiledMma,
acc: cute.Tensor,
tCrA: cute.Tensor,
tCrB: cute.Tensor,
tCsA: cute.Tensor,
tCsB: cute.Tensor,
smem_thr_copy_A: cute.TiledCopy,
smem_thr_copy_B: cute.TiledCopy,
hook_fn: Optional[Callable] = None,
A_in_regs: cutlass.Constexpr[bool] = False,
B_in_regs: cutlass.Constexpr[bool] = False,
swap_AB: cutlass.Constexpr[bool] = False,
) -> None:
if cutlass.const_expr(swap_AB):
gemm(
tiled_mma,
acc,
tCrB,
tCrA,
tCsB,
tCsA,
smem_thr_copy_B,
smem_thr_copy_A,
hook_fn,
A_in_regs=B_in_regs,
B_in_regs=A_in_regs,
swap_AB=False,
)
else:
tCrA_copy_view = smem_thr_copy_A.retile(tCrA)
tCrB_copy_view = smem_thr_copy_B.retile(tCrB)
if cutlass.const_expr(not A_in_regs):
cute.copy(smem_thr_copy_A, tCsA[None, None, 0], tCrA_copy_view[None, None, 0])
if cutlass.const_expr(not B_in_regs):
cute.copy(smem_thr_copy_B, tCsB[None, None, 0], tCrB_copy_view[None, None, 0])
for k in cutlass.range_constexpr(cute.size(tCsA.shape[2])):
if k < cute.size(tCsA.shape[2]) - 1:
if cutlass.const_expr(not A_in_regs):
cute.copy(
smem_thr_copy_A, tCsA[None, None, k + 1], tCrA_copy_view[None, None, k + 1]
)
if cutlass.const_expr(not B_in_regs):
cute.copy(
smem_thr_copy_B, tCsB[None, None, k + 1], tCrB_copy_view[None, None, k + 1]
)
cute.gemm(tiled_mma, acc, tCrA[None, None, k], tCrB[None, None, k], acc)
if cutlass.const_expr(k == 0 and hook_fn is not None):
hook_fn()
@cute.jit
def gemm_rs(
tiled_mma: cute.TiledMma,
acc: cute.Tensor,
tCrA: cute.Tensor,
tCrB: cute.Tensor,
tCsB: cute.Tensor,
smem_thr_copy_B: cute.TiledCopy,
hook_fn: Optional[Callable] = None,
) -> None:
tCrB_copy_view = smem_thr_copy_B.retile(tCrB)
cute.copy(smem_thr_copy_B, tCsB[None, None, 0], tCrB_copy_view[None, None, 0])
for k in cutlass.range_constexpr(cute.size(tCrA.shape[2])):
if cutlass.const_expr(k < cute.size(tCrA.shape[2]) - 1):
cute.copy(smem_thr_copy_B, tCsB[None, None, k + 1], tCrB_copy_view[None, None, k + 1])
cute.gemm(tiled_mma, acc, tCrA[None, None, k], tCrB[None, None, k], acc)
if cutlass.const_expr(k == 0 and hook_fn is not None):
hook_fn()
+71
View File
@@ -0,0 +1,71 @@
import cutlass
import cutlass.cute as cute
from cutlass import Int32
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import llvm
@dsl_user_op
def ld_acquire(lock_ptr: cute.Pointer, *, loc=None, ip=None) -> cutlass.Int32:
lock_ptr_i64 = lock_ptr.toint(loc=loc, ip=ip).ir_value()
state = llvm.inline_asm(
T.i32(),
[lock_ptr_i64],
"ld.global.acquire.gpu.b32 $0, [$1];",
"=r,l",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
return cutlass.Int32(state)
@dsl_user_op
def red_relaxed(
lock_ptr: cute.Pointer, val: cutlass.Constexpr[Int32], *, loc=None, ip=None
) -> None:
lock_ptr_i64 = lock_ptr.toint(loc=loc, ip=ip).ir_value()
llvm.inline_asm(
None,
[lock_ptr_i64, Int32(val).ir_value(loc=loc, ip=ip)],
"red.relaxed.gpu.global.add.s32 [$0], $1;",
"l,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
@dsl_user_op
def red_release(
lock_ptr: cute.Pointer, val: cutlass.Constexpr[Int32], *, loc=None, ip=None
) -> None:
lock_ptr_i64 = lock_ptr.toint(loc=loc, ip=ip).ir_value()
llvm.inline_asm(
None,
[lock_ptr_i64, Int32(val).ir_value(loc=loc, ip=ip)],
"red.release.gpu.global.add.s32 [$0], $1;",
"l,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
@cute.jit
def wait_eq(lock_ptr: cute.Pointer, thread_idx: int | Int32, flag_offset: int, val: Int32) -> None:
flag_ptr = lock_ptr + flag_offset
if thread_idx == 0:
read_val = Int32(0)
while read_val != val:
read_val = ld_acquire(flag_ptr)
@cute.jit
def arrive_inc(
lock_ptr: cute.Pointer, thread_idx: int | Int32, flag_offset: int, val: cutlass.Constexpr[Int32]
) -> None:
flag_ptr = lock_ptr + flag_offset
if thread_idx == 0:
red_release(flag_ptr, val)
# red_relaxed(flag_ptr, val)
+243
View File
@@ -0,0 +1,243 @@
"""Shared benchmark utilities: attention_ref, cuDNN helpers, flops calculation."""
import math
import torch
try:
import cudnn
except ImportError:
cudnn = None
# ── FLOPS calculation ────────────────────────────────────────────────────────
def flops(
batch,
nheads,
seqlen_q,
seqlen_k,
headdim,
headdim_v,
causal=False,
window_size=(None, None),
has_qv=False,
):
if causal:
avg_seqlen = (max(0, seqlen_k - seqlen_q) + seqlen_k) / 2
else:
if window_size == (None, None):
avg_seqlen = seqlen_k
else:
row_idx = torch.arange(seqlen_q, device="cuda")
col_left = (
torch.maximum(row_idx + seqlen_k - seqlen_q - window_size[0], torch.tensor(0))
if window_size[0] is not None
else torch.zeros_like(row_idx)
)
col_right = (
torch.minimum(
row_idx + seqlen_k - seqlen_q + window_size[1], torch.tensor(seqlen_k - 1)
)
if window_size[1] is not None
else torch.full_like(row_idx, seqlen_k - 1)
)
avg_seqlen = (col_right - col_left + 1).float().mean().item()
eff_headdim = headdim + headdim_v if has_qv else headdim
return batch * nheads * 2 * seqlen_q * avg_seqlen * (eff_headdim + headdim_v)
# ── Bandwidth calculation ────────────────────────────────────────────────────
def bandwidth_fwd_bytes(
batch,
nheads,
nheads_kv,
seqlen_q,
seqlen_k,
headdim,
headdim_v,
dtype_bytes=2,
has_qv=False,
shared_kv=False,
):
"""HBM traffic for one attention pass: read Q,K,V + write O."""
q = batch * nheads * seqlen_q * headdim
qv = batch * nheads * seqlen_q * headdim_v if has_qv else 0
k = batch * nheads_kv * seqlen_k * headdim
v = batch * nheads_kv * seqlen_k * headdim_v if not shared_kv else 0
o = batch * nheads * seqlen_q * headdim_v
return (q + qv + k + v + o) * dtype_bytes
def bandwidth_bwd_bytes(
batch, nheads, nheads_kv, seqlen_q, seqlen_k, headdim, headdim_v, dtype_bytes=2
):
"""HBM traffic for one attention pass: read Q,K,V,dO + write dQ,dK,dV."""
q = batch * nheads * seqlen_q * headdim
k = batch * nheads_kv * seqlen_k * headdim
v = batch * nheads_kv * seqlen_k * headdim_v
do = batch * nheads * seqlen_q * headdim_v
dq = q
dk = k
dv = v
return (q + k + v + do + dq + dk + dv) * dtype_bytes
# ── Reference attention ─────────────────────────────────────────────────────
_attention_ref_mask_cache = {}
def attention_ref(q, k, v, causal=False):
"""Standard attention reference implementation.
Args:
q, k, v: (batch, seqlen, nheads, headdim) tensors.
causal: whether to apply causal mask.
"""
softmax_scale = 1.0 / math.sqrt(q.shape[-1])
scores = torch.einsum("bthd,bshd->bhts", q * softmax_scale, k)
if causal:
if scores.shape[-2] not in _attention_ref_mask_cache:
mask = torch.tril(
torch.ones(scores.shape[-2:], device=scores.device, dtype=torch.bool), diagonal=0
)
_attention_ref_mask_cache[scores.shape[-2]] = mask
else:
mask = _attention_ref_mask_cache[scores.shape[-2]]
scores = scores.masked_fill(mask, float("-inf"))
attn = torch.softmax(scores, dim=-1)
return torch.einsum("bhts,bshd->bthd", attn, v)
# ── cuDNN graph helpers ─────────────────────────────────────────────────────
_TORCH_TO_CUDNN_DTYPE = {
torch.float16: "HALF",
torch.bfloat16: "BFLOAT16",
torch.float32: "FLOAT",
torch.int32: "INT32",
torch.int64: "INT64",
}
def _build_cudnn_graph(io_dtype, tensors, build_fn):
"""Build a cuDNN graph. Returns (graph, variant_pack, workspace)."""
assert cudnn is not None, "cuDNN is not available"
cudnn_dtype = getattr(cudnn.data_type, _TORCH_TO_CUDNN_DTYPE[io_dtype])
graph = cudnn.pygraph(
io_data_type=cudnn_dtype,
intermediate_data_type=cudnn.data_type.FLOAT,
compute_data_type=cudnn.data_type.FLOAT,
)
graph_tensors = {name: graph.tensor_like(t.detach()) for name, t in tensors.items()}
variant_pack = build_fn(graph, graph_tensors)
graph.validate()
graph.build_operation_graph()
graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK])
graph.check_support()
graph.build_plans()
workspace = torch.empty(graph.get_workspace_size(), device="cuda", dtype=torch.uint8)
return graph, variant_pack, workspace
def cudnn_fwd_setup(q, k, v, causal=False, window_size_left=None):
"""Build a cuDNN forward SDPA graph.
Args:
q, k, v: (batch, nheads, seqlen, headdim) tensors (cuDNN layout).
causal: whether to apply causal mask.
window_size_left: sliding window size (None for no window).
Returns:
(fwd_fn, o_gpu, stats_gpu) where fwd_fn is a zero-arg callable.
"""
b, nheads, seqlen_q, headdim = q.shape
headdim_v = v.shape[-1]
o_gpu = torch.empty(b, nheads, seqlen_q, headdim_v, dtype=q.dtype, device=q.device)
stats_gpu = torch.empty(b, nheads, seqlen_q, 1, dtype=torch.float32, device=q.device)
def build(graph, gt):
o, stats = graph.sdpa(
name="sdpa",
q=gt["q"],
k=gt["k"],
v=gt["v"],
is_inference=False,
attn_scale=1.0 / math.sqrt(headdim),
use_causal_mask=causal or window_size_left is not None,
sliding_window_length=window_size_left
if window_size_left is not None and not causal
else None,
)
o.set_output(True).set_dim(o_gpu.shape).set_stride(o_gpu.stride())
stats.set_output(True).set_data_type(cudnn.data_type.FLOAT)
return {gt["q"]: q, gt["k"]: k, gt["v"]: v, o: o_gpu, stats: stats_gpu}
graph, variant_pack, workspace = _build_cudnn_graph(q.dtype, {"q": q, "k": k, "v": v}, build)
def fwd_fn():
graph.execute(variant_pack, workspace)
return o_gpu
return fwd_fn, o_gpu, stats_gpu
def cudnn_bwd_setup(q, k, v, o, g, lse, causal=False, window_size_left=None):
"""Build a cuDNN backward SDPA graph.
Args:
q, k, v, o, g, lse: (batch, nheads, seqlen, dim) tensors (cuDNN layout).
causal: whether to apply causal mask.
window_size_left: sliding window size (None for no window).
Returns:
bwd_fn: zero-arg callable that returns (dq, dk, dv).
"""
headdim = q.shape[-1]
dq_gpu, dk_gpu, dv_gpu = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
def build(graph, gt):
dq, dk, dv = graph.sdpa_backward(
name="sdpa_backward",
q=gt["q"],
k=gt["k"],
v=gt["v"],
o=gt["o"],
dO=gt["g"],
stats=gt["lse"],
attn_scale=1.0 / math.sqrt(headdim),
use_causal_mask=causal or window_size_left is not None,
sliding_window_length=window_size_left
if window_size_left is not None and not causal
else None,
use_deterministic_algorithm=False,
)
dq.set_output(True).set_dim(dq_gpu.shape).set_stride(dq_gpu.stride())
dk.set_output(True).set_dim(dk_gpu.shape).set_stride(dk_gpu.stride())
dv.set_output(True).set_dim(dv_gpu.shape).set_stride(dv_gpu.stride())
return {
gt["q"]: q,
gt["k"]: k,
gt["v"]: v,
gt["o"]: o,
gt["g"]: g,
gt["lse"]: lse,
dq: dq_gpu,
dk: dk_gpu,
dv: dv_gpu,
}
graph, variant_pack, workspace = _build_cudnn_graph(
q.dtype,
{"q": q, "k": k, "v": v, "o": o, "g": g, "lse": lse},
build,
)
def bwd_fn():
graph.execute(variant_pack, workspace)
return dq_gpu, dk_gpu, dv_gpu
return bwd_fn
+268
View File
@@ -0,0 +1,268 @@
# Copyright (c) 2023, Tri Dao.
"""Useful functions for writing test code."""
import torch
import torch.utils.benchmark as benchmark
def benchmark_forward(
fn, *inputs, repeats=10, desc="", verbose=True, amp=False, amp_dtype=torch.float16, **kwinputs
):
"""Use Pytorch Benchmark on the forward pass of an arbitrary function."""
if verbose:
print(desc, "- Forward pass")
def amp_wrapper(*inputs, **kwinputs):
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
fn(*inputs, **kwinputs)
t = benchmark.Timer(
stmt="fn_amp(*inputs, **kwinputs)",
globals={"fn_amp": amp_wrapper, "inputs": inputs, "kwinputs": kwinputs},
num_threads=torch.get_num_threads(),
)
m = t.timeit(repeats)
if verbose:
print(m)
return t, m
def benchmark_backward(
fn,
*inputs,
grad=None,
repeats=10,
desc="",
verbose=True,
amp=False,
amp_dtype=torch.float16,
**kwinputs,
):
"""Use Pytorch Benchmark on the backward pass of an arbitrary function."""
if verbose:
print(desc, "- Backward pass")
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
y = fn(*inputs, **kwinputs)
if type(y) is tuple:
y = y[0]
if grad is None:
grad = torch.randn_like(y)
else:
if grad.shape != y.shape:
raise RuntimeError("Grad shape does not match output shape")
def f(*inputs, y, grad):
# Set .grad to None to avoid extra operation of gradient accumulation
for x in inputs:
if isinstance(x, torch.Tensor):
x.grad = None
y.backward(grad, retain_graph=True)
t = benchmark.Timer(
stmt="f(*inputs, y=y, grad=grad)",
globals={"f": f, "inputs": inputs, "y": y, "grad": grad},
num_threads=torch.get_num_threads(),
)
m = t.timeit(repeats)
if verbose:
print(m)
return t, m
def benchmark_combined(
fn,
*inputs,
grad=None,
repeats=10,
desc="",
verbose=True,
amp=False,
amp_dtype=torch.float16,
**kwinputs,
):
"""Use Pytorch Benchmark on the forward+backward pass of an arbitrary function."""
if verbose:
print(desc, "- Forward + Backward pass")
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
y = fn(*inputs, **kwinputs)
if type(y) is tuple:
y = y[0]
if grad is None:
grad = torch.randn_like(y)
else:
if grad.shape != y.shape:
raise RuntimeError("Grad shape does not match output shape")
def f(grad, *inputs, **kwinputs):
for x in inputs:
if isinstance(x, torch.Tensor):
x.grad = None
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
y = fn(*inputs, **kwinputs)
if type(y) is tuple:
y = y[0]
y.backward(grad, retain_graph=True)
t = benchmark.Timer(
stmt="f(grad, *inputs, **kwinputs)",
globals={"f": f, "fn": fn, "inputs": inputs, "grad": grad, "kwinputs": kwinputs},
num_threads=torch.get_num_threads(),
)
m = t.timeit(repeats)
if verbose:
print(m)
return t, m
def benchmark_fwd_bwd(
fn,
*inputs,
grad=None,
repeats=10,
desc="",
verbose=True,
amp=False,
amp_dtype=torch.float16,
**kwinputs,
):
"""Use Pytorch Benchmark on the forward+backward pass of an arbitrary function."""
return (
benchmark_forward(
fn,
*inputs,
repeats=repeats,
desc=desc,
verbose=verbose,
amp=amp,
amp_dtype=amp_dtype,
**kwinputs,
),
benchmark_backward(
fn,
*inputs,
grad=grad,
repeats=repeats,
desc=desc,
verbose=verbose,
amp=amp,
amp_dtype=amp_dtype,
**kwinputs,
),
)
def benchmark_all(
fn,
*inputs,
grad=None,
repeats=10,
desc="",
verbose=True,
amp=False,
amp_dtype=torch.float16,
**kwinputs,
):
"""Use Pytorch Benchmark on the forward+backward pass of an arbitrary function."""
return (
benchmark_forward(
fn,
*inputs,
repeats=repeats,
desc=desc,
verbose=verbose,
amp=amp,
amp_dtype=amp_dtype,
**kwinputs,
),
benchmark_backward(
fn,
*inputs,
grad=grad,
repeats=repeats,
desc=desc,
verbose=verbose,
amp=amp,
amp_dtype=amp_dtype,
**kwinputs,
),
benchmark_combined(
fn,
*inputs,
grad=grad,
repeats=repeats,
desc=desc,
verbose=verbose,
amp=amp,
amp_dtype=amp_dtype,
**kwinputs,
),
)
def pytorch_profiler(
fn,
*inputs,
trace_filename=None,
backward=False,
amp=False,
amp_dtype=torch.float16,
cpu=False,
verbose=True,
**kwinputs,
):
"""Wrap benchmark functions in Pytorch profiler to see CUDA information."""
if backward:
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
out = fn(*inputs, **kwinputs)
if type(out) is tuple:
out = out[0]
g = torch.randn_like(out)
for _ in range(30): # Warm up
if backward:
for x in inputs:
if isinstance(x, torch.Tensor):
x.grad = None
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
out = fn(*inputs, **kwinputs)
if type(out) is tuple:
out = out[0]
# Backward should be done outside autocast
if backward:
out.backward(g, retain_graph=True)
activities = ([torch.profiler.ProfilerActivity.CPU] if cpu else []) + [
torch.profiler.ProfilerActivity.CUDA
]
with torch.profiler.profile(
activities=activities,
record_shapes=True,
# profile_memory=True,
with_stack=True,
) as prof:
if backward:
for x in inputs:
if isinstance(x, torch.Tensor):
x.grad = None
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
out = fn(*inputs, **kwinputs)
if type(out) is tuple:
out = out[0]
if backward:
out.backward(g, retain_graph=True)
if verbose:
# print(prof.key_averages().table(sort_by="self_cuda_time_total", row_limit=50))
print(prof.key_averages().table(row_limit=50))
if trace_filename is not None:
prof.export_chrome_trace(trace_filename)
def benchmark_memory(fn, *inputs, desc="", verbose=True, **kwinputs):
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
torch.cuda.synchronize()
fn(*inputs, **kwinputs)
torch.cuda.synchronize()
mem = torch.cuda.max_memory_allocated() / ((2**20) * 1000)
if verbose:
print(f"{desc} max memory: {mem}GB")
torch.cuda.empty_cache()
return mem
@@ -0,0 +1,434 @@
# Benchmark FP8 attention for FA4 (CuTe-DSL) on SM100.
#
# Run (recommended):
# python -m flash_attn.cute.benchmark_flash_attention_fp8
#
# Notes:
# - This is intended to be used while bringing up FP8 support for SM100.
# - FP8 correctness depends on descales + max-offset scaling being implemented in the SM100 kernel.
# This script optionally checks output vs a BF16 PyTorch baseline on dequantized FP8 inputs.
#
# Adapted from: `hopper/benchmark_flash_attention_fp8.py`
from __future__ import annotations
import argparse
import inspect
import math
import time
from typing import Iterable
import torch
from einops import rearrange
from flash_attn.cute.benchmark import benchmark_forward
from flash_attn.cute.interface import _flash_attn_fwd as flash_attn_cute_fwd
try:
import cudnn
except ImportError:
cudnn = None
def _torch_float8_dtype(name: str) -> torch.dtype:
if name in ("fp8", "fp8_e4m3", "fp8_e4m3fn"):
return torch.float8_e4m3fn
if name in ("fp8_e5m2", "fp8_e5m2fn"):
return torch.float8_e5m2
raise ValueError(f"Unsupported fp8 dtype name: {name}")
def _parse_int_list(csv: str) -> list[int]:
out: list[int] = []
for part in csv.split(","):
part = part.strip()
if not part:
continue
out.append(int(part))
return out
def attention_pytorch(qkv: torch.Tensor, causal: bool) -> torch.Tensor:
"""
qkv: (batch, seqlen, 3, nheads, headdim)
out: (batch, seqlen, nheads, headdim)
"""
batch_size, seqlen, _, nheads, d = qkv.shape
q, k, v = qkv.unbind(dim=2)
q = rearrange(q, "b t h d -> (b h) t d")
k = rearrange(k, "b s h d -> (b h) d s")
softmax_scale = 1.0 / math.sqrt(d)
scores = torch.empty(batch_size * nheads, seqlen, seqlen, dtype=qkv.dtype, device=qkv.device)
scores = rearrange(
torch.baddbmm(scores, q, k, beta=0, alpha=softmax_scale), "(b h) t s -> b h t s", h=nheads
)
if causal:
causal_mask = torch.triu(torch.full((seqlen, seqlen), -10000.0, device=scores.device), 1)
scores = scores + causal_mask.to(dtype=scores.dtype)
attention = torch.softmax(scores, dim=-1)
output = torch.einsum("bhts,bshd->bthd", attention, v)
return output.to(dtype=qkv.dtype)
def flops(batch: int, seqlen: int, headdim: int, nheads: int, causal: bool) -> int:
# Matches the hopper benchmark’s convention.
return 4 * batch * seqlen**2 * nheads * headdim // (2 if causal else 1)
def efficiency(flop: int, seconds: float) -> float:
return (flop / seconds / 1e12) if not math.isnan(seconds) else 0.0
def time_fwd(fn, *args, repeats: int, **kwargs) -> float:
time.sleep(1) # reduce residual throttling effects between benchmarks
_, m = benchmark_forward(fn, *args, repeats=repeats, verbose=False, **kwargs)
return float(m.mean)
def convert_to_cudnn_type(torch_type):
if torch_type == torch.float16:
return cudnn.data_type.HALF
if torch_type == torch.bfloat16:
return cudnn.data_type.BFLOAT16
if torch_type == torch.float32:
return cudnn.data_type.FLOAT
if torch_type == torch.int32:
return cudnn.data_type.INT32
if torch_type == torch.int64:
return cudnn.data_type.INT64
if torch_type == torch.float8_e4m3fn:
return cudnn.data_type.FP8_E4M3
if torch_type == torch.float8_e5m2:
return cudnn.data_type.FP8_E5M2
raise ValueError("Unsupported tensor data type.")
def cudnn_sdpa_fp8_setup(qkv: torch.Tensor, seqlen_q: int, seqlen_k: int, causal: bool):
"""Minimal cudnn.fp8 sdpa runner (optional)."""
assert cudnn is not None, "cudnn python bindings not available"
b, _, _, nheads, headdim = qkv.shape
o_gpu = torch.zeros(b, seqlen_q, nheads, headdim, dtype=qkv.dtype, device=qkv.device)
o_gpu_transposed = torch.as_strided(
o_gpu,
[b, nheads, seqlen_q, headdim],
[nheads * seqlen_q * headdim, headdim, nheads * headdim, 1],
)
amax_s_gpu = torch.empty(1, 1, 1, 1, dtype=torch.float32, device=qkv.device)
amax_o_gpu = torch.empty(1, 1, 1, 1, dtype=torch.float32, device=qkv.device)
graph = cudnn.pygraph(
io_data_type=convert_to_cudnn_type(qkv.dtype),
intermediate_data_type=cudnn.data_type.FLOAT,
compute_data_type=cudnn.data_type.FLOAT,
)
new_q = torch.as_strided(
qkv,
[b, nheads, seqlen_q, headdim],
[seqlen_q * nheads * headdim * 3, headdim, headdim * nheads * 3, 1],
storage_offset=0,
)
q = graph.tensor(
name="Q",
dim=list(new_q.shape),
stride=list(new_q.stride()),
data_type=convert_to_cudnn_type(qkv.dtype),
)
new_k = torch.as_strided(
qkv,
[b, nheads, seqlen_k, headdim],
[seqlen_k * nheads * headdim * 3, headdim, headdim * nheads * 3, 1],
storage_offset=nheads * headdim,
)
k = graph.tensor(
name="K",
dim=list(new_k.shape),
stride=list(new_k.stride()),
data_type=convert_to_cudnn_type(qkv.dtype),
)
new_v = torch.as_strided(
qkv,
[b, nheads, seqlen_k, headdim],
[seqlen_k * nheads * headdim * 3, headdim, headdim * nheads * 3, 1],
storage_offset=nheads * headdim * 2,
)
v = graph.tensor(
name="V",
dim=list(new_v.shape),
stride=list(new_v.stride()),
data_type=convert_to_cudnn_type(qkv.dtype),
)
def _scale_tensor():
return graph.tensor(dim=[1, 1, 1, 1], stride=[1, 1, 1, 1], data_type=cudnn.data_type.FLOAT)
default_scale_gpu = torch.ones(1, 1, 1, 1, dtype=torch.float32, device="cuda")
descale_q = _scale_tensor()
descale_k = _scale_tensor()
descale_v = _scale_tensor()
descale_s = _scale_tensor()
scale_s = _scale_tensor()
scale_o = _scale_tensor()
o, _, amax_s, amax_o = graph.sdpa_fp8(
q=q,
k=k,
v=v,
descale_q=descale_q,
descale_k=descale_k,
descale_v=descale_v,
descale_s=descale_s,
scale_s=scale_s,
scale_o=scale_o,
is_inference=True,
attn_scale=1.0 / math.sqrt(headdim),
use_causal_mask=causal,
name="sdpa",
)
o.set_output(True).set_dim(o_gpu_transposed.shape).set_stride(o_gpu_transposed.stride())
amax_s.set_output(False).set_dim(amax_s_gpu.shape).set_stride(amax_s_gpu.stride())
amax_o.set_output(False).set_dim(amax_o_gpu.shape).set_stride(amax_o_gpu.stride())
graph.validate()
graph.build_operation_graph()
graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK])
graph.check_support()
graph.build_plans()
variant_pack = {
q: new_q,
k: new_k,
v: new_v,
descale_q: default_scale_gpu,
descale_k: default_scale_gpu,
descale_v: default_scale_gpu,
descale_s: default_scale_gpu,
scale_s: default_scale_gpu,
scale_o: default_scale_gpu,
o: o_gpu_transposed,
amax_s: amax_s_gpu,
amax_o: amax_o_gpu,
}
workspace = torch.empty(graph.get_workspace_size(), device="cuda", dtype=torch.uint8)
def run():
graph.execute(variant_pack, workspace)
return o_gpu
return run
def _maybe_pass_descales(callable_, **kwargs):
sig = inspect.signature(callable_)
return {k: v for k, v in kwargs.items() if k in sig.parameters}
def main(argv: Iterable[str] | None = None) -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--repeats", type=int, default=30)
parser.add_argument("--dim", type=int, default=2048)
parser.add_argument("--headdims", default="64,128")
parser.add_argument("--dtype", default="fp8_e4m3fn")
parser.add_argument("--seed", type=int, default=0)
parser.add_argument(
"--check",
action=argparse.BooleanOptionalAction,
default=True,
help="Enable correctness checks vs BF16 PyTorch baseline.",
)
parser.add_argument(
"--check-quantization-only",
action="store_true",
help="Check FP8 kernel vs dequantized-FP8 baseline (quantization error only).",
)
parser.add_argument("--atol-bf16", type=float, default=0.10)
parser.add_argument("--rtol-bf16", type=float, default=0.10)
parser.add_argument("--atol-fp8", type=float, default=0.50)
parser.add_argument("--rtol-fp8", type=float, default=0.50)
parser.add_argument("--run-cudnn", action="store_true")
args = parser.parse_args(list(argv) if argv is not None else None)
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required")
major, minor = torch.cuda.get_device_capability()
if major != 10:
raise RuntimeError(
f"This benchmark is for SM100 (compute capability 10.x). Got {major}.{minor}."
)
torch.manual_seed(args.seed)
device = "cuda"
fp8_dtype = _torch_float8_dtype(args.dtype)
headdim_vals = _parse_int_list(args.headdims)
bs_seqlen_vals = [(32, 512), (16, 1024), (8, 2048), (4, 4096), (2, 8192), (1, 16384)]
methods = ["Pytorch", "FA4-CuTe-BF16", "FA4-CuTe-FP8"] + (
["cuDNN-FP8"] if args.run_cudnn and cudnn is not None else []
)
fp8_failures = []
for headdim in headdim_vals:
for causal in (False, True):
for batch, seqlen in bs_seqlen_vals:
torch.cuda.empty_cache()
nheads = args.dim // headdim
if args.dim % headdim != 0:
raise ValueError(f"--dim must be divisible by headdim ({args.dim=} {headdim=})")
q_bf16 = torch.randn(
batch, seqlen, nheads, headdim, device=device, dtype=torch.bfloat16
)
k_bf16 = torch.randn(
batch, seqlen, nheads, headdim, device=device, dtype=torch.bfloat16
)
v_bf16 = torch.randn(
batch, seqlen, nheads, headdim, device=device, dtype=torch.bfloat16
)
qkv_bf16 = torch.stack([q_bf16, k_bf16, v_bf16], dim=2)
times = {}
speeds = {}
out_ref_bf16 = None
try:
out_ref_bf16 = attention_pytorch(qkv_bf16, causal=causal) # warmup / reference
t = time_fwd(attention_pytorch, qkv_bf16, causal=causal, repeats=args.repeats)
times["Pytorch"] = t
except RuntimeError as e:
if "out of memory" in str(e).lower():
times["Pytorch"] = float("nan")
out_ref_bf16 = None
else:
raise
# FA4 / CuTe BF16 baseline
try:
softmax_scale = headdim**-0.5
out_fa4_bf16, *_ = flash_attn_cute_fwd(
q_bf16, k_bf16, v_bf16, softmax_scale=softmax_scale, causal=causal
) # warmup / compile
t = time_fwd(
flash_attn_cute_fwd,
q_bf16,
k_bf16,
v_bf16,
softmax_scale=softmax_scale,
causal=causal,
repeats=args.repeats,
)
times["FA4-CuTe-BF16"] = t
if args.check and out_ref_bf16 is not None:
torch.testing.assert_close(
out_fa4_bf16,
out_ref_bf16,
atol=args.atol_bf16,
rtol=args.rtol_bf16,
)
except Exception as e:
# Treat as fatal: BF16 kernel should be usable for basic sanity checking.
raise RuntimeError("FA4-CuTe BF16 baseline failed") from e
# FA4 / CuTe FP8
q_fp8 = q_bf16.to(fp8_dtype)
k_fp8 = k_bf16.to(fp8_dtype)
v_fp8 = v_bf16.to(fp8_dtype)
# Placeholder descales (FA3-style: per-(batch, kv_head)).
q_descale = torch.ones(batch, nheads, device=device, dtype=torch.float32)
k_descale = torch.ones(batch, nheads, device=device, dtype=torch.float32)
v_descale = torch.ones(batch, nheads, device=device, dtype=torch.float32)
# Optional: FP8 reference baseline (dequantized FP8 -> PyTorch) for quantization-error-only checks
out_ref_fp8 = None
if args.check and args.check_quantization_only:
try:
# Dequantize FP8 inputs back to BF16 (applying descales)
q_ref_fp8 = (q_fp8.to(torch.bfloat16) * q_descale[:, None, :, None]).to(
torch.bfloat16
)
k_ref_fp8 = (k_fp8.to(torch.bfloat16) * k_descale[:, None, :, None]).to(
torch.bfloat16
)
v_ref_fp8 = (v_fp8.to(torch.bfloat16) * v_descale[:, None, :, None]).to(
torch.bfloat16
)
qkv_ref_fp8 = torch.stack([q_ref_fp8, k_ref_fp8, v_ref_fp8], dim=2)
out_ref_fp8 = attention_pytorch(qkv_ref_fp8, causal=causal)
except RuntimeError as e:
if "out of memory" in str(e).lower():
out_ref_fp8 = None
else:
raise
fa4_kwargs = dict(softmax_scale=softmax_scale, causal=causal)
fa4_kwargs.update(
_maybe_pass_descales(
flash_attn_cute_fwd,
q_descale=q_descale,
k_descale=k_descale,
v_descale=v_descale,
)
)
try:
# Warmup/compile (will raise until FP8 is implemented)
out_fa4_fp8, *_ = flash_attn_cute_fwd(q_fp8, k_fp8, v_fp8, **fa4_kwargs)
t = time_fwd(
flash_attn_cute_fwd,
q_fp8,
k_fp8,
v_fp8,
repeats=args.repeats,
**fa4_kwargs,
)
times["FA4-CuTe-FP8"] = t
if args.check:
# Choose baseline: quantization-only (dequantized FP8) or full (BF16)
if args.check_quantization_only:
ref_baseline = out_ref_fp8
else:
ref_baseline = out_ref_bf16
if ref_baseline is not None:
torch.testing.assert_close(
out_fa4_fp8,
ref_baseline,
atol=args.atol_fp8,
rtol=args.rtol_fp8,
)
except Exception as e:
fp8_failures.append((causal, headdim, batch, seqlen, repr(e)))
times["FA4-CuTe-FP8"] = float("nan")
if args.run_cudnn and cudnn is not None:
qkv_fp8 = qkv_bf16.to(fp8_dtype)
runner = cudnn_sdpa_fp8_setup(qkv_fp8, seqlen, seqlen, causal=causal)
_ = runner() # warmup
t = time_fwd(lambda: runner(), repeats=args.repeats)
times["cuDNN-FP8"] = t
print(f"### causal={causal}, headdim={headdim}, batch={batch}, seqlen={seqlen} ###")
for method in methods:
t = times.get(method, float("nan"))
speeds[method] = efficiency(flops(batch, seqlen, headdim, nheads, causal), t)
if math.isnan(t):
print(f"{method} fwd: (skipped)")
else:
print(f"{method} fwd: {speeds[method]:.2f} TFLOPs/s, {t * 1e3:.3f} ms")
if math.isnan(times.get("FA4-CuTe-FP8", float("nan"))):
print("FA4-CuTe-FP8 status: FAILED")
if fp8_failures:
print(f"\nFP8 failures: {len(fp8_failures)} (showing first 5)")
for causal, headdim, batch, seqlen, err in fp8_failures[:5]:
print(f"- causal={causal} headdim={headdim} batch={batch} seqlen={seqlen}: {err}")
return 1
return 0
if __name__ == "__main__":
raise SystemExit(main())
File diff suppressed because it is too large Load Diff
+156
View File
@@ -0,0 +1,156 @@
# Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
from typing import Tuple, Optional
from dataclasses import dataclass
import cutlass
import cutlass.cute as cute
from cutlass import Int32, const_expr
from flash_attn.cute.seqlen_info import SeqlenInfoQK, SeqlenInfoQKNewK
@dataclass(frozen=True)
class BlockInfo:
tile_m: cutlass.Constexpr[int]
tile_n: cutlass.Constexpr[int]
is_causal: cutlass.Constexpr[bool]
is_local: cutlass.Constexpr[bool] = False
is_split_kv: cutlass.Constexpr[bool] = False
window_size_left: Optional[Int32] = None
window_size_right: Optional[Int32] = None
qhead_per_kvhead_packgqa: cutlass.Constexpr[int] = 1
@cute.jit
def get_n_block_min_max(
self,
seqlen_info: SeqlenInfoQK,
m_block: Int32,
split_idx: Int32 = 0,
num_splits: Int32 = 1,
) -> Tuple[Int32, Int32]:
n_block_max = cute.ceil_div(seqlen_info.seqlen_k, self.tile_n)
if const_expr(self.is_causal or (self.is_local and self.window_size_right is not None)):
m_idx_max = (m_block + 1) * self.tile_m
if const_expr(self.qhead_per_kvhead_packgqa > 1):
m_idx_max = cute.ceil_div(m_idx_max, self.qhead_per_kvhead_packgqa)
n_idx = m_idx_max + seqlen_info.seqlen_k - seqlen_info.seqlen_q
n_idx_right = n_idx if const_expr(self.is_causal) else n_idx + self.window_size_right
n_block_max = min(n_block_max, cute.ceil_div(n_idx_right, self.tile_n))
n_block_min = 0
if const_expr(self.is_local and self.window_size_left is not None):
m_idx_min = m_block * self.tile_m
if const_expr(self.qhead_per_kvhead_packgqa > 1):
m_idx_min = m_idx_min // self.qhead_per_kvhead_packgqa
n_idx = m_idx_min + seqlen_info.seqlen_k - seqlen_info.seqlen_q
n_idx_left = n_idx - self.window_size_left
n_block_min = cutlass.max(n_idx_left // self.tile_n, 0)
if cutlass.const_expr(self.is_split_kv):
num_n_blocks_per_split = (
Int32(0)
if n_block_max <= n_block_min
else (n_block_max - n_block_min + num_splits - 1) // num_splits
)
n_block_min = n_block_min + split_idx * num_n_blocks_per_split
n_block_max = cutlass.min(n_block_min + num_n_blocks_per_split, n_block_max)
return n_block_min, n_block_max
@cute.jit
def get_m_block_min_max(self, seqlen_info: SeqlenInfoQK, n_block: Int32) -> Tuple[Int32, Int32]:
m_block_max = cute.ceil_div(seqlen_info.seqlen_q, self.tile_m)
m_block_min = 0
if const_expr(self.is_causal or (self.is_local and self.window_size_right is not None)):
n_idx_min = n_block * self.tile_n
m_idx = n_idx_min + seqlen_info.seqlen_q - seqlen_info.seqlen_k
m_idx_right = m_idx if const_expr(self.is_causal) else m_idx - self.window_size_right
m_block_min = max(m_block_min, m_idx_right // self.tile_m)
if const_expr(self.is_local and self.window_size_left is not None):
n_idx_max = (n_block + 1) * self.tile_n
m_idx = n_idx_max + seqlen_info.seqlen_q - seqlen_info.seqlen_k
m_idx_left = m_idx + self.window_size_left
m_block_max = min(m_block_max, cute.ceil_div(m_idx_left, self.tile_m))
return m_block_min, m_block_max
@cute.jit
def get_n_block_k_new_min_max(
self,
seqlen_info: SeqlenInfoQKNewK,
m_block: Int32,
split_idx: Int32 = 0,
num_splits: Int32 = 1,
) -> Tuple[Int32, Int32]:
"""Get the block range for new K tokens (append KV).
First computes the full n_block range via get_n_block_min_max, then maps
those blocks into the new-K index space by subtracting seqlen_k_og.
"""
n_block_min, n_block_max = self.get_n_block_min_max(
seqlen_info,
m_block,
split_idx,
num_splits,
)
idx_k_new_min = cutlass.max(n_block_min * self.tile_n - seqlen_info.seqlen_k_og, 0)
idx_k_new_max = cutlass.min(
n_block_max * self.tile_n - seqlen_info.seqlen_k_og, seqlen_info.seqlen_k_new
)
n_block_new_min = idx_k_new_min // self.tile_n
n_block_new_max = (
cute.ceil_div(idx_k_new_max, self.tile_n)
if idx_k_new_max > idx_k_new_min
else n_block_new_min
)
return n_block_new_min, n_block_new_max
@cute.jit
def get_n_block_min_causal_local_mask(
self,
seqlen_info: SeqlenInfoQK,
m_block: Int32,
n_block_min: Int32,
) -> Int32:
"""If we have separate iterations with causal or local masking at the start, where do we stop"""
m_idx_min = m_block * self.tile_m
if const_expr(self.qhead_per_kvhead_packgqa > 1):
m_idx_min = m_idx_min // self.qhead_per_kvhead_packgqa
n_idx = m_idx_min + seqlen_info.seqlen_k - seqlen_info.seqlen_q
n_idx_right = (
n_idx
if const_expr(not self.is_local or self.window_size_right is None)
else n_idx + self.window_size_right
)
return cutlass.max(n_block_min, n_idx_right // self.tile_n)
@cute.jit
def get_n_block_min_before_local_mask(
self,
seqlen_info: SeqlenInfoQK,
m_block: Int32,
n_block_min: Int32,
) -> Int32:
"""If we have separate iterations with local masking at the end, where do we stop the non-masked iterations"""
if const_expr(not self.is_local or self.window_size_left is None):
return n_block_min
else:
m_idx_max = (m_block + 1) * self.tile_m
if const_expr(self.qhead_per_kvhead_packgqa > 1):
m_idx_max = cute.ceil_div(m_idx_max, self.qhead_per_kvhead_packgqa)
n_idx = m_idx_max + seqlen_info.seqlen_k - seqlen_info.seqlen_q
n_idx_left = n_idx - self.window_size_left
return cutlass.max(n_block_min, cute.ceil_div(n_idx_left, self.tile_n))
@cute.jit
def get_n_block_max_for_m_block(
self,
seqlen_info: SeqlenInfoQK,
m_block: Int32,
) -> Int32:
n_block_max = cute.ceil_div(seqlen_info.seqlen_k, self.tile_n)
if const_expr(self.is_causal or self.window_size_right is not None):
m_idx_max = (m_block + 1) * self.tile_m
if const_expr(self.qhead_per_kvhead_packgqa > 1):
m_idx_max = cute.ceil_div(m_idx_max, self.qhead_per_kvhead_packgqa)
n_idx_right = m_idx_max + seqlen_info.seqlen_k - seqlen_info.seqlen_q
if const_expr(self.window_size_right is not None):
n_idx_right += self.window_size_right
n_block_max = min(n_block_max, cute.ceil_div(n_idx_right, self.tile_n))
return n_block_max
File diff suppressed because it is too large Load Diff
+676
View File
@@ -0,0 +1,676 @@
"""
Block-sparsity utilities for FlexAttention
"""
from typing import Callable, NamedTuple, Tuple
import cutlass.cute as cute
import torch
from flash_attn.cute.cute_dsl_utils import get_broadcast_dims, to_cute_tensor
def ceildiv(a: int, b: int) -> int:
return (a + b - 1) // b
class BlockSparseTensors(NamedTuple):
mask_block_cnt: cute.Tensor
mask_block_idx: cute.Tensor
full_block_cnt: cute.Tensor | None = None
full_block_idx: cute.Tensor | None = None
cu_total_m_blocks: cute.Tensor | None = None
cu_block_idx_offsets: cute.Tensor | None = None
dq_write_order: cute.Tensor | None = None
dq_write_order_full: cute.Tensor | None = None
def __new_from_mlir_values__(self, values):
new_fields = []
idx = 0
for original in self:
if original is None:
new_fields.append(None)
else:
new_fields.append(values[idx])
idx += 1
return BlockSparseTensors(*new_fields)
class BlockSparseTensorsTorch(NamedTuple):
mask_block_cnt: torch.Tensor
mask_block_idx: torch.Tensor
full_block_cnt: torch.Tensor | None = None
full_block_idx: torch.Tensor | None = None
cu_total_m_blocks: torch.Tensor | None = None
cu_block_idx_offsets: torch.Tensor | None = None
block_size: tuple[int, int] | None = None
dq_write_order: torch.Tensor | None = None
dq_write_order_full: torch.Tensor | None = None
spt: bool | None = None
def _ordered_to_dense_simple(
num_blocks: torch.Tensor,
indices: torch.Tensor,
num_cols: int,
) -> torch.Tensor:
"""Convert ordered sparse representation to dense binary matrix.
Args:
num_blocks: [B, H, num_rows] count of valid entries per row
indices: [B, H, num_rows, max_entries] column indices (valid entries packed left)
num_cols: total number of columns
Returns:
dense: [B, H, num_rows, num_cols] binary int32 matrix
"""
B, H, num_rows, max_entries = indices.shape
device = indices.device
dense = torch.zeros(B, H, num_rows, num_cols + 1, dtype=torch.int32, device=device)
col_range = torch.arange(max_entries, device=device)
valid = col_range[None, None, None, :] < num_blocks[:, :, :, None]
safe_indices = torch.where(valid, indices.long(), num_cols)
row_idx = torch.arange(num_rows, device=device)[None, None, :, None].expand_as(indices)
b_idx = torch.arange(B, device=device)[:, None, None, None].expand_as(indices)
h_idx = torch.arange(H, device=device)[None, :, None, None].expand_as(indices)
dense[b_idx, h_idx, row_idx, safe_indices] = 1
return dense[:, :, :, :num_cols]
def compute_dq_write_order(
fwd_mask_cnt: torch.Tensor,
fwd_mask_idx: torch.Tensor,
fwd_full_cnt: torch.Tensor | None,
fwd_full_idx: torch.Tensor | None,
bwd_mask_cnt: torch.Tensor,
bwd_mask_idx: torch.Tensor,
bwd_full_cnt: torch.Tensor | None,
bwd_full_idx: torch.Tensor | None,
spt: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Compute dQ write-order metadata for deterministic block-sparse backward.
For each (n_block, i) in the backward iteration, computes the semaphore
lock value: the rank of n_block in the combined (partial + full) sorted
contributor list for the target m_block.
Lock values are assigned in ascending n_block order (or descending if spt=True)
to guarantee deadlock-freedom with the CTA scheduling order.
Args:
fwd_mask_cnt: [B, H, num_m_blocks] partial contributor counts per m_block
fwd_mask_idx: [B, H, num_m_blocks, max_kv] partial contributor n_block indices (ascending)
fwd_full_cnt: [B, H, num_m_blocks] full contributor counts per m_block (optional)
fwd_full_idx: [B, H, num_m_blocks, max_kv] full contributor n_block indices (optional)
bwd_mask_cnt: [B, H, num_n_blocks] partial iteration counts per n_block
bwd_mask_idx: [B, H, num_n_blocks, max_q] partial iteration m_block indices
bwd_full_cnt: [B, H, num_n_blocks] full iteration counts per n_block (optional)
bwd_full_idx: [B, H, num_n_blocks, max_q] full iteration m_block indices (optional)
spt: if True, reverse ordering (highest n_block gets lock_value=0)
Returns:
(dq_write_order, dq_write_order_full): tensors parallel to bwd_mask_idx
and bwd_full_idx respectively, containing lock values.
"""
device = fwd_mask_idx.device
B, H, num_m, max_kv_partial = fwd_mask_idx.shape
_, _, num_n, max_q_partial = bwd_mask_idx.shape
has_full = fwd_full_cnt is not None and fwd_full_idx is not None
dense_partial = _ordered_to_dense_simple(fwd_mask_cnt, fwd_mask_idx, num_n)
if has_full:
dense_full = _ordered_to_dense_simple(fwd_full_cnt, fwd_full_idx, num_n)
dense = (dense_partial + dense_full).clamp(max=1)
else:
dense = dense_partial
cumsum = dense.cumsum(dim=-1)
rank_table = (cumsum - dense).to(torch.int32)
if spt:
total_per_m = cumsum[:, :, :, -1:]
rank_table = (total_per_m - 1 - rank_table).to(torch.int32)
def _gather_write_order(bwd_idx, bwd_cnt):
b_i = torch.arange(B, device=device)[:, None, None, None].expand_as(bwd_idx)
h_i = torch.arange(H, device=device)[None, :, None, None].expand_as(bwd_idx)
n_i = torch.arange(bwd_idx.shape[2], device=device)[None, None, :, None].expand_as(bwd_idx)
m_vals = bwd_idx.long().clamp(0, num_m - 1)
return rank_table[b_i, h_i, m_vals, n_i].to(torch.int32)
dq_write_order = _gather_write_order(bwd_mask_idx, bwd_mask_cnt)
dq_write_order_full = None
if has_full and bwd_full_cnt is not None and bwd_full_idx is not None:
dq_write_order_full = _gather_write_order(bwd_full_idx, bwd_full_cnt)
return dq_write_order, dq_write_order_full
def compute_dq_write_order_from_block_mask(
block_mask,
spt: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
(
_seq_q,
_seq_k,
kv_mask_cnt,
kv_mask_idx,
full_kv_cnt,
full_kv_idx,
q_mask_cnt,
q_mask_idx,
full_q_cnt,
full_q_idx,
*_,
) = block_mask.as_tuple()
return compute_dq_write_order(
kv_mask_cnt,
kv_mask_idx,
full_kv_cnt,
full_kv_idx,
q_mask_cnt,
q_mask_idx,
full_q_cnt,
full_q_idx,
spt=spt,
)
def get_sparse_q_block_size(
tensors: BlockSparseTensorsTorch | None,
seqlen_q: int,
) -> int | None:
"""Return the Q sparse block size, or None when sparsity is unset or ambiguous."""
if tensors is None:
return None
if tensors.block_size is not None:
return tensors.block_size[0]
num_m_blocks = tensors.mask_block_idx.shape[2]
min_block_size = ceildiv(seqlen_q, num_m_blocks)
max_block_size = seqlen_q if num_m_blocks == 1 else (seqlen_q - 1) // (num_m_blocks - 1)
if min_block_size != max_block_size:
return None
return min_block_size
def _expand_sparsity_tensor(
tensor: torch.Tensor,
expected_shape: Tuple[int, ...],
tensor_name: str,
context: str | None,
hint: str | Callable[[], str] | None,
) -> torch.Tensor:
"""Check if we need to expand the tensor to expected shape, and do so if possible."""
needs_expand = tensor.shape != expected_shape
if not needs_expand:
return tensor
can_expand = all(map(lambda cur, tgt: cur == tgt or cur == 1, tensor.shape, expected_shape))
if not can_expand:
context_clause = f" ({context})" if context else ""
resolved_hint = hint() if callable(hint) else hint
hint_clause = f" Hint: {resolved_hint}" if resolved_hint else ""
raise ValueError(
f"{tensor_name}{context_clause} with shape {tensor.shape} cannot be expanded to expected shape {expected_shape}."
f"{hint_clause}"
)
return tensor.expand(*expected_shape)
def _check_and_expand_block(
name: str,
cnt: torch.Tensor | None,
idx: torch.Tensor | None,
expected_count_shape: Tuple[int, ...],
expected_index_shape: Tuple[int, ...],
context: str | None,
hint: str | Callable[[], str] | None,
) -> Tuple[torch.Tensor | None, torch.Tensor | None]:
if (cnt is None) != (idx is None):
raise ValueError(
f"{name}_block_cnt and {name}_block_idx must both be provided or both be None"
)
if cnt is None or idx is None:
return None, None
if cnt.dtype != torch.int32 or idx.dtype != torch.int32:
raise ValueError(f"{name}_block tensors must have dtype torch.int32")
if cnt.device != idx.device:
raise ValueError(f"{name}_block_cnt and {name}_block_idx must be on the same device")
if not cnt.is_cuda or not idx.is_cuda:
raise ValueError(f"{name}_block tensors must live on CUDA")
expanded_cnt = _expand_sparsity_tensor(
cnt, expected_count_shape, f"{name}_block_cnt", context, hint
)
# [Note] Allow Compact block sparse indices
# Allow the last dimension (n_blocks) of idx to be <= expected, since
# FA4 only accesses indices 0..cnt-1 per query tile. This enables compact
# index tensors that avoid O(N^2) memory at long sequence lengths.
if idx.ndim == 4 and idx.shape[3] <= expected_index_shape[3]:
expected_index_shape = (*expected_index_shape[:3], idx.shape[3])
expanded_idx = _expand_sparsity_tensor(
idx, expected_index_shape, f"{name}_block_idx", context, hint
)
return expanded_cnt, expanded_idx
def _check_and_expand_metadata_tensor(
name: str,
tensor: torch.Tensor | None,
expected_shape: Tuple[int, ...],
context: str | None,
hint: str | Callable[[], str] | None,
device: torch.device,
) -> torch.Tensor | None:
if tensor is None:
return None
if tensor.dtype != torch.int32:
raise ValueError(f"{name} must have dtype torch.int32")
if tensor.device != device:
raise ValueError(f"{name} must be on the same device as block sparse tensors")
if not tensor.is_cuda:
raise ValueError(f"{name} must live on CUDA")
return _expand_sparsity_tensor(tensor, expected_shape, name, context, hint)
def get_block_sparse_expected_shapes(
batch_size: int,
num_head: int,
seqlen_q: int,
seqlen_k: int,
m_block_size: int,
n_block_size: int,
q_stage: int,
) -> Tuple[Tuple[int, int, int], Tuple[int, int, int, int]]:
"""Return (expected_count_shape, expected_index_shape) for block sparse normalization."""
m_block_size_effective = q_stage * m_block_size
expected_m_blocks = ceildiv(seqlen_q, m_block_size_effective)
expected_n_blocks = ceildiv(seqlen_k, n_block_size)
expected_count_shape = (batch_size, num_head, expected_m_blocks)
expected_index_shape = (batch_size, num_head, expected_m_blocks, expected_n_blocks)
return expected_count_shape, expected_index_shape
def infer_block_sparse_expected_shapes(
tensors: BlockSparseTensorsTorch,
*,
batch_size: int,
num_head: int,
seqlen_q: int,
seqlen_k: int,
m_block_size: int,
n_block_size: int,
q_stage: int,
context: str,
sparse_block_size_q: int | None = None,
sparse_block_size_kv: int | None = None,
) -> Tuple[Tuple[int, int, int], Tuple[int, int, int, int], int]:
"""Infer shapes and scaling for block-sparse tensors.
Expectations:
- mask_block_cnt is (B, H, M) and mask_block_idx is (B, H, M, N).
- Batch/head dims may be 1 for broadcast, or match the requested sizes.
- sparse_block_size_kv must match tile_n.
- sparse_block_size_q must be a multiple of q_stage * tile_m.
- If sparse_block_size_q is omitted and seqlen_q/num_m_blocks is ambiguous,
the caller must provide block_size to disambiguate. TODO will make this required in a future PR.
"""
base_m_block = q_stage * m_block_size
base_n_block = n_block_size
if sparse_block_size_kv is None:
sparse_block_size_kv = base_n_block
if sparse_block_size_kv != base_n_block:
raise ValueError(f"Block sparse tensors{context} require BLOCK_SIZE_KV={base_n_block}.")
if tensors.mask_block_idx is None:
raise ValueError("mask_block_cnt and mask_block_idx must be provided for block sparsity.")
num_m_blocks = tensors.mask_block_idx.shape[2]
if sparse_block_size_q is None:
sparse_block_size_q = get_sparse_q_block_size(tensors, seqlen_q)
if sparse_block_size_q is None and base_m_block != 1:
raise ValueError(
f"Block sparse tensors{context} require explicit sparse_block_size[0] "
f"to disambiguate block size for seqlen_q={seqlen_q} and num_m_blocks={num_m_blocks}."
)
if sparse_block_size_q is None:
sparse_block_size_q = ceildiv(seqlen_q, num_m_blocks)
if sparse_block_size_q % base_m_block != 0:
raise ValueError(
f"Block sparse tensors{context} have block size {sparse_block_size_q}, "
f"which must be a multiple of {base_m_block}."
)
expected_m_blocks = ceildiv(seqlen_q, sparse_block_size_q)
expected_n_blocks = ceildiv(seqlen_k, sparse_block_size_kv)
q_subtile_factor = sparse_block_size_q // base_m_block
expected_count_shape = (batch_size, num_head, expected_m_blocks)
expected_index_shape = (batch_size, num_head, expected_m_blocks, expected_n_blocks)
mask_block_cnt = tensors.mask_block_cnt
mask_block_idx = tensors.mask_block_idx
if mask_block_cnt is None or mask_block_idx is None:
raise ValueError("mask_block_cnt and mask_block_idx must be provided for block sparsity.")
if mask_block_cnt.ndim != 3 or mask_block_idx.ndim != 4:
raise ValueError(
f"Block sparse tensors{context} must have shapes (B, H, M) and (B, H, M, N)."
)
for dim_name, cur, tgt in (
("batch", mask_block_cnt.shape[0], expected_count_shape[0]),
("head", mask_block_cnt.shape[1], expected_count_shape[1]),
):
if cur != tgt and cur != 1:
raise ValueError(f"Block sparse tensors{context} {dim_name} dim must be {tgt} or 1.")
for dim_name, cur, tgt in (
("batch", mask_block_idx.shape[0], expected_index_shape[0]),
("head", mask_block_idx.shape[1], expected_index_shape[1]),
):
if cur != tgt and cur != 1:
raise ValueError(f"Block sparse tensors{context} {dim_name} dim must be {tgt} or 1.")
if mask_block_cnt.shape[2] != mask_block_idx.shape[2]:
raise ValueError(f"Block sparse tensors{context} must share the same m-block dimension.")
# [Note] Allow Compact block sparse indices: FA4 only accesses indices 0..cnt-1
# per query tile, so idx.shape[3] can be <= expected_n_blocks.
if mask_block_idx.shape[3] > expected_n_blocks:
raise ValueError(
f"Block sparse tensors{context} n-block dimension must be <= {expected_n_blocks}."
)
if expected_m_blocks != num_m_blocks:
raise ValueError(
f"Block sparse tensors{context} m-block dimension {num_m_blocks} does not match "
f"sparse_block_size_q={sparse_block_size_q}. "
f"Set BlockSparseTensorsTorch.block_size to match the BlockMask BLOCK_SIZE."
)
return expected_count_shape, expected_index_shape, q_subtile_factor
def get_block_sparse_expected_shapes_bwd(
batch_size: int,
num_head: int,
seqlen_q: int,
seqlen_k: int,
m_block_size: int,
n_block_size: int,
q_subtile_factor: int,
) -> Tuple[Tuple[int, int, int], Tuple[int, int, int, int]]:
"""Return (expected_count_shape, expected_index_shape) for backward block sparse normalization.
Backward uses Q-direction indexing (transposed from forward), where shapes are
indexed by N-blocks first, then M-blocks. The sparse_block_size_q is determined
by q_subtile_factor * m_block_size.
"""
sparse_block_size_q = q_subtile_factor * m_block_size
expected_m_blocks = ceildiv(seqlen_q, sparse_block_size_q)
expected_n_blocks = ceildiv(seqlen_k, n_block_size)
expected_count_shape = (batch_size, num_head, expected_n_blocks)
expected_index_shape = (batch_size, num_head, expected_n_blocks, expected_m_blocks)
return expected_count_shape, expected_index_shape
def normalize_block_sparse_tensors(
tensors: BlockSparseTensorsTorch,
*,
expected_count_shape: Tuple[int, ...],
expected_index_shape: Tuple[int, ...],
context: str | None = None,
hint: str | Callable[[], str] | None = None,
) -> BlockSparseTensorsTorch:
if tensors.mask_block_cnt is None or tensors.mask_block_idx is None:
raise ValueError("mask_block_cnt and mask_block_idx must be provided for block sparsity.")
mask_cnt, mask_idx = _check_and_expand_block(
"mask",
tensors.mask_block_cnt,
tensors.mask_block_idx,
expected_count_shape,
expected_index_shape,
context,
hint,
)
if mask_cnt is None or mask_idx is None:
raise ValueError("mask_block_cnt and mask_block_idx must be provided for block sparsity.")
full_cnt, full_idx = _check_and_expand_block(
"full",
tensors.full_block_cnt,
tensors.full_block_idx,
expected_count_shape,
expected_index_shape,
context,
hint,
)
if full_cnt is not None and mask_cnt.device != full_cnt.device:
raise ValueError("All block sparse tensors must be on the same device")
dq_write_order = _check_and_expand_metadata_tensor(
"dq_write_order",
tensors.dq_write_order,
tuple(mask_idx.shape),
context,
hint,
mask_cnt.device,
)
dq_write_order_full = _check_and_expand_metadata_tensor(
"dq_write_order_full",
tensors.dq_write_order_full,
tuple(full_idx.shape) if full_idx is not None else expected_index_shape,
context,
hint,
mask_cnt.device,
)
spt = tensors.spt
if spt is not None and not isinstance(spt, bool):
raise ValueError("spt must be a bool when provided")
if spt is not None and dq_write_order is None:
raise ValueError("spt requires dq_write_order to be provided")
return BlockSparseTensorsTorch(
mask_block_cnt=mask_cnt,
mask_block_idx=mask_idx,
full_block_cnt=full_cnt,
full_block_idx=full_idx,
cu_total_m_blocks=tensors.cu_total_m_blocks,
cu_block_idx_offsets=tensors.cu_block_idx_offsets,
block_size=tensors.block_size,
dq_write_order=dq_write_order,
dq_write_order_full=dq_write_order_full,
spt=spt,
)
def is_block_sparsity_enabled(tensors: BlockSparseTensorsTorch) -> bool:
return any(t is not None for t in (tensors.full_block_cnt, tensors.mask_block_cnt))
def get_block_sparse_broadcast_pattern(
tensors: BlockSparseTensorsTorch,
) -> Tuple[Tuple[bool, ...], ...] | None:
"""Return broadcast pattern for block sparse tensors by checking actual strides.
Returns a tuple of broadcast patterns (one per tensor) where each pattern
is a tuple of bools indicating which dims have stride=0.
This is used in compile keys to ensure kernels are recompiled when
broadcast patterns change, since CuTe's mark_layout_dynamic() keeps
stride=0 as static.
The tensors should already be expanded/normalized before calling this function.
Returns None if block sparsity is not enabled.
"""
if not is_block_sparsity_enabled(tensors):
return None
patterns = []
for tensor in (
tensors.mask_block_cnt,
tensors.mask_block_idx,
tensors.full_block_cnt,
tensors.full_block_idx,
tensors.dq_write_order,
tensors.dq_write_order_full,
):
if tensor is not None:
patterns.append(get_broadcast_dims(tensor))
else:
patterns.append(None)
return tuple(patterns)
def normalize_block_sparse_config(
tensors: BlockSparseTensorsTorch,
*,
batch_size: int,
num_head: int,
seqlen_q: int,
seqlen_k: int,
block_size: tuple[int, int],
q_stage: int,
) -> tuple[BlockSparseTensorsTorch, Tuple[Tuple[bool, ...], ...] | None, int]:
"""Validate the block-sparse config, infer expected shapes, and normalize.
Handles both fixed-length (3D `[B, H, M]` / 4D `[B, H, M, N]`) and varlen
(2D `[H, total_m_blocks]` / `[H, total_n_blocks]`) layouts. Varlen is
detected by `tensors.cu_total_m_blocks is not None` and forces
`q_subtile_factor == 1` (TODO: potentially remove this restriction).
"""
m_block_size, n_block_size = block_size
if tensors.block_size is None:
sparse_block_size_q, sparse_block_size_kv = None, n_block_size
else:
sparse_block_size_q, sparse_block_size_kv = tensors.block_size
if sparse_block_size_kv != n_block_size:
raise ValueError(
f"Block sparsity requires sparse_block_size[1]={n_block_size} to match tile_n."
)
if tensors.cu_total_m_blocks is not None:
base_m_block = q_stage * m_block_size
if sparse_block_size_q is not None and sparse_block_size_q != base_m_block:
raise ValueError(
f"Varlen block sparsity requires sparse_block_size[0]={base_m_block} "
f"(= q_stage * tile_m); got {sparse_block_size_q}."
)
total_m_blocks = tensors.mask_block_cnt.shape[-1]
total_n_blocks = tensors.mask_block_idx.shape[-1]
expected_count_shape = (num_head, total_m_blocks)
expected_index_shape = (num_head, total_n_blocks)
q_subtile_factor = 1
else:
expected_count_shape, expected_index_shape, q_subtile_factor = (
infer_block_sparse_expected_shapes(
tensors,
batch_size=batch_size,
num_head=num_head,
seqlen_q=seqlen_q,
seqlen_k=seqlen_k,
m_block_size=m_block_size,
n_block_size=n_block_size,
q_stage=q_stage,
context="forward",
sparse_block_size_q=sparse_block_size_q,
sparse_block_size_kv=sparse_block_size_kv,
)
)
normalized_tensors = normalize_block_sparse_tensors(
tensors,
expected_count_shape=expected_count_shape,
expected_index_shape=expected_index_shape,
)
return (
normalized_tensors,
get_block_sparse_broadcast_pattern(normalized_tensors),
q_subtile_factor,
)
def normalize_block_sparse_config_bwd(
tensors: BlockSparseTensorsTorch,
*,
batch_size: int,
num_head: int,
seqlen_q: int,
seqlen_k: int,
block_size: tuple[int, int],
q_subtile_factor: int,
) -> tuple[BlockSparseTensorsTorch, Tuple[Tuple[bool, ...], ...] | None]:
m_block_size, n_block_size = block_size
if tensors.block_size is None:
sparse_block_size_q, sparse_block_size_kv = q_subtile_factor * m_block_size, n_block_size
else:
sparse_block_size_q, sparse_block_size_kv = tensors.block_size
if sparse_block_size_q != q_subtile_factor * m_block_size:
raise ValueError(
f"Block sparsity expects sparse_block_size_q={q_subtile_factor * m_block_size} "
f"for q_subtile_factor={q_subtile_factor}."
)
if sparse_block_size_kv != n_block_size:
raise ValueError(
f"Block sparsity expects sparse_block_size[1]={n_block_size} to match tile_n."
)
expected_count_shape, expected_index_shape = get_block_sparse_expected_shapes_bwd(
batch_size,
num_head,
seqlen_q,
seqlen_k,
m_block_size,
n_block_size,
q_subtile_factor,
)
normalized_tensors = normalize_block_sparse_tensors(
tensors,
expected_count_shape=expected_count_shape,
expected_index_shape=expected_index_shape,
context="_flash_attn_bwd",
hint=lambda: (
f"Backward expects Q-direction block-sparse tensors (q_mask_cnt/q_mask_idx, "
f"and optionally full_q_cnt/full_q_idx). Regenerate the backward BlockMask with "
f"BLOCK_SIZE=({q_subtile_factor * m_block_size}, {n_block_size})."
),
)
return normalized_tensors, get_block_sparse_broadcast_pattern(normalized_tensors)
def to_cute_block_sparse_tensors(
tensors: BlockSparseTensorsTorch, enable_tvm_ffi: bool = True
) -> BlockSparseTensors | None:
"""Convert torch block sparsity tensors to CuTe tensors, optionally for tvm ffi"""
if not is_block_sparsity_enabled(tensors):
return None
mask_block_cnt_tensor, mask_block_idx_tensor = [
to_cute_tensor(t, assumed_align=4, leading_dim=-1, enable_tvm_ffi=enable_tvm_ffi)
for t in (tensors.mask_block_cnt, tensors.mask_block_idx)
]
full_block_cnt_tensor, full_block_idx_tensor = [
to_cute_tensor(t, assumed_align=4, leading_dim=-1, enable_tvm_ffi=enable_tvm_ffi)
if t is not None
else None
for t in (tensors.full_block_cnt, tensors.full_block_idx)
]
cu_total_m_blocks_tensor, cu_block_idx_offsets_tensor = [
to_cute_tensor(t, assumed_align=4, leading_dim=0, enable_tvm_ffi=enable_tvm_ffi)
if t is not None
else None
for t in (tensors.cu_total_m_blocks, tensors.cu_block_idx_offsets)
]
dq_write_order_tensor, dq_write_order_full_tensor = [
to_cute_tensor(t, assumed_align=4, leading_dim=-1, enable_tvm_ffi=enable_tvm_ffi)
if t is not None
else None
for t in (tensors.dq_write_order, tensors.dq_write_order_full)
]
return BlockSparseTensors(
mask_block_cnt_tensor,
mask_block_idx_tensor,
full_block_cnt_tensor,
full_block_idx_tensor,
cu_total_m_blocks_tensor,
cu_block_idx_offsets_tensor,
dq_write_order_tensor,
dq_write_order_full_tensor,
)
def fast_sampling(mask_mod):
"""Convenience decorator to mark mask_mod as safe for 5-point fast sampling"""
mask_mod.use_fast_sampling = True
return mask_mod
+281
View File
@@ -0,0 +1,281 @@
# Manage Ahead-of-Time (AOT) compiled kernels
import fcntl
import hashlib
import os
import pickle
import sys
import tempfile
import time
from functools import lru_cache
from getpass import getuser
from pathlib import Path
from typing import Hashable, TypeAlias
import ctypes
import cutlass
import cutlass.cute as cute
import tvm_ffi
from cutlass.cutlass_dsl import JitCompiledFunction
from flash_attn.cute.fa_logging import fa_log
# Pre-load cute DSL runtime libraries with RTLD_GLOBAL so that their symbols
# (e.g. _cudaLibraryLoadData) are visible to .so modules loaded later via dlopen.
# Upstream cute.runtime.load_module loads these without RTLD_GLOBAL, which causes
# "undefined symbol" errors when loading cached kernels from disk.
for _lib_path in cute.runtime.find_runtime_libraries(enable_tvm_ffi=False):
if Path(_lib_path).exists():
ctypes.CDLL(_lib_path, mode=ctypes.RTLD_GLOBAL)
CompileKeyType: TypeAlias = tuple[Hashable, ...]
CallableFunction: TypeAlias = JitCompiledFunction | tvm_ffi.Function
# Enable cache via `FLASH_ATTENTION_CUTE_DSL_CACHE_ENABLED=1`
CUTE_DSL_CACHE_ENABLED: bool = os.getenv("FLASH_ATTENTION_CUTE_DSL_CACHE_ENABLED", "0") == "1"
# Customize cache dir via `FLASH_ATTENTION_CUTE_DSL_CACHE_DIR`, default is
# `/tmp/${USER}/flash_attention_cute_dsl_cache``
CUTE_DSL_CACHE_DIR: str | None = os.getenv("FLASH_ATTENTION_CUTE_DSL_CACHE_DIR", None)
def get_cache_path() -> Path:
if CUTE_DSL_CACHE_DIR is not None:
cache_dir = Path(CUTE_DSL_CACHE_DIR)
else:
cache_dir = Path(tempfile.gettempdir()) / getuser() / "flash_attention_cute_dsl_cache"
cache_dir.mkdir(parents=True, exist_ok=True)
return cache_dir
@lru_cache(maxsize=1)
def _compute_source_fingerprint() -> str:
"""
Hash all CuTe Python sources plus runtime ABI stamps into a short fingerprint.
The fingerprint changes whenever:
- Any .py file under flash_attn/cute is added, removed, renamed, or modified.
- The Python minor version changes (e.g. 3.13 -> 3.14).
- The cutlass or tvm_ffi package version changes.
Computed once per process and cached.
"""
cute_root = Path(__file__).resolve().parent
h = hashlib.sha256()
h.update(f"py{sys.version_info.major}.{sys.version_info.minor}".encode())
h.update(f"cutlass={cutlass.__version__}".encode())
h.update(f"tvm_ffi={tvm_ffi.__version__}".encode())
for src in sorted(cute_root.rglob("*.py")):
if not src.is_file():
continue
h.update(src.relative_to(cute_root).as_posix().encode())
content = src.read_bytes()
h.update(len(content).to_bytes(8, "little"))
h.update(content)
return h.hexdigest()
class FileLock:
"""Context manager for advisory file locks using fcntl.flock.
Supports exclusive (write) and shared (read) locks.
Always blocks with polling until the lock is acquired or timeout is reached.
Usage:
with FileLock(lock_path, exclusive=True, timeout=15, label="abc"):
# do work under lock
"""
def __init__(
self,
lock_path: Path,
exclusive: bool,
timeout: float = 15,
label: str = "",
):
"""
Args:
lock_path: Path to the lock file on disk.
exclusive: True for exclusive (write) lock, False for shared (read) lock.
timeout: Max seconds to wait for lock acquisition before raising RuntimeError.
label: Optional human-readable label for error messages.
"""
self.lock_path: Path = lock_path
self.exclusive: bool = exclusive
self.timeout: float = timeout
self.label: str = label
self._fd: int = -1
@property
def _lock_label(self) -> str:
kind = "exclusive" if self.exclusive else "shared"
return f"{kind} {self.label}" if self.label else kind
def __enter__(self) -> "FileLock":
open_flags = os.O_WRONLY | os.O_CREAT if self.exclusive else os.O_RDONLY | os.O_CREAT
lock_type = fcntl.LOCK_EX if self.exclusive else fcntl.LOCK_SH
self._fd = os.open(str(self.lock_path), open_flags)
deadline = time.monotonic() + self.timeout
acquired = False
while time.monotonic() < deadline:
try:
fcntl.flock(self._fd, lock_type | fcntl.LOCK_NB)
acquired = True
break
except OSError:
time.sleep(0.1)
if not acquired:
os.close(self._fd)
self._fd = None
raise RuntimeError(
f"Timed out after {self.timeout}s waiting for "
f"{self._lock_label} lock: {self.lock_path}"
)
return self
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
if self._fd is not None:
fcntl.flock(self._fd, fcntl.LOCK_UN)
os.close(self._fd)
self._fd = None
class JITCache:
"""
In-memory cache for compiled functions.
"""
def __init__(self):
self.cache: dict[CompileKeyType, CallableFunction] = {}
def __setitem__(self, key: CompileKeyType, fn: JitCompiledFunction) -> None:
self.cache[key] = fn
def __getitem__(self, key: CompileKeyType) -> CallableFunction:
return self.cache[key]
def __contains__(self, key: CompileKeyType) -> bool:
return key in self.cache
def clear(self) -> None:
"""
Clear in-memory cache of compiled functions
"""
self.cache.clear()
class JITPersistentCache(JITCache):
"""
In-memory cache for compiled functions, which is also backed by persistent storage.
Use cutedsl ahead-of-time (AOT) compilation, only supporting enable_tvm_ffi=True
"""
EXPORT_FUNCTION_PREFIX = "func"
LOCK_TIMEOUT_SECONDS = 15
def __init__(self, cache_path: Path):
super().__init__()
cache_path.mkdir(parents=True, exist_ok=True)
self.cache_path: Path = cache_path
def __setitem__(self, key: CompileKeyType, fn: JitCompiledFunction) -> None:
JITCache.__setitem__(self, key, fn)
self._try_export_to_storage(key, fn)
def __getitem__(self, key: CompileKeyType) -> CallableFunction:
# Use __contains__ to try populating in-memory cache with persistent storage
self.__contains__(key)
return JITCache.__getitem__(self, key)
def __contains__(self, key: CompileKeyType) -> bool:
# Checks in-memory cache first, then tries loading from storage.
# When returning True, guarantees the in-memory cache is populated.
if JITCache.__contains__(self, key):
return True
return self._try_load_from_storage(key)
def _try_load_from_storage(self, key: CompileKeyType) -> bool:
"""
Try to load a function from persistent storage into in-memory cache.
Returns True if loaded successfully, False if not found on disk.
Holds a shared lock during loading to prevent concurrent writes.
"""
sha256_hex = self._key_to_hash(key)
obj_path = self.cache_path / f"{sha256_hex}.o"
with FileLock(
self._lock_path(sha256_hex),
exclusive=False,
timeout=self.LOCK_TIMEOUT_SECONDS,
label=sha256_hex,
):
if obj_path.exists():
fa_log(1, f"Loading compiled function from disk: {obj_path}")
m = cute.runtime.load_module(str(obj_path), enable_tvm_ffi=True)
fn = getattr(m, self.EXPORT_FUNCTION_PREFIX)
JITCache.__setitem__(self, key, fn)
return True
else:
fa_log(1, f"Cache miss on disk for key hash {sha256_hex}")
return False
def _try_export_to_storage(self, key: CompileKeyType, fn: JitCompiledFunction) -> None:
"""Export a compiled function to persistent storage under exclusive lock."""
sha256_hex = self._key_to_hash(key)
with FileLock(
self._lock_path(sha256_hex),
exclusive=True,
timeout=self.LOCK_TIMEOUT_SECONDS,
label=sha256_hex,
):
obj_path = self.cache_path / f"{sha256_hex}.o"
if obj_path.exists():
# Another process already exported.
fa_log(1, f"Skipping export, already on disk: {obj_path}")
return
fa_log(1, f"Exporting compiled function to disk: {obj_path}")
fn.export_to_c(
object_file_path=str(obj_path),
function_name=self.EXPORT_FUNCTION_PREFIX,
)
fa_log(1, f"Successfully exported compiled function to disk: {obj_path}")
def _key_to_hash(self, key: CompileKeyType) -> str:
return hashlib.sha256(pickle.dumps(key)).hexdigest()
def _lock_path(self, sha256_hex: str) -> Path:
return self.cache_path / f"{sha256_hex}.lock"
def clear(self) -> None:
"""
Not only clear the in-memory cache. Also purge persistent compilation cache.
"""
fa_log(1, f"Clearing persistent cache at {self.cache_path}")
super().clear()
for child in self.cache_path.iterdir():
child.unlink()
def get_jit_cache(name: str | None = None) -> JITCache:
"""
JIT cache factory.
`name` is an optional identifier to create subdirectories to manage cache.
When persistent caching is enabled, artifacts are namespaced under a
source fingerprint directory so that code or dependency changes
automatically invalidate stale entries.
"""
if CUTE_DSL_CACHE_ENABLED:
path = get_cache_path() / _compute_source_fingerprint()
if name:
path = path / name
fa_log(1, f"Creating persistent JIT cache at {path}")
return JITPersistentCache(path)
else:
fa_log(1, "Persistent cache disabled, using in-memory JIT cache")
return JITCache()
@@ -0,0 +1,551 @@
from functools import partial
from typing import Callable, Optional, Tuple
import cutlass
import cutlass.cute as cute
import torch
from cutlass import Boolean, Int8, Int32, const_expr
from flash_attn.cute.block_sparsity import (
BlockSparseTensors,
BlockSparseTensorsTorch,
to_cute_block_sparse_tensors,
)
from flash_attn.cute.block_sparse_utils import get_curr_blocksparse_tensors
from flash_attn.cute.testing import is_fake_mode
from flash_attn.cute.cute_dsl_utils import (
get_aux_tensor_metadata,
to_cute_aux_tensor,
to_cute_tensor,
)
from flash_attn.cute.utils import (
get_batch_from_cu_tensor,
hash_callable,
scalar_to_ssa,
ssa_to_scalar,
)
from flash_attn.cute.mask import call_mask_mod
from flash_attn.cute.seqlen_info import SeqlenInfoQK
from flash_attn.cute.utils import AuxData
class BlockSparsityKernel:
"""Block sparsity kernel for FlexAttention.
This kernel computes `mask_mod` for every token of each block
to determine if an n block is full, masked, or neither.
Writes block counts and indices to a BlockSparseTensors object.
When use_fast_sampling=True, uses 5-point sampling (4 corners + center)
which is much faster but only suitable for masks where this is sufficient.
TODO:
- optimize mask_mod evaluation
- transposed tensors for bwd pass
"""
def __init__(
self,
mask_mod: Callable,
tile_mn: Tuple[int, int],
compute_full_blocks: bool = True,
use_aux_tensors: bool = False,
use_fast_sampling: bool = False,
):
self.mask_mod = mask_mod
self.tile_mn = tile_mn
self.compute_full_blocks = compute_full_blocks
self.use_aux_tensors = use_aux_tensors
self.use_fast_sampling = use_fast_sampling
@cute.jit
def __call__(
self,
blocksparse_tensors: BlockSparseTensors,
seqlen_q: Int32,
seqlen_k: Int32,
mCuSeqlensQ: Optional[cute.Tensor] = None,
mCuSeqlensK: Optional[cute.Tensor] = None,
mSeqUsedQ: Optional[cute.Tensor] = None,
mSeqUsedK: Optional[cute.Tensor] = None,
aux_data: AuxData = AuxData(),
):
mask_cnt, mask_idx, full_cnt, full_idx, mCuTotalMBlocks, mCuBlockIdxOffsets, *_ = (
blocksparse_tensors
)
self.is_varlen_q = const_expr(mCuSeqlensQ is not None)
if const_expr(self.compute_full_blocks):
assert full_cnt is not None and full_idx is not None, (
"full block tensors must be provided when computing full blocks"
)
if const_expr(not self.is_varlen_q):
batch_size, num_heads, num_m_blocks, _ = mask_idx.shape
total_m_blocks = batch_size * num_m_blocks
else:
assert const_expr(mCuTotalMBlocks is not None), (
"mCuTotalMBlocks must be provided when varlen q"
)
num_heads, total_m_blocks = mask_cnt.shape # num_m_blocks is total_m_blocks
batch_size = mCuSeqlensQ.shape[0] - 1
if const_expr(self.use_fast_sampling):
num_threads = 5
self.num_warps = 1
else:
num_threads = self.tile_mn[0]
self.num_warps = (num_threads + 32 - 1) // 32
if const_expr(not self.is_varlen_q):
grid = [num_m_blocks, num_heads, batch_size]
else:
grid = [total_m_blocks, num_heads, 1]
self.kernel(
blocksparse_tensors,
seqlen_q,
seqlen_k,
batch_size,
mCuSeqlensQ,
mCuSeqlensK,
mSeqUsedQ,
mSeqUsedK,
mCuTotalMBlocks,
mCuBlockIdxOffsets,
aux_data,
).launch(grid=grid, block=[num_threads, 1, 1])
@cute.kernel
def kernel(
self,
blocksparse_tensors: BlockSparseTensors,
seqlen_q: Int32,
seqlen_k: Int32,
batch_size: Int32,
mCuSeqlensQ: Optional[cute.Tensor] = None,
mCuSeqlensK: Optional[cute.Tensor] = None,
mSeqUsedQ: Optional[cute.Tensor] = None,
mSeqUsedK: Optional[cute.Tensor] = None,
mCuTotalMBlocks: Optional[cute.Tensor] = None,
mCuBlockIdxOffsets: Optional[cute.Tensor] = None,
aux_data: AuxData = AuxData(),
):
tidx, _, _ = cute.arch.thread_idx()
warp_idx = cute.arch.warp_idx()
lane_id = cute.arch.lane_idx()
ssa = partial(scalar_to_ssa, dtype=Int32)
@cute.struct
class SharedStorage:
reduction_buffer_smem: cute.struct.Align[
cute.struct.MemRange[cutlass.Int8, 2 * self.num_warps], 1024
]
smem = cutlass.utils.SmemAllocator()
storage = smem.allocate(SharedStorage, 16)
reduction_buffer = storage.reduction_buffer_smem.get_tensor(
cute.make_layout((self.num_warps, 2))
)
SeqlenInfoCls = partial(
SeqlenInfoQK.create,
seqlen_q_static=seqlen_q,
seqlen_k_static=seqlen_k,
mCuSeqlensQ=mCuSeqlensQ,
mCuSeqlensK=mCuSeqlensK,
mSeqUsedQ=mSeqUsedQ,
mSeqUsedK=mSeqUsedK,
mCuTotalMBlocks=mCuTotalMBlocks,
mCuBlockIdxOffsets=mCuBlockIdxOffsets,
tile_m=self.tile_mn[0],
tile_n=self.tile_mn[1],
)
if const_expr(not self.is_varlen_q):
m_block, head_idx, batch_idx = cute.arch.block_idx()
else:
global_m_block, head_idx, _ = cute.arch.block_idx()
batch_idx = get_batch_from_cu_tensor(global_m_block, mCuTotalMBlocks)
m_block = global_m_block - mCuTotalMBlocks[batch_idx]
seqlen = SeqlenInfoCls(batch_idx)
seqlen_q = seqlen.seqlen_q
seqlen_k = seqlen.seqlen_k
global_m_block = seqlen.m_block_offset + m_block
num_n_blocks = (seqlen_k + self.tile_mn[1] - 1) // self.tile_mn[1]
_, curr_mask_idx, _, curr_full_idx = get_curr_blocksparse_tensors(
batch_idx, head_idx, m_block, blocksparse_tensors, seqlen
)
num_mask_blocks = Int32(0)
num_full_blocks = Int32(0)
m_base = m_block * self.tile_mn[0]
if const_expr(self.use_fast_sampling):
# Loop-invariant per-thread q_idx for the 5 sample points
# (tidx 0, 1: top corners; 2, 3: bottom corners; 4: center).
q_idx_sample = m_base
if tidx == 2 or tidx == 3:
q_idx_sample = cutlass.min(m_base + self.tile_mn[0] - 1, seqlen_q - 1)
elif tidx == 4:
q_idx_sample = m_base + cutlass.min(seqlen_q - m_base, self.tile_mn[0]) // 2
else:
q_idx_thread = m_base + tidx
thread_in_bounds = Boolean(tidx < self.tile_mn[0] and q_idx_thread < seqlen_q)
for n_block in cutlass.range(num_n_blocks):
n_base = n_block * self.tile_mn[1]
if const_expr(self.use_fast_sampling):
# 5-point sampling (4 corners + center). Interior n_blocks
# (n_base + tile_n <= seqlen_k) skip the OOB clamp on the right /
# center samples.
is_interior = (n_base + self.tile_mn[1]) <= seqlen_k
n_right = Int32(0)
n_mid = Int32(0)
if is_interior:
n_right = n_base + self.tile_mn[1] - 1
n_mid = n_base + self.tile_mn[1] // 2
else:
n_right = cutlass.min(n_base + self.tile_mn[1] - 1, seqlen_k - 1)
n_mid = n_base + cutlass.min(seqlen_k - n_base, self.tile_mn[1]) // 2
kv_idx = n_base
if tidx == 1 or tidx == 3:
kv_idx = n_right
elif tidx == 4:
kv_idx = n_mid
thread_result = Boolean(False)
thread_is_valid = Boolean(False)
if tidx < 5:
thread_is_valid = Boolean(True)
thread_result = ssa_to_scalar(
call_mask_mod(
self.mask_mod,
ssa(batch_idx),
ssa(head_idx),
ssa(q_idx_sample),
ssa(kv_idx),
seqlen,
aux_data,
)
)
has_unmasked = cute.arch.vote_any_sync(thread_result & thread_is_valid)
has_masked = cute.arch.vote_any_sync(Boolean(not thread_result) & thread_is_valid)
else:
# Full path. Interior blocks (n_base + tile_n <= seqlen_k) drop the
# per-element bound check; the boundary block (at most one) keeps it.
thread_has_unmasked = Boolean(False)
thread_has_masked = Boolean(False)
kv_idx = Int32(0)
is_interior = (n_base + self.tile_mn[1]) <= seqlen_k
if is_interior:
if thread_in_bounds:
for c in cutlass.range(self.tile_mn[1], unroll_full=True):
mask_val = ssa_to_scalar(
call_mask_mod(
self.mask_mod,
ssa(batch_idx),
ssa(head_idx),
ssa(q_idx_thread),
ssa(n_base + c),
seqlen,
aux_data,
)
)
thread_has_unmasked |= Boolean(mask_val)
thread_has_masked |= Boolean(not mask_val)
else:
if thread_in_bounds:
for c in cutlass.range(self.tile_mn[1], unroll_full=True):
kv_idx = n_base + c
if kv_idx < seqlen_k:
mask_val = ssa_to_scalar(
call_mask_mod(
self.mask_mod,
ssa(batch_idx),
ssa(head_idx),
ssa(q_idx_thread),
ssa(kv_idx),
seqlen,
aux_data,
)
)
thread_has_unmasked |= Boolean(mask_val)
thread_has_masked |= Boolean(not mask_val)
warp_unmasked = cute.arch.vote_any_sync(thread_has_unmasked & thread_in_bounds)
warp_masked = cute.arch.vote_any_sync(thread_has_masked & thread_in_bounds)
if lane_id == 0:
reduction_buffer[warp_idx, 0] = Int8(1) if warp_unmasked else Int8(0)
reduction_buffer[warp_idx, 1] = Int8(1) if warp_masked else Int8(0)
cute.arch.sync_threads()
# Cross-warp OR via warp 0; thread 0 (lane 0 of warp 0) holds the result.
has_unmasked = Boolean(False)
has_masked = Boolean(False)
if warp_idx == 0:
lane_unmasked = Boolean(False)
lane_masked = Boolean(False)
if lane_id < self.num_warps:
lane_unmasked = reduction_buffer[lane_id, 0] != Int8(0)
lane_masked = reduction_buffer[lane_id, 1] != Int8(0)
has_unmasked = cute.arch.vote_any_sync(lane_unmasked)
has_masked = cute.arch.vote_any_sync(lane_masked)
# Only thread 0 updates the output arrays (common to both paths)
if tidx == 0:
# Block classification based on what we found:
# - If has_masked and has_unmasked: partial block (needs masking)
# - If only has_unmasked: full block (no masking needed)
# - If only has_masked: skip this block entirely
is_partial = Boolean(has_masked and has_unmasked)
is_full = Boolean(has_unmasked and (not has_masked))
if is_partial:
curr_mask_idx[num_mask_blocks] = n_block
num_mask_blocks += 1
elif is_full and const_expr(self.compute_full_blocks):
curr_full_idx[num_full_blocks] = n_block
num_full_blocks += 1
# Only thread 0 writes back the counts
if tidx == 0:
mask_cnt, _, full_cnt, *_ = blocksparse_tensors
if const_expr(self.is_varlen_q):
mask_cnt[head_idx, global_m_block] = num_mask_blocks
if const_expr(self.compute_full_blocks):
full_cnt[head_idx, global_m_block] = num_full_blocks
else:
mask_cnt[batch_idx, head_idx, m_block] = num_mask_blocks
if const_expr(self.compute_full_blocks):
full_cnt[batch_idx, head_idx, m_block] = num_full_blocks
def compute_block_sparsity(
tile_m,
tile_n,
batch_size,
num_heads,
seqlen_q,
seqlen_k,
mask_mod: Callable,
aux_tensors: Optional[list],
device,
aux_scalars: Optional[tuple] = None,
cu_seqlens_q: Optional[torch.Tensor] = None,
cu_seqlens_k: Optional[torch.Tensor] = None,
seqused_q: Optional[torch.Tensor] = None,
seqused_k: Optional[torch.Tensor] = None,
cu_total_m_blocks: Optional[torch.Tensor] = None,
cu_block_idx_offsets: Optional[torch.Tensor] = None,
compute_full_blocks: bool = True,
use_fast_sampling: bool = False,
) -> BlockSparseTensorsTorch:
"""
Computes block sparsity for a given `mask_mod`.
Args:
tile_m: The tile size for the m dimension.
tile_n: The tile size for the n dimension.
batch_size: The batch size.
num_heads: The number of heads.
seqlen_q: The sequence length for the query.
seqlen_k: The sequence length for the key.
mask_mod: The `mask_mod` callable to use.
aux_tensors: A list of auxiliary tensors.
device: The device to use.
cu_seqlens_q: Cumulative q sequence lengths for varlen
cu_seqlens_k: Cumulative k sequence lengths for varlen
seqused_q: Per-batch effective q sequence lengths
seqused_k: Per-batch effective k sequence lengths
cu_total_m_blocks: Cumulative total m blocks tensor for varlen q
cu_block_idx_offsets: Cumulative offsets into the packed mask_block_idx /
full_block_idx tensors per batch (== cumsum of M_b * N_b).
compute_full_blocks: Whether to compute full blocks. If False, only partially-masked blocks are computed.
use_fast_sampling: Whether to use 5-point sampling (4 corners + center). This is much faster, but only suitable for masks where this check is sufficient.
Returns:
BlockSparseTensorsTorch
"""
aux_scalars = tuple(aux_scalars) if aux_scalars else None
# Check if mask_mod is marked as suitable for 5-point sampling
use_fast_sampling = getattr(mask_mod, "use_fast_sampling", use_fast_sampling)
num_m_blocks = (seqlen_q + tile_m - 1) // tile_m
num_n_blocks = (seqlen_k + tile_n - 1) // tile_n
if cu_seqlens_q is not None:
assert cu_total_m_blocks is not None, "total m blocks must be provided when varlen q"
total_m_blocks = cu_total_m_blocks[-1].item()
if cu_block_idx_offsets is None and (cu_seqlens_k is not None or seqused_k is not None):
# Derive cu_block_idx_offsets from per-batch K seqlens.
cu_block_idx_offsets_list = [0]
for batch_idx in range(batch_size):
batch_seqlen_q = cu_seqlens_q[batch_idx + 1].item() - cu_seqlens_q[batch_idx].item()
if cu_seqlens_k is not None:
batch_seqlen_k = (
cu_seqlens_k[batch_idx + 1].item() - cu_seqlens_k[batch_idx].item()
)
else:
batch_seqlen_k = seqused_k[batch_idx].item()
num_m_blocks_batch = (batch_seqlen_q + tile_m - 1) // tile_m
num_n_blocks_batch = (batch_seqlen_k + tile_n - 1) // tile_n
cu_block_idx_offsets_list.append(
cu_block_idx_offsets_list[-1] + num_m_blocks_batch * num_n_blocks_batch
)
cu_block_idx_offsets = torch.tensor(
cu_block_idx_offsets_list, dtype=torch.int32, device=device
)
if cu_block_idx_offsets is not None:
total_n_blocks = cu_block_idx_offsets[-1].item()
else:
# Uniform-K varlen-Q: every batch has the same K seqlen.
total_n_blocks = total_m_blocks * num_n_blocks
mask_block_cnt = torch.zeros((num_heads, total_m_blocks), device=device, dtype=torch.int32)
mask_block_idx = torch.zeros((num_heads, total_n_blocks), device=device, dtype=torch.int32)
full_block_cnt = (
torch.zeros((num_heads, total_m_blocks), device=device, dtype=torch.int32)
if compute_full_blocks
else None
)
full_block_idx = (
torch.zeros((num_heads, total_n_blocks), device=device, dtype=torch.int32)
if compute_full_blocks
else None
)
else:
total_m_blocks = batch_size * num_m_blocks
total_n_blocks = batch_size * num_m_blocks * num_n_blocks
mask_block_cnt = torch.zeros(
(batch_size, num_heads, num_m_blocks), device=device, dtype=torch.int32
)
mask_block_idx = torch.zeros(
(batch_size, num_heads, num_m_blocks, num_n_blocks), device=device, dtype=torch.int32
)
full_block_cnt = (
torch.zeros((batch_size, num_heads, num_m_blocks), device=device, dtype=torch.int32)
if compute_full_blocks
else None
)
full_block_idx = (
torch.zeros(
(batch_size, num_heads, num_m_blocks, num_n_blocks),
device=device,
dtype=torch.int32,
)
if compute_full_blocks
else None
)
blocksparse_tensors_torch = BlockSparseTensorsTorch(
mask_block_cnt=mask_block_cnt,
mask_block_idx=mask_block_idx,
full_block_cnt=full_block_cnt,
full_block_idx=full_block_idx,
cu_total_m_blocks=cu_total_m_blocks,
cu_block_idx_offsets=cu_block_idx_offsets,
block_size=(tile_m, tile_n),
)
mask_mod_hash = hash_callable(mask_mod)
if aux_tensors is not None:
aux_tensor_metadata = get_aux_tensor_metadata(aux_tensors)
else:
aux_tensor_metadata = None
aux_scalar_metadata = tuple(type(s) for s in aux_scalars) if aux_scalars is not None else None
compile_key = (
tile_m,
tile_n,
mask_mod_hash,
aux_tensor_metadata,
aux_scalar_metadata,
compute_full_blocks,
cu_seqlens_q is None,
cu_seqlens_k is None,
seqused_q is None,
seqused_k is None,
aux_tensors is not None,
use_fast_sampling,
)
if compile_key not in compute_block_sparsity.compile_cache:
(
cu_seqlens_q_tensor,
cu_seqlens_k_tensor,
seqused_q_tensor,
seqused_k_tensor,
) = [
to_cute_tensor(t, assumed_align=4, leading_dim=0) if t is not None else None
for t in (
cu_seqlens_q,
cu_seqlens_k,
seqused_q,
seqused_k,
)
]
blocksparse_tensors = to_cute_block_sparse_tensors(
blocksparse_tensors_torch, enable_tvm_ffi=True
)
if aux_tensors is not None:
cute_aux_tensors = [to_cute_aux_tensor(buf) for buf in aux_tensors]
else:
cute_aux_tensors = None
kernel = BlockSparsityKernel(
mask_mod,
tile_mn=(tile_m, tile_n),
compute_full_blocks=compute_full_blocks,
use_aux_tensors=aux_tensors is not None,
use_fast_sampling=use_fast_sampling,
)
compute_block_sparsity.compile_cache[compile_key] = cute.compile(
kernel,
blocksparse_tensors,
seqlen_q,
seqlen_k,
cu_seqlens_q_tensor,
cu_seqlens_k_tensor,
seqused_q_tensor,
seqused_k_tensor,
AuxData(cute_aux_tensors, aux_scalars),
options="--enable-tvm-ffi",
)
if not is_fake_mode():
compute_block_sparsity.compile_cache[compile_key](
(
blocksparse_tensors_torch.mask_block_cnt,
blocksparse_tensors_torch.mask_block_idx,
blocksparse_tensors_torch.full_block_cnt,
blocksparse_tensors_torch.full_block_idx,
blocksparse_tensors_torch.cu_total_m_blocks,
blocksparse_tensors_torch.cu_block_idx_offsets,
blocksparse_tensors_torch.dq_write_order,
blocksparse_tensors_torch.dq_write_order_full,
),
seqlen_q,
seqlen_k,
cu_seqlens_q,
cu_seqlens_k,
seqused_q,
seqused_k,
AuxData(aux_tensors, aux_scalars),
)
return blocksparse_tensors_torch
compute_block_sparsity.compile_cache = {}
+372
View File
@@ -0,0 +1,372 @@
# Copyright (c) 2025, Wentao Guo, Ted Zadouri, Tri Dao.
import math
from typing import Optional, Type, Callable
import cutlass
import cutlass.cute as cute
from cutlass import Float32, Int32, const_expr
from cutlass.cute.nvgpu import cpasync
import cutlass.utils.blackwell_helpers as sm100_utils
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import llvm
import cutlass.pipeline
@dsl_user_op
def cvt_copy(
atom: cute.CopyAtom,
src: cute.Tensor,
dst: cute.Tensor,
*,
pred: Optional[cute.Tensor] = None,
loc=None,
ip=None,
**kwargs,
) -> None:
assert isinstance(src.iterator, cute.Pointer) and src.memspace == cute.AddressSpace.rmem
if const_expr(src.element_type != dst.element_type):
src_cvt = cute.make_fragment_like(src, dst.element_type, loc=loc, ip=ip)
src_cvt.store(src.load().to(dst.element_type))
src = src_cvt
cute.copy(atom, src, dst, pred=pred, loc=loc, ip=ip, **kwargs)
@dsl_user_op
def load_s2r(src: cute.Tensor, *, loc=None, ip=None) -> cute.Tensor:
dst = cute.make_fragment_like(src, src.element_type, loc=loc, ip=ip)
cute.autovec_copy(src, dst, loc=loc, ip=ip)
return dst
@dsl_user_op
def get_copy_atom(
dtype: Type[cutlass.Numeric], num_copy_elems: int, is_async: bool = False, *, loc=None, ip=None
) -> cute.CopyAtom:
num_copy_bits = const_expr(min(128, num_copy_elems * dtype.width))
copy_op = cpasync.CopyG2SOp() if is_async else cute.nvgpu.CopyUniversalOp()
return cute.make_copy_atom(copy_op, dtype, num_bits_per_copy=num_copy_bits)
@dsl_user_op
def make_tmem_copy(
tmem_copy_atom: cute.CopyAtom, num_wg: int = 1, *, loc=None, ip=None
) -> cute.CopyAtom:
num_dp, num_bits, num_rep, _ = sm100_utils.get_tmem_copy_properties(tmem_copy_atom)
assert num_dp == 32
assert num_bits == 32
tiler_mn = (cute.make_layout((128 * num_rep * num_wg // 32, 32), stride=(32, 1)),)
layout_tv = cute.make_layout(
((32, 4, num_wg), (num_rep, 32)), stride=((0, 1, 4 * num_rep), (4, 4 * num_rep * num_wg))
)
return cute.make_tiled_copy(tmem_copy_atom, layout_tv, tiler_mn)
@dsl_user_op
def copy(
src: cute.Tensor,
dst: cute.Tensor,
*,
pred: Optional[cute.Tensor] = None,
num_copy_elems: int = 1,
is_async: bool = False,
loc=None,
ip=None,
**kwargs,
) -> None:
copy_atom = get_copy_atom(src.element_type, num_copy_elems, is_async)
cute.copy(copy_atom, src, dst, pred=pred, loc=loc, ip=ip, **kwargs)
def tiled_copy_1d(
dtype: Type[cutlass.Numeric], num_threads: int, num_copy_elems: int = 1, is_async: bool = False
) -> cute.TiledCopy:
num_copy_bits = num_copy_elems * dtype.width
copy_op = cpasync.CopyG2SOp() if is_async else cute.nvgpu.CopyUniversalOp()
copy_atom = cute.make_copy_atom(copy_op, dtype, num_bits_per_copy=num_copy_bits)
thr_layout = cute.make_layout(num_threads)
val_layout = cute.make_layout(num_copy_elems)
return cute.make_tiled_copy_tv(copy_atom, thr_layout, val_layout)
def tiled_copy_2d(
dtype: Type[cutlass.Numeric], major_mode_size: int, num_threads: int, is_async: bool = False
) -> cute.TiledCopy:
num_copy_bits = math.gcd(major_mode_size, 128 // dtype.width) * dtype.width
copy_elems = num_copy_bits // dtype.width
copy_op = cpasync.CopyG2SOp() if is_async else cute.nvgpu.CopyUniversalOp()
copy_atom = cute.make_copy_atom(copy_op, dtype, num_bits_per_copy=num_copy_bits)
gmem_threads_per_row = major_mode_size // copy_elems
assert num_threads % gmem_threads_per_row == 0
thr_layout = cute.make_ordered_layout(
(num_threads // gmem_threads_per_row, gmem_threads_per_row),
order=(1, 0),
)
val_layout = cute.make_layout((1, copy_elems))
return cute.make_tiled_copy_tv(copy_atom, thr_layout, val_layout)
@dsl_user_op
def atomic_add_fp32x4(
a: Float32, b: Float32, c: Float32, d: Float32, gmem_ptr: cute.Pointer, *, loc=None, ip=None
) -> None:
gmem_ptr_i64 = gmem_ptr.toint(loc=loc, ip=ip).ir_value()
# cache_hint = cutlass.Int64(0x12F0000000000000)
llvm.inline_asm(
None,
[
gmem_ptr_i64,
Float32(a).ir_value(loc=loc, ip=ip),
Float32(b).ir_value(loc=loc, ip=ip),
Float32(c).ir_value(loc=loc, ip=ip),
Float32(d).ir_value(loc=loc, ip=ip),
],
# [gmem_ptr_i64, Float32(a).ir_value(loc=loc, ip=ip), cache_hint.ir_value()],
"{\n\t"
# ".reg .b128 abcd;\n\t"
# "mov.b128 abcd, {$1, $2, $3, $4};\n\t"
".reg .v4 .f32 abcd;\n\t"
# "mov.b128 abcd, {$1, $2, $3, $4};\n\t"
"mov.f32 abcd.x, $1;\n\t"
"mov.f32 abcd.y, $2;\n\t"
"mov.f32 abcd.z, $3;\n\t"
"mov.f32 abcd.w, $4;\n\t"
"red.global.add.v4.f32 [$0], abcd;\n\t"
# "red.global.add.L2::cache_hint.v4.f32 [$0], abcd, 0x14F0000000000000;\n\t"
"}\n",
# "red.global.add.L2::cache_hint.f32 [$0], $1, 0x12F0000000000000;",
# "red.global.add.L2::cache_hint.f32 [$0], $1, $2;",
"l,f,f,f,f",
# "l,f,l",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
@dsl_user_op
def set_block_rank(
smem_ptr: cute.Pointer, peer_cta_rank_in_cluster: Int32, *, loc=None, ip=None
) -> Int32:
"""Map the given smem pointer to the address at another CTA rank in the cluster."""
smem_ptr_i32 = smem_ptr.toint(loc=loc, ip=ip).ir_value()
return Int32(
llvm.inline_asm(
T.i32(),
[smem_ptr_i32, peer_cta_rank_in_cluster.ir_value()],
"mapa.shared::cluster.u32 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def store_shared_remote_fp32x4(
a: Float32,
b: Float32,
c: Float32,
d: Float32,
smem_ptr: cute.Pointer,
mbar_ptr: cute.Pointer,
peer_cta_rank_in_cluster: Int32,
*,
loc=None,
ip=None,
) -> None:
remote_smem_ptr_i32 = set_block_rank(
smem_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip
).ir_value()
remote_mbar_ptr_i32 = set_block_rank(
mbar_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip
).ir_value()
llvm.inline_asm(
None,
[
remote_smem_ptr_i32,
remote_mbar_ptr_i32,
Float32(a).ir_value(loc=loc, ip=ip),
Float32(b).ir_value(loc=loc, ip=ip),
Float32(c).ir_value(loc=loc, ip=ip),
Float32(d).ir_value(loc=loc, ip=ip),
],
"{\n\t"
".reg .v4 .f32 abcd;\n\t"
"mov.f32 abcd.x, $2;\n\t"
"mov.f32 abcd.y, $3;\n\t"
"mov.f32 abcd.z, $4;\n\t"
"mov.f32 abcd.w, $5;\n\t"
"st.async.shared::cluster.mbarrier::complete_tx::bytes.v4.f32 [$0], abcd, [$1];\n\t"
"}\n",
"r,r,f,f,f,f",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
@dsl_user_op
def cpasync_bulk_s2cluster(
smem_src_ptr: cute.Pointer,
smem_dst_ptr: cute.Pointer,
mbar_ptr: cute.Pointer,
size: int | Int32,
peer_cta_rank_in_cluster: Int32,
*,
loc=None,
ip=None,
):
smem_src_ptr_i32 = smem_src_ptr.toint(loc=loc, ip=ip).ir_value()
smem_dst_ptr_i32 = set_block_rank(
smem_dst_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip
).ir_value()
mbar_ptr_i32 = set_block_rank(mbar_ptr, peer_cta_rank_in_cluster, loc=loc, ip=ip).ir_value()
llvm.inline_asm(
None,
[
smem_dst_ptr_i32,
smem_src_ptr_i32,
mbar_ptr_i32,
Int32(size).ir_value(loc=loc, ip=ip),
],
"cp.async.bulk.shared::cluster.shared::cta.mbarrier::complete_tx::bytes [$0], [$1], $3, [$2];",
"r,r,r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
@dsl_user_op
def cpasync_bulk_g2s(
gmem_ptr: cute.Pointer,
smem_ptr: cute.Pointer,
tma_bar_ptr: cute.Pointer,
size: int | Int32,
*,
loc=None,
ip=None,
):
gmem_ptr_i64 = gmem_ptr.toint(loc=loc, ip=ip).ir_value()
smem_ptr_i32 = smem_ptr.toint(loc=loc, ip=ip).ir_value()
mbar_ptr_i32 = tma_bar_ptr.toint(loc=loc, ip=ip).ir_value()
llvm.inline_asm(
None,
[gmem_ptr_i64, smem_ptr_i32, mbar_ptr_i32, Int32(size).ir_value()],
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [$1], [$0], $3, [$2];",
"l,r,r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
@dsl_user_op
def cpasync_reduce_bulk_add_f32(
smem_ptr: cute.Pointer,
gmem_ptr: cute.Pointer,
store_bytes: int | Int32,
*,
loc=None,
ip=None,
):
smem_ptr_i32 = smem_ptr.toint(loc=loc, ip=ip).ir_value()
# cache_hint = cutlass.Int64(0x14F0000000000000) # EVICT_LAST
llvm.inline_asm(
None,
[gmem_ptr.llvm_ptr, smem_ptr_i32, Int32(store_bytes).ir_value()],
"cp.reduce.async.bulk.global.shared::cta.bulk_group.add.f32 [$0], [$1], $2;",
"l,r,r",
# [gmem_ptr.llvm_ptr, smem_ptr_i32, Int32(store_bytes).ir_value(), cache_hint.ir_value()],
# "cp.reduce.async.bulk.global.shared::cta.bulk_group.L2::cache_hint.add.f32 [$0], [$1], $2, $3;",
# "l,r,r,l",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
def cpasync_bulk_get_copy_fn(
src_tensor: cute.Tensor,
dst_tensor: cute.Tensor,
single_stage: bool = False,
**kwargs,
) -> Callable:
# src_is_smem = const_expr(
# isinstance(src_tensor.iterator, cute.Pointer)
# and src_tensor.memspace == cute.AddressSpace.smem
# )
group_rank_src = const_expr(cute.rank(src_tensor) - (1 if not single_stage else 0))
group_rank_dst = const_expr(cute.rank(dst_tensor) - (1 if not single_stage else 0))
# ((atom_v, rest_v), STAGE), ((atom_v, rest_v), RestK)
src = cute.group_modes(src_tensor, 0, group_rank_src)
dst = cute.group_modes(dst_tensor, 0, group_rank_dst)
def copy_bulk(src_idx, dst_idx, **new_kwargs):
size = const_expr(cute.size(src.shape[:-1]) * src.element_type.width // 8)
cpasync_bulk_g2s(
src[None, src_idx].iterator,
dst[None, dst_idx].iterator,
size=size,
**new_kwargs,
**kwargs,
)
def copy_bulk_single_stage(**new_kwargs):
size = const_expr(cute.size(src.shape) * src.element_type.width // 8)
cpasync_bulk_g2s(src.iterator, dst.iterator, size=size, **new_kwargs, **kwargs)
return copy_bulk if const_expr(not single_stage) else copy_bulk_single_stage
def tma_get_copy_fn(
atom: cute.CopyAtom,
cta_coord: cute.Coord,
cta_layout: cute.Layout,
src_tensor: cute.Tensor,
dst_tensor: cute.Tensor,
filter_zeros: bool = False,
single_stage: bool = False,
**kwargs,
) -> Callable:
src_is_smem = const_expr(
isinstance(src_tensor.iterator, cute.Pointer)
and src_tensor.memspace == cute.AddressSpace.smem
)
smem_tensor, gmem_tensor = (src_tensor, dst_tensor) if src_is_smem else (dst_tensor, src_tensor)
group_rank_smem = const_expr(cute.rank(smem_tensor) - (1 if not single_stage else 0))
group_rank_gmem = const_expr(cute.rank(gmem_tensor) - (1 if not single_stage else 0))
# ((atom_v, rest_v), STAGE), ((atom_v, rest_v), RestK)
s, g = cpasync.tma_partition(
atom,
cta_coord,
cta_layout,
cute.group_modes(smem_tensor, 0, group_rank_smem),
cute.group_modes(gmem_tensor, 0, group_rank_gmem),
)
if const_expr(filter_zeros):
s = cute.filter_zeros(s)
g = cute.filter_zeros(g)
src, dst = (s, g) if src_is_smem else (g, s)
def copy_tma(src_idx, dst_idx, **new_kwargs):
cute.copy(atom, src[None, src_idx], dst[None, dst_idx], **new_kwargs, **kwargs)
def copy_tma_single_stage(**new_kwargs):
cute.copy(atom, src, dst, **new_kwargs, **kwargs)
return (copy_tma if const_expr(not single_stage) else copy_tma_single_stage), s, g
def tma_producer_copy_fn(copy: Callable, pipeline: cutlass.pipeline.PipelineAsync):
def copy_fn(src_idx, producer_state: cutlass.pipeline.PipelineState, **new_kwargs):
copy(
src_idx=src_idx,
dst_idx=producer_state.index,
tma_bar_ptr=pipeline.producer_get_barrier(producer_state),
**new_kwargs,
)
return copy_fn
+151
View File
@@ -0,0 +1,151 @@
"""
System ptxas replacement for CUTLASS DSL.
Environment variables:
CUTE_DSL_PTXAS_PATH - Path to ptxas (e.g., /usr/local/cuda/bin/ptxas)
CUTE_DSL_PTXAS_VERBOSE - Set to 1 for verbose output
"""
import os
import sys
import re
import ctypes
import subprocess
from pathlib import Path
import cutlass
CUTE_DSL_PTXAS_PATH = os.environ.get("CUTE_DSL_PTXAS_PATH", None)
VERBOSE = os.environ.get("CUTE_DSL_PTXAS_VERBOSE", "0") == "1"
_original_load_cuda_library = None
_user_wanted_ptx = False # True if user originally set CUTE_DSL_KEEP_PTX=1
def _log(msg):
if VERBOSE:
print(f"[ptxas] {msg}", file=sys.stderr)
def _get_ptx(compiled_func) -> tuple[str, Path] | None:
"""Find and read PTX file, stripping null bytes."""
func_name = getattr(compiled_func, "function_name", None)
if not func_name:
return None
dump_dir = os.environ.get("CUTE_DSL_DUMP_DIR", Path.cwd())
for ptx_path in Path(dump_dir).glob(f"*{func_name}*.ptx"):
content = ptx_path.read_text().rstrip("\x00")
if ".entry " in content and content.rstrip().endswith("}"):
_log(f"Found PTX: {ptx_path}")
return content, ptx_path
return None
def _compile_ptx(ptx_path: Path, ptx_content: str) -> bytes:
"""Compile PTX to cubin using system ptxas."""
# Extract arch from PTX
match = re.search(r"\.target\s+(sm_\d+[a-z]?)", ptx_content)
arch = match.group(1) if match else "sm_90a"
# Write stripped content back if needed
if ptx_path.read_text() != ptx_content:
ptx_path.write_text(ptx_content)
# Compile
cubin_tmp = ptx_path.with_suffix(".cubin.tmp")
try:
assert CUTE_DSL_PTXAS_PATH is not None
result = subprocess.run(
[CUTE_DSL_PTXAS_PATH, f"-arch={arch}", "-O3", "-o", str(cubin_tmp), str(ptx_path)],
capture_output=True,
text=True,
)
if result.returncode != 0:
raise RuntimeError(f"ptxas failed: {result.stderr}")
cubin_data = cubin_tmp.read_bytes()
_log(f"Compiled {ptx_path.name} -> {len(cubin_data)} bytes ({arch})")
# Save cubin if CUTE_DSL_KEEP_CUBIN is set
if os.environ.get("CUTE_DSL_KEEP_CUBIN", "0") == "1":
cubin_out = ptx_path.with_suffix(".cubin")
cubin_out.write_bytes(cubin_data)
_log(f"Saved: {cubin_out}")
return cubin_data
finally:
cubin_tmp.unlink(missing_ok=True)
def _patched_load_cuda_library(self):
"""Replacement for _load_cuda_library that uses system ptxas."""
result = _get_ptx(self)
if not result:
_log("PTX not found, falling back to embedded ptxas")
return _original_load_cuda_library(self)
ptx_content, ptx_path = result
try:
cubin = _compile_ptx(ptx_path, ptx_content)
except Exception as e:
_log(f"Compilation failed ({e}), falling back to embedded ptxas")
return _original_load_cuda_library(self)
# Load cubin
import cuda.bindings.runtime as cuda_runtime
err, library = cuda_runtime.cudaLibraryLoadData(cubin, None, None, 0, None, None, 0)
if err != cuda_runtime.cudaError_t.cudaSuccess:
_log(f"cudaLibraryLoadData failed ({err}), falling back to embedded ptxas")
return _original_load_cuda_library(self)
# Register kernels on all devices
_, cuda_load_to_device = self._get_cuda_init_and_load()
lib_ptr = ctypes.c_void_p(int(library))
dev_id = ctypes.c_int32(0)
err_val = ctypes.c_int32(0)
args = (ctypes.c_void_p * 3)(
ctypes.cast(ctypes.pointer(lib_ptr), ctypes.c_void_p),
ctypes.cast(ctypes.pointer(dev_id), ctypes.c_void_p),
ctypes.cast(ctypes.pointer(err_val), ctypes.c_void_p),
)
for dev in range(self.num_devices):
dev_id.value = dev
cuda_load_to_device(args)
if err_val.value != 0:
_log("cuda_load_to_device failed, falling back to embedded ptxas")
return _original_load_cuda_library(self)
_log(f"Loaded kernel from {ptx_path.name}")
# Delete PTX if user didn't originally want it kept
if not _user_wanted_ptx:
ptx_path.unlink(missing_ok=True)
return [cuda_runtime.cudaLibrary_t(lib_ptr.value)]
def patch():
"""Install system ptxas hook. Call before importing cutlass."""
global _original_load_cuda_library, _user_wanted_ptx
assert CUTE_DSL_PTXAS_PATH is not None
if not os.path.isfile(CUTE_DSL_PTXAS_PATH) or not os.access(CUTE_DSL_PTXAS_PATH, os.X_OK):
raise RuntimeError(f"ptxas not found: {CUTE_DSL_PTXAS_PATH}")
# Track if user originally wanted PTX kept
_user_wanted_ptx = os.environ.get("CUTE_DSL_KEEP_PTX", "0") == "1"
# os.environ['CUTE_DSL_KEEP_PTX'] = '1'
assert os.environ.get("CUTE_DSL_KEEP_PTX", "0") == "1", (
"Require CUTE_DSL_KEEP_PTX=1 to use system's ptxas"
)
cls = cutlass.cutlass_dsl.cuda_jit_executor.CudaDialectJitCompiledFunction
_original_load_cuda_library = cls._load_cuda_library
cls._load_cuda_library = _patched_load_cuda_library
_log("Patch applied")
return
+158
View File
@@ -0,0 +1,158 @@
# Copyright (c) 2025, Tri Dao.
from typing import Tuple
from functools import lru_cache
import torch
try:
from triton.tools.disasm import extract
except ImportError:
extract = None
import cutlass
import cutlass.cute as cute
from cutlass.cutlass_dsl import NumericMeta
from cutlass.cute.runtime import from_dlpack
StaticTypes = (cutlass.Constexpr, NumericMeta, int, bool, str, float, type(None))
load_cubin_module_data_og = cutlass.base_dsl.runtime.cuda.load_cubin_module_data
cute_compile_og = cute.compile
torch2cute_dtype_map = {
torch.float16: cutlass.Float16,
torch.bfloat16: cutlass.BFloat16,
torch.float32: cutlass.Float32,
torch.float8_e4m3fn: cutlass.Float8E4M3FN,
torch.float8_e5m2: cutlass.Float8E5M2,
}
@lru_cache
def get_max_active_clusters(cluster_size):
return cutlass.utils.HardwareInfo().get_max_active_clusters(cluster_size=cluster_size)
@lru_cache
def get_device_capacity(device: torch.device = None) -> Tuple[int, int]:
return torch.cuda.get_device_capability(device)
def assume_strides_aligned(t):
"""Assume all strides except the last are divisible by 128 bits.
Python int strides (e.g., stride=0 from GQA expand) are kept as-is
since they're static and don't need alignment assumptions.
"""
divby = 128 // t.element_type.width
strides = tuple(s if isinstance(s, int) else cute.assume(s, divby=divby) for s in t.stride[:-1])
return (*strides, t.stride[-1])
def assume_tensor_aligned(t):
"""Rebuild a tensor with 128-bit aligned stride assumptions. Passes through None."""
if t is None:
return None
return cute.make_tensor(t.iterator, cute.make_layout(t.shape, stride=assume_strides_aligned(t)))
def to_cute_tensor(t, assumed_align=16, leading_dim=-1, fully_dynamic=False, enable_tvm_ffi=True):
"""Convert torch tensor to cute tensor for TVM FFI. leading_dim=-1 defaults to t.ndim-1."""
if t is None:
return None
# NOTE: torch 2.9.1 doesn't support fp8 via DLPack but 2.11.0 nightly does
# currently export raw bytes as uint8 and tell cutlass correct type
# can directly export as fp8 when torch supports it
if t.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
tensor = from_dlpack(
t.view(torch.uint8).detach(),
assumed_align=assumed_align,
enable_tvm_ffi=enable_tvm_ffi,
)
tensor.element_type = (
cutlass.Float8E4M3FN if t.dtype == torch.float8_e4m3fn else cutlass.Float8E5M2
)
else:
tensor = from_dlpack(t.detach(), assumed_align=assumed_align, enable_tvm_ffi=enable_tvm_ffi)
if fully_dynamic:
return tensor.mark_layout_dynamic()
if leading_dim == -1:
leading_dim = t.ndim - 1
return tensor.mark_layout_dynamic(leading_dim=leading_dim)
def to_cute_aux_tensor(t, enable_tvm_ffi=True):
"""Convert torch tensor to cute tensor for TVM FFI, tailored to FlexAttention aux tensors.
This allows the user to specify alignment and leading dimension for aux tensors used in
custom score_mod callables.
"""
assumed_align: int = getattr(t, "__assumed_align__", None)
leading_dim: int = getattr(t, "__leading_dim__", None)
fully_dynamic: bool = leading_dim is None
return to_cute_tensor(
t,
assumed_align=assumed_align,
leading_dim=leading_dim,
fully_dynamic=fully_dynamic,
enable_tvm_ffi=enable_tvm_ffi,
)
def get_aux_tensor_metadata(aux_tensors):
return tuple(
(
getattr(t, "__assumed_align__", 0),
getattr(t, "__leading_dim__", -1),
hasattr(t, "__leading_dim__"),
)
for t in aux_tensors
)
def get_broadcast_dims(tensor: torch.Tensor) -> Tuple[bool, ...]:
"""Return tuple of bools indicating which dims have stride=0 (broadcast).
This is useful for compile keys since CuTe's mark_layout_dynamic() keeps
stride=0 as static, meaning kernels compiled with different broadcast
patterns are not interchangeable.
"""
return tuple(s == 0 for s in tensor.stride())
# credit: monellz (https://github.com/NVIDIA/cutlass/issues/2658#issuecomment-3630564264)
def dump_kernel_attributes(compiled_kernel):
from cuda.bindings import driver
from cutlass.utils import HardwareInfo
import torch
device_id = torch.cuda.current_device()
hardware_info = HardwareInfo(device_id=device_id)
cubin_data = compiled_kernel.artifacts.CUBIN
assert cubin_data is not None, "cubin_data is None, need '--keep-cubin' option when compiling"
cuda_library = hardware_info._checkCudaErrors(
driver.cuLibraryLoadData(cubin_data, None, None, 0, None, None, 0)
)
kernels = hardware_info._checkCudaErrors(driver.cuLibraryEnumerateKernels(1, cuda_library))
kernel = hardware_info._checkCudaErrors(driver.cuKernelGetFunction(kernels[0]))
# more metrics: https://docs.nvidia.com/cuda/cuda-driver-api/group__CUDA__EXEC.html#group__CUDA__EXEC_1g5e92a1b0d8d1b82cb00dcfb2de15961b
local_size_bytes = hardware_info._checkCudaErrors(
driver.cuFuncGetAttribute(
driver.CUfunction_attribute.CU_FUNC_ATTRIBUTE_LOCAL_SIZE_BYTES,
kernel,
)
)
num_regs = hardware_info._checkCudaErrors(
driver.cuFuncGetAttribute(
driver.CUfunction_attribute.CU_FUNC_ATTRIBUTE_NUM_REGS,
kernel,
)
)
print("--- Kernel Info ---")
print(f"local_size_bytes: {local_size_bytes}")
print(f"num_regs: {num_regs}")
print("--- End Kernel Info ---")
+97
View File
@@ -0,0 +1,97 @@
# Copyright (c) 2025, Tri Dao.
"""Unified FlashAttention logging controlled by a single ``FA_LOG_LEVEL`` env var.
Host-side messages go through Python ``logging`` (logger name ``flash_attn``).
A default ``StreamHandler`` is attached automatically when ``FA_LOG_LEVEL >= 1``
so that standalone scripts get output without extra setup; applications that
configure their own logging can remove or replace it via the standard API.
FA_LOG_LEVEL mapping::
0 off nothing logged
1 host host-side summaries only (no kernel printf)
2 kernel host + curated kernel traces
3 max host + all kernel traces (noisy, perf hit)
Set via environment variable::
FA_LOG_LEVEL=1 python train.py
Device-side ``cute.printf`` calls are compile-time eliminated via
``cutlass.const_expr`` when the log level is below the callsite threshold,
so there is zero performance cost when device logging is off.
Changing the log level after kernel compilation requires a recompile
(the level participates in the forward compile key).
"""
import logging
import os
import sys
import cutlass.cute as cute
from cutlass import const_expr
_LOG_LEVEL_NAMES = {"off": 0, "host": 1, "kernel": 2, "max": 3}
def _parse_log_level(raw: str) -> int:
if raw in _LOG_LEVEL_NAMES:
return _LOG_LEVEL_NAMES[raw]
try:
level = int(raw)
except ValueError:
return 0
return max(0, min(level, 3))
_fa_log_level: int = _parse_log_level(os.environ.get("FA_LOG_LEVEL", "0"))
_logger = logging.getLogger("flash_attn")
_logger.addHandler(logging.NullHandler())
_default_handler: logging.Handler | None = None
def _configure_default_handler() -> None:
global _default_handler
if _fa_log_level >= 1:
if _default_handler is None:
_default_handler = logging.StreamHandler(sys.stdout)
_default_handler.setFormatter(logging.Formatter("[FA] %(message)s"))
_logger.addHandler(_default_handler)
_logger.setLevel(logging.DEBUG)
else:
if _default_handler is not None:
_logger.removeHandler(_default_handler)
_default_handler = None
_logger.setLevel(logging.WARNING)
_configure_default_handler()
def get_fa_log_level() -> int:
return _fa_log_level
def set_fa_log_level(level: int | str) -> None:
"""Set the FA log level programmatically.
Host logging takes effect immediately. Device logging changes only
affect kernels compiled after this call (new compile-key selection).
"""
global _fa_log_level
if isinstance(level, str):
level = _parse_log_level(level)
_fa_log_level = max(0, min(int(level), 3))
_configure_default_handler()
def fa_log(level: int, msg: str):
if _fa_log_level >= level:
_logger.info(msg)
def fa_printf(level: int, fmt, *args):
if const_expr(_fa_log_level >= level):
cute.printf(fmt, *args)
+21
View File
@@ -0,0 +1,21 @@
# Copyright (c) 2025, Tri Dao.
import cutlass
import cutlass.cute as cute
from cutlass import Int32
@cute.jit
def clz(x: Int32) -> Int32:
# for i in cutlass.range_constexpr(32):
# if (1 << (31 - i)) & x:
# return Int32(i)
# return Int32(32)
# Early exit is not supported yet
res = Int32(32)
done = False
for i in cutlass.range(32):
if ((1 << (31 - i)) & x) and not done:
res = Int32(i)
done = True
return res
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,587 @@
# Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
# A reimplementation of https://github.com/Dao-AILab/flash-attention/blob/main/hopper/flash_bwd_postprocess_kernel.h
# from Cutlass C++ to Cute-DSL.
import math
from typing import Callable, Optional, Type
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
import cutlass.utils.hopper_helpers as sm90_utils_basic
import cutlass.utils.blackwell_helpers as sm100_utils_basic
from cutlass.cute.nvgpu import cpasync, warp, warpgroup
from cutlass import Float32, const_expr
from cutlass.utils import LayoutEnum
from quack import copy_utils
from quack import layout_utils
from quack import sm90_utils
from flash_attn.cute import utils
from flash_attn.cute.cute_dsl_utils import assume_tensor_aligned
from flash_attn.cute import ampere_helpers as sm80_utils
from flash_attn.cute.seqlen_info import SeqlenInfoQK
import cutlass.cute.nvgpu.tcgen05 as tcgen05
from quack.cute_dsl_utils import ParamsBase
from flash_attn.cute.tile_scheduler import (
SingleTileScheduler,
SingleTileVarlenScheduler,
TileSchedulerArguments,
)
class FlashAttentionBackwardPostprocess:
def __init__(
self,
dtype: Type[cutlass.Numeric],
head_dim: int,
arch: int,
tile_m: int = 128,
num_threads: int = 256,
AtomLayoutMdQ: int = 1,
dQ_swapAB: bool = False,
use_2cta_instrs: bool = False,
cluster_size: int = 1, # for varlen offsets
):
"""
:param head_dim: head dimension
:type head_dim: int
:param tile_m: m block size
:type tile_m: int
"""
self.dtype = dtype
self.tile_m = tile_m
assert arch // 10 in [8, 9, 10, 11, 12], (
"Only Ampere (8.x), Hopper (9.x), and Blackwell (10.x, 11.x, 12.x) are supported"
)
self.arch = arch
# padding head_dim to a multiple of 32 as k_block_size
hdim_multiple_of = 32
self.tile_hdim = int(math.ceil(head_dim / hdim_multiple_of) * hdim_multiple_of)
self.check_hdim_oob = head_dim != self.tile_hdim
self.num_threads = num_threads
self.AtomLayoutMdQ = AtomLayoutMdQ
self.dQ_swapAB = dQ_swapAB
self.use_2cta_instrs = use_2cta_instrs and arch // 10 in [10, 11] and head_dim != 64
self.cluster_size = cluster_size
@staticmethod
def can_implement(dtype, head_dim, tile_m, num_threads) -> bool:
"""Check if the kernel can be implemented with the given parameters.
:param dtype: data type
:type dtype: cutlass.Numeric
:param head_dim: head dimension
:type head_dim: int
:param tile_m: m block size
:type tile_m: int
:return: True if the kernel can be implemented, False otherwise
:rtype: bool
"""
if dtype not in [cutlass.Float16, cutlass.BFloat16]:
return False
if head_dim % 8 != 0:
return False
if num_threads % 32 != 0:
return False
return True
def _get_tiled_mma(self):
if const_expr(self.arch // 10 in [8, 12]):
num_mma_warps = self.num_threads // 32
atom_layout_dQ = (
(self.AtomLayoutMdQ, num_mma_warps // self.AtomLayoutMdQ, 1)
if const_expr(not self.dQ_swapAB)
else (num_mma_warps // self.AtomLayoutMdQ, self.AtomLayoutMdQ, 1)
)
tiled_mma = cute.make_tiled_mma(
warp.MmaF16BF16Op(self.dtype, Float32, (16, 8, 16)),
atom_layout_dQ,
permutation_mnk=(atom_layout_dQ[0] * 16, atom_layout_dQ[1] * 16, 16),
)
elif const_expr(self.arch // 10 == 9):
num_wg_mma = self.num_threads // 128
atom_layout_dQ = (self.AtomLayoutMdQ, num_wg_mma // self.AtomLayoutMdQ)
tiler_mn_dQ = (self.tile_m // atom_layout_dQ[0], self.tile_hdim // atom_layout_dQ[1])
tiled_mma = sm90_utils_basic.make_trivial_tiled_mma(
self.dtype,
self.dtype,
warpgroup.OperandMajorMode.K, # These don't matter, we only care about the accum
warpgroup.OperandMajorMode.K,
Float32,
atom_layout_mnk=(atom_layout_dQ if not self.dQ_swapAB else atom_layout_dQ[::-1])
+ (1,),
tiler_mn=tiler_mn_dQ if not self.dQ_swapAB else tiler_mn_dQ[::-1],
)
else:
cta_group = tcgen05.CtaGroup.ONE
tiled_mma = sm100_utils_basic.make_trivial_tiled_mma(
self.dtype,
tcgen05.OperandMajorMode.MN, # dS_major_mode
tcgen05.OperandMajorMode.MN, # Kt_major_mode
Float32,
cta_group,
(self.tile_m, self.tile_hdim),
)
if const_expr(self.arch // 10 in [8, 9, 12]):
assert self.num_threads == tiled_mma.size
return tiled_mma
def _setup_attributes(self):
# ///////////////////////////////////////////////////////////////////////////////
# GMEM Tiled copy:
# ///////////////////////////////////////////////////////////////////////////////
# Thread layouts for copies
universal_copy_bits = 128
async_copy_elems_accum = universal_copy_bits // Float32.width
atom_async_copy_accum = cute.make_copy_atom(
cpasync.CopyG2SOp(cache_mode=cpasync.LoadCacheMode.GLOBAL),
Float32,
num_bits_per_copy=universal_copy_bits,
)
# We don't do bound checking for the gmem -> smem load so we just assert here.
assert (self.tile_m * self.tile_hdim // async_copy_elems_accum) % self.num_threads == 0
self.g2s_tiled_copy_dQaccum = cute.make_tiled_copy_tv(
atom_async_copy_accum,
cute.make_layout(self.num_threads),
cute.make_layout(async_copy_elems_accum),
)
num_s2r_copy_elems = 1 if const_expr(self.arch // 10 in [8, 12]) else 4
if const_expr(self.arch // 10 in [8, 12]):
self.s2r_tiled_copy_dQaccum = copy_utils.tiled_copy_1d(
Float32, self.num_threads, num_s2r_copy_elems
)
self.sdQaccum_layout = cute.make_layout(self.tile_m * self.tile_hdim)
elif const_expr(self.arch // 10 == 9):
num_threads_per_warp_group = 128
num_wg_mma = self.num_threads // 128
self.s2r_tiled_copy_dQaccum = cute.make_tiled_copy_tv(
cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), Float32, num_bits_per_copy=128),
cute.make_layout((num_threads_per_warp_group, num_wg_mma)), # thr_layout
cute.make_layout(128 // Float32.width), # val_layout
)
self.sdQaccum_layout = cute.make_layout(
(self.tile_m * self.tile_hdim // num_wg_mma, num_wg_mma)
)
else:
self.dQ_reduce_ncol = 32
dQaccum_reduce_stage = self.tile_hdim // self.dQ_reduce_ncol
assert self.num_threads == 128 # TODO: currently hard-coded
self.s2r_tiled_copy_dQaccum = copy_utils.tiled_copy_1d(
Float32, self.num_threads, num_s2r_copy_elems
)
self.sdQaccum_layout = cute.make_layout(
(self.tile_m * self.tile_hdim // dQaccum_reduce_stage, dQaccum_reduce_stage)
)
num_copy_elems = 128 // self.dtype.width
threads_per_row = math.gcd(128, self.tile_hdim) // num_copy_elems
self.gmem_tiled_copy_dQ = copy_utils.tiled_copy_2d(
self.dtype, threads_per_row, self.num_threads, num_copy_elems
)
# ///////////////////////////////////////////////////////////////////////////////
# Shared memory layout: dQ
# ///////////////////////////////////////////////////////////////////////////////
# We can't just use kHeadDim here. E.g. if MMA shape is 64 x 96 but split across 2 WGs,
# then setting kBlockKSmem to 32 will cause "Static shape_div failure".
# We want to treat it as 64 x 48, so kBlockKSmem should be 16.
mma_shape_n = self.tiled_mma.get_tile_size(1)
if const_expr(self.arch // 10 in [8, 12]):
sdQ_layout_atom = sm80_utils.get_smem_layout_atom(self.dtype, mma_shape_n)
self.sdQ_layout = cute.tile_to_shape(
sdQ_layout_atom, (self.tile_m, self.tile_hdim), (0, 1)
)
elif const_expr(self.arch // 10 == 9):
wg_d_dQ = num_wg_mma // self.AtomLayoutMdQ
self.sdQ_layout = sm90_utils.make_smem_layout(
self.dtype,
LayoutEnum.ROW_MAJOR,
(self.tile_m, self.tile_hdim),
major_mode_size=self.tile_hdim // wg_d_dQ,
)
else:
# TODO: this is hard-coded for hdim 128
self.sdQ_layout = sm100_utils_basic.make_smem_layout_epi(
self.dtype, LayoutEnum.ROW_MAJOR, (self.tile_m, self.tile_hdim), 1
)
@cute.jit
def __call__(
self,
mdQaccum: cute.Tensor,
mdQ: cute.Tensor,
scale: cutlass.Float32,
mCuSeqlensQ: Optional[cute.Tensor],
mSeqUsedQ: Optional[cute.Tensor],
# Always keep stream as the last parameter (EnvStream: obtained implicitly via TVM FFI).
stream: cuda.CUstream = None,
):
# Get the data type and check if it is fp16 or bf16
if const_expr(mdQ.element_type not in [cutlass.Float16, cutlass.BFloat16]):
raise TypeError("Only Float16 or BFloat16 is supported")
if const_expr(mdQaccum is not None):
if const_expr(mdQaccum.element_type not in [cutlass.Float32]):
raise TypeError("dQaccum tensor must be Float32")
mdQaccum, mdQ = [assume_tensor_aligned(t) for t in (mdQaccum, mdQ)]
self.tiled_mma = self._get_tiled_mma()
self._setup_attributes()
smem_size = max(
cute.size_in_bytes(cutlass.Float32, self.sdQaccum_layout),
cute.size_in_bytes(self.dtype, self.sdQ_layout),
)
if const_expr(mCuSeqlensQ is not None):
TileScheduler = SingleTileVarlenScheduler
num_head = mdQ.shape[1]
num_batch = mCuSeqlensQ.shape[0] - 1
num_block = cute.ceil_div(mdQ.shape[0], self.tile_m)
else:
TileScheduler = SingleTileScheduler
num_head = mdQ.shape[2]
num_batch = mdQ.shape[0]
num_block = cute.ceil_div(mdQ.shape[1], self.tile_m)
tile_sched_args = TileSchedulerArguments(
num_block=num_block,
num_head=num_head,
num_batch=num_batch,
num_splits=1,
seqlen_k=0,
headdim=mdQ.shape[2],
headdim_v=0,
total_q=mdQ.shape[0],
tile_shape_mn=(self.tile_m, 1),
mCuSeqlensQ=mCuSeqlensQ,
mSeqUsedQ=mSeqUsedQ,
)
tile_sched_params = TileScheduler.to_underlying_arguments(tile_sched_args)
grid_dim = TileScheduler.get_grid_shape(tile_sched_params)
# grid_dim: (m_block, num_head, batch_size)
self.kernel(
mdQaccum,
mdQ,
mCuSeqlensQ,
mSeqUsedQ,
scale,
self.tiled_mma,
self.dQ_swapAB,
self.sdQaccum_layout,
self.sdQ_layout,
self.g2s_tiled_copy_dQaccum,
self.s2r_tiled_copy_dQaccum,
self.gmem_tiled_copy_dQ,
tile_sched_params,
TileScheduler,
).launch(
grid=grid_dim,
block=[self.num_threads, 1, 1],
smem=smem_size,
stream=stream,
)
@cute.kernel
def kernel(
self,
mdQaccum: cute.Tensor,
mdQ: cute.Tensor,
mCuSeqlensQ: Optional[cute.Tensor],
mSeqUsedQ: Optional[cute.Tensor],
scale: cutlass.Float32,
tiled_mma: cute.TiledMma,
dQ_swapAB: cutlass.Constexpr,
sdQaccum_layout: cute.Layout,
sdQ_layout: cute.ComposedLayout,
g2s_tiled_copy_dQaccum: cute.TiledCopy,
s2r_tiled_copy_dQaccum: cute.TiledCopy,
gmem_tiled_copy_dQ: cute.TiledCopy,
tile_sched_params: ParamsBase,
TileScheduler: cutlass.Constexpr[Callable],
):
# ///////////////////////////////////////////////////////////////////////////////
# Get shared memory buffer
# ///////////////////////////////////////////////////////////////////////////////
smem = cutlass.utils.SmemAllocator()
sdQaccum = smem.allocate_tensor(cutlass.Float32, sdQaccum_layout, byte_alignment=1024)
sdQaccum_flat = cute.make_tensor(sdQaccum.iterator, cute.make_layout(cute.size(sdQaccum)))
if const_expr(self.arch // 10 in [8, 9, 12]):
sdQ = cute.make_tensor(cute.recast_ptr(sdQaccum.iterator, dtype=self.dtype), sdQ_layout)
else:
# extra stage dimension
sdQ = cute.make_tensor(
cute.recast_ptr(sdQaccum.iterator, sdQ_layout.inner, dtype=self.dtype),
sdQ_layout.outer,
)[None, None, 0]
sdQt = layout_utils.transpose_view(sdQ)
# Thread index, block index
tidx, _, _ = cute.arch.thread_idx()
tile_scheduler = TileScheduler.create(tile_sched_params)
work_tile = tile_scheduler.initial_work_tile_info()
m_block, head_idx, batch_idx, _ = work_tile.tile_idx
if work_tile.is_valid_tile:
# ///////////////////////////////////////////////////////////////////////////////
# Get the appropriate tiles for this thread block.
# ///////////////////////////////////////////////////////////////////////////////
seqlen = SeqlenInfoQK.create(
batch_idx,
mdQ.shape[1],
0,
mCuSeqlensQ=mCuSeqlensQ,
mCuSeqlensK=None,
mSeqUsedQ=mSeqUsedQ,
mSeqUsedK=None,
tile_m=self.tile_m * self.cluster_size,
)
if const_expr(not seqlen.has_cu_seqlens_q):
mdQ_cur = mdQ[batch_idx, None, head_idx, None]
mdQaccum_cur = mdQaccum[batch_idx, head_idx, None]
head_dim = mdQ.shape[3]
else:
padded_offset_q = seqlen.padded_offset_q
mdQ_cur = cute.domain_offset((seqlen.offset_q, 0), mdQ[None, head_idx, None])
mdQaccum_cur = cute.domain_offset(
(padded_offset_q * self.tile_hdim,), mdQaccum[head_idx, None]
)
head_dim = mdQ.shape[2]
# HACK: Compiler doesn't seem to recognize that padding
# by padded_offset_q * self.tile_hdim keeps alignment
# since statically divisible by 4
mdQaccum_cur_ptr = cute.make_ptr(
dtype=mdQaccum_cur.element_type,
value=mdQaccum_cur.iterator.toint(),
mem_space=mdQaccum_cur.iterator.memspace,
assumed_align=mdQaccum.iterator.alignment,
)
mdQaccum_cur = cute.make_tensor(mdQaccum_cur_ptr, mdQaccum_cur.layout)
gdQaccum = cute.local_tile(mdQaccum_cur, (self.tile_m * self.tile_hdim,), (m_block,))
gdQ = cute.local_tile(mdQ_cur, (self.tile_m, self.tile_hdim), (m_block, 0))
seqlen_q = seqlen.seqlen_q
seqlen_q_rounded = cute.round_up(seqlen_q, self.tile_m)
if const_expr(self.arch // 10 in [10, 11] and self.use_2cta_instrs):
# 2-CTA: remap dQaccum layout into TMEM view before writing sdQ
num_reduce_threads = self.num_threads
thr_mma_dsk = tiled_mma.get_slice(tidx)
dQacc_shape = thr_mma_dsk.partition_shape_C((self.tile_m, self.tile_hdim))
tdQtdQ = thr_mma_dsk.make_fragment_C(dQacc_shape)
tdQtdQ = cute.make_tensor(tdQtdQ.iterator, tdQtdQ.layout)
tmem_load_atom = cute.make_copy_atom(
tcgen05.copy.Ld32x32bOp(tcgen05.copy.Repetition(self.dQ_reduce_ncol)), Float32
)
tiled_tmem_ld = tcgen05.make_tmem_copy(tmem_load_atom, tdQtdQ)
thr_tmem_ld = tiled_tmem_ld.get_slice(tidx)
cdQ = cute.make_identity_tensor((self.tile_m, self.tile_hdim))
tdQcdQ = thr_mma_dsk.partition_C(cdQ)
tdQcdQ_tensor = cute.make_tensor(tdQcdQ.iterator, tdQcdQ.layout)
tdQrdQ = thr_tmem_ld.partition_D(tdQcdQ_tensor)
tiled_copy_accum = s2r_tiled_copy_dQaccum
g2s_thr_copy = tiled_copy_accum.get_slice(tidx)
# S -> R
tdQrdQ_fp32 = cute.make_rmem_tensor(tdQrdQ.shape, cutlass.Float32)
tdQrdQ_s2r = cute.make_tensor(tdQrdQ_fp32.iterator, tdQrdQ_fp32.shape)
smem_copy_atom = sm100_utils_basic.get_smem_store_op(
LayoutEnum.ROW_MAJOR, self.dtype, cutlass.Float32, tiled_tmem_ld
)
r2s_tiled_copy = cute.make_tiled_copy(
smem_copy_atom,
layout_tv=tiled_tmem_ld.layout_dst_tv_tiled,
tiler_mn=tiled_tmem_ld.tiler_mn,
)
tdQsdQ_r2s = thr_tmem_ld.partition_D(thr_mma_dsk.partition_C(sdQ))
tdQrdQ_r2s = cute.make_rmem_tensor(tdQsdQ_r2s.shape, self.dtype)
num_stages = cute.size(tdQrdQ_fp32, mode=[1])
stage_stride = self.dQ_reduce_ncol
row_groups = 2
assert num_stages % row_groups == 0
assert num_reduce_threads % row_groups == 0
stage_groups = num_stages // row_groups
threads_per_row_group = num_reduce_threads // row_groups
stage_loads = tuple((row_group, row_group) for row_group in range(row_groups))
stage_iters = tuple(
(row_group, row_group * threads_per_row_group)
for row_group in range(row_groups)
)
s2r_lane = tidx % threads_per_row_group
s2r_buf = tidx // threads_per_row_group
gdQaccum_layout_g2s = cute.make_layout(
shape=(self.tile_m * self.dQ_reduce_ncol, 1), stride=(1, 0)
)
sdQaccum_g2s = g2s_thr_copy.partition_D(sdQaccum)
# G -> S
for stage_group in cutlass.range_constexpr(stage_groups):
for stage_offset, smem_buf in stage_loads:
stage_idx = stage_group + stage_offset * stage_groups
gdQaccum_stage = cute.local_tile(
gdQaccum,
(self.tile_m * self.dQ_reduce_ncol,),
(stage_idx,),
)
gdQaccum_stage_g2s = cute.make_tensor(
gdQaccum_stage.iterator,
gdQaccum_layout_g2s,
)
tdQgdQ = g2s_thr_copy.partition_S(gdQaccum_stage_g2s)
cute.copy(
g2s_thr_copy,
tdQgdQ[None, None, 0],
sdQaccum_g2s[None, None, smem_buf],
)
cute.arch.fence_view_async_shared()
cute.arch.barrier(barrier_id=6, number_of_threads=num_reduce_threads)
# S -> R
for stage_offset, lane_offset in stage_iters:
stage_idx = stage_group + stage_offset * stage_groups
s2r_src_tidx = s2r_lane + lane_offset
s2r_thr_copy = tiled_copy_accum.get_slice(s2r_src_tidx)
sdQaccum_src = s2r_thr_copy.partition_S(sdQaccum)[None, None, s2r_buf]
tdQrdQ_s2r_cpy = tdQrdQ_s2r[None, stage_idx, None, None]
tdQrdQ_r2s_cpy = cute.make_tensor(
tdQrdQ_s2r_cpy.iterator, cute.make_layout(sdQaccum_src.shape)
)
cute.copy(s2r_thr_copy, sdQaccum_src, tdQrdQ_r2s_cpy)
cute.arch.fence_view_async_shared()
cute.arch.barrier(barrier_id=7, number_of_threads=num_reduce_threads)
# R -> S
stage_lo = stage_idx % stage_stride
stage_hi = stage_idx // stage_stride
tdQrdQ_r2s_cpy = cute.make_tensor(
cute.recast_ptr(tdQrdQ_r2s_cpy.iterator),
tdQrdQ_r2s[((None, 0), (stage_lo, stage_hi), 0, 0)].shape,
)
dQ_vec = tdQrdQ_r2s_cpy.load() * scale
tdQrdQ_r2s[((None, 0), (stage_lo, stage_hi), 0, 0)].store(
dQ_vec.to(self.dtype)
)
# R -> S
cute.copy(
r2s_tiled_copy,
tdQrdQ_r2s[None, None, None, 0],
tdQsdQ_r2s[None, None, None, 0],
)
cute.arch.fence_view_async_shared()
cute.arch.barrier(barrier_id=8, number_of_threads=num_reduce_threads)
else:
# Step 1: load dQaccum from gmem to smem
g2s_thr_copy_dQaccum = g2s_tiled_copy_dQaccum.get_slice(tidx)
tdQgdQaccum = g2s_thr_copy_dQaccum.partition_S(gdQaccum)
tdQsdQaccumg2s = g2s_thr_copy_dQaccum.partition_D(sdQaccum_flat)
cute.copy(g2s_tiled_copy_dQaccum, tdQgdQaccum, tdQsdQaccumg2s)
cute.arch.cp_async_commit_group()
cute.arch.cp_async_wait_group(0)
cute.arch.barrier()
# Step 2: load dQ from smem to rmem
s2r_thr_copy_dQaccum = s2r_tiled_copy_dQaccum.get_slice(tidx)
tdQsdQaccum = s2r_thr_copy_dQaccum.partition_S(sdQaccum)
tile_shape = (self.tile_m, self.tile_hdim)
acc = None
tiled_copy_t2r = None
if const_expr(self.arch // 10 in [8, 9, 12]):
acc_shape = tiled_mma.partition_shape_C(
tile_shape if const_expr(not dQ_swapAB) else tile_shape[::-1]
)
acc = cute.make_rmem_tensor(acc_shape, cutlass.Float32)
assert cute.size(acc) == cute.size(tdQsdQaccum)
else:
thr_mma = tiled_mma.get_slice(0) # 1-CTA
dQacc_shape = tiled_mma.partition_shape_C((self.tile_m, self.tile_hdim))
tdQtdQ = tiled_mma.make_fragment_C(dQacc_shape)
tdQcdQ = thr_mma.partition_C(
cute.make_identity_tensor((self.tile_m, self.tile_hdim))
)
tmem_load_atom = cute.make_copy_atom(
tcgen05.copy.Ld32x32bOp(tcgen05.copy.Repetition(self.dQ_reduce_ncol)),
Float32,
)
tiled_copy_t2r = tcgen05.make_tmem_copy(tmem_load_atom, tdQtdQ)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tdQrdQ_t2r_shape = thr_copy_t2r.partition_D(tdQcdQ).shape
acc = cute.make_rmem_tensor(tdQrdQ_t2r_shape, Float32)
tdQrdQaccum = cute.make_tensor(acc.iterator, cute.make_layout(tdQsdQaccum.shape))
cute.autovec_copy(tdQsdQaccum, tdQrdQaccum)
# Convert tdQrdQaccum from fp32 to fp16/bf16
rdQ = cute.make_fragment_like(acc, self.dtype)
rdQ.store((acc.load() * scale).to(self.dtype))
# Step 3: Copy dQ from register to smem
cute.arch.barrier() # make sure all threads have finished loading dQaccum
if const_expr(self.arch // 10 in [8, 9, 12]):
copy_atom_r2s_dQ = utils.get_smem_store_atom(
self.arch, self.dtype, transpose=self.dQ_swapAB
)
tiled_copy_r2s_dQ = cute.make_tiled_copy_C(copy_atom_r2s_dQ, tiled_mma)
else:
# copy_atom_r2s_dQ = sm100_utils_basic.get_smem_store_op(
# LayoutEnum.ROW_MAJOR, self.dtype, Float32, tiled_copy_t2r,
# )
# tiled_copy_r2s_dQ = cute.make_tiled_copy_D(copy_atom_r2s_dQ, tiled_copy_t2r)
thr_layout_r2s_dQ = cute.make_layout((self.num_threads, 1)) # 128 threads
val_layout_r2s_dQ = cute.make_layout((1, 128 // self.dtype.width))
copy_atom_r2s_dQ = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
self.dtype,
num_bits_per_copy=128,
)
tiled_copy_r2s_dQ = cute.make_tiled_copy_tv(
copy_atom_r2s_dQ, thr_layout_r2s_dQ, val_layout_r2s_dQ
)
thr_copy_r2s_dQ = tiled_copy_r2s_dQ.get_slice(tidx)
cdQ = cute.make_identity_tensor((self.tile_m, self.tile_hdim))
if const_expr(self.arch // 10 in [8, 9, 12]):
taccdQrdQ = thr_copy_r2s_dQ.retile(rdQ)
else:
taccdQcdQ_shape = thr_copy_r2s_dQ.partition_S(cdQ).shape
taccdQrdQ = cute.make_tensor(rdQ.iterator, taccdQcdQ_shape)
taccdQsdQ = thr_copy_r2s_dQ.partition_D(
sdQ if const_expr(not self.dQ_swapAB) else sdQt
)
cute.copy(thr_copy_r2s_dQ, taccdQrdQ, taccdQsdQ)
# Step 4: Copy dQ from smem to register to prepare for coalesced write to gmem
cute.arch.barrier() # make sure all smem stores are done
gmem_thr_copy_dQ = gmem_tiled_copy_dQ.get_slice(tidx)
tdQgdQ = gmem_thr_copy_dQ.partition_S(gdQ)
tdQsdQ = gmem_thr_copy_dQ.partition_D(sdQ)
tdQrdQ = cute.make_fragment_like(tdQsdQ, self.dtype)
# TODO: check OOB when reading from smem if kBlockM isn't evenly tiled
cute.autovec_copy(tdQsdQ, tdQrdQ)
# Step 5: Copy dQ from register to gmem
tdQcdQ = gmem_thr_copy_dQ.partition_S(cdQ)
tdQpdQ = utils.predicate_k(tdQcdQ, limit=head_dim)
for rest_m in cutlass.range(cute.size(tdQrdQ.shape[1]), unroll_full=True):
if tdQcdQ[0, rest_m, 0][0] < seqlen_q - m_block * self.tile_m:
cute.copy(
gmem_tiled_copy_dQ,
tdQrdQ[None, rest_m, None],
tdQgdQ[None, rest_m, None],
pred=tdQpdQ[None, rest_m, None],
)
@@ -0,0 +1,465 @@
# Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
# A reimplementation of https://github.com/Dao-AILab/flash-attention/blob/main/hopper/flash_bwd_preprocess_kernel.h
# from Cutlass C++ to Cute-DSL.
#
# Computes D_i = (dO_i * O_i).sum(dim=-1), optionally adjusted for LSE gradient:
# D'_i = D_i - dLSE_i
# This works because in the backward pass:
# dS_ij = P_ij * (dP_ij - D_i) [standard]
# When LSE is differentiable, d(loss)/d(S_ij) gets an extra term dLSE_i * P_ij
# (since d(LSE_i)/d(S_ij) = P_ij), giving:
# dS_ij = P_ij * (dP_ij - D_i) + dLSE_i * P_ij
# = P_ij * (dP_ij - (D_i - dLSE_i))
# So the main backward kernel is unchanged; we just replace D with D' = D - dLSE here.
import math
import operator
from functools import partial
from typing import Callable, Type, Optional
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
from cutlass import Float32, const_expr
from cutlass.cutlass_dsl import Arch, BaseDSL
from quack import copy_utils, layout_utils
from flash_attn.cute import utils
from flash_attn.cute.seqlen_info import SeqlenInfo
from quack.cute_dsl_utils import ParamsBase
from flash_attn.cute.tile_scheduler import (
SingleTileScheduler,
SingleTileVarlenScheduler,
TileSchedulerArguments,
)
from flash_attn.cute.pack_gqa import pack_gqa_layout
class FlashAttentionBackwardPreprocess:
def __init__(
self,
dtype: Type[cutlass.Numeric],
head_dim: int,
head_dim_v: int,
tile_m: int = 128,
num_threads: int = 256,
use_padded_offsets: bool = True,
nheads_major: bool = False,
pack_gqa: bool = False,
qhead_per_kvhead: int = 1,
nheads_kv: int = 1,
):
"""
All contiguous dimensions must be at least 16 bytes aligned which indicates the head dimension
should be a multiple of 8.
:param head_dim: head dimension
:type head_dim: int
:param tile_m: m block size
:type tile_m: int
:param num_threads: number of threads
:type num_threads: int
"""
self.use_pdl = BaseDSL._get_dsl().get_arch_enum() >= Arch.sm_90a
self.dtype = dtype
self.tile_m = tile_m
# padding head_dim to a multiple of 32 as k_block_size
hdim_multiple_of = 32
self.head_dim_padded = int(math.ceil(head_dim / hdim_multiple_of) * hdim_multiple_of)
self.head_dim_v_padded = int(math.ceil(head_dim_v / hdim_multiple_of) * hdim_multiple_of)
self.check_hdim_v_oob = head_dim_v != self.head_dim_v_padded
self.num_threads = num_threads
self.use_padded_offsets = use_padded_offsets
self.nheads_major = nheads_major
self.pack_gqa = pack_gqa
self.qhead_per_kvhead = qhead_per_kvhead
self.nheads_kv = nheads_kv
@staticmethod
def can_implement(dtype, head_dim, tile_m, num_threads) -> bool:
"""Check if the kernel can be implemented with the given parameters.
:param dtype: data type
:type dtype: cutlass.Numeric
:param head_dim: head dimension
:type head_dim: int
:param tile_m: m block size
:type tile_m: int
:param num_threads: number of threads
:type num_threads: int
:return: True if the kernel can be implemented, False otherwise
:rtype: bool
"""
if dtype not in [cutlass.Float16, cutlass.BFloat16]:
return False
if head_dim % 8 != 0:
return False
if num_threads % 32 != 0:
return False
if num_threads < tile_m: # For multiplying lse with log2
return False
return True
def _setup_attributes(self):
# ///////////////////////////////////////////////////////////////////////////////
# GMEM Tiled copy:
# ///////////////////////////////////////////////////////////////////////////////
# Thread layouts for copies
# We want kBlockKGmem to be a power of 2 so that when we do the summing,
# it's just between threads in the same warp
gmem_k_block_size = (
128
if self.head_dim_v_padded % 128 == 0
else (
64
if self.head_dim_v_padded % 64 == 0
else (32 if self.head_dim_v_padded % 32 == 0 else 16)
)
)
num_copy_elems = 128 // self.dtype.width
threads_per_row = gmem_k_block_size // num_copy_elems
self.gmem_tiled_copy_O = copy_utils.tiled_copy_2d(
self.dtype, threads_per_row, self.num_threads, num_copy_elems
)
universal_copy_bits = 128
num_copy_elems_dQaccum = universal_copy_bits // Float32.width
assert (
self.tile_m * self.head_dim_padded // num_copy_elems_dQaccum
) % self.num_threads == 0
self.gmem_tiled_copy_dQaccum = copy_utils.tiled_copy_1d(
Float32, self.num_threads, num_copy_elems_dQaccum
)
@cute.jit
def __call__(
self,
mO: cute.Tensor, # (batch, seqlen, nheads, head_dim_v) or (total_q, nheads, head_dim_v)
mdO: cute.Tensor, # same shape as mO
mPdPsum: cute.Tensor, # (batch, nheads, seqlen_padded) or (nheads, total_q_padded)
mLSE: Optional[cute.Tensor], # (batch, nheads, seqlen) or (nheads, total_q)
mLSElog2: Optional[cute.Tensor], # same shape as mPdPsum
# (batch, nheads, seqlen_padded * head_dim_v) or (nheads, total_q_padded * head_dim_v)
mdQaccum: Optional[cute.Tensor],
mCuSeqlensQ: Optional[cute.Tensor], # (batch + 1,)
mSeqUsedQ: Optional[cute.Tensor], # (batch,)
mdLSE: Optional[cute.Tensor], # (batch, nheads, seqlen) or (nheads, total_q)
mRowMax: Optional[cute.Tensor], # (b, s, n, h) or (t, n, h)
mScaleP: Optional[cute.Tensor], # == mRowMax
softmax_scale: Float32,
# Always keep stream as the last parameter (EnvStream: obtained implicitly via TVM FFI).
stream: cuda.CUstream = None,
):
# Get the data type and check if it is fp16 or bf16
if const_expr(not (mO.element_type == mdO.element_type)):
raise TypeError("All tensors must have the same data type")
if const_expr(mO.element_type not in [cutlass.Float16, cutlass.BFloat16]):
raise TypeError("Only Float16 or BFloat16 is supported")
if const_expr(mPdPsum.element_type not in [Float32]):
raise TypeError("PdPsum tensor must be Float32")
if const_expr(mdQaccum is not None):
assert self.nheads_major is False
assert self.pack_gqa is False
assert self.use_padded_offsets is True
if const_expr(mdQaccum.element_type not in [Float32]):
raise TypeError("dQaccum tensor must be Float32")
if const_expr(mLSE is not None):
if const_expr(mLSE.element_type not in [Float32]):
raise TypeError("LSE tensor must be Float32")
if const_expr(mLSElog2 is not None):
if const_expr(mLSElog2.element_type not in [Float32]):
raise TypeError("LSElog2 tensor must be Float32")
if const_expr(mdLSE is not None):
if const_expr(mdLSE.element_type not in [Float32]):
raise TypeError("dLSE tensor must be Float32")
if const_expr(mScaleP is not None):
assert self.nheads_major is True
assert self.pack_gqa is True
assert mRowMax is not None
if const_expr(mScaleP.element_type not in [Float32]):
raise TypeError("ScaleP tensor must be Float32")
if const_expr(mRowMax.element_type not in [Float32]):
raise TypeError("RowMax tensor must be Float32")
self._setup_attributes()
# (b, s, h, d) -> (s, d, h, b) or
# (total, h, d) -> (total, d, h)
QO_layout_transpose = [1, 3, 2, 0] if const_expr(mCuSeqlensQ is None) else [0, 2, 1]
mO, mdO = [
cute.make_tensor(mX.iterator, cute.select(mX.layout, mode=QO_layout_transpose))
for mX in (mO, mdO)
]
if const_expr(not self.nheads_major):
# (batch, nheads, seqlen) -> (seqlen, nheads, batch) or
# (nheads, total_q) -> (total_q, nheads)
transpose = [2, 1, 0] if const_expr(mCuSeqlensQ is None) else [1, 0]
else:
# (batch, seqlen, nheads) -> (seqlen, nheads, batch) or
# (total_q, nheads) -> (total_q, nheads)
transpose = [1, 2, 0] if const_expr(mCuSeqlensQ is None) else [0, 1]
mPdPsum, mLSE, mLSElog2, mdLSE, mdQaccum = [
layout_utils.select(mX, transpose) if mX is not None else None
for mX in (mPdPsum, mLSE, mLSElog2, mdLSE, mdQaccum)
]
# (b, s, n, h) => (s, n, h, b) or
# (total, n, h) == (total, n, h)
rowmax_layout_transpose = [1, 2, 3, 0] if const_expr(mCuSeqlensQ is None) else [0, 1, 2]
if const_expr(mRowMax is not None):
mRowMax = layout_utils.select(mRowMax, rowmax_layout_transpose)
if const_expr(mScaleP is not None):
mScaleP = layout_utils.select(mScaleP, rowmax_layout_transpose)
# pack gqa
if const_expr(self.pack_gqa):
mO, mdO, mRowMax, mScaleP = [
pack_gqa_layout(mX, self.qhead_per_kvhead, self.nheads_kv, head_idx=2)
if mX is not None
else None
for mX in (mO, mdO, mRowMax, mScaleP)
]
mPdPsum, mLSE, mLSElog2, mdLSE = [
pack_gqa_layout(mX, self.qhead_per_kvhead, self.nheads_kv, head_idx=1)
if mX is not None
else None
for mX in (mPdPsum, mLSE, mLSElog2, mdLSE)
]
# mO: (s, d, h, b) or (total, d, h)
if const_expr(mCuSeqlensQ is not None):
TileScheduler = SingleTileVarlenScheduler
num_head = mO.shape[2]
num_batch = mCuSeqlensQ.shape[0] - 1
else:
TileScheduler = SingleTileScheduler
num_head = mO.shape[2]
num_batch = mO.shape[3]
tile_sched_args = TileSchedulerArguments(
num_block=cute.ceil_div(mO.shape[0], self.tile_m),
num_head=num_head,
num_batch=num_batch,
num_splits=1,
seqlen_k=0,
headdim=0,
headdim_v=mO.shape[1],
total_q=cute.size(mO.shape[0])
if const_expr(mCuSeqlensQ is not None)
else cute.size(mO.shape[0]) * cute.size(mO.shape[3]),
tile_shape_mn=(self.tile_m, 1),
mCuSeqlensQ=mCuSeqlensQ,
mSeqUsedQ=mSeqUsedQ,
qhead_per_kvhead_packgqa=self.qhead_per_kvhead if const_expr(self.pack_gqa) else 1,
)
tile_sched_params = TileScheduler.to_underlying_arguments(tile_sched_args)
grid_dim = TileScheduler.get_grid_shape(tile_sched_params)
LOG2_E = math.log2(math.e)
softmax_scale_log2 = softmax_scale * LOG2_E
self.kernel(
mO,
mdO,
mPdPsum,
mLSE,
mLSElog2,
mdQaccum,
mCuSeqlensQ,
mSeqUsedQ,
mdLSE,
mRowMax,
mScaleP,
softmax_scale_log2,
self.gmem_tiled_copy_O,
self.gmem_tiled_copy_dQaccum,
tile_sched_params,
TileScheduler,
).launch(
grid=grid_dim,
block=[self.num_threads, 1, 1],
stream=stream,
use_pdl=self.use_pdl,
)
@cute.kernel
def kernel(
self,
mO: cute.Tensor,
mdO: cute.Tensor,
mPdPsum: cute.Tensor,
mLSE: Optional[cute.Tensor],
mLSElog2: Optional[cute.Tensor],
mdQaccum: Optional[cute.Tensor],
mCuSeqlensQ: Optional[cute.Tensor],
mSeqUsedQ: Optional[cute.Tensor],
mdLSE: Optional[cute.Tensor],
mRowMax: Optional[cute.Tensor],
mScaleP: Optional[cute.Tensor],
softmax_scale_log2: Float32,
gmem_tiled_copy_O: cute.TiledCopy,
gmem_tiled_copy_dQaccum: cute.TiledCopy,
tile_sched_params: ParamsBase,
TileScheduler: cutlass.Constexpr[Callable],
):
# Thread index, block index
tidx, _, _ = cute.arch.thread_idx()
tile_scheduler = TileScheduler.create(tile_sched_params)
work_tile = tile_scheduler.initial_work_tile_info()
m_block, head_idx, batch_idx, _ = work_tile.tile_idx
# This kernel is launched with use_pdl=True, so the GPU may start executing it in
# "prologue" mode while the previous stream kernel is still running. We must wait
# before touching any upstream GMEM outputs (mO, mdO, mLSE); otherwise we risk
# reading a partially-written dout, which silently corrupts dpsum = sum(O * dO) and
# propagates to dQ/dK via dS = P * (dP - dpsum).
if const_expr(self.use_pdl):
cute.arch.griddepcontrol_wait()
if work_tile.is_valid_tile:
# ///////////////////////////////////////////////////////////////////////////////
# Get the appropriate tiles for this thread block.
# ///////////////////////////////////////////////////////////////////////////////
seqlen_static = mO.shape[0] if const_expr(not self.pack_gqa) else mO.shape[0][1]
seqlen = SeqlenInfo.create(
batch_idx, seqlen_static, mCuSeqlensQ, mSeqUsedQ, tile=self.tile_m
)
# (seqlen, dv)
mO_cur, mdO_cur = [
seqlen.offset_batch(mX, batch_idx, dim=3)[None, None, head_idx] for mX in (mO, mdO)
]
mPdPsum_cur = seqlen.offset_batch(
mPdPsum, batch_idx, dim=2, padded=self.use_padded_offsets
)[None, head_idx]
headdim_v = mO_cur.shape[1]
seqlen_q = (
seqlen.seqlen
if const_expr(not self.pack_gqa)
else seqlen.seqlen * self.qhead_per_kvhead
)
seqlen_q_rounded = cute.round_up(seqlen_q, self.tile_m)
seqlen_limit = seqlen_q - m_block * self.tile_m
lse = None
if const_expr(mLSE is not None):
mLSE_cur = seqlen.offset_batch(mLSE, batch_idx, dim=2)[None, head_idx]
gLSE = cute.local_tile(mLSE_cur, (self.tile_m,), (m_block,))
lse = Float32.inf
if tidx < seqlen_limit:
lse = gLSE[tidx]
blk_shape = (self.tile_m, self.head_dim_v_padded)
gO = cute.local_tile(mO_cur, blk_shape, (m_block, 0))
gdO = cute.local_tile(mdO_cur, blk_shape, (m_block, 0))
gmem_thr_copy_O = gmem_tiled_copy_O.get_slice(tidx)
# (CPY_Atom, CPY_M, CPY_K)
tOgO = gmem_thr_copy_O.partition_S(gO)
tOgdO = gmem_thr_copy_O.partition_S(gdO)
cO = cute.make_identity_tensor(blk_shape)
tOcO = gmem_thr_copy_O.partition_S(cO)
t0OcO = gmem_thr_copy_O.get_slice(0).partition_S(cO)
tOpO = None
if const_expr(self.check_hdim_v_oob):
tOpO = copy_utils.predicate_k(tOcO, limit=headdim_v)
# Each copy will use the same predicate
copy = partial(copy_utils.copy, pred=tOpO)
tOrO = cute.make_rmem_tensor_like(tOgO)
tOrdO = cute.make_rmem_tensor_like(tOgdO)
if const_expr(self.check_hdim_v_oob):
tOrO.fill(0.0)
tOrdO.fill(0.0)
assert tOgO.shape == tOgdO.shape
for m in cutlass.range(cute.size(tOrO.shape[1]), unroll_full=True):
# Instead of using tOcO, we using t0OcO and subtract the offset from the limit.
# This is bc the entries of t0OcO are known at compile time.
if t0OcO[0, m, 0][0] < seqlen_limit - tOcO[0][0]:
copy(tOgO[None, m, None], tOrO[None, m, None])
copy(tOgdO[None, m, None], tOrdO[None, m, None])
# O and dO loads are done; signal that the next kernel can start.
# Correctness is ensured by griddepcontrol_wait() in bwd_sm90 before it reads our outputs.
if const_expr(self.use_pdl):
cute.arch.griddepcontrol_launch_dependents()
# Sum across the "k" dimension
pdpsum = (tOrO.load().to(Float32) * tOrdO.load().to(Float32)).reduce(
cute.ReductionOp.ADD, init_val=0.0, reduction_profile=(0, None, 1)
)
threads_per_row = gmem_tiled_copy_O.layout_src_tv_tiled[0].shape[0]
assert cute.arch.WARP_SIZE % threads_per_row == 0
pdpsum = utils.warp_reduce(pdpsum, operator.add, width=threads_per_row)
PdP_sum = cute.make_rmem_tensor(cute.size(tOrO, mode=[1]), Float32)
PdP_sum.store(pdpsum)
# If dLSE is provided, compute D' = D - dLSE (see module docstring for derivation).
gdLSE = None
if const_expr(mdLSE is not None):
mdLSE_cur = seqlen.offset_batch(mdLSE, batch_idx, dim=2)[None, head_idx]
gdLSE = cute.local_tile(mdLSE_cur, (self.tile_m,), (m_block,))
# Write PdPsum from rmem -> gmem
gPdPsum = cute.local_tile(mPdPsum_cur, (self.tile_m,), (m_block,))
# Only the thread corresponding to column 0 writes out the PdPsum to gmem
if tOcO[0, 0, 0][1] == 0:
for m in cutlass.range(cute.size(PdP_sum), unroll_full=True):
row = tOcO[0, m, 0][0]
PdPsum_val = 0.0
if row < seqlen_limit:
PdPsum_val = PdP_sum[m]
if const_expr(mdLSE is not None):
PdPsum_val -= gdLSE[row]
gPdPsum[row] = PdPsum_val
# Clear dQaccum
if const_expr(mdQaccum is not None):
mdQaccum_cur = seqlen.offset_batch(
mdQaccum,
batch_idx,
dim=2,
padded=self.use_padded_offsets,
multiple=self.head_dim_padded,
)[None, head_idx]
blkdQaccum_shape = (self.tile_m * self.head_dim_padded,)
gdQaccum = cute.local_tile(mdQaccum_cur, blkdQaccum_shape, (m_block,))
gmem_thr_copy_dQaccum = gmem_tiled_copy_dQaccum.get_slice(tidx)
tdQgdQaccum = gmem_thr_copy_dQaccum.partition_S(gdQaccum)
zero = cute.make_rmem_tensor_like(tdQgdQaccum)
zero.fill(0.0)
cute.copy(gmem_tiled_copy_dQaccum, zero, tdQgdQaccum)
LOG2_E = math.log2(math.e)
lse_log2 = lse * LOG2_E if lse != -Float32.inf else 0.0
if const_expr(mLSElog2 is not None):
mLSElog2_cur = seqlen.offset_batch(
mLSElog2, batch_idx, dim=2, padded=self.use_padded_offsets
)[None, head_idx]
gLSElog2 = cute.local_tile(mLSElog2_cur, (self.tile_m,), (m_block,))
LOG2_E = math.log2(math.e)
if tidx < seqlen_q_rounded - m_block * self.tile_m:
gLSElog2[tidx] = lse_log2
if const_expr(mRowMax is not None):
assert mLSE is not None
# (s, n)
mRowMax_cur, mScaleP_cur = [
seqlen.offset_batch(mX, batch_idx, dim=3)[None, None, head_idx]
for mX in (mRowMax, mScaleP)
]
# (tile_m, n)
gRowMax, gScaleP = [
cute.local_tile(mX, (self.tile_m,), (m_block, None))
for mX in (mRowMax_cur, mScaleP_cur)
]
assert self.tile_m <= self.num_threads
if const_expr(self.tile_m == self.num_threads) or tidx < self.tile_m:
for n in cutlass.range(gRowMax.shape[1], unroll=4):
row_max = gRowMax[tidx, n]
scale = 0.0
if row_max != -Float32.inf and lse != -Float32.inf:
scale = softmax_scale_log2 * row_max - lse_log2
scale = cute.math.exp2(scale, fastmath=True)
gScaleP[tidx, n] = scale
File diff suppressed because it is too large Load Diff
+55
View File
@@ -0,0 +1,55 @@
# Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
# SM120 (Blackwell GeForce / DGX Spark) backward pass.
#
# SM120 uses the same SM80-era MMA instructions (mma.sync.aligned.m16n8k16) but has
# a smaller shared memory capacity (99 KB vs 163 KB on SM80). This module subclasses
# FlashAttentionBackwardSm80 and overrides the SMEM capacity check accordingly.
import cutlass
import cutlass.utils as utils_basic
from flash_attn.cute.flash_bwd import FlashAttentionBackwardSm80
class FlashAttentionBackwardSm120(FlashAttentionBackwardSm80):
@staticmethod
def can_implement(
dtype,
head_dim,
head_dim_v,
m_block_size,
n_block_size,
num_stages_Q,
num_stages_dO,
num_threads,
is_causal,
V_in_regs=False,
) -> bool:
"""Check if the kernel can be implemented on SM120.
Same logic as SM80 but uses SM120's shared memory capacity (99 KB).
"""
if dtype not in [cutlass.Float16, cutlass.BFloat16]:
return False
if head_dim % 8 != 0:
return False
if head_dim_v % 8 != 0:
return False
if n_block_size % 16 != 0:
return False
if num_threads % 32 != 0:
return False
# Shared memory usage: Q tile + dO tile + K tile + V tile
smem_usage_Q = m_block_size * head_dim * num_stages_Q * 2
smem_usage_dO = m_block_size * head_dim_v * num_stages_dO * 2
smem_usage_K = n_block_size * head_dim * 2
smem_usage_V = n_block_size * head_dim_v * 2
smem_usage_QV = (
(smem_usage_Q + smem_usage_V) if not V_in_regs else max(smem_usage_Q, smem_usage_V)
)
smem_usage = smem_usage_QV + smem_usage_dO + smem_usage_K
# SM120 has 99 KB shared memory (vs 163 KB on SM80)
smem_capacity = utils_basic.get_smem_capacity_in_bytes("sm_120")
if smem_usage > smem_capacity:
return False
return True
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+698
View File
@@ -0,0 +1,698 @@
# Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
# A reimplementation of https://github.com/Dao-AILab/flash-attention/blob/main/hopper/flash_fwd_combine_kernel.h
# from Cutlass C++ to Cute-DSL.
import math
from typing import Type, Optional
from functools import partial
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync
from cutlass import Float32, Int32, Boolean, const_expr
from flash_attn.cute import utils
from flash_attn.cute.cute_dsl_utils import assume_tensor_aligned
from flash_attn.cute.seqlen_info import SeqlenInfo
from cutlass.cute import FastDivmodDivisor
class FlashAttentionForwardCombine:
def __init__(
self,
dtype: Type[cutlass.Numeric],
dtype_partial: Type[cutlass.Numeric],
head_dim: int,
tile_m: int = 8,
k_block_size: int = 64,
log_max_splits: int = 4,
num_threads: int = 256,
stages: int = 4,
):
"""
Forward combine kernel for split attention computation.
:param dtype: output data type
:param dtype_partial: partial accumulation data type
:param head_dim: head dimension
:param tile_m: m block size
:param k_block_size: k block size
:param log_max_splits: log2 of maximum splits
:param num_threads: number of threads
:param varlen: whether using variable length sequences
:param stages: number of pipeline stages
"""
self.dtype = dtype
self.dtype_partial = dtype_partial
self.head_dim = head_dim
self.tile_m = tile_m
self.k_block_size = k_block_size
self.max_splits = 1 << log_max_splits
self.num_threads = num_threads
self.is_even_k = head_dim % k_block_size == 0
self.stages = stages
@staticmethod
def can_implement(
dtype,
dtype_partial,
head_dim,
tile_m,
k_block_size,
log_max_splits,
num_threads,
) -> bool:
"""Check if the kernel can be implemented with the given parameters."""
if dtype not in [cutlass.Float16, cutlass.BFloat16, cutlass.Float32]:
return False
if dtype_partial not in [cutlass.Float16, cutlass.BFloat16, Float32]:
return False
if head_dim % 8 != 0:
return False
if num_threads % 32 != 0:
return False
if tile_m % 8 != 0:
return False
max_splits = 1 << log_max_splits
if max_splits > 256:
return False
if (tile_m * max_splits) % num_threads != 0:
return False
return True
def _setup_attributes(self):
# GMEM copy setup for O partial
universal_copy_bits = 128
async_copy_elems = universal_copy_bits // self.dtype_partial.width
assert self.k_block_size % async_copy_elems == 0
k_block_gmem = (
128 if self.k_block_size % 128 == 0 else (64 if self.k_block_size % 64 == 0 else 32)
)
gmem_threads_per_row = k_block_gmem // async_copy_elems
assert self.num_threads % gmem_threads_per_row == 0
# Async copy atom for O partial load
atom_async_copy_partial = cute.make_copy_atom(
cpasync.CopyG2SOp(cache_mode=cpasync.LoadCacheMode.GLOBAL),
self.dtype_partial,
num_bits_per_copy=universal_copy_bits,
)
tOpartial_layout = cute.make_ordered_layout(
(self.num_threads // gmem_threads_per_row, gmem_threads_per_row),
order=(1, 0),
)
vOpartial_layout = cute.make_layout((1, async_copy_elems)) # 4 vals per load
self.gmem_tiled_copy_O_partial = cute.make_tiled_copy_tv(
atom_async_copy_partial, tOpartial_layout, vOpartial_layout
)
# GMEM copy setup for final O (use universal copy for store)
atom_universal_copy = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
self.dtype,
num_bits_per_copy=async_copy_elems * self.dtype.width,
)
self.gmem_tiled_copy_O = cute.make_tiled_copy_tv(
atom_universal_copy,
tOpartial_layout,
vOpartial_layout, # 4 vals per store
)
# LSE copy setup with async copy (alignment = 1)
lse_copy_bits = Float32.width # 1 element per copy, width is in bits
m_block_smem = (
128
if self.tile_m % 128 == 0
else (
64
if self.tile_m % 64 == 0
else (32 if self.tile_m % 32 == 0 else (16 if self.tile_m % 16 == 0 else 8))
)
)
gmem_threads_per_row_lse = m_block_smem
assert self.num_threads % gmem_threads_per_row_lse == 0
# Async copy atom for LSE load
atom_async_copy_lse = cute.make_copy_atom(
cpasync.CopyG2SOp(cache_mode=cpasync.LoadCacheMode.ALWAYS),
Float32,
num_bits_per_copy=lse_copy_bits,
)
tLSE_layout = cute.make_ordered_layout(
(self.num_threads // gmem_threads_per_row_lse, gmem_threads_per_row_lse),
order=(1, 0),
)
vLSE_layout = cute.make_layout(1)
self.gmem_tiled_copy_LSE = cute.make_tiled_copy_tv(
atom_async_copy_lse, tLSE_layout, vLSE_layout
)
# ///////////////////////////////////////////////////////////////////////////////
# Shared memory
# ///////////////////////////////////////////////////////////////////////////////
# Shared memory to register copy for LSE
self.smem_threads_per_col_lse = self.num_threads // m_block_smem
assert 32 % self.smem_threads_per_col_lse == 0 # Must divide warp size
s2r_layout_atom_lse = cute.make_ordered_layout(
(self.smem_threads_per_col_lse, self.num_threads // self.smem_threads_per_col_lse),
order=(0, 1),
)
self.s2r_tiled_copy_LSE = cute.make_tiled_copy_tv(
cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), Float32),
s2r_layout_atom_lse,
cute.make_layout(1),
)
# LSE shared memory layout with swizzling to avoid bank conflicts
# This works for kBlockMSmem = 8, 16, 32, 64, 128, no bank conflicts
if const_expr(m_block_smem == 8):
smem_lse_swizzle = cute.make_swizzle(5, 0, 5)
elif const_expr(m_block_smem == 16):
smem_lse_swizzle = cute.make_swizzle(4, 0, 4)
else:
smem_lse_swizzle = cute.make_swizzle(3, 2, 3)
smem_layout_atom_lse = cute.make_composed_layout(
smem_lse_swizzle, 0, cute.make_ordered_layout((8, m_block_smem), order=(1, 0))
)
self.smem_layout_lse = cute.tile_to_shape(
smem_layout_atom_lse, (self.max_splits, self.tile_m), (0, 1)
)
# O partial shared memory layout (simple layout for pipeline stages)
self.smem_layout_o = cute.make_ordered_layout(
(self.tile_m, self.k_block_size, self.stages), order=(1, 0, 2)
)
@cute.jit
def __call__(
self,
mO_partial: cute.Tensor,
mLSE_partial: cute.Tensor,
mO: cute.Tensor,
mLSE: Optional[cute.Tensor] = None,
cu_seqlens: Optional[cute.Tensor] = None,
seqused: Optional[cute.Tensor] = None,
num_splits_dynamic_ptr: Optional[cute.Tensor] = None,
varlen_batch_idx: Optional[cute.Tensor] = None,
semaphore_to_reset: Optional[cute.Tensor] = None,
# Always keep stream as the last parameter (EnvStream: obtained implicitly via TVM FFI).
stream: cuda.CUstream = None,
):
# Type checking
if const_expr(not (mO_partial.element_type == self.dtype_partial)):
raise TypeError("O partial tensor must match dtype_partial")
if const_expr(not (mO.element_type == self.dtype)):
raise TypeError("O tensor must match dtype")
if const_expr(mLSE_partial.element_type not in [Float32]):
raise TypeError("LSE partial tensor must be Float32")
if const_expr(mLSE is not None and mLSE.element_type not in [Float32]):
raise TypeError("LSE tensor must be Float32")
# Shape validation - input tensors are in user format, need to be converted to kernel format
if const_expr(len(mO_partial.shape) not in [4, 5]):
raise ValueError(
"O partial tensor must have 4 or 5 dimensions: (num_splits, batch, seqlen, nheads, headdim) or (num_splits, total_q, nheads, headdim)"
)
if const_expr(len(mLSE_partial.shape) not in [3, 4]):
raise ValueError(
"LSE partial tensor must have 3 or 4 dimensions: (num_splits, batch, seqlen, nheads) or (num_splits, total_q, nheads)"
)
if const_expr(len(mO.shape) not in [3, 4]):
raise ValueError(
"O tensor must have 3 or 4 dimensions: (batch, seqlen, nheads, headdim) or (total_q, nheads, headdim)"
)
if const_expr(mLSE is not None and len(mLSE.shape) not in [2, 3]):
raise ValueError(
"LSE tensor must have 2 or 3 dimensions: (batch, seqlen, nheads) or (total_q, nheads)"
)
mO_partial, mO = [assume_tensor_aligned(t) for t in (mO_partial, mO)]
# (num_splits, b, seqlen, h, d) -> (seqlen, d, num_splits, h, b)
# or (num_splits, total_q, h, d) -> (total_q, d, num_splits, h)
O_partial_layout_transpose = (
[2, 4, 0, 3, 1] if const_expr(cu_seqlens is None) else [1, 3, 0, 2]
)
# (b, seqlen, h, d) -> (seqlen, d, h, b) or (total_q, h, d) -> (total_q, d, h)
mO_partial = cute.make_tensor(
mO_partial.iterator, cute.select(mO_partial.layout, mode=O_partial_layout_transpose)
)
O_layout_transpose = [1, 3, 2, 0] if const_expr(cu_seqlens is None) else [0, 2, 1]
mO = cute.make_tensor(mO.iterator, cute.select(mO.layout, mode=O_layout_transpose))
# (num_splits, b, seqlen, h) -> (seqlen, num_splits, h, b)
# or (num_splits, total_q, h) -> (total_q, num_splits, h)
LSE_partial_layout_transpose = [2, 0, 3, 1] if const_expr(cu_seqlens is None) else [1, 0, 2]
mLSE_partial = cute.make_tensor(
mLSE_partial.iterator,
cute.select(mLSE_partial.layout, mode=LSE_partial_layout_transpose),
)
# (b, seqlen, h) -> (seqlen, h, b) or (total_q, h) -> (total_q, h)
LSE_layout_transpose = [1, 2, 0] if const_expr(cu_seqlens is None) else [0, 1]
mLSE = (
cute.make_tensor(mLSE.iterator, cute.select(mLSE.layout, mode=LSE_layout_transpose))
if mLSE is not None
else None
)
# Determine if we have variable length sequences
varlen = const_expr(cu_seqlens is not None or seqused is not None)
self._setup_attributes()
@cute.struct
class SharedStorage:
sLSE: cute.struct.Align[
cute.struct.MemRange[Float32, cute.cosize(self.smem_layout_lse)], 128
]
sMaxValidSplit: cute.struct.Align[cute.struct.MemRange[Int32, self.tile_m], 128]
sO: cute.struct.Align[
cute.struct.MemRange[self.dtype_partial, cute.cosize(self.smem_layout_o)], 128
]
smem_size = SharedStorage.size_in_bytes()
# Grid dimensions: (ceil_div(seqlen, m_block), ceil_div(head_dim, k_block), num_head * batch)
seqlen = mO_partial.shape[0]
num_head = mO_partial.shape[3]
batch_size = (
mO_partial.shape[4]
if const_expr(cu_seqlens is None)
else Int32(cu_seqlens.shape[0] - 1)
)
# Create FastDivmodDivisor objects for efficient division
seqlen_divmod = FastDivmodDivisor(seqlen)
head_divmod = FastDivmodDivisor(num_head)
grid_dim = (
cute.ceil_div(seqlen * num_head, self.tile_m),
cute.ceil_div(self.head_dim, self.k_block_size),
batch_size,
)
self.kernel(
mO_partial,
mLSE_partial,
mO,
mLSE,
cu_seqlens,
seqused,
num_splits_dynamic_ptr,
varlen_batch_idx,
semaphore_to_reset,
SharedStorage,
self.smem_layout_lse,
self.smem_layout_o,
self.gmem_tiled_copy_O_partial,
self.gmem_tiled_copy_O,
self.gmem_tiled_copy_LSE,
self.s2r_tiled_copy_LSE,
seqlen_divmod,
head_divmod,
varlen,
).launch(
grid=grid_dim,
block=[self.num_threads, 1, 1],
smem=smem_size,
stream=stream,
)
@cute.kernel
def kernel(
self,
mO_partial: cute.Tensor,
mLSE_partial: cute.Tensor,
mO: cute.Tensor,
mLSE: Optional[cute.Tensor],
cu_seqlens: Optional[cute.Tensor],
seqused: Optional[cute.Tensor],
num_splits_dynamic_ptr: Optional[cute.Tensor],
varlen_batch_idx: Optional[cute.Tensor],
semaphore_to_reset: Optional[cute.Tensor],
SharedStorage: cutlass.Constexpr,
smem_layout_lse: cute.Layout | cute.ComposedLayout,
smem_layout_o: cute.Layout,
gmem_tiled_copy_O_partial: cute.TiledCopy,
gmem_tiled_copy_O: cute.TiledCopy,
gmem_tiled_copy_LSE: cute.TiledCopy,
s2r_tiled_copy_LSE: cute.TiledCopy,
seqlen_divmod: FastDivmodDivisor,
head_divmod: FastDivmodDivisor,
varlen: cutlass.Constexpr[bool],
):
# Thread and block indices
tidx, _, _ = cute.arch.thread_idx()
m_block, k_block, maybe_virtual_batch = cute.arch.block_idx()
# Map virtual batch index to real batch index (for persistent tile schedulers)
batch_idx = (
varlen_batch_idx[maybe_virtual_batch]
if const_expr(varlen_batch_idx is not None)
else maybe_virtual_batch
)
# ///////////////////////////////////////////////////////////////////////////////
# Get shared memory buffer
# ///////////////////////////////////////////////////////////////////////////////
smem = cutlass.utils.SmemAllocator()
storage = smem.allocate(SharedStorage)
sLSE = storage.sLSE.get_tensor(smem_layout_lse)
sMaxValidSplit = storage.sMaxValidSplit.get_tensor((self.tile_m,))
sO = storage.sO.get_tensor(smem_layout_o)
# Handle semaphore reset — wait for dependent grids first
if const_expr(semaphore_to_reset is not None):
if (
tidx == 0
and m_block == cute.arch.grid_dim()[0] - 1
and k_block == cute.arch.grid_dim()[1] - 1
and maybe_virtual_batch == cute.arch.grid_dim()[2] - 1
):
cute.arch.griddepcontrol_wait()
semaphore_to_reset[0] = 0
# Get number of splits (use maybe_virtual_batch for per-batch-slot splits)
num_splits = (
num_splits_dynamic_ptr[maybe_virtual_batch]
if const_expr(num_splits_dynamic_ptr is not None)
else mLSE_partial.shape[1]
)
# Handle variable length sequences using SeqlenInfo
seqlen_info = SeqlenInfo.create(
batch_idx=batch_idx,
seqlen_static=mO_partial.shape[0],
cu_seqlens=cu_seqlens,
seqused=seqused,
# Don't need to pass in tile size since we won't use offset_padded
)
seqlen, offset = seqlen_info.seqlen, seqlen_info.offset
# Extract number of heads (head index will be determined dynamically)
num_head = mO_partial.shape[3]
max_idx = seqlen * num_head
# Early exit for single split if dynamic
if (const_expr(num_splits_dynamic_ptr is None) or num_splits > 1) and (
const_expr(not varlen) or m_block * self.tile_m < max_idx
):
# Wait for dependent grids (e.g., the main attention kernel that produces O_partial/LSE_partial)
cute.arch.griddepcontrol_wait()
# ===============================
# Step 1: Load LSE_partial from gmem to shared memory
# ===============================
mLSE_partial_cur = seqlen_info.offset_batch(mLSE_partial, batch_idx, dim=3)
mLSE_partial_copy = cute.tiled_divide(mLSE_partial_cur, (1,))
gmem_thr_copy_LSE = gmem_tiled_copy_LSE.get_slice(tidx)
tLSEsLSE = gmem_thr_copy_LSE.partition_D(sLSE)
# Create identity tensor for coordinate tracking
cLSE = cute.make_identity_tensor((self.max_splits, self.tile_m))
tLSEcLSE = gmem_thr_copy_LSE.partition_S(cLSE)
# Load LSE partial values
for m in cutlass.range(cute.size(tLSEcLSE, mode=[2]), unroll_full=True):
mi = tLSEcLSE[0, 0, m][1] # Get m coordinate
idx = m_block * self.tile_m + mi
if idx < max_idx:
# Calculate actual sequence position and head using FastDivmodDivisor
if const_expr(not varlen):
head_idx, m_idx = divmod(idx, seqlen_divmod)
else:
head_idx = idx // seqlen
m_idx = idx - head_idx * seqlen
mLSE_partial_cur_copy = mLSE_partial_copy[None, m_idx, None, head_idx]
for s in cutlass.range(cute.size(tLSEcLSE, mode=[1]), unroll_full=True):
si = tLSEcLSE[0, s, 0][0] # Get split coordinate
if si < num_splits:
cute.copy(
gmem_thr_copy_LSE,
mLSE_partial_cur_copy[None, si],
tLSEsLSE[None, s, m],
)
else:
tLSEsLSE[None, s, m].fill(-Float32.inf)
# Don't need to zero out the rest of the LSEs, as we will not write the output to gmem
cute.arch.cp_async_commit_group()
# ===============================
# Step 2: Load O_partial for pipeline stages
# ===============================
gmem_thr_copy_O_partial = gmem_tiled_copy_O_partial.get_slice(tidx)
cO = cute.make_identity_tensor((self.tile_m, self.k_block_size))
tOcO = gmem_thr_copy_O_partial.partition_D(cO)
tOsO_partial = gmem_thr_copy_O_partial.partition_D(sO)
mO_partial_cur = seqlen_info.offset_batch(mO_partial, batch_idx, dim=4)
# Precompute these values to avoid recomputing them in the loop
num_rows = const_expr(cute.size(tOcO, mode=[1]))
tOmidx = cute.make_rmem_tensor(num_rows, cutlass.Int32)
tOhidx = cute.make_rmem_tensor(num_rows, cutlass.Int32)
tOrOptr = cute.make_rmem_tensor(num_rows, cutlass.Int64)
for m in cutlass.range(num_rows, unroll_full=True):
mi = tOcO[0, m, 0][0] # m coordinate
idx = m_block * self.tile_m + mi
if const_expr(not varlen):
tOhidx[m], tOmidx[m] = divmod(idx, seqlen_divmod)
else:
tOhidx[m] = idx // seqlen
tOmidx[m] = idx - tOhidx[m] * seqlen
tOrOptr[m] = utils.elem_pointer(
mO_partial_cur, (tOmidx[m], k_block * self.k_block_size, 0, tOhidx[m])
).toint()
if idx >= max_idx:
tOhidx[m] = -1
tOpO = None
if const_expr(not self.is_even_k):
tOpO = cute.make_rmem_tensor(cute.size(tOcO, mode=[2]), Boolean)
for k in cutlass.range(cute.size(tOpO), unroll_full=True):
tOpO[k] = tOcO[0, 0, k][1] < mO_partial.shape[1] - k_block * self.k_block_size
# if cute.arch.thread_idx()[0] == 0 and k_block == 1: cute.print_tensor(tOpO)
load_O_partial = partial(
self.load_O_partial,
gmem_tiled_copy_O_partial,
tOrOptr,
tOsO_partial,
tOhidx,
tOpO,
tOcO,
mO_partial_cur.layout,
)
# Load first few stages of O_partial
for stage in cutlass.range(self.stages - 1, unroll_full=True):
if stage < num_splits:
load_O_partial(stage, stage)
cute.arch.cp_async_commit_group()
# ===============================
# Step 3: Load and transpose LSE from smem to registers
# ===============================
# Wait for LSE and initial O partial stages to complete
cute.arch.cp_async_wait_group(self.stages - 1)
cute.arch.sync_threads()
# if cute.arch.thread_idx()[0] == 0:
# # cute.print_tensor(sLSE)
# for i in range(64):
# cute.printf("sLSE[%d, 0] = %f", i, sLSE[i, 0])
# cute.arch.sync_threads()
s2r_thr_copy_LSE = s2r_tiled_copy_LSE.get_slice(tidx)
ts2rsLSE = s2r_thr_copy_LSE.partition_S(sLSE)
ts2rrLSE = cute.make_rmem_tensor_like(ts2rsLSE)
cute.copy(s2r_tiled_copy_LSE, ts2rsLSE, ts2rrLSE)
# ===============================
# Step 4: Compute final LSE along split dimension
# ===============================
lse_sum = cute.make_rmem_tensor(cute.size(ts2rrLSE, mode=[2]), Float32)
ts2rcLSE = s2r_thr_copy_LSE.partition_D(cLSE)
# We compute the max valid split for each row to short-circuit the computation later
max_valid_split = cute.make_rmem_tensor(cute.size(ts2rrLSE, mode=[2]), Int32)
assert cute.size(ts2rrLSE, mode=[0]) == 1
# Compute max, scales, and final LSE for each row
for m in cutlass.range(cute.size(ts2rrLSE, mode=[2]), unroll_full=True):
# Find max LSE value across splits
threads_per_col = const_expr(self.smem_threads_per_col_lse)
lse_max = cute.arch.warp_reduction_max(
ts2rrLSE[None, None, m]
.load()
.reduce(cute.ReductionOp.MAX, init_val=-Float32.inf, reduction_profile=0),
threads_in_group=threads_per_col,
)
# if cute.arch.thread_idx()[0] == 0: cute.printf(lse_max)
# Find max valid split index
max_valid_idx = -1
for s in cutlass.range(cute.size(ts2rrLSE, mode=[1]), unroll_full=True):
if ts2rrLSE[0, s, m] != -Float32.inf:
max_valid_idx = ts2rcLSE[0, s, 0][0] # Get split coordinate
# if cute.arch.thread_idx()[0] < 32: cute.printf(max_valid_idx)
max_valid_split[m] = cute.arch.warp_reduction_max(
max_valid_idx, threads_in_group=threads_per_col
)
# Compute exp scales and sum
lse_max_cur = (
0.0 if lse_max == -Float32.inf else lse_max
) # In case all local LSEs are -inf
LOG2_E = math.log2(math.e)
lse_sum_cur = 0.0
for s in cutlass.range(cute.size(ts2rrLSE, mode=[1]), unroll_full=True):
scale = cute.math.exp2(
ts2rrLSE[0, s, m] * LOG2_E - (lse_max_cur * LOG2_E), fastmath=True
)
lse_sum_cur += scale
ts2rrLSE[0, s, m] = scale # Store scale for later use
lse_sum_cur = cute.arch.warp_reduction_sum(
lse_sum_cur, threads_in_group=threads_per_col
)
lse_sum[m] = cute.math.log(lse_sum_cur, fastmath=True) + lse_max
# Normalize scales
inv_sum = (
0.0 if (lse_sum_cur == 0.0 or lse_sum_cur != lse_sum_cur) else 1.0 / lse_sum_cur
)
ts2rrLSE[None, None, m].store(ts2rrLSE[None, None, m].load() * inv_sum)
# Store the scales exp(lse - lse_logsum) back to smem
cute.copy(s2r_tiled_copy_LSE, ts2rrLSE, ts2rsLSE)
# Store max valid split to smem
for m in cutlass.range(cute.size(ts2rrLSE, mode=[2]), unroll_full=True):
if ts2rcLSE[0, 0, m][0] == 0: # Only thread responsible for s=0 writes
mi = ts2rcLSE[0, 0, m][1]
if mi < self.tile_m:
sMaxValidSplit[mi] = max_valid_split[m]
# ===============================
# Step 5: Store final LSE to gmem
# ===============================
if const_expr(mLSE is not None):
if const_expr(cu_seqlens is None):
mLSE_cur = mLSE[None, None, batch_idx]
else:
mLSE_cur = cute.domain_offset((offset, 0), mLSE)
if k_block == 0: # Only first k_block writes LSE when mLSE is provided
for m in cutlass.range(cute.size(ts2rrLSE, mode=[2]), unroll_full=True):
if ts2rcLSE[0, 0, m][0] == 0: # Only thread responsible for s=0 writes
mi = ts2rcLSE[0, 0, m][1]
idx = m_block * self.tile_m + mi
if idx < max_idx:
if const_expr(not varlen):
head_idx, m_idx = divmod(idx, seqlen_divmod)
else:
head_idx = idx // seqlen
m_idx = idx - head_idx * seqlen
mLSE_cur[m_idx, head_idx] = lse_sum[m]
# ===============================
# Step 6: Read O_partial and accumulate final O
# ===============================
cute.arch.sync_threads()
# Get max valid split for this thread
thr_max_valid_split = sMaxValidSplit[tOcO[0, 0, 0][0]]
for m in cutlass.range(1, cute.size(tOcO, mode=[1]), unroll_full=True):
thr_max_valid_split = max(thr_max_valid_split, sMaxValidSplit[tOcO[0, m, 0][0]])
tOrO_partial = cute.make_rmem_tensor_like(tOsO_partial[None, None, None, 0])
tOrO = cute.make_rmem_tensor_like(tOrO_partial, Float32)
tOrO.fill(0.0)
stage_load = self.stages - 1
stage_compute = 0
# Main accumulation loop
for s in cutlass.range(thr_max_valid_split + 1, unroll=4):
# Get scales for this split
scale = cute.make_rmem_tensor(num_rows, Float32)
for m in cutlass.range(num_rows, unroll_full=True):
scale[m] = sLSE[s, tOcO[0, m, 0][0]] # Get scale from smem
# Load next stage if needed
split_to_load = s + self.stages - 1
if split_to_load <= thr_max_valid_split:
load_O_partial(split_to_load, stage_load)
cute.arch.cp_async_commit_group()
stage_load = 0 if stage_load == self.stages - 1 else stage_load + 1
# Wait for the current stage to be ready
cute.arch.cp_async_wait_group(self.stages - 1)
# We don't need __syncthreads() because each thread is just reading its own data from smem
# Copy from smem to registers
cute.autovec_copy(tOsO_partial[None, None, None, stage_compute], tOrO_partial)
stage_compute = 0 if stage_compute == self.stages - 1 else stage_compute + 1
# Accumulate scaled partial results
for m in cutlass.range(num_rows, unroll_full=True):
if tOhidx[m] >= 0 and scale[m] > 0.0:
tOrO[None, m, None].store(
tOrO[None, m, None].load()
+ scale[m] * tOrO_partial[None, m, None].load().to(Float32)
)
# ===============================
# Step 7: Write final O to gmem
# ===============================
rO = cute.make_rmem_tensor_like(tOrO, self.dtype)
rO.store(tOrO.load().to(self.dtype))
mO_cur = seqlen_info.offset_batch(mO, batch_idx, dim=3)
if const_expr(cu_seqlens is None):
mO_cur = mO[None, None, None, batch_idx]
else:
mO_cur = cute.domain_offset((offset, 0, 0), mO)
mO_cur = utils.domain_offset_aligned((0, k_block * self.k_block_size, 0), mO_cur)
elems_per_store = const_expr(cute.size(gmem_tiled_copy_O.layout_tv_tiled[1]))
# mO_cur_copy = cute.tiled_divide(mO_cur, (1, elems_per_store,))
gmem_thr_copy_O = gmem_tiled_copy_O.get_slice(tidx)
# Write final results
for m in cutlass.range(num_rows, unroll_full=True):
if tOhidx[m] >= 0:
mO_cur_copy = cute.tiled_divide(
mO_cur[tOmidx[m], None, tOhidx[m]], (elems_per_store,)
)
for k in cutlass.range(cute.size(tOcO, mode=[2]), unroll_full=True):
k_idx = tOcO[0, 0, k][1] // elems_per_store
if const_expr(self.is_even_k) or tOpO[k]:
cute.copy(gmem_thr_copy_O, rO[None, m, k], mO_cur_copy[None, k_idx])
@cute.jit
def load_O_partial(
self,
gmem_tiled_copy_O_partial: cute.TiledCopy,
tOrOptr: cute.Tensor,
tOsO_partial: cute.Tensor,
tOhidx: cute.Tensor,
tOpO: Optional[cute.Tensor],
tOcO: cute.Tensor,
mO_cur_partial_layout: cute.Layout,
split: Int32,
stage: Int32,
) -> None:
elems_per_load = const_expr(cute.size(gmem_tiled_copy_O_partial.layout_tv_tiled[1]))
tOsO_partial_cur = tOsO_partial[None, None, None, stage]
for m in cutlass.range(cute.size(tOcO, [1]), unroll_full=True):
if tOhidx[m] >= 0:
o_gmem_ptr = cute.make_ptr(
tOsO_partial.element_type, tOrOptr[m], cute.AddressSpace.gmem, assumed_align=16
)
mO_partial_cur = cute.make_tensor(
o_gmem_ptr, cute.slice_(mO_cur_partial_layout, (0, None, None, 0))
)
mO_partial_cur_copy = cute.tiled_divide(mO_partial_cur, (elems_per_load,))
for k in cutlass.range(cute.size(tOcO, mode=[2]), unroll_full=True):
k_idx = tOcO[0, 0, k][1] // elems_per_load
if const_expr(tOpO is None) or tOpO[k]:
cute.copy(
gmem_tiled_copy_O_partial,
mO_partial_cur_copy[None, k_idx, split],
tOsO_partial_cur[None, m, k],
)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+59
View File
@@ -0,0 +1,59 @@
# Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
# SM120 (Blackwell GeForce / DGX Spark) forward pass.
#
# SM120 uses the same SM80-era MMA instructions (mma.sync.aligned.m16n8k16) but has
# a smaller shared memory capacity (99 KB vs 163 KB on SM80). This module subclasses
# FlashAttentionForwardSm80 and overrides the SMEM capacity check accordingly.
import cutlass
import cutlass.utils as utils_basic
from flash_attn.cute.flash_fwd import FlashAttentionForwardSm80
class FlashAttentionForwardSm120(FlashAttentionForwardSm80):
# Keep arch = 80 to use CpAsync code paths (no TMA for output).
# The compilation target is determined by the GPU at compile time, not this field.
arch = 80
@staticmethod
def can_implement(
dtype,
head_dim,
head_dim_v,
tile_m,
tile_n,
num_stages,
num_threads,
is_causal,
Q_in_regs=False,
) -> bool:
"""Check if the kernel can be implemented on SM120.
Same logic as SM80 but uses SM120's shared memory capacity (99 KB).
"""
if dtype not in [cutlass.Float16, cutlass.BFloat16]:
return False
if head_dim % 8 != 0:
return False
if head_dim_v % 8 != 0:
return False
if tile_n % 16 != 0:
return False
if num_threads % 32 != 0:
return False
# Shared memory usage: Q tile + (K tile + V tile)
smem_usage_Q = tile_m * head_dim * 2
smem_usage_K = tile_n * head_dim * num_stages * 2
smem_usage_V = tile_n * head_dim_v * num_stages * 2
smem_usage_QV = (
(smem_usage_Q + smem_usage_V) if not Q_in_regs else max(smem_usage_Q, smem_usage_V)
)
smem_usage = smem_usage_QV + smem_usage_K
# SM120 has 99 KB shared memory (vs 163 KB on SM80)
smem_capacity = utils_basic.get_smem_capacity_in_bytes("sm_120")
if smem_usage > smem_capacity:
return False
if (tile_m * 2) % num_threads != 0:
return False
return True
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+296
View File
@@ -0,0 +1,296 @@
# Copyright (c) 2025, Tri Dao.
# Ported Cutlass code from C++ to Python:
# https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/mma_sm100_desc.hpp
# https://github.com/NVIDIA/cutlass/blob/main/include/cute/atom/mma_traits_sm100.hpp
from enum import IntEnum
import cutlass
import cutlass.cute as cute
# ---------------------------------------------------------------------------
# Enumerations that match the HW encodings (values MUST stay identical)
# ---------------------------------------------------------------------------
class Major(IntEnum): # matrix “layout” in the ISA docs
K = 0
MN = 1
class ScaleIn(IntEnum): # negate flags
One = 0
Neg = 1
class Saturate(IntEnum):
False_ = 0
True_ = 1
class CFormat(IntEnum): # 2-bit field (bits 4-5)
F16 = 0
F32 = 1
S32 = 2
class F16F32Format(IntEnum): # 3-bit field (A/B element type)
F16 = 0
BF16 = 1
TF32 = 2
class S8Format(IntEnum):
UINT8 = 0
INT8 = 1
class MXF8F6F4Format(IntEnum):
E4M3 = 0
E5M2 = 1
E2M3 = 3
E3M2 = 4
E2M1 = 5
class MaxShift(IntEnum):
NoShift = 0
MaxShift8 = 1
MaxShift16 = 2
MaxShift32 = 3
# ---------------------------------------------------------------------------
# CUTLASS-type → encoding helpers
# ---------------------------------------------------------------------------
def to_UMMA_format(cutlass_type) -> int:
"""
Map a CUTLASS scalar class to the 3-bit encoding for Matrix A/B.
"""
if cutlass_type is cutlass.Int8:
return S8Format.INT8
# Unsigned 8-bit (if available in your CUTLASS build)
if cutlass_type is cutlass.Uint8:
return S8Format.UINT8
# FP-16 / BF-16
if cutlass_type is cutlass.Float16:
return F16F32Format.F16
if cutlass_type is cutlass.BFloat16:
return F16F32Format.BF16
# TensorFloat-32 (8-bit exponent, 10-bit mantissa packed in 19 bits)
if cutlass_type is cutlass.TFloat32:
return F16F32Format.TF32
# Float-8 / Float-6 / Float-4 – add whenever CUTLASS exposes them
if cutlass_type is cutlass.Float8E4M3FN:
return MXF8F6F4Format.E4M3
if cutlass_type is cutlass.Float8E5M2:
return MXF8F6F4Format.E5M2
raise TypeError(f"Unsupported CUTLASS scalar type for A/B: {cutlass_type!r}")
def to_C_format(cutlass_type) -> int:
"""
Map a CUTLASS scalar class to the 2-bit accumulator encoding.
"""
if cutlass_type is cutlass.Float16:
return CFormat.F16
if cutlass_type is cutlass.Float32:
return CFormat.F32
if cutlass_type is cutlass.Int32:
return CFormat.S32
raise TypeError(f"Unsupported CUTLASS scalar type for accumulator: {cutlass_type!r}")
# ---------------------------------------------------------------------------
# The constructor – accepts only CUTLASS scalar classes
# ---------------------------------------------------------------------------
def make_instr_desc(
a_type, # CUTLASS scalar class, e.g. cutlass.Int8
b_type,
c_type,
M: int, # 64, 128 or 256
N: int, # 8 … 256 (multiple of 8)
a_major: Major,
b_major: Major,
a_neg: ScaleIn = ScaleIn.One,
b_neg: ScaleIn = ScaleIn.One,
c_sat: Saturate = Saturate.False_,
is_sparse: bool = False,
max_shift: MaxShift = MaxShift.NoShift,
) -> int:
"""
Build the 32-bit instruction descriptor for Blackwell MMA.
All matrix/accumulator **types must be CUTLASS scalar classes** –
passing integers is forbidden.
"""
# --- encode element formats -------------------------------------------------
a_fmt = int(to_UMMA_format(a_type))
b_fmt = int(to_UMMA_format(b_type))
c_fmt = int(to_C_format(c_type))
# --- range checks on M/N -----------------------------------------------------
if M not in (64, 128, 256):
raise ValueError("M must be 64, 128 or 256")
if N < 8 or N > 256 or (N & 7):
raise ValueError("N must be a multiple of 8 in the range 8…256")
m_dim = M >> 4 # 5-bit field
n_dim = N >> 3 # 6-bit field
# fmt: off
# --- pack the bit-fields -----------------------------------------------------
desc = 0
desc |= (0 & 0x3) << 0 # sparse_id2 (always 0 here)
desc |= (int(is_sparse) & 0x1) << 2 # sparse_flag
desc |= (int(c_sat) & 0x1) << 3 # saturate
desc |= (c_fmt & 0x3) << 4 # c_format
desc |= (a_fmt & 0x7) << 7 # a_format
desc |= (b_fmt & 0x7) << 10 # b_format
desc |= (int(a_neg) & 0x1) << 13 # a_negate
desc |= (int(b_neg) & 0x1) << 14 # b_negate
desc |= (int(a_major) & 0x1) << 15 # a_major
desc |= (int(b_major) & 0x1) << 16 # b_major
desc |= (n_dim & 0x3F) << 17 # n_dim (6 bits)
desc |= (m_dim & 0x1F) << 24 # m_dim (5 bits)
desc |= (int(max_shift) & 0x3) << 30 # max_shift (2 bits)
# fmt: on
return desc & 0xFFFF_FFFF # ensure 32-bit result
def mma_op_to_idesc(op: cute.nvgpu.tcgen05.mma.MmaOp):
return make_instr_desc(
op.a_dtype,
op.b_dtype,
op.acc_dtype,
op.shape_mnk[0],
op.shape_mnk[1],
Major.K if op.a_major_mode == cute.nvgpu.tcgen05.mma.OperandMajorMode.K else Major.MN,
Major.K if op.b_major_mode == cute.nvgpu.tcgen05.mma.OperandMajorMode.K else Major.MN,
)
class LayoutType(IntEnum): # occupies the top-3 bits [61:64)
SWIZZLE_NONE = 0 # (a.k.a. “INTERLEAVE” in older docs)
SWIZZLE_128B_BASE32B = 1
SWIZZLE_128B = 2
SWIZZLE_64B = 4
SWIZZLE_32B = 6
# values 3,5,7 are reserved / illegal for UMMA
# ---------------------------------------------------------------------------
# Helpers – figure out the SWIZZLE_* family from the tensor layout
# ---------------------------------------------------------------------------
def _layout_type(swizzle: cute.Swizzle) -> LayoutType:
B, M, S = swizzle.num_bits, swizzle.num_base, swizzle.num_shift
if M == 4: # Swizzle<*,4,3>
if S != 3:
raise ValueError("Unexpected swizzle shift – want S==3 for M==4")
return {
0: LayoutType.SWIZZLE_NONE,
1: LayoutType.SWIZZLE_32B,
2: LayoutType.SWIZZLE_64B,
3: LayoutType.SWIZZLE_128B,
}[B] # KeyError ⇒ invalid B→ raise
if M == 5: # Swizzle<2,5,2> (the only legal triple for M==5)
if (B, S) != (2, 2):
raise ValueError("Only Swizzle<2,5,2> supported for 128B_BASE32B")
return LayoutType.SWIZZLE_128B_BASE32B
# Any other (M,B,S) triple is not a UMMA-legal shared-memory layout
raise ValueError("Unsupported swizzle triple for UMMA smem descriptor")
def make_smem_desc_base(layout: cute.Layout, swizzle: cute.Swizzle, major: Major) -> int:
"""
Convert a 2-D *shared-memory* Cute layout into the Blackwell 64-bit
smem-descriptor, without the smem start address.
layout must correspond to layout of an uint128 tensor.
"""
# ------------------------------------------------------------------ meta
layout_type = _layout_type(swizzle) # resolve SWIZZLE_* family
VERSION = 1 # bits 46–47
LBO_MODE = 0 # bit 52
BASE_OFFSET = 0 # bits 49–51 (CUTLASS always 0)
# ---------------------------------------------------------- strides (units: uint128_t = 16 B)
swizzle_atom_mn_size = {
LayoutType.SWIZZLE_NONE: 1,
LayoutType.SWIZZLE_32B: 2,
LayoutType.SWIZZLE_64B: 4,
LayoutType.SWIZZLE_128B: 8,
LayoutType.SWIZZLE_128B_BASE32B: 8,
}[layout_type]
if major is Major.MN:
swizzle_atom_k_size = 4 if layout_type is LayoutType.SWIZZLE_128B_BASE32B else 8
canonical_layout = cute.logical_divide(layout, (swizzle_atom_mn_size, swizzle_atom_k_size))
if not cute.is_congruent(canonical_layout, ((1, 1), (1, 1))):
raise ValueError("Not a canonical UMMA_MN Layout: Expected profile failure.")
stride_00 = canonical_layout.stride[0][0]
if layout_type is not LayoutType.SWIZZLE_NONE and stride_00 != 1:
raise ValueError("Not a canonical UMMA_MN Layout: Expected stride failure.")
stride_10 = canonical_layout.stride[1][0]
if stride_10 != swizzle_atom_mn_size:
raise ValueError("Not a canonical UMMA_MN Layout: Expected stride failure.")
stride_01, stride_11 = canonical_layout.stride[0][1], canonical_layout.stride[1][1]
if layout_type is LayoutType.SWIZZLE_NONE:
stride_byte_offset, leading_byte_offset = stride_01, stride_11
else:
stride_byte_offset, leading_byte_offset = stride_11, stride_01
else:
if layout_type == LayoutType.SWIZZLE_128B_BASE32B:
raise ValueError("SWIZZLE_128B_BASE32B is invalid for Major-K")
if not cute.size(layout.shape[0]) % 8 == 0:
raise ValueError("Not a canonical UMMA_K Layout: Expected MN-size multiple of 8.")
canonical_layout = cute.logical_divide(layout, (8, 2))
if not cute.is_congruent(canonical_layout, ((1, 1), (1, 1))):
raise ValueError("Not a canonical UMMA_K Layout: Expected profile failure.")
stride_00 = canonical_layout.stride[0][0]
if stride_00 != swizzle_atom_mn_size:
raise ValueError("Not a canonical UMMA_K Layout: Expected stride failure.")
stride_10 = canonical_layout.stride[1][0]
if layout_type is not LayoutType.SWIZZLE_NONE and stride_10 != 1:
raise ValueError("Not a canonical UMMA_K Layout: Expected stride failure.")
stride_01 = canonical_layout.stride[0][1]
stride_byte_offset, leading_byte_offset = stride_01, stride_10
# ------------------------------------------------------------------ pack
desc = 0
# leading_byte_offset_ [16:30)
desc |= (leading_byte_offset & 0x3FFF) << 16
# stride_byte_offset_ [32:46)
desc |= (stride_byte_offset & 0x3FFF) << 32
# version_ [46:48)
desc |= (VERSION & 0x3) << 46
# base_offset_ [49:52)
desc |= (BASE_OFFSET & 0x7) << 49
# lbo_mode_ [52:53)
desc |= (LBO_MODE & 0x1) << 52
# layout_type_ [61:64)
desc |= (int(layout_type) & 0x7) << 61
return desc & 0xFFFF_FFFF_FFFF_FFFF # force 64-bit width
def make_smem_desc_start_addr(start_addr: cute.Pointer) -> cutlass.Int32:
# 14 bits, remove 4 LSB (bits 0-13 in desc)
return (start_addr.toint() & 0x3FFFF) >> 4
def smem_desc_base_from_tensor(sA: cute.Tensor, major: Major) -> int:
sA_swizzle = sA.iterator.type.swizzle_type
return make_smem_desc_base(
cute.recast_layout(128, sA.element_type.width, sA.layout[0]),
sA_swizzle,
major,
)
+68
View File
@@ -0,0 +1,68 @@
# Copyright (c) 2025, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
import enum
class NamedBarrierFwd(enum.IntEnum):
Epilogue = enum.auto() # starts from 1 as barrier 0 is reserved for sync_threads()
WarpSchedulerWG1 = enum.auto()
WarpSchedulerWG2 = enum.auto()
WarpSchedulerWG3 = enum.auto()
PFull = enum.auto()
PEmpty = enum.auto()
class NamedBarrierFwdSm100(enum.IntEnum):
Epilogue = enum.auto() # starts from 1 as barrier 0 is reserved for sync_threads()
TmemPtr = enum.auto()
SoftmaxStatsW0 = enum.auto()
SoftmaxStatsW1 = enum.auto()
SoftmaxStatsW2 = enum.auto()
SoftmaxStatsW3 = enum.auto()
SoftmaxStatsW4 = enum.auto()
SoftmaxStatsW5 = enum.auto()
SoftmaxStatsW6 = enum.auto()
SoftmaxStatsW7 = enum.auto()
# Reserve one independent SMEM-P visibility barrier per compute slot. A
# single shared barrier can mix arrivals from the two softmax warpgroups.
SoftmaxSmemP0 = enum.auto()
SoftmaxSmemP1 = enum.auto()
CorrectionScale = enum.auto()
class NamedBarrierBwd(enum.IntEnum):
Epilogue = enum.auto()
WarpSchedulerWG1 = enum.auto()
WarpSchedulerWG2 = enum.auto()
WarpSchedulerWG3 = enum.auto()
PdS = enum.auto()
dQFullWG0 = enum.auto()
dQFullWG1 = enum.auto()
dQFullWG2 = enum.auto()
dQEmptyWG0 = enum.auto()
dQEmptyWG1 = enum.auto()
dQEmptyWG2 = enum.auto()
class NamedBarrierBwdSm100(enum.IntEnum):
EpilogueWG1 = enum.auto()
EpilogueWG2 = enum.auto()
Compute = enum.auto()
dQaccReduce = enum.auto()
TmemPtr = enum.auto()
class NamedBarrierFwdSm100_MLA2CTA(enum.IntEnum):
Epilogue = enum.auto()
TmemPtr = enum.auto()
Cpasync = enum.auto()
Softmax = enum.auto()
SoftmaxStatsFull = enum.auto()
SoftmaxStatsEmpty = enum.auto()
class NamedBarrierBwdSm100_MLA2CTA(enum.IntEnum):
Epilogue = enum.auto()
TmemPtr = enum.auto()
Cpasync = enum.auto()
Softmax = enum.auto()
+263
View File
@@ -0,0 +1,263 @@
# Copyright (c) 2025, Tri Dao.
from dataclasses import dataclass
from typing import Union, Tuple
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync
from quack import layout_utils
import flash_attn.cute.utils as utils
def pack_gqa_layout(T, qhead_per_kvhead, nheads_kv, head_idx):
"""Reshape a tensor to fold qhead_per_kvhead into the seqlen dimension (mode 0).
The head dimension is at mode ``head_idx``. Modes before it (1..head_idx-1)
are kept as-is (e.g. headdim for Q/O tensors), and modes after it are kept
as-is (e.g. batch).
For Q/O tensors (head_idx=2):
(seqlen_q, headdim, nheads, batch, ...) -> ((qhead_per_kvhead, seqlen_q), headdim, nheads_kv, batch, ...)
For LSE tensors (head_idx=1):
(seqlen_q, nheads, batch, ...) -> ((qhead_per_kvhead, seqlen_q), nheads_kv, batch, ...)
"""
head_stride = T.stride[head_idx]
shape_packed = (
(qhead_per_kvhead, T.shape[0]),
*[T.shape[i] for i in range(1, head_idx)],
nheads_kv,
*[T.shape[i] for i in range(head_idx + 1, len(T.shape))],
)
stride_packed = (
(head_stride, T.stride[0]),
*[T.stride[i] for i in range(1, head_idx)],
head_stride * qhead_per_kvhead,
*[T.stride[i] for i in range(head_idx + 1, len(T.shape))],
)
return cute.make_tensor(T.iterator, cute.make_layout(shape_packed, stride=stride_packed))
def make_packgqa_tiled_tma_atom(
op: cute.atom.CopyOp,
gmem_tensor: cute.Tensor,
smem_layout: Union[cute.Layout, cute.ComposedLayout],
cta_tiler: Tuple[int, int],
qhead_per_kvhead: int,
head_idx: int,
):
# This packing and unpacking of the layout is so that we keep the same TMA dimension as usual.
# e.g. for (seqlen, d, nheads, b) layout, we still have 4D TMA after packing to
# ((nheads, seqlen), d, b).
# If we instead pack directly to ((qhead_per_kvhead, seqlen), d, nheads_kv, b) we'd have 5D TMA.
# Pack headdim and seqlen dim into 1: (seqlen, d, nheads, b) -> ((nheads, seqlen), d, b)
gmem_tensor = layout_utils.select(
gmem_tensor, [head_idx, *range(head_idx), *range(head_idx + 1, cute.rank(gmem_tensor))]
)
gmem_tensor = cute.group_modes(gmem_tensor, 0, 2)
assert cta_tiler[0] % qhead_per_kvhead == 0, (
"CTA tile size in the seqlen dimension must be divisible by qhead_per_kvhead"
)
tma_atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
gmem_tensor,
smem_layout,
((qhead_per_kvhead, cta_tiler[0] // qhead_per_kvhead), cta_tiler[1]), # No mcast
)
# Unpack from ((nheads, seqlen), d, b) -> ((qhead_per_kvhead, seqlen), d, nheads_kv, b)
T = tma_tensor
shape_packed = (
(qhead_per_kvhead, T.shape[0][1]),
*[T.shape[i] for i in range(1, head_idx)],
T.shape[0][0] // qhead_per_kvhead,
*[T.shape[i] for i in range(head_idx, len(T.shape))],
)
stride_packed = (
*[T.stride[i] for i in range(head_idx)],
T.stride[0][0] * qhead_per_kvhead,
*[T.stride[i] for i in range(head_idx, len(T.shape))],
)
tma_tensor = cute.make_tensor(T.iterator, cute.make_layout(shape_packed, stride=stride_packed))
return tma_atom, tma_tensor
def unpack_gqa_layout(T, qhead_per_kvhead, head_idx):
"""Reverse of pack_gqa_layout: unfold qhead_per_kvhead from the seqlen dimension (mode 0).
The head dimension is at mode ``head_idx``. Modes before it (1..head_idx-1)
are kept as-is (e.g. headdim for Q/O tensors), and modes after it are kept
as-is (e.g. batch).
For Q/O tensors (head_idx=2):
((qhead_per_kvhead, seqlen_q), headdim, nheads_kv, batch, ...) -> (seqlen_q, headdim, nheads, batch, ...)
For LSE tensors (head_idx=1):
((qhead_per_kvhead, seqlen_q), nheads_kv, batch, ...) -> (seqlen_q, nheads, batch, ...)
"""
seqlen_stride = T.stride[0][1]
head_stride = T.stride[0][0]
shape_unpacked = (
T.shape[0][1],
*[T.shape[i] for i in range(1, head_idx)],
T.shape[head_idx] * qhead_per_kvhead,
*[T.shape[i] for i in range(head_idx + 1, len(T.shape))],
)
stride_unpacked = (
seqlen_stride,
*[T.stride[i] for i in range(1, head_idx)],
head_stride,
*[T.stride[i] for i in range(head_idx + 1, len(T.shape))],
)
return cute.make_tensor(T.iterator, cute.make_layout(shape_unpacked, stride=stride_unpacked))
@dataclass
class PackGQA:
m_block_size: cutlass.Constexpr[int]
head_dim_padded: cutlass.Constexpr[int]
check_hdim_oob: cutlass.Constexpr[bool]
qhead_per_kvhead: cutlass.Constexpr[bool]
@cute.jit
def compute_ptr(
self,
tensor: cute.Tensor,
cRows: cute.Tensor,
tidx: cutlass.Int32,
block: cutlass.Int32,
threads_per_row: cutlass.Constexpr[int],
num_threads: cutlass.Constexpr[int],
):
num_ptr_per_thread = cute.ceil_div(cute.size(cRows), threads_per_row)
tPrPtr = cute.make_rmem_tensor(num_ptr_per_thread, cutlass.Int64)
for i in cutlass.range_constexpr(num_ptr_per_thread):
row = i * num_threads + cRows[tidx % threads_per_row][0]
idx = block * self.m_block_size + row
m_idx = idx // self.qhead_per_kvhead
h_idx = idx - m_idx * self.qhead_per_kvhead
tPrPtr[i] = utils.elem_pointer(tensor, ((h_idx, m_idx),)).toint()
return tPrPtr
@cute.jit
def load_Q(
self,
mQ: cute.Tensor, # ((qhead_per_kvhead, seqlen_q), headdim)
sQ: cute.Tensor, # (m_block_size, head_dim_padded)
gmem_tiled_copy: cute.TiledCopy,
tidx: cutlass.Int32,
block: cutlass.Int32,
seqlen: cutlass.Int32,
):
gmem_thr_copy = gmem_tiled_copy.get_slice(tidx)
cQ = cute.make_identity_tensor((self.m_block_size, self.head_dim_padded))
tQsQ = gmem_thr_copy.partition_D(sQ)
tQcQ = gmem_thr_copy.partition_S(cQ)
t0QcQ = gmem_thr_copy.get_slice(0).partition_S(cQ)
tQpQ = utils.predicate_k(tQcQ, limit=mQ.shape[1])
tQcQ_row = tQcQ[0, None, 0]
threads_per_row = gmem_tiled_copy.layout_tv_tiled.shape[0][0]
assert cute.arch.WARP_SIZE % threads_per_row == 0, "threads_per_row must divide WARP_SIZE"
num_threads = gmem_tiled_copy.size
tPrQPtr = self.compute_ptr(mQ[None, 0], tQcQ_row, tidx, block, threads_per_row, num_threads)
for m in cutlass.range_constexpr(cute.size(tQsQ.shape[1])):
q_ptr_i64 = utils.shuffle_sync(
tPrQPtr[m // threads_per_row], m % threads_per_row, width=threads_per_row
)
q_gmem_ptr = cute.make_ptr(
mQ.element_type, q_ptr_i64, cute.AddressSpace.gmem, assumed_align=16
)
if (
t0QcQ[0, m, 0][0]
< seqlen * self.qhead_per_kvhead - block * self.m_block_size - tQcQ_row[0][0]
):
mQ_cur = cute.make_tensor(q_gmem_ptr, (self.head_dim_padded,))
elems_per_load = cute.size(tQsQ.shape[0][0])
mQ_cur_copy = cute.tiled_divide(mQ_cur, (elems_per_load,))
for k in cutlass.range_constexpr(cute.size(tQsQ.shape[2])):
ki = tQcQ[0, 0, k][1] // elems_per_load
cute.copy(
gmem_thr_copy,
mQ_cur_copy[None, ki],
tQsQ[None, m, k],
pred=tQpQ[None, m, k] if cutlass.const_expr(self.check_hdim_oob) else None,
)
# We don't need to clear the sQ smem tiles since we'll only write out the valid outputs
@cute.jit
def store_LSE(
self,
mLSE: cute.Tensor, # (qhead_per_kvhead, seqlen_q)
tLSErLSE: cute.Tensor, # (m_block_size, head_dim_padded)
tiled_mma: cute.TiledMma,
tidx: cutlass.Int32,
block: cutlass.Int32,
seqlen: cutlass.Int32,
):
thr_mma = tiled_mma.get_slice(tidx)
caccO = cute.make_identity_tensor((self.m_block_size, self.head_dim_padded))
taccOcO = thr_mma.partition_C(caccO)
taccOcO_row = layout_utils.reshape_acc_to_mn(taccOcO)[None, 0]
assert cute.size(tLSErLSE) == cute.size(taccOcO_row)
threads_per_row = tiled_mma.tv_layout_C.shape[0][0]
assert cute.arch.WARP_SIZE % threads_per_row == 0, "threads_per_row must divide WARP_SIZE"
assert cute.size(tLSErLSE) <= threads_per_row
num_threads = tiled_mma.size
tPrLSEPtr = self.compute_ptr(mLSE, taccOcO_row, tidx, block, threads_per_row, num_threads)
for m in cutlass.range_constexpr(cute.size(tLSErLSE)):
lse_ptr_i64 = utils.shuffle_sync(
tPrLSEPtr[m // threads_per_row],
m % threads_per_row,
width=threads_per_row,
)
lse_gmem_ptr = cute.make_ptr(
mLSE.element_type, lse_ptr_i64, cute.AddressSpace.gmem, assumed_align=4
)
row = block * self.m_block_size + taccOcO_row[m][0]
# Only the thread corresponding to column 0 writes out the lse to gmem
if taccOcO[0][1] == 0 and row < seqlen * self.qhead_per_kvhead:
mLSE_copy = cute.make_tensor(lse_gmem_ptr, (1,))
mLSE_copy[0] = tLSErLSE[m]
@cute.jit
def store_O(
self,
mO: cute.Tensor, # ((qhead_per_kvhead, seqlen_q), headdim)
tOrO: cute.Tensor, # (m_block_size, head_dim_padded) split across threads according to gmem_tiled_copy
gmem_tiled_copy: cute.TiledCopy,
tidx: cutlass.Int32,
block: cutlass.Int32,
seqlen: cutlass.Int32,
):
gmem_thr_copy = gmem_tiled_copy.get_slice(tidx)
cO = cute.make_identity_tensor((self.m_block_size, self.head_dim_padded))
tOcO = gmem_thr_copy.partition_S(cO)
t0OcO = gmem_thr_copy.get_slice(0).partition_S(cO)
tOpO = utils.predicate_k(tOcO, limit=mO.shape[1])
tOcO_row = tOcO[0, None, 0]
threads_per_row = gmem_tiled_copy.layout_tv_tiled.shape[0][0]
assert cute.arch.WARP_SIZE % threads_per_row == 0, "threads_per_row must divide WARP_SIZE"
num_threads = gmem_tiled_copy.size
tPrOPtr = self.compute_ptr(mO[None, 0], tOcO_row, tidx, block, threads_per_row, num_threads)
for m in cutlass.range_constexpr(cute.size(tOrO.shape[1])):
o_ptr_i64 = utils.shuffle_sync(
tPrOPtr[m // threads_per_row], m % threads_per_row, width=threads_per_row
)
o_gmem_ptr = cute.make_ptr(
mO.element_type, o_ptr_i64, cute.AddressSpace.gmem, assumed_align=16
)
if (
t0OcO[0, m, 0][0]
< seqlen * self.qhead_per_kvhead - block * self.m_block_size - tOcO_row[0][0]
):
mO_cur = cute.make_tensor(o_gmem_ptr, (self.head_dim_padded,))
elems_per_load = cute.size(tOrO.shape[0][0])
mO_cur_copy = cute.tiled_divide(mO_cur, (elems_per_load,))
for k in cutlass.range_constexpr(cute.size(tOrO.shape[2])):
ki = tOcO[0, 0, k][1] // elems_per_load
cute.copy(
gmem_thr_copy,
tOrO[None, m, k],
mO_cur_copy[None, ki],
pred=tOpO[None, m, k] if cutlass.const_expr(self.check_hdim_oob) else None,
)
+247
View File
@@ -0,0 +1,247 @@
from typing import Type
from dataclasses import dataclass
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync
from cutlass import Int32, const_expr
from flash_attn.cute import utils
from quack.cute_dsl_utils import ParamsBase
from cutlass.cute import FastDivmodDivisor
import math
@dataclass
class PagedKVManager(ParamsBase):
mPageTable: cute.Tensor
mK_paged: cute.Tensor
mV_paged: cute.Tensor
thread_idx: Int32
page_size_divmod: FastDivmodDivisor
seqlen_k: Int32
leftpad_k: Int32
n_block_size: Int32
num_threads: cutlass.Constexpr[Int32]
head_dim_padded: cutlass.Constexpr[Int32]
head_dim_v_padded: cutlass.Constexpr[Int32]
arch: cutlass.Constexpr[Int32]
v_gmem_transposed: cutlass.Constexpr[bool]
gmem_threads_per_row: cutlass.Constexpr[Int32]
page_entry_per_thread: Int32
async_copy_elems: Int32
gmem_tiled_copy_KV: cute.TiledCopy
gmem_thr_copy_KV: cute.TiledCopy
tPrPage: cute.Tensor
tPrPageOffset: cute.Tensor
tKpK: cute.Tensor
tVpV: cute.Tensor
@staticmethod
def create(
mPageTable: cute.Tensor,
mK_paged: cute.Tensor,
mV_paged: cute.Tensor,
page_size_divmod: FastDivmodDivisor,
bidb: Int32,
bidh: Int32,
thread_idx: Int32,
seqlen_k: Int32,
leftpad_k: Int32,
n_block_size: cutlass.Constexpr[Int32],
head_dim_padded: cutlass.Constexpr[Int32],
head_dim_v_padded: cutlass.Constexpr[Int32],
num_threads: cutlass.Constexpr[Int32],
dtype: Type[cutlass.Numeric],
arch: cutlass.Constexpr[int] = 100,
):
# SM100 transposes V in gmem to (dv, page_size, num_pages);
# SM90 keeps V as (page_size, dv, num_pages), same layout as K.
v_gmem_transposed = arch != 90
universal_copy_bits = 128
async_copy_elems = universal_copy_bits // dtype.width
dtype_bytes = dtype.width // 8
gmem_k_block_size = math.gcd(
head_dim_padded,
head_dim_v_padded,
128 // dtype_bytes,
)
assert gmem_k_block_size % async_copy_elems == 0
gmem_threads_per_row = gmem_k_block_size // async_copy_elems
assert cute.arch.WARP_SIZE % gmem_threads_per_row == 0
atom_async_copy = cute.make_copy_atom(
cpasync.CopyG2SOp(cache_mode=cpasync.LoadCacheMode.GLOBAL),
dtype,
num_bits_per_copy=universal_copy_bits,
)
thr_layout = cute.make_ordered_layout(
(num_threads // gmem_threads_per_row, gmem_threads_per_row),
order=(1, 0),
)
val_layout = cute.make_layout((1, async_copy_elems))
gmem_tiled_copy_KV = cute.make_tiled_copy_tv(atom_async_copy, thr_layout, val_layout)
gmem_thr_copy_KV = gmem_tiled_copy_KV.get_slice(thread_idx)
page_entry_per_thread = n_block_size // num_threads
tPrPage = cute.make_rmem_tensor((page_entry_per_thread,), Int32)
tPrPageOffset = cute.make_rmem_tensor((page_entry_per_thread,), Int32)
mPageTable = mPageTable[bidb, None]
mK_paged = mK_paged[None, None, bidh, None]
mV_paged = mV_paged[None, None, bidh, None]
cK = cute.make_identity_tensor((n_block_size, head_dim_padded))
tKcK = gmem_thr_copy_KV.partition_S(cK)
tKpK = utils.predicate_k(tKcK, limit=mK_paged.shape[1])
if const_expr(head_dim_padded == head_dim_v_padded):
tVpV = tKpK
else:
cV = cute.make_identity_tensor((n_block_size, head_dim_v_padded))
tVcV = gmem_thr_copy_KV.partition_S(cV)
# When V is transposed in gmem, dv is shape[0]; otherwise dv is shape[1] (same as K)
V_limit = cute.size(mV_paged.shape[0 if v_gmem_transposed else 1])
tVpV = utils.predicate_k(tVcV, limit=V_limit)
return PagedKVManager(
mPageTable,
mK_paged,
mV_paged,
thread_idx,
page_size_divmod,
seqlen_k,
leftpad_k,
n_block_size,
num_threads,
head_dim_padded,
head_dim_v_padded,
arch,
v_gmem_transposed,
gmem_threads_per_row,
page_entry_per_thread,
async_copy_elems,
gmem_tiled_copy_KV,
gmem_thr_copy_KV,
tPrPage,
tPrPageOffset,
tKpK,
tVpV,
)
@cute.jit
def load_page_table(self, n_block: Int32):
for i in cutlass.range(self.page_entry_per_thread, unroll=1):
row = (
i * self.num_threads
+ (self.thread_idx % self.gmem_threads_per_row)
* (self.num_threads // self.gmem_threads_per_row)
+ (self.thread_idx // self.gmem_threads_per_row)
)
row_idx = n_block * self.n_block_size + row
page_idx, page_offset = divmod(row_idx + self.leftpad_k, self.page_size_divmod)
is_valid = (
(i + 1) * self.num_threads <= self.n_block_size or row < self.n_block_size
) and row_idx < self.seqlen_k
page = self.mPageTable[page_idx] if is_valid else 0
self.tPrPage[i] = page
self.tPrPageOffset[i] = page_offset
@cute.jit
def compute_X_ptr(self, K_or_V: str, d_offset: int = 0):
tPrXPtr = cute.make_rmem_tensor((self.page_entry_per_thread,), cutlass.Int64)
mX = self.mK_paged if const_expr(K_or_V == "K") else self.mV_paged
# K is always (page_size, d, num_pages). V matches K when not transposed,
# but is (dv, page_size, num_pages) when transposed (SM100).
transposed = const_expr(K_or_V == "V" and self.v_gmem_transposed)
for i in cutlass.range(self.page_entry_per_thread, unroll=1):
page = self.tPrPage[i]
page_offset = self.tPrPageOffset[i]
if const_expr(transposed):
tPrXPtr[i] = utils.elem_pointer(mX, (d_offset, page_offset, page)).toint()
else:
tPrXPtr[i] = utils.elem_pointer(mX, (page_offset, d_offset, page)).toint()
return tPrXPtr
@cute.jit
def _flatten_smem_sm100(self, sX: cute.Tensor, K_or_V: str):
"""Flatten SM100 smem ((a,b), cta_split, k) to (a,(b,k)); transpose V to (d,page_size)."""
sX_pi = cute.make_tensor(
sX.iterator,
cute.make_layout(
(sX.shape[0][0], (sX.shape[0][1], sX.shape[2])),
stride=(sX.stride[0][0], (sX.stride[0][1], sX.stride[2])),
),
)
if const_expr(K_or_V == "V"):
sX_pi = cute.make_tensor(sX_pi.iterator, cute.select(sX_pi.layout, mode=[1, 0]))
return sX_pi
@cute.jit
def _copy_row_async(
self,
tXsX: cute.Tensor,
tXcX: cute.Tensor,
mX_paged_cur_copy: cute.Tensor,
m: Int32,
should_load: cute.Tensor,
):
"""Issue cp.async copies for one row across all k-tiles."""
for k in cutlass.range_constexpr(cute.size(tXsX, mode=[2])):
ki = tXcX[0, 0, k][1] // self.async_copy_elems
mX_paged_cur_copy_ki = mX_paged_cur_copy[None, ki]
tXsX_k = tXsX[None, m, k]
mX_paged_cur_copy_ki = cute.make_tensor(mX_paged_cur_copy_ki.iterator, tXsX_k.layout)
cute.copy(
self.gmem_tiled_copy_KV,
mX_paged_cur_copy_ki,
tXsX_k,
pred=should_load,
)
@cute.jit
def load_KV(self, n_block: Int32, sX: cute.Tensor, K_or_V: str):
assert K_or_V in ("K", "V")
tPrXPtr = self.compute_X_ptr(K_or_V)
if const_expr(self.arch == 90):
# SM90: sX is already stage-sliced by caller (sK[None, None, stage]).
# Flatten hierarchical modes to get (n_block_size, head_dim).
sX_pi = cute.group_modes(sX, 0, 1)
# SM90 does NOT transpose V here (it's transposed via utils.transpose_view before MMA)
else:
sX_pi = self._flatten_smem_sm100(sX, K_or_V)
head_dim = self.head_dim_v_padded if const_expr(K_or_V == "V") else self.head_dim_padded
cX = cute.make_identity_tensor((self.n_block_size, head_dim))
tXsX = self.gmem_thr_copy_KV.partition_D(sX_pi)
tXcX = self.gmem_thr_copy_KV.partition_S(cX)
tXc0X = self.gmem_thr_copy_KV.get_slice(0).partition_S(cX)
seqlenk_row_limit = (
self.seqlen_k - n_block * self.n_block_size - tXcX[0][0] if n_block >= 0 else 0
)
for m in cutlass.range_constexpr(cute.size(tXsX, mode=[1])):
row_valid = tXc0X[0, m, 0][0] < seqlenk_row_limit
should_load = cute.make_fragment_like(tXsX[(0, None), m, 0], cute.Boolean)
should_load.fill(row_valid)
x_ptr_i64 = utils.shuffle_sync(
tPrXPtr[m // self.gmem_threads_per_row],
m % self.gmem_threads_per_row,
width=self.gmem_threads_per_row,
)
x_gmem_ptr = cute.make_ptr(
self.mK_paged.element_type, x_ptr_i64, cute.AddressSpace.gmem, assumed_align=16
)
mX_paged_cur = cute.make_tensor(x_gmem_ptr, cute.make_layout((head_dim,)))
mX_paged_cur_copy = cute.tiled_divide(mX_paged_cur, (self.async_copy_elems,))
self._copy_row_async(tXsX, tXcX, mX_paged_cur_copy, m, should_load)
+402
View File
@@ -0,0 +1,402 @@
# Copyright (c) 2025, Tri Dao.
from typing import Optional
from dataclasses import dataclass
import cutlass.cute as cute
from cutlass import Boolean, Int32, const_expr
from cutlass.cutlass_dsl import if_generate, dsl_user_op
from cutlass.pipeline import PipelineState
from cutlass.pipeline import PipelineUserType
from cutlass.pipeline import NamedBarrier as NamedBarrierOg
from cutlass.pipeline import PipelineAsync as PipelineAsyncOg
from cutlass.pipeline import PipelineCpAsync as PipelineCpAsyncOg
from cutlass.pipeline import PipelineTmaAsync as PipelineTmaAsyncOg
from cutlass.pipeline import PipelineTmaUmma as PipelineTmaUmmaOg
from cutlass.pipeline import PipelineUmmaAsync as PipelineUmmaAsyncOg
from cutlass.pipeline import PipelineAsyncUmma as PipelineAsyncUmmaOg
def _override_create(parent_cls, child_cls):
"""Create a static factory that constructs parent_cls then re-classes to child_cls."""
@staticmethod
def create(*args, **kwargs):
obj = parent_cls.create(*args, **kwargs)
# Can't assign to __class__ directly since the dataclass is frozen
object.__setattr__(obj, "__class__", child_cls)
return obj
return create
def _make_state(index: Int32, phase: Int32) -> PipelineState:
"""Construct a PipelineState from index and phase (count/stages unused by callers)."""
return PipelineState(stages=0, count=Int32(0), index=index, phase=phase)
class PipelineStateSimple:
"""
Pipeline state contains an index and phase bit corresponding to the current position in the circular buffer.
Use a single Int32 to store both the index and phase bit, then we use divmod to get the
index and phase. If stages is a power of 2, divmod turns into bit twiddling.
"""
def __init__(self, stages: int, phase_index: Int32):
self._stages = stages
self._phase_index = phase_index
def clone(self) -> "PipelineStateSimple":
return PipelineStateSimple(self.stages, self._phase_index)
@property
def stages(self) -> int:
return self._stages
@property
def index(self) -> Int32:
if const_expr(self._stages == 1):
return Int32(0)
else:
return self._phase_index % self._stages
@property
def phase(self) -> Int32:
# PTX docs say that the phase parity needs to be 0 or 1, so by right we need to
# take modulo 2. But in practice just passing the phase in without modulo works fine.
if const_expr(self._stages == 1):
return self._phase_index
else:
return self._phase_index // self._stages
def advance(self):
if const_expr(self._stages == 1):
self._phase_index ^= 1
else:
self._phase_index += 1
def __extract_mlir_values__(self):
phase_index = self._phase_index
return [phase_index.ir_value()]
def __new_from_mlir_values__(self, values):
return PipelineStateSimple(self.stages, Int32(values[0]))
def make_pipeline_state(type: PipelineUserType, stages: int):
"""
Creates a pipeline state. Producers are assumed to start with an empty buffer and have a flipped phase bit of 1.
"""
if type is PipelineUserType.Producer:
return PipelineStateSimple(stages, Int32(stages))
elif type is PipelineUserType.Consumer:
return PipelineStateSimple(stages, Int32(0))
else:
assert False, "Error: invalid PipelineUserType specified for make_pipeline_state."
# ── Shared helpers ───────────────────────────────────────────────────────────
def _call_with_elect_one(parent_method, self, state, elect_one, syncwarp, loc, ip):
"""Optionally wrap a parent pipeline method call in sync_warp + elect_one."""
if const_expr(elect_one):
if const_expr(syncwarp):
cute.arch.sync_warp()
with cute.arch.elect_one():
parent_method(self, state, loc=loc, ip=ip)
else:
parent_method(self, state, loc=loc, ip=ip)
# ── Mixin: _w_index / _w_index_phase variants that delegate to parent ───────
# Each parent class has PipelineState-based methods (producer_acquire, producer_commit,
# consumer_wait, consumer_release). The _w_index_phase variants just construct a
# PipelineState from (index, phase) and delegate.
class _PipelineIndexPhaseMixin:
"""Mixin providing _w_index_phase / _w_index methods that delegate to PipelineState-based parents."""
@dsl_user_op
def producer_acquire_w_index_phase(
self,
index: Int32,
phase: Int32,
try_acquire_token: Optional[Boolean] = None,
*,
loc=None,
ip=None,
):
state = _make_state(index, phase)
# Call the parent's producer_acquire (which takes PipelineState)
self.producer_acquire(state, try_acquire_token, loc=loc, ip=ip)
@dsl_user_op
def producer_commit_w_index(self, index: Int32, *, loc=None, ip=None):
state = _make_state(index, Int32(0))
self.producer_commit(state, loc=loc, ip=ip)
@dsl_user_op
def consumer_wait_w_index_phase(
self,
index: Int32,
phase: Int32,
try_wait_token: Optional[Boolean] = None,
*,
loc=None,
ip=None,
):
state = _make_state(index, phase)
self.consumer_wait(state, try_wait_token, loc=loc, ip=ip)
@dsl_user_op
def consumer_release_w_index(self, index: Int32, *, loc=None, ip=None):
state = _make_state(index, Int32(0))
self.consumer_release(state, loc=loc, ip=ip)
# ── NamedBarrier ─────────────────────────────────────────────────────────────
@dataclass(frozen=True)
class NamedBarrier(NamedBarrierOg):
create = _override_create(NamedBarrierOg, None) # patched below
@dsl_user_op
def arrive_w_index(self, index: Int32, *, loc=None, ip=None) -> None:
"""
The aligned flavor of arrive is used when all threads in the CTA will execute the
same instruction. See PTX documentation.
"""
cute.arch.barrier_arrive(
barrier_id=self.barrier_id + index,
number_of_threads=self.num_threads,
loc=loc,
ip=ip,
)
@dsl_user_op
def arrive_and_wait_w_index(self, index: Int32, *, loc=None, ip=None) -> None:
cute.arch.barrier(
barrier_id=self.barrier_id + index,
number_of_threads=self.num_threads,
loc=loc,
ip=ip,
)
NamedBarrier.create = _override_create(NamedBarrierOg, NamedBarrier)
# ── PipelineAsync ────────────────────────────────────────────────────────────
@dataclass(frozen=True)
class PipelineAsync(_PipelineIndexPhaseMixin, PipelineAsyncOg):
"""
PipelineAsync with optional elect_one for producer_commit and consumer_release.
When elect_one_*=True (set at create time), only one elected thread per warp
signals the barrier arrive. This is useful when the mask count is set to 1 per warp.
Args (to create):
elect_one_commit: If True, only elected thread signals producer_commit.
syncwarp_before_commit: If True (default), issue syncwarp before elect_one.
elect_one_release: If True, only elected thread signals consumer_release.
syncwarp_before_release: If True (default), issue syncwarp before elect_one.
Set syncwarp to False when threads are already converged (e.g. after wgmma wait_group).
"""
_elect_one_commit: bool = False
_syncwarp_before_commit: bool = True
_elect_one_release: bool = False
_syncwarp_before_release: bool = True
@staticmethod
def create(
*args,
elect_one_commit: bool = False,
syncwarp_before_commit: bool = True,
elect_one_release: bool = False,
syncwarp_before_release: bool = True,
**kwargs,
):
obj = PipelineAsyncOg.create(*args, **kwargs)
object.__setattr__(obj, "__class__", PipelineAsync)
object.__setattr__(obj, "_elect_one_commit", elect_one_commit)
object.__setattr__(obj, "_syncwarp_before_commit", syncwarp_before_commit)
object.__setattr__(obj, "_elect_one_release", elect_one_release)
object.__setattr__(obj, "_syncwarp_before_release", syncwarp_before_release)
return obj
@dsl_user_op
def producer_commit(self, state: PipelineState, *, loc=None, ip=None):
_call_with_elect_one(
PipelineAsyncOg.producer_commit,
self,
state,
self._elect_one_commit,
self._syncwarp_before_commit,
loc,
ip,
)
@dsl_user_op
def consumer_release(self, state: PipelineState, *, loc=None, ip=None):
_call_with_elect_one(
PipelineAsyncOg.consumer_release,
self,
state,
self._elect_one_release,
self._syncwarp_before_release,
loc,
ip,
)
# _w_index variants inherited from _PipelineIndexPhaseMixin, which delegate
# to producer_commit / consumer_release above.
# ── PipelineCpAsync ──────────────────────────────────────────────────────────
@dataclass(frozen=True)
class PipelineCpAsync(_PipelineIndexPhaseMixin, PipelineCpAsyncOg):
_elect_one_release: bool = False
_syncwarp_before_release: bool = True
@staticmethod
def create(
*args,
elect_one_release: bool = False,
syncwarp_before_release: bool = True,
**kwargs,
):
obj = PipelineCpAsyncOg.create(*args, **kwargs)
object.__setattr__(obj, "__class__", PipelineCpAsync)
object.__setattr__(obj, "_elect_one_release", elect_one_release)
object.__setattr__(obj, "_syncwarp_before_release", syncwarp_before_release)
return obj
@dsl_user_op
def consumer_release(self, state: PipelineState, *, loc=None, ip=None):
_call_with_elect_one(
PipelineCpAsyncOg.consumer_release,
self,
state,
self._elect_one_release,
self._syncwarp_before_release,
loc,
ip,
)
# _w_index variants inherited from _PipelineIndexPhaseMixin.
# ── PipelineTmaAsync ────────────────────────────────────────────────────────
@dataclass(frozen=True)
class PipelineTmaAsync(_PipelineIndexPhaseMixin, PipelineTmaAsyncOg):
"""Override producer_acquire to take in extra_tx_count parameter."""
@dsl_user_op
def producer_acquire(
self,
state: PipelineState,
try_acquire_token: Optional[Boolean] = None,
extra_tx_count: int = 0,
*,
loc=None,
ip=None,
):
"""
TMA producer commit conditionally waits on buffer empty and sets the transaction barrier for leader threadblocks.
"""
if_generate(
try_acquire_token is None or try_acquire_token == 0,
lambda: self.sync_object_empty.wait(state.index, state.phase, loc=loc, ip=ip),
loc=loc,
ip=ip,
)
if const_expr(extra_tx_count == 0):
self.sync_object_full.arrive(state.index, self.producer_mask, loc=loc, ip=ip)
else:
tx_count = self.sync_object_full.tx_count + extra_tx_count
self.sync_object_full.arrive_and_expect_tx(state.index, tx_count, loc=loc, ip=ip)
PipelineTmaAsync.create = _override_create(PipelineTmaAsyncOg, PipelineTmaAsync)
# ── PipelineTmaUmma ─────────────────────────────────────────────────────────
@dataclass(frozen=True)
class PipelineTmaUmma(_PipelineIndexPhaseMixin, PipelineTmaUmmaOg):
"""Override producer_acquire to take in extra_tx_count parameter."""
@dsl_user_op
def producer_acquire(
self,
state: PipelineState,
try_acquire_token: Optional[Boolean] = None,
extra_tx_count: int = 0,
*,
loc=None,
ip=None,
):
"""
TMA producer commit conditionally waits on buffer empty and sets the transaction barrier for leader threadblocks.
"""
if_generate(
try_acquire_token is None or try_acquire_token == 0,
lambda: self.sync_object_empty.wait(state.index, state.phase, loc=loc, ip=ip),
loc=loc,
ip=ip,
)
if const_expr(extra_tx_count == 0):
if_generate(
self.is_leader_cta,
lambda: self.sync_object_full.arrive(
state.index, self.producer_mask, loc=loc, ip=ip
),
loc=loc,
ip=ip,
)
else:
tx_count = self.sync_object_full.tx_count + extra_tx_count
if_generate(
self.is_leader_cta,
lambda: self.sync_object_full.arrive_and_expect_tx(
state.index, tx_count, loc=loc, ip=ip
),
loc=loc,
ip=ip,
)
PipelineTmaUmma.create = _override_create(PipelineTmaUmmaOg, PipelineTmaUmma)
# ── PipelineUmmaAsync ───────────────────────────────────────────────────────
@dataclass(frozen=True)
class PipelineUmmaAsync(_PipelineIndexPhaseMixin, PipelineUmmaAsyncOg):
pass
PipelineUmmaAsync.create = _override_create(PipelineUmmaAsyncOg, PipelineUmmaAsync)
# ── PipelineAsyncUmma ───────────────────────────────────────────────────────
@dataclass(frozen=True)
class PipelineAsyncUmma(_PipelineIndexPhaseMixin, PipelineAsyncUmmaOg):
pass
PipelineAsyncUmma.create = _override_create(PipelineAsyncUmmaOg, PipelineAsyncUmma)
+68
View File
@@ -0,0 +1,68 @@
[build-system]
requires = ["setuptools>=75"]
build-backend = "setuptools.build_meta"
[project]
name = "flash-attn-4"
version = "4.0.0b20.dev2+fastvideo.82d6441"
description = "Flash Attention CUTE (CUDA Template Engine) implementation"
readme = "README.md"
requires-python = ">=3.10"
license = "BSD-3-Clause"
authors = [
{name = "Tri Dao"},
]
classifiers = [
"Development Status :: 3 - Alpha",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
]
dependencies = [
"nvidia-cutlass-dsl==4.6.0.dev0",
"torch",
"einops",
"typing_extensions",
"apache-tvm-ffi>=0.1.5,<0.2",
"torch-c-dlpack-ext",
"quack-kernels>=0.5.0",
]
[project.optional-dependencies]
cu13 = ["nvidia-cutlass-dsl[cu13]==4.6.0.dev0"]
dev = [
"pytest",
"pytest-xdist",
"ruff",
]
[project.urls]
Homepage = "https://github.com/Dao-AILab/flash-attention"
Repository = "https://github.com/Dao-AILab/flash-attention"
[tool.setuptools]
packages = ["flash_attn.cute"]
package-dir = {"flash_attn.cute" = "."}
[[tool.uv.index]]
name = "pytorch-cu130"
url = "https://download.pytorch.org/whl/cu130"
explicit = true
[tool.uv.sources]
torch = [
{ index = "pytorch-cu130", marker = "extra == 'cu13'" },
]
[tool.ruff]
line-length = 100
[tool.ruff.lint]
ignore = [
"E731", # do not assign a lambda expression, use a def
"E741", # Do not use variables named 'I', 'O', or 'l'
"F841", # local variable is assigned to but never used
"D102", # Missing docstring in public methods
]
+302
View File
@@ -0,0 +1,302 @@
from typing import Optional
from dataclasses import dataclass
import cutlass
import cutlass.cute as cute
from cutlass import Int32, const_expr
from quack import copy_utils
"""
This consolidates all the info related to sequence length. This is so that we can do all
the gmem reads once at the beginning of each tile, rather than having to repeat these reads
to compute various things like n_block_min, n_block_max, etc.
"""
@dataclass(frozen=True)
class SeqlenInfo:
offset: Int32
offset_padded: Int32
seqlen: Int32
has_cu_seqlens: cutlass.Constexpr[bool] = False
@staticmethod
def create(
batch_idx: Int32,
seqlen_static: Int32,
cu_seqlens: Optional[cute.Tensor] = None,
seqused: Optional[cute.Tensor] = None,
tile: cutlass.Constexpr[int] = 128,
):
offset = 0 if const_expr(cu_seqlens is None) else cu_seqlens[batch_idx]
offset_padded = (
0
if const_expr(cu_seqlens is None)
# Add divby so that the compiler knows the alignment when moving by offset_padded
else cute.assume((offset + batch_idx * tile) // tile * tile, divby=tile)
)
if const_expr(seqused is not None):
seqlen = seqused[batch_idx]
elif const_expr(cu_seqlens is not None):
seqlen = cu_seqlens[batch_idx + 1] - cu_seqlens[batch_idx]
else:
seqlen = seqlen_static
return SeqlenInfo(offset, offset_padded, seqlen, has_cu_seqlens=cu_seqlens is not None)
def offset_batch(
self,
mT: cute.Tensor,
batch_idx: Int32,
dim: int,
padded: cutlass.Constexpr[bool] = False,
multiple: int = 1,
) -> cute.Tensor:
"""Offset a tensor by batch index. batch dim is at position `dim`, seqlen is at dim=0."""
if const_expr(not self.has_cu_seqlens):
idx = (None,) * dim + (batch_idx,) + (None,) * (cute.rank(mT) - 1 - dim)
return mT[idx]
else:
off = multiple * (self.offset if const_expr(not padded) else self.offset_padded)
offset = off if const_expr(cute.rank(mT.shape[0]) == 1) else (0, off)
idx = (offset,) + (None,) * (cute.rank(mT) - 1)
return cute.domain_offset(idx, mT)
@dataclass(frozen=True)
class SeqlenInfoQK:
offset_q: Int32
offset_k: Int32
padded_offset_q: Int32
padded_offset_k: Int32
seqlen_q: Int32
seqlen_k: Int32
m_block_offset: Int32
block_idx_offset: Int32
num_n_blocks: Int32
has_cu_seqlens_q: cutlass.Constexpr[bool]
has_cu_seqlens_k: cutlass.Constexpr[bool]
has_seqused_q: cutlass.Constexpr[bool]
has_seqused_k: cutlass.Constexpr[bool]
@staticmethod
def create(
batch_idx: Int32,
seqlen_q_static: Int32,
seqlen_k_static: Int32,
mCuSeqlensQ: Optional[cute.Tensor] = None,
mCuSeqlensK: Optional[cute.Tensor] = None,
mSeqUsedQ: Optional[cute.Tensor] = None,
mSeqUsedK: Optional[cute.Tensor] = None,
mCuTotalMBlocks: Optional[cute.Tensor] = None,
mCuBlockIdxOffsets: Optional[cute.Tensor] = None,
tile_m: cutlass.Constexpr[Int32] = 128,
tile_n: cutlass.Constexpr[Int32] = 128,
):
offset_q = 0 if const_expr(mCuSeqlensQ is None) else mCuSeqlensQ[batch_idx]
offset_k = 0 if const_expr(mCuSeqlensK is None) else mCuSeqlensK[batch_idx]
padded_offset_q = (
0
if const_expr(mCuSeqlensQ is None)
else cute.assume((offset_q + batch_idx * tile_m) // tile_m * tile_m, divby=tile_m)
)
padded_offset_k = (
0
if const_expr(mCuSeqlensK is None)
else cute.assume((offset_k + batch_idx * tile_n) // tile_n * tile_n, divby=tile_n)
)
if const_expr(mSeqUsedQ is not None):
seqlen_q = mSeqUsedQ[batch_idx]
else:
seqlen_q = (
seqlen_q_static
if const_expr(mCuSeqlensQ is None)
else mCuSeqlensQ[batch_idx + 1] - offset_q
)
if const_expr(mSeqUsedK is not None):
seqlen_k = mSeqUsedK[batch_idx]
else:
seqlen_k = (
seqlen_k_static
if const_expr(mCuSeqlensK is None)
else mCuSeqlensK[batch_idx + 1] - offset_k
)
m_block_offset = 0 if const_expr(mCuTotalMBlocks is None) else mCuTotalMBlocks[batch_idx]
num_n_blocks = (seqlen_k + tile_n - 1) // tile_n
block_idx_offset = (
mCuBlockIdxOffsets[batch_idx]
if const_expr(mCuBlockIdxOffsets is not None)
else m_block_offset * num_n_blocks
)
return SeqlenInfoQK(
offset_q,
offset_k,
padded_offset_q,
padded_offset_k,
seqlen_q,
seqlen_k,
m_block_offset,
block_idx_offset,
num_n_blocks,
has_cu_seqlens_q=mCuSeqlensQ is not None,
has_cu_seqlens_k=mCuSeqlensK is not None,
has_seqused_q=mSeqUsedQ is not None,
has_seqused_k=mSeqUsedK is not None,
)
def offset_batch_Q(
self,
mQ: cute.Tensor,
batch_idx: Int32,
dim: int,
padded: cutlass.Constexpr[bool] = False,
ragged: cutlass.Constexpr[bool] = False,
) -> cute.Tensor:
"""Seqlen must be the first dimension of mQ"""
if const_expr(not ragged):
if const_expr(not self.has_cu_seqlens_q):
idx = (None,) * dim + (batch_idx,) + (None,) * (cute.rank(mQ) - 1 - dim)
return mQ[idx]
else:
offset_q = self.offset_q if const_expr(not padded) else self.padded_offset_q
offset_q = offset_q if const_expr(cute.rank(mQ.shape[0]) == 1) else (None, offset_q)
idx = (offset_q,) + (None,) * (cute.rank(mQ) - 1)
return cute.domain_offset(idx, mQ)
else:
if const_expr(not self.has_cu_seqlens_q):
offset_q = 0
idx = (None,) * dim + (batch_idx,) + (None,) * (cute.rank(mQ) - 1 - dim)
mQ = mQ[idx]
else:
offset_q = self.offset_q if const_expr(not padded) else self.padded_offset_q
if const_expr(cute.rank(mQ.shape[0]) == 1):
return copy_utils.offset_ragged_tensor(
mQ, offset_q, self.seqlen_q, ragged_dim=0, ptr_shift=True
)
else: # PackGQA
assert cute.rank(mQ.shape[0]) == 2
# Unpack before calling offset_ragged_tensor, then pack
idx = ((None, None),) + (None,) * (cute.rank(mQ) - 1)
mQ = mQ[idx]
mQ = copy_utils.offset_ragged_tensor(
mQ, offset_q, self.seqlen_q, ragged_dim=1, ptr_shift=True
)
return cute.group_modes(mQ, 0, 2)
def offset_batch_K(
self,
mK: cute.Tensor,
batch_idx: Int32,
dim: int,
padded: cutlass.Constexpr[bool] = False,
ragged: cutlass.Constexpr[bool] = False,
multiple: int = 1,
) -> cute.Tensor:
"""Seqlen must be the first dimension of mK"""
if const_expr(not ragged):
if const_expr(not self.has_cu_seqlens_k):
idx = (None,) * dim + (batch_idx,) + (None,) * (cute.rank(mK) - 1 - dim)
return mK[idx]
else:
offset_k = self.offset_k if const_expr(not padded) else self.padded_offset_k
offset_k *= multiple
idx = (offset_k,) + (None,) * (cute.rank(mK) - 1)
return cute.domain_offset(idx, mK)
else:
if const_expr(not self.has_cu_seqlens_k):
offset_k = 0
idx = (None,) * dim + (batch_idx,) + (None,) * (cute.rank(mK) - 1 - dim)
mK = mK[idx]
else:
offset_k = self.offset_k if const_expr(not padded) else self.padded_offset_k
offset_k *= multiple
return copy_utils.offset_ragged_tensor(
mK, offset_k, self.seqlen_k, ragged_dim=0, ptr_shift=True
)
@dataclass(frozen=True)
class SeqlenInfoQKNewK:
"""Sequence length info for append-KV with left-padding and new K support.
Extends SeqlenInfoQK with:
- leftpad_k: left padding for K (tokens to skip at the start of the KV cache)
- offset_k_new: offset into the new K tensor
- seqlen_k_og: original K length (before appending new K), excluding leftpad
- seqlen_k_new: length of new K to append
- seqlen_k: total K length (seqlen_k_og + seqlen_k_new)
- seqlen_rotary: position for rotary embedding computation
"""
leftpad_k: Int32
offset_q: Int32
offset_k: Int32
offset_k_new: Int32
seqlen_q: Int32
seqlen_k_og: Int32
seqlen_k_new: Int32
seqlen_k: Int32
seqlen_rotary: Int32
@staticmethod
def create(
batch_idx: Int32,
seqlen_q_static: Int32,
seqlen_k_static: Int32,
shape_K_new_0: Int32,
mCuSeqlensQ: Optional[cute.Tensor] = None,
mCuSeqlensK: Optional[cute.Tensor] = None,
mCuSeqlensKNew: Optional[cute.Tensor] = None,
mSeqUsedQ: Optional[cute.Tensor] = None,
mSeqUsedK: Optional[cute.Tensor] = None,
mLeftpadK: Optional[cute.Tensor] = None,
mSeqlensRotary: Optional[cute.Tensor] = None,
):
leftpad_k = 0 if const_expr(mLeftpadK is None) else mLeftpadK[batch_idx]
offset_q = 0 if const_expr(mCuSeqlensQ is None) else mCuSeqlensQ[batch_idx]
if const_expr(mCuSeqlensK is not None):
offset_k = mCuSeqlensK[batch_idx] + leftpad_k
else:
offset_k = leftpad_k if const_expr(mCuSeqlensQ is not None) else 0
offset_k_new = 0 if const_expr(mCuSeqlensKNew is None) else mCuSeqlensKNew[batch_idx]
# seqlen_q
if const_expr(mSeqUsedQ is not None):
seqlen_q = mSeqUsedQ[batch_idx]
elif const_expr(mCuSeqlensQ is not None):
seqlen_q = mCuSeqlensQ[batch_idx + 1] - mCuSeqlensQ[batch_idx]
else:
seqlen_q = seqlen_q_static
# seqlen_k_og: original K length (excluding leftpad)
if const_expr(mSeqUsedK is not None):
seqlen_k_og = mSeqUsedK[batch_idx] - leftpad_k
elif const_expr(mCuSeqlensK is not None):
seqlen_k_og = mCuSeqlensK[batch_idx + 1] - mCuSeqlensK[batch_idx] - leftpad_k
else:
seqlen_k_og = (
seqlen_k_static - leftpad_k
if const_expr(mCuSeqlensQ is not None)
else seqlen_k_static
)
# seqlen_k_new
if const_expr(mCuSeqlensKNew is None):
seqlen_k_new = 0 if const_expr(mCuSeqlensQ is None) else shape_K_new_0
else:
seqlen_k_new = mCuSeqlensKNew[batch_idx + 1] - mCuSeqlensKNew[batch_idx]
seqlen_k = seqlen_k_og if const_expr(mCuSeqlensQ is None) else seqlen_k_og + seqlen_k_new
# seqlen_rotary: defaults to seqlen_k_og + leftpad_k unless explicitly provided
if const_expr(mSeqlensRotary is not None):
seqlen_rotary = mSeqlensRotary[batch_idx]
else:
seqlen_rotary = seqlen_k_og + leftpad_k
return SeqlenInfoQKNewK(
leftpad_k,
offset_q,
offset_k,
offset_k_new,
seqlen_q,
seqlen_k_og,
seqlen_k_new,
seqlen_k,
seqlen_rotary,
)
@@ -0,0 +1,298 @@
# Copyright (c) 2025, Siyu Wang, Shengbin Di, Yuxi Chi, Johnsonms, Linfeng Zheng, Haoyan Huang, Lanbo Li, Yun Zhong, Man Yuan, Minmin Sun, Yong Li, Wei Lin.
"""Fused multi-head attention (FMHA) backward for the SM100 architecture using CUTE DSL.
Constraints:
* Supported head dimensions: 256 only
* mma_tiler_mn must be 64,64
* Batch size must be the same for Q, K, and V tensors
"""
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
from cutlass.cute.typing import Int32
from flash_attn.cute.sm100_hd256_2cta_fmha_backward_dqkernel import (
BlackwellFusedMultiHeadAttentionBackwardDQKernel,
)
from flash_attn.cute.sm100_hd256_2cta_fmha_backward_dkdvkernel import (
BlackwellFusedMultiHeadAttentionBackwardDKDVKernel,
)
from flash_attn.cute.cute_dsl_utils import assume_tensor_aligned
from flash_attn.cute.utils import AuxData
def _as_bshkrd_tensor(
tensor: cute.Tensor,
h_k: Int32,
h_r: Int32,
varlen: bool,
) -> cute.Tensor:
"""Normalize (B,S,H,D)/(S,H,D) tensors to (B,S,H_k,H_r,D) view."""
if cutlass.const_expr(cute.rank(tensor.layout) == 5):
return tensor
if cutlass.const_expr(cute.rank(tensor.layout) == 4):
return cute.make_tensor(
tensor.iterator,
cute.make_layout(
(tensor.shape[0], tensor.shape[1], h_k, h_r, tensor.shape[3]),
stride=(
tensor.stride[0],
tensor.stride[1],
tensor.stride[2] * h_r,
tensor.stride[2],
tensor.stride[3],
),
),
)
assert cutlass.const_expr(cute.rank(tensor.layout) == 3), "Expected rank-3 varlen tensor"
assert cutlass.const_expr(varlen), "Rank-3 input is only valid for varlen backward"
return cute.make_tensor(
tensor.iterator,
cute.make_layout(
(1, tensor.shape[0], h_k, h_r, tensor.shape[2]),
stride=(
0,
tensor.stride[0],
tensor.stride[1] * h_r,
tensor.stride[1],
tensor.stride[2],
),
),
)
def _as_shhb_tensor(
tensor: cute.Tensor,
h_k: Int32,
h_r: Int32,
b: Int32,
varlen: bool,
) -> cute.Tensor:
"""Normalize (B,H,S)/(H,S) tensors to (S, ((H_r, H_k), B)) view."""
if cutlass.const_expr(cute.rank(tensor.layout) == 3):
return cute.make_tensor(
tensor.iterator,
cute.make_layout(
(tensor.shape[2], ((h_r, h_k), tensor.shape[0])),
stride=(
tensor.stride[2],
((tensor.stride[1], tensor.stride[1] * h_r), tensor.stride[0]),
),
),
)
assert cutlass.const_expr(cute.rank(tensor.layout) == 2), "Expected rank-2 varlen tensor"
assert cutlass.const_expr(varlen), "Rank-2 input is only valid for varlen backward"
return cute.make_tensor(
tensor.iterator,
cute.make_layout(
(tensor.shape[1], ((h_r, h_k), b)),
stride=(
tensor.stride[1],
((tensor.stride[0], tensor.stride[0] * h_r), 0),
),
),
)
class BlackwellFusedMultiHeadAttentionBackward:
"""FMHA backward class for executing CuTeDSL kernel."""
def __init__(
self,
head_dim: int,
head_dim_v: int | None = None,
is_causal: bool = False,
is_local: bool = False,
qhead_per_kvhead: cutlass.Constexpr[int] = 1,
is_persistent: bool = False,
deterministic: bool = False,
cluster_size: int = 1,
use_2cta_instrs: bool = False,
score_mod: cutlass.Constexpr | None = None,
score_mod_bwd: cutlass.Constexpr | None = None,
mask_mod: cutlass.Constexpr | None = None,
has_aux_tensors: cutlass.Constexpr = False,
q_subtile_factor: cutlass.Constexpr[int] = 1,
tile_m_dq: int = 128,
tile_n_dq: int = 128,
tile_m_dkdv: int = 128,
tile_n_dkdv: int = 64,
window_size_left: int | None = None,
window_size_right: int | None = None,
use_clc_scheduler: bool = False,
):
"""Initialization."""
head_dim_v = head_dim if head_dim_v is None else head_dim_v
assert head_dim == 256 and head_dim_v == 256, (
"SM100 dedicated backward kernel only supports (head_dim, head_dim_v) = (256, 256)"
)
assert not is_local, "SM100 backward with head_dim=256 does not support local attention"
assert tile_m_dq == 128 and tile_n_dq == 128, (
"SM100 dedicated backward kernel only supports tile_m_dq=128 and tile_n_dq=128"
)
assert tile_m_dkdv == 128 and tile_n_dkdv == 64, (
"SM100 dedicated backward kernel only supports tile_m_dkdv=128 and tile_n_dkdv=64"
)
assert score_mod is None and score_mod_bwd is None and mask_mod is None, (
"SM100 backward with head_dim=256 does not support score_mod/mask_mod"
)
assert not deterministic, (
"SM100 backward with head_dim=256 does not support deterministic mode"
)
assert not has_aux_tensors, "SM100 backward with head_dim=256 does not support aux_tensors"
assert cluster_size in (1, 2), (
"SM100 backward with head_dim=256 only supports cluster_size in {1, 2}"
)
assert use_2cta_instrs, "SM100 backward with head_dim=256 requires use_2cta_instrs=True"
# q_subtile_factor is accepted for interface parity with FlashAttentionBackwardSm100,
# but this dedicated kernel uses fixed internal behavior.
self.acc_dtype = cutlass.Float32
self.is_causal = is_causal
self.window_size_left = (
None if (window_size_left is None or window_size_left < 0) else window_size_left
)
self.window_size_right = (
None if (window_size_right is None or window_size_right < 0) else window_size_right
)
self.tile_m_dq = tile_m_dq
self.tile_n_dq = tile_n_dq
self.tile_m_dkdv = tile_m_dkdv
self.tile_n_dkdv = tile_n_dkdv
self.use_clc_scheduler = use_clc_scheduler
self.dq_kernel = BlackwellFusedMultiHeadAttentionBackwardDQKernel(
self.acc_dtype,
(self.tile_m_dq, self.tile_n_dq, 256),
self.is_causal,
self.window_size_left,
self.window_size_right,
False, # is_persistent
False, # split_head
use_clc_scheduler=self.use_clc_scheduler,
)
self.dkdv_kernel = BlackwellFusedMultiHeadAttentionBackwardDKDVKernel(
self.acc_dtype,
(self.tile_m_dkdv, self.tile_n_dkdv, 256),
self.is_causal,
self.window_size_left,
self.window_size_right,
use_clc_scheduler=self.use_clc_scheduler,
)
@cute.jit
def __call__(
self,
Q: cute.Tensor,
K: cute.Tensor,
V: cute.Tensor,
dO: cute.Tensor,
lse_log2: cute.Tensor,
dpsum: cute.Tensor,
dQ_accum: cute.Tensor | None,
dK: cute.Tensor,
dV: cute.Tensor,
scale_softmax: cutlass.Float32,
cumulative_s_q: cute.Tensor | None,
cumulative_s_k: cute.Tensor | None,
seqused_q: cute.Tensor | None = None,
seqused_k: cute.Tensor | None = None,
window_size_left: Int32 | None = None,
window_size_right: Int32 | None = None,
dQ_semaphore: cute.Tensor | None = None,
dK_semaphore: cute.Tensor | None = None,
dV_semaphore: cute.Tensor | None = None,
aux_data: AuxData = AuxData(),
block_sparse_tensors: cute.Tensor | None = None,
stream: cuda.CUstream = None,
):
"""Host function to launch CuTeDSL kernel."""
assert seqused_q is None and seqused_k is None, (
"SM100 backward with head_dim=256 does not support seqused_q/seqused_k"
)
assert window_size_left is None and window_size_right is None, (
"SM100 backward with head_dim=256 uses constructor-provided window sizes"
)
assert dQ_semaphore is None and dK_semaphore is None and dV_semaphore is None, (
"SM100 backward with head_dim=256 does not use semaphores"
)
assert block_sparse_tensors is None, (
"SM100 backward with head_dim=256 does not support block sparse tensors"
)
assert aux_data.tensors is None or len(aux_data.tensors) == 0, (
"SM100 backward with head_dim=256 does not support aux_tensors"
)
assert aux_data.scalars is None or len(aux_data.scalars) == 0, (
"SM100 backward with head_dim=256 does not support aux_scalars"
)
assert dQ_accum is not None, (
"SM100 backward with head_dim=256 expects dQ tensor at dQ_accum slot"
)
dQ = dQ_accum
varlen = cumulative_s_q is not None or cumulative_s_k is not None
q_rank = cute.rank(Q.layout)
k_rank = cute.rank(K.layout)
if cutlass.const_expr(q_rank == 5):
h_q = Q.shape[2] * Q.shape[3]
elif cutlass.const_expr(q_rank == 4):
h_q = Q.shape[2]
else:
h_q = Q.shape[1]
if cutlass.const_expr(k_rank == 5):
h_k = K.shape[2]
elif cutlass.const_expr(k_rank == 4):
h_k = K.shape[2]
else:
h_k = K.shape[1]
h_r = h_q // h_k
if cutlass.const_expr(cumulative_s_q is not None):
b = cumulative_s_q.shape[0] - 1
elif cutlass.const_expr(cumulative_s_k is not None):
b = cumulative_s_k.shape[0] - 1
else:
b = Q.shape[0]
Q, K, V, dQ, dK, dV, dO = [assume_tensor_aligned(t) for t in (Q, K, V, dQ, dK, dV, dO)]
Q = _as_bshkrd_tensor(Q, h_k, h_r, varlen)
K = _as_bshkrd_tensor(K, h_k, 1, varlen)
V = _as_bshkrd_tensor(V, h_k, 1, varlen)
dQ = _as_bshkrd_tensor(dQ, h_k, h_r, varlen)
dK = _as_bshkrd_tensor(dK, h_k, 1, varlen)
dV = _as_bshkrd_tensor(dV, h_k, 1, varlen)
dO = _as_bshkrd_tensor(dO, h_k, h_r, varlen)
scaled_LSE = _as_shhb_tensor(lse_log2, h_k, h_r, b, varlen)
sum_OdO = _as_shhb_tensor(dpsum, h_k, h_r, b, varlen)
# Keep original order: dQ first, then dKdV.
self.dq_kernel(
Q,
K,
V,
dQ,
dO,
scaled_LSE,
sum_OdO,
cumulative_s_q,
cumulative_s_k,
scale_softmax,
stream,
)
self.dkdv_kernel(
Q,
K,
V,
dK,
dV,
dO,
scaled_LSE,
sum_OdO,
cumulative_s_q,
cumulative_s_k,
scale_softmax,
stream,
)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+402
View File
@@ -0,0 +1,402 @@
"""Search feasible SM90 fwd/bwd attention configs for given (head_dim, head_dim_v).
Enumerates tile sizes, swap modes, atom layouts, and staging options.
Checks GMMA divisibility, register budget, and shared memory budget.
Usage:
python flash_attn/cute/sm90_config_search.py --headdim 128
python flash_attn/cute/sm90_config_search.py --mode fwd --headdim 192-128
python flash_attn/cute/sm90_config_search.py --mode bwd --headdim 192 --tile-n 64,96
"""
import math
# H100 hardware limits
SMEM_LIMIT = 224 * 1024 # 228 KB minus ~3 KB for LSE, dPsum, mbarriers
REG_LIMITS = {2: 216, 3: 128} # per-WG budget: 2WG=240-24, 3WG=160-32
THREADS_PER_WG = 128
def _divisors(n):
return [d for d in range(1, n + 1) if n % d == 0]
def _acc_regs(M, N, num_wg):
"""Accumulator registers per thread per WG."""
return M * N // (num_wg * THREADS_PER_WG)
def _check_mma(M, N, num_wg, atom_layout_m, swap_AB):
"""Check MMA feasibility. Returns regs per WG, or None if infeasible.
GMMA atom M=64. Swap exchanges (M, N) and atom layout.
Requires: M divisible by (atom_layout_m * 64), N by (atom_layout_n * 8).
"""
if swap_AB:
M, N = N, M
atom_layout_m = num_wg // atom_layout_m
atom_layout_n = num_wg // atom_layout_m
if M % (atom_layout_m * 64) != 0 or N % (atom_layout_n * 8) != 0:
return None
return _acc_regs(M, N, num_wg)
def _mma_traffic(M_eff, N_eff, K_red, num_wg, wg_n, is_rs=False):
"""Total SMEM read traffic for one MMA (all WGs combined).
num_instr = (M_eff / 64) * wg_n instructions total.
Each reads A(64, K_red) and B(N_eff/wg_n, K_red) from smem (bf16).
"""
num_instr = (M_eff // 64) * wg_n
A_per = 64 * K_red * 2 if not is_rs else 0
B_per = (N_eff // wg_n) * K_red * 2
return num_instr * (A_per + B_per)
# ============================================================================
# Backward
# ============================================================================
def _check_bwd_config(
hdim,
hdimv,
tile_m,
tile_n,
num_wg,
SdP_swapAB,
dKV_swapAB,
dQ_swapAB,
AtomLayoutMSdP,
AtomLayoutNdKV,
AtomLayoutMdQ,
):
reg_limit = REG_LIMITS[num_wg]
# MMA feasibility
regs_SdP = _check_mma(tile_m, tile_n, num_wg, AtomLayoutMSdP, SdP_swapAB)
regs_dK = _check_mma(tile_n, hdim, num_wg, AtomLayoutNdKV, dKV_swapAB)
regs_dV = _check_mma(tile_n, hdimv, num_wg, AtomLayoutNdKV, dKV_swapAB)
regs_dQ = _check_mma(tile_m, hdim, num_wg, AtomLayoutMdQ, dQ_swapAB)
if any(r is None for r in (regs_SdP, regs_dK, regs_dV, regs_dQ)):
return None
# Peak regs: max(S+dP, dQ) + dK + dV
total_regs = max(2 * regs_SdP, regs_dQ) + regs_dK + regs_dV
if total_regs > reg_limit:
return None
# SMEM
mma_dkv_is_rs = (
AtomLayoutMSdP == 1 and AtomLayoutNdKV == num_wg and SdP_swapAB and not dKV_swapAB
)
Q_stage, PdS_stage = 2, 1
for dO_stage in (2, 1):
sQ = tile_m * hdim * 2 * Q_stage
sK = tile_n * hdim * 2
sV = tile_n * hdimv * 2
sdO = tile_m * hdimv * 2 * dO_stage
sPdS = tile_m * tile_n * 2 * PdS_stage
sP = sPdS if not mma_dkv_is_rs else 0
sdQaccum = tile_m * hdim * 4
smem = sQ + sK + sV + sdO + sP + sPdS + sdQaccum
if smem <= SMEM_LIMIT:
break
else:
return None
# SMEM traffic
def _swap(a, b, s):
return (b, a) if s else (a, b)
def _wg_n(al_m, s):
return al_m if s else num_wg // al_m
M_s, N_s = _swap(tile_m, tile_n, SdP_swapAB)
wn_SdP = _wg_n(AtomLayoutMSdP, SdP_swapAB)
traffic_S = _mma_traffic(M_s, N_s, hdim, num_wg, wn_SdP)
traffic_dP = _mma_traffic(M_s, N_s, hdimv, num_wg, wn_SdP)
wn_dKV = _wg_n(AtomLayoutNdKV, dKV_swapAB)
M_dv, N_dv = _swap(tile_n, hdimv, dKV_swapAB)
traffic_dV = _mma_traffic(M_dv, N_dv, tile_m, num_wg, wn_dKV, is_rs=mma_dkv_is_rs)
M_dk, N_dk = _swap(tile_n, hdim, dKV_swapAB)
traffic_dK = _mma_traffic(M_dk, N_dk, tile_m, num_wg, wn_dKV, is_rs=mma_dkv_is_rs)
M_dq, N_dq = _swap(tile_m, hdim, dQ_swapAB)
wn_dQ = _wg_n(AtomLayoutMdQ, dQ_swapAB)
traffic_dQ = _mma_traffic(M_dq, N_dq, tile_n, num_wg, wn_dQ)
traffic_P_store = tile_m * tile_n * 2 if not mma_dkv_is_rs else 0
traffic_dS_store = tile_m * tile_n * 2
traffic_dQ_smem = tile_m * hdim * 4 * 2 # store + TMA load
smem_traffic = (
traffic_S
+ traffic_dP
+ traffic_dV
+ traffic_dK
+ traffic_dQ
+ traffic_P_store
+ traffic_dS_store
+ traffic_dQ_smem
)
return dict(
tile_m=tile_m,
tile_n=tile_n,
num_wg=num_wg,
Q_stage=Q_stage,
dO_stage=dO_stage,
PdS_stage=PdS_stage,
SdP_swapAB=SdP_swapAB,
dKV_swapAB=dKV_swapAB,
dQ_swapAB=dQ_swapAB,
AtomLayoutMSdP=AtomLayoutMSdP,
AtomLayoutNdKV=AtomLayoutNdKV,
AtomLayoutMdQ=AtomLayoutMdQ,
mma_dkv_is_rs=mma_dkv_is_rs,
regs_SdP=regs_SdP,
regs_dK=regs_dK,
regs_dV=regs_dV,
regs_dQ=regs_dQ,
total_regs=total_regs,
reg_limit=reg_limit,
smem_bytes=smem,
smem_kb=smem / 1024,
smem_traffic=smem_traffic,
smem_traffic_kb=smem_traffic / 1024,
smem_traffic_per_block=smem_traffic / (tile_m * tile_n),
)
def find_feasible_bwd_configs(
head_dim,
head_dim_v=None,
tile_m_choices=(64, 80, 96, 112, 128),
tile_n_choices=(64, 80, 96, 112, 128),
):
if head_dim_v is None:
head_dim_v = head_dim
hdim = int(math.ceil(head_dim / 32) * 32)
hdimv = int(math.ceil(head_dim_v / 32) * 32)
results = []
for num_wg in (2, 3):
divs = _divisors(num_wg)
for tile_m in tile_m_choices:
for tile_n in tile_n_choices:
for SdP_swap in (False, True):
if (tile_n if SdP_swap else tile_m) % 64 != 0:
continue
for dKV_swap in (False, True):
if not dKV_swap and tile_n % 64 != 0:
continue
if dKV_swap and (hdim % 64 != 0 or hdimv % 64 != 0):
continue
for dQ_swap in (False, True):
if (hdim if dQ_swap else tile_m) % 64 != 0:
continue
for a1 in divs:
for a2 in divs:
for a3 in divs:
cfg = _check_bwd_config(
hdim,
hdimv,
tile_m,
tile_n,
num_wg,
SdP_swap,
dKV_swap,
dQ_swap,
a1,
a2,
a3,
)
if cfg is not None:
results.append(cfg)
results.sort(key=lambda c: (-c["tile_n"], -c["tile_m"], c["smem_traffic_per_block"]))
return results
def print_bwd_configs(configs, max_results=20):
if not configs:
print("No feasible configs found!")
return
n = min(len(configs), max_results)
print(f"Found {len(configs)} feasible configs (showing top {n}):\n")
hdr = (
f"{'wg':>2} {'tm':>3} {'tn':>3} "
f"{'SdP':>3} {'dKV':>3} {'dQ':>3} "
f"{'aSdP':>4} {'adKV':>4} {'adQ':>4} "
f"{'Qs':>2} {'dOs':>3} "
f"{'rS':>3} {'rdK':>3} {'rdV':>3} {'rdQ':>3} {'tot':>4}/{'':<3} "
f"{'smem':>5} {'traffic':>7} {'tr/blk':>6}"
)
print(hdr)
print("-" * len(hdr))
B = lambda b: "T" if b else "F"
for c in configs[:max_results]:
print(
f"{c['num_wg']:>2} {c['tile_m']:>3} {c['tile_n']:>3} "
f"{B(c['SdP_swapAB']):>3} {B(c['dKV_swapAB']):>3} {B(c['dQ_swapAB']):>3} "
f"{c['AtomLayoutMSdP']:>4} {c['AtomLayoutNdKV']:>4} {c['AtomLayoutMdQ']:>4} "
f"{c['Q_stage']:>2} {c['dO_stage']:>3} "
f"{c['regs_SdP']:>3} {c['regs_dK']:>3} {c['regs_dV']:>3} {c['regs_dQ']:>3} "
f"{c['total_regs']:>4}/{c['reg_limit']:<3} "
f"{c['smem_kb']:>4.0f}K "
f"{c['smem_traffic_kb']:>6.0f}K "
f"{c['smem_traffic_per_block']:>6.1f}"
)
# ============================================================================
# Forward
# ============================================================================
def _check_fwd_config(hdim, hdimv, tile_n, num_wg, pv_is_rs, overlap_wg):
reg_limit = REG_LIMITS[num_wg]
tile_m = num_wg * 64
if tile_n % 8 != 0:
return None
regs_S = _acc_regs(tile_m, tile_n, num_wg)
regs_O = _acc_regs(tile_m, hdimv, num_wg)
regs_P = regs_S // 2 # bf16 = half of f32
if overlap_wg:
total_regs = regs_S + regs_P + regs_O
else:
total_regs = regs_S + regs_O
if total_regs > reg_limit:
return None
# SMEM: 1 stage Q, 2 stages K/V, O overlaps Q, sP if not RS
sQ = tile_m * hdim * 2
sK = tile_n * hdim * 2 * 2
sV = tile_n * hdimv * 2 * 2
sO = tile_m * hdimv * 2
sP = tile_m * tile_n * 2 if not pv_is_rs else 0
smem = max(sQ, sO) + sK + sV + sP
if smem > SMEM_LIMIT:
return None
# SMEM traffic: num_instr = num_wg (all WGs in M, wg_n=1)
traffic_S = num_wg * (64 * hdim * 2 + tile_n * hdim * 2)
A_pv = 64 * tile_n * 2 if not pv_is_rs else 0
traffic_O = num_wg * (A_pv + hdimv * tile_n * 2)
traffic_P_store = tile_m * tile_n * 2 if not pv_is_rs else 0
smem_traffic = traffic_S + traffic_O + traffic_P_store
return dict(
tile_m=tile_m,
tile_n=tile_n,
num_wg=num_wg,
pv_is_rs=pv_is_rs,
overlap_wg=overlap_wg,
regs_S=regs_S,
regs_O=regs_O,
regs_P=regs_P,
total_regs=total_regs,
reg_limit=reg_limit,
smem_bytes=smem,
smem_kb=smem / 1024,
smem_traffic=smem_traffic,
smem_traffic_kb=smem_traffic / 1024,
smem_traffic_per_block=smem_traffic / (tile_m * tile_n),
)
def find_feasible_fwd_configs(
head_dim, head_dim_v=None, tile_n_choices=(64, 80, 96, 112, 128, 144, 160, 176, 192)
):
if head_dim_v is None:
head_dim_v = head_dim
hdim = int(math.ceil(head_dim / 32) * 32)
hdimv = int(math.ceil(head_dim_v / 32) * 32)
results = []
for num_wg in (2, 3):
for tile_n in tile_n_choices:
for pv_is_rs in (True, False):
for overlap_wg in (True, False):
cfg = _check_fwd_config(hdim, hdimv, tile_n, num_wg, pv_is_rs, overlap_wg)
if cfg is not None:
results.append(cfg)
results.sort(key=lambda c: (-c["tile_n"], c["smem_traffic_per_block"]))
return results
def print_fwd_configs(configs, max_results=20):
if not configs:
print("No feasible configs found!")
return
n = min(len(configs), max_results)
print(f"Found {len(configs)} feasible configs (showing top {n}):\n")
hdr = (
f"{'wg':>2} {'tm':>3} {'tn':>3} "
f"{'RS':>2} {'olap':>4} "
f"{'rS':>3} {'rP':>3} {'rO':>3} {'tot':>4}/{'':<3} "
f"{'smem':>5} {'traffic':>7} {'tr/blk':>6}"
)
print(hdr)
print("-" * len(hdr))
B = lambda b: "T" if b else "F"
for c in configs[:max_results]:
print(
f"{c['num_wg']:>2} {c['tile_m']:>3} {c['tile_n']:>3} "
f"{B(c['pv_is_rs']):>2} {B(c['overlap_wg']):>4} "
f"{c['regs_S']:>3} {c['regs_P']:>3} {c['regs_O']:>3} "
f"{c['total_regs']:>4}/{c['reg_limit']:<3} "
f"{c['smem_kb']:>4.0f}K "
f"{c['smem_traffic_kb']:>6.0f}K "
f"{c['smem_traffic_per_block']:>6.1f}"
)
# ============================================================================
# CLI
# ============================================================================
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="Search feasible SM90 MMA configs")
parser.add_argument("--mode", choices=["fwd", "bwd", "both"], default="both")
parser.add_argument(
"--headdim", type=str, default="128", help="Head dim, or hdim-hdimv (e.g. 192-128)"
)
parser.add_argument("--tile-m", type=str, default="64,80,96,112,128", help="Bwd tile_m choices")
parser.add_argument(
"--tile-n",
type=str,
default=None,
help="tile_n choices (default: fwd up to 192, bwd up to 128)",
)
parser.add_argument("-n", "--num-results", type=int, default=30)
args = parser.parse_args()
parts = args.headdim.split("-")
hdim = int(parts[0])
hdimv = int(parts[1]) if len(parts) > 1 else hdim
TN_FWD = "64,80,96,112,128,144,160,176,192"
TN_BWD = "64,80,96,112,128"
if args.mode in ("fwd", "both"):
tn = tuple(int(x) for x in (args.tile_n or TN_FWD).split(","))
print(f"=== FWD configs: hdim={hdim}, hdimv={hdimv} ===\n")
print_fwd_configs(find_feasible_fwd_configs(hdim, hdimv, tn), args.num_results)
print()
if args.mode in ("bwd", "both"):
tm = tuple(int(x) for x in args.tile_m.split(","))
tn = tuple(int(x) for x in (args.tile_n or TN_BWD).split(","))
print(f"=== BWD configs: hdim={hdim}, hdimv={hdimv} ===\n")
print_bwd_configs(find_feasible_bwd_configs(hdim, hdimv, tm, tn), args.num_results)
+740
View File
@@ -0,0 +1,740 @@
# Copyright (c) 2025, Tri Dao.
import math
import operator
from typing import Tuple
from dataclasses import dataclass
import cutlass
import cutlass.cute as cute
from cutlass import Float32, Boolean
from quack import layout_utils
import flash_attn.cute.utils as utils
from quack.cute_dsl_utils import ParamsBase
from flash_attn.cute.seqlen_info import SeqlenInfoQK
from flash_attn.cute.utils import AuxData
@cute.jit
def call_score_mod(
score_mod: cutlass.Constexpr,
score,
batch_idx,
head_idx,
q_idx,
kv_idx,
seqlen_info,
aux_data: AuxData,
):
aux_tensors = aux_data.tensors if aux_data.tensors is not None else ()
# Compatibility shim for pre-aux_scalars score_mod callables.
if cutlass.const_expr(aux_data.scalars is not None):
return score_mod(
score,
batch_idx,
head_idx,
q_idx=q_idx,
kv_idx=kv_idx,
seqlen_info=seqlen_info,
aux_tensors=aux_tensors,
aux_scalars=aux_data.scalars,
)
return score_mod(
score,
batch_idx,
head_idx,
q_idx=q_idx,
kv_idx=kv_idx,
seqlen_info=seqlen_info,
aux_tensors=aux_tensors,
)
@cute.jit
def call_score_mod_bwd(
score_mod_bwd: cutlass.Constexpr,
grad,
score,
batch_idx,
head_idx,
q_idx,
kv_idx,
seqlen_info,
aux_data: AuxData,
):
aux_tensors = aux_data.tensors if aux_data.tensors is not None else ()
# Compatibility shim for pre-aux_scalars score_mod_bwd callables.
if cutlass.const_expr(aux_data.scalars is not None):
return score_mod_bwd(
grad,
score,
batch_idx,
head_idx,
q_idx=q_idx,
kv_idx=kv_idx,
seqlen_info=seqlen_info,
aux_tensors=aux_tensors,
aux_scalars=aux_data.scalars,
)
return score_mod_bwd(
grad,
score,
batch_idx,
head_idx,
q_idx=q_idx,
kv_idx=kv_idx,
seqlen_info=seqlen_info,
aux_tensors=aux_tensors,
)
@dataclass
class Softmax(ParamsBase):
scale_log2: Float32
num_rows: cutlass.Constexpr[int]
row_max: cute.Tensor
row_sum: cute.Tensor
arch: cutlass.Constexpr[int] = 80
softmax_scale: Float32 | None = None
@staticmethod
def create(
scale_log2: Float32,
num_rows: cutlass.Constexpr[int],
arch: cutlass.Constexpr[int] = 80,
softmax_scale: Float32 | None = None,
):
row_max = cute.make_rmem_tensor(num_rows, Float32)
row_sum = cute.make_rmem_tensor(num_rows, Float32)
return Softmax(scale_log2, num_rows, row_max, row_sum, arch, softmax_scale)
def reset(self) -> None:
self.row_max.fill(-Float32.inf)
self.row_sum.fill(0.0)
def _compute_row_max(
self, acc_S_row: cute.TensorSSA, init_val: float | Float32 | None = None
) -> Float32:
return utils.fmax_reduce(acc_S_row, init_val, arch=self.arch)
def _compute_row_sum(
self, acc_S_row_exp: cute.TensorSSA, init_val: float | Float32 | None = None
) -> Float32:
return utils.fadd_reduce(acc_S_row_exp, init_val, arch=self.arch)
@cute.jit
def online_softmax(
self,
acc_S: cute.Tensor,
is_first: cutlass.Constexpr[bool] = False,
check_inf: cutlass.Constexpr[bool] = True,
) -> cute.Tensor:
"""Apply online softmax and return the row_scale to rescale O.
:param acc_S: acc_S tensor
:type acc_S: cute.Tensor
:param is_first: is first n_block
:type is_first: cutlass.Constexpr
"""
# Change acc_S to M,N layout view.
acc_S_mn = layout_utils.reshape_acc_to_mn(acc_S)
row_scale = cute.make_fragment_like(self.row_max, Float32)
row_max = self.row_max
row_sum = self.row_sum
scale_log2 = self.scale_log2
arch = self.arch
# Each iteration processes one row of acc_S
for r in cutlass.range(cute.size(row_max), unroll_full=True):
acc_S_row = acc_S_mn[r, None].load() # (n_block_size)
row_max_cur = utils.fmax_reduce(
acc_S_row,
init_val=row_max[r] if cutlass.const_expr(not is_first) else None,
arch=arch,
)
row_max_cur = cute.arch.warp_reduction_max(row_max_cur, threads_in_group=4)
# Update row_max before changing row_max_cur to safe value for -inf
row_max_prev = row_max[r]
row_max[r] = row_max_cur
if cutlass.const_expr(check_inf):
row_max_cur = 0.0 if row_max_cur == -Float32.inf else row_max_cur
if cutlass.const_expr(is_first):
row_max_cur_scaled = row_max_cur * scale_log2
acc_S_row_exp = cute.math.exp2(
acc_S_row * scale_log2 - row_max_cur_scaled, fastmath=True
)
acc_S_row_sum = utils.fadd_reduce(acc_S_row_exp, init_val=None, arch=arch)
row_scale[r] = 1.0
else:
row_max_cur_scaled = row_max_cur * scale_log2
acc_S_row_exp = cute.math.exp2(
acc_S_row * scale_log2 - row_max_cur_scaled, fastmath=True
)
# row_scale[r] = cute.math.exp2(row_max_prev * self.scale_log2 - row_max_cur_scaled)
row_scale[r] = cute.math.exp2(
(row_max_prev - row_max_cur) * scale_log2, fastmath=True
)
acc_S_row_sum = utils.fadd_reduce(
acc_S_row_exp, init_val=row_sum[r] * row_scale[r], arch=arch
)
row_sum[r] = acc_S_row_sum
acc_S_mn[r, None].store(acc_S_row_exp)
return row_scale
@cute.jit
def finalize(
self, final_scale: Float32 = 1.0, sink_val: Float32 | cute.Tensor | None = None
) -> cute.Tensor:
"""Finalize the online softmax by computing the scale and logsumexp."""
if cutlass.const_expr(sink_val is not None and isinstance(sink_val, cute.Tensor)):
assert cute.size(sink_val) == cute.size(self.row_sum)
row_sum = self.row_sum
row_max = self.row_max
scale_log2 = self.scale_log2
# quad reduction for row_sum as we didn't do it during each iteration of online softmax
row_sum.store(utils.warp_reduce(row_sum.load(), operator.add, width=4))
row_scale = cute.make_fragment_like(row_max, Float32)
for r in cutlass.range(cute.size(row_sum), unroll_full=True):
if cutlass.const_expr(sink_val is not None):
sink_val_cur = sink_val if not isinstance(sink_val, cute.Tensor) else sink_val[r]
LOG2_E = math.log2(math.e)
row_sum[r] += cute.math.exp2(
sink_val_cur * LOG2_E - row_max[r] * scale_log2, fastmath=True
)
# if row_sum is zero or nan, set acc_O_mn_row to 1.0
acc_O_mn_row_is_zero_or_nan = row_sum[r] == 0.0 or row_sum[r] != row_sum[r]
row_scale[r] = (
cute.arch.rcp_approx(row_sum[r] if not acc_O_mn_row_is_zero_or_nan else 1.0)
) * final_scale
row_sum_cur = row_sum[r]
LN2 = math.log(2.0)
row_sum[r] = (
(row_max[r] * scale_log2 + cute.math.log2(row_sum_cur, fastmath=True)) * LN2
if not acc_O_mn_row_is_zero_or_nan
else -Float32.inf
)
return row_scale
@cute.jit
def rescale_O(self, acc_O: cute.Tensor, row_scale: cute.Tensor) -> None:
"""Scale each row of acc_O by the given scale tensor.
:param acc_O: input tensor
:type acc_O: cute.Tensor
:param row_scale: row_scale tensor
:type row_scale: cute.Tensor
"""
acc_O_mn = layout_utils.reshape_acc_to_mn(acc_O)
assert cute.size(row_scale) == cute.size(acc_O_mn, mode=[0])
for r in cutlass.range(cute.size(row_scale), unroll_full=True):
acc_O_mn[r, None].store(acc_O_mn[r, None].load() * row_scale[r])
@dataclass
class SoftmaxSm100(Softmax):
rescale_threshold: cutlass.Constexpr[float] = 0.0
max_offset: cutlass.Constexpr[int] = 0
row_pair_stride: cutlass.Constexpr[int] = 0
@staticmethod
def create(
scale_log2: Float32,
rescale_threshold: cutlass.Constexpr[float] = 0.0,
softmax_scale: Float32 | None = None,
max_offset: cutlass.Constexpr[int] = 0,
row_pair_stride: cutlass.Constexpr[int] = 0,
):
num_rows = 1
arch = 100
row_max = cute.make_rmem_tensor(num_rows, Float32)
row_sum = cute.make_rmem_tensor(num_rows, Float32)
return SoftmaxSm100(
scale_log2,
num_rows,
row_max,
row_sum,
arch,
softmax_scale,
rescale_threshold=rescale_threshold,
max_offset=max_offset,
row_pair_stride=row_pair_stride,
)
@cute.jit
def compute_row_max_local(self, acc_S_row: cute.TensorSSA, is_first: Boolean) -> Float32:
if cutlass.const_expr(is_first):
row_max_new = self._compute_row_max(acc_S_row)
else:
row_max_old = self.row_max[0]
row_max_new = self._compute_row_max(acc_S_row, init_val=row_max_old)
return row_max_new
@cute.jit
def update_row_max_from_local(
self,
row_max_new: Float32,
is_first: Boolean,
) -> Tuple[Float32, Float32]:
if cutlass.const_expr(self.row_pair_stride != 0):
row_max_new = utils.fmax(
row_max_new,
cute.arch.shuffle_sync_bfly(
row_max_new,
offset=self.row_pair_stride,
mask=-1,
mask_and_clamp=31,
),
)
if cutlass.const_expr(is_first):
row_max_safe = row_max_new if row_max_new != -cutlass.Float32.inf else 0.0
acc_scale = 0.0
else:
row_max_old = self.row_max[0]
row_max_safe = row_max_new if row_max_new != -cutlass.Float32.inf else 0.0
acc_scale_ = (row_max_old - row_max_safe) * self.scale_log2
acc_scale = cute.math.exp2(acc_scale_)
if cutlass.const_expr(self.rescale_threshold > 0.0):
if acc_scale_ >= -self.rescale_threshold:
row_max_new = row_max_old
row_max_safe = row_max_old
acc_scale = 1.0
self.row_max[0] = row_max_new
return row_max_safe, acc_scale
@cute.jit
def update_row_max(self, acc_S_row: cute.TensorSSA, is_first: int) -> Tuple[Float32, Float32]:
if cutlass.const_expr(is_first):
row_max_new = self._compute_row_max(acc_S_row)
if cutlass.const_expr(self.row_pair_stride != 0):
row_max_new = utils.fmax(
row_max_new,
cute.arch.shuffle_sync_bfly(
row_max_new,
offset=self.row_pair_stride,
mask=-1,
mask_and_clamp=31,
),
)
row_max_safe = row_max_new if row_max_new != -cutlass.Float32.inf else 0.0
acc_scale = 0.0
else:
row_max_old = self.row_max[0]
row_max_new = self._compute_row_max(acc_S_row, init_val=row_max_old)
if cutlass.const_expr(self.row_pair_stride != 0):
row_max_new = utils.fmax(
row_max_new,
cute.arch.shuffle_sync_bfly(
row_max_new,
offset=self.row_pair_stride,
mask=-1,
mask_and_clamp=31,
),
)
row_max_safe = row_max_new if row_max_new != -cutlass.Float32.inf else 0.0
acc_scale_ = (row_max_old - row_max_safe) * self.scale_log2
acc_scale = cute.math.exp2(acc_scale_, fastmath=True)
if cutlass.const_expr(self.rescale_threshold > 0.0):
if acc_scale_ >= -self.rescale_threshold:
row_max_new = row_max_old
row_max_safe = row_max_old
acc_scale = 1.0
self.row_max[0] = row_max_new
return row_max_safe, acc_scale
def update_row_sum(
self, acc_S_row_exp: cute.TensorSSA, row_scale: Float32, is_first: int = False
) -> None:
if cutlass.const_expr(self.row_pair_stride != 0):
local_sum = self._compute_row_sum(acc_S_row_exp)
row_sum_new = local_sum + cute.arch.shuffle_sync_bfly(
local_sum,
offset=self.row_pair_stride,
mask=-1,
mask_and_clamp=31,
)
if cutlass.const_expr(not is_first):
row_sum_new = row_sum_new + self.row_sum[0] * row_scale
self.row_sum[0] = row_sum_new
else:
init_val = self.row_sum[0] * row_scale if cutlass.const_expr(not is_first) else None
self.row_sum[0] = self._compute_row_sum(acc_S_row_exp, init_val=init_val)
@cute.jit
def scale_subtract_rowmax(
self,
acc_S_row: cute.Tensor,
row_max: Float32,
):
assert cute.size(acc_S_row.shape) % 2 == 0, "acc_S_row must have an even number of elements"
row_max_scaled = row_max * self.scale_log2
max_offset = Float32(self.max_offset)
bias = max_offset - row_max_scaled
for i in cutlass.range(0, cute.size(acc_S_row.shape), 2, unroll_full=True):
acc_S_row[i], acc_S_row[i + 1] = cute.arch.fma_packed_f32x2(
(acc_S_row[i], acc_S_row[i + 1]),
(self.scale_log2, self.scale_log2),
(bias, bias),
)
@cute.jit
def apply_exp2_convert(
self,
acc_S_row: cute.Tensor,
acc_S_row_converted: cute.Tensor,
ex2_emu_freq: cutlass.Constexpr[int] = 0,
ex2_emu_res: cutlass.Constexpr[int] = 4,
ex2_emu_start_frg: cutlass.Constexpr[int] = 0,
):
assert cute.size(acc_S_row.shape) % 2 == 0, "acc_S_row must have an even number of elements"
frg_tile = 32
assert frg_tile % 2 == 0
frg_cnt = cute.size(acc_S_row) // frg_tile
assert cute.size(acc_S_row) % frg_tile == 0
acc_S_row_frg = cute.logical_divide(acc_S_row, cute.make_layout(frg_tile))
acc_S_row_converted_frg = cute.logical_divide(
acc_S_row_converted, cute.make_layout(frg_tile)
)
for j in cutlass.range_constexpr(frg_cnt):
for k in cutlass.range_constexpr(0, cute.size(acc_S_row_frg, mode=[0]), 2):
# acc_S_row_frg[k, j] = cute.math.exp2(acc_S_row_frg[k, j], fastmath=True)
# acc_S_row_frg[k + 1, j] = cute.math.exp2(acc_S_row_frg[k + 1, j], fastmath=True)
if cutlass.const_expr(ex2_emu_freq == 0):
acc_S_row_frg[k, j] = cute.math.exp2(acc_S_row_frg[k, j], fastmath=True)
acc_S_row_frg[k + 1, j] = cute.math.exp2(acc_S_row_frg[k + 1, j], fastmath=True)
else:
if cutlass.const_expr(
k % ex2_emu_freq < ex2_emu_freq - ex2_emu_res
or j >= frg_cnt - 1
or j < ex2_emu_start_frg
):
acc_S_row_frg[k, j] = cute.math.exp2(acc_S_row_frg[k, j], fastmath=True)
acc_S_row_frg[k + 1, j] = cute.math.exp2(
acc_S_row_frg[k + 1, j], fastmath=True
)
else:
# acc_S_row_frg[k, j], acc_S_row_frg[k + 1, j] = utils.e2e_asm2(acc_S_row_frg[k, j], acc_S_row_frg[k + 1, j])
acc_S_row_frg[k, j], acc_S_row_frg[k + 1, j] = utils.ex2_emulation_2(
acc_S_row_frg[k, j], acc_S_row_frg[k + 1, j]
)
acc_S_row_converted_frg[None, j].store(
acc_S_row_frg[None, j].load().to(acc_S_row_converted.element_type)
)
@cute.jit
def scale_apply_exp2_convert(
self,
acc_S_row: cute.Tensor,
row_max: Float32,
acc_S_row_converted: cute.Tensor,
):
assert cute.size(acc_S_row.shape) % 2 == 0, "acc_S_row must have an even number of elements"
minus_row_max_scaled = -row_max * self.scale_log2
for i in cutlass.range_constexpr(0, cute.size(acc_S_row.shape), 2):
acc_S_row[i], acc_S_row[i + 1] = cute.arch.fma_packed_f32x2(
(acc_S_row[i], acc_S_row[i + 1]),
(self.scale_log2, self.scale_log2),
(minus_row_max_scaled, minus_row_max_scaled),
)
# for i in cutlass.range_constexpr(0, cute.size(acc_S_row.shape), 2):
# acc_S_row[i], acc_S_row[i + 1] = cute.arch.fma_packed_f32x2(
# (acc_S_row[i], acc_S_row[i + 1]),
# (self.scale_log2, self.scale_log2),
# (minus_row_max_scaled, minus_row_max_scaled),
# )
# acc_S_row[i] = cute.math.exp2(acc_S_row[i], fastmath=True)
# acc_S_row[i + 1] = cute.math.exp2(acc_S_row[i + 1], fastmath=True)
frg_tile = 32
assert frg_tile % 2 == 0
frg_cnt = cute.size(acc_S_row) // frg_tile
assert cute.size(acc_S_row) % frg_tile == 0
acc_S_row_frg = cute.logical_divide(acc_S_row, cute.make_layout(frg_tile))
acc_S_row_converted_frg = cute.logical_divide(
acc_S_row_converted, cute.make_layout(frg_tile)
)
for j in cutlass.range_constexpr(frg_cnt):
for k in cutlass.range_constexpr(0, cute.size(acc_S_row_frg, mode=[0]), 2):
# acc_S_row_frg[k, j], acc_S_row_frg[k + 1, j] = (
# cute.arch.fma_packed_f32x2(
# (acc_S_row_frg[k, j], acc_S_row_frg[k + 1, j]),
# (self.scale_log2, self.scale_log2),
# (minus_row_max_scaled, minus_row_max_scaled),
# )
# )
# acc_S_row_frg[k, j] = cute.math.exp2(acc_S_row_frg[k, j], fastmath=True)
# acc_S_row_frg[k + 1, j] = cute.math.exp2(acc_S_row_frg[k + 1, j], fastmath=True)
acc_S_row_frg[k, j] = cute.math.exp2(acc_S_row_frg[k, j], fastmath=True)
acc_S_row_frg[k + 1, j] = cute.math.exp2(acc_S_row_frg[k + 1, j], fastmath=True)
acc_S_row_converted_frg[None, j].store(
acc_S_row_frg[None, j].load().to(acc_S_row_converted.element_type)
)
@cute.jit
def floor_if_packed(
q_idx,
qhead_per_kvhead: cutlass.Constexpr[int],
) -> cute.Tensor:
"""Convert q_idx to packed format for Pack-GQA."""
if cutlass.const_expr(qhead_per_kvhead == 1):
return q_idx
return q_idx // qhead_per_kvhead
@cute.jit
def apply_score_mod_inner(
score_tensor,
index_tensor,
score_mod: cutlass.Constexpr,
batch_idx,
head_idx,
softmax_scale,
vec_size: cutlass.Constexpr,
qk_acc_dtype: cutlass.Constexpr,
aux_data: AuxData,
fastdiv_mods,
seqlen_info: SeqlenInfoQK,
constant_q_idx: cutlass.Constexpr,
qhead_per_kvhead: cutlass.Constexpr[int] = 1,
transpose_indices: cutlass.Constexpr[bool] = False,
):
"""Shared implementation for applying score modification.
Args:
score_tensor: The scores to modify (acc_S for flash_fwd, tSrS_t2r for sm100)
index_tensor: Index positions (tScS for flash_fwd, tScS_t2r for sm100)
score_mod: The score modification function to apply
batch_idx: Batch index
head_idx: Head index
softmax_scale: Scale to apply
vec_size: Vector size for processing elements
qk_acc_dtype: Data type for accumulator
aux_tensors: Optional aux_tensors for FlexAttention
aux_scalars: Optional runtime scalar captures for FlexAttention
fastdiv_mods: Tuple of (seqlen_q_divmod, seqlen_k_divmod) for wrapping
seqlen_info: Sequence length info
constant_q_idx: If provided, use this constant for all q_idx values
If None, compute q_idx per-element
qhead_per_kvhead_packgqa: Pack-GQA replication factor. Divide q_idx by this
when greater than 1 so score mods see logical heads.
transpose_indices: If True, swap q_idx/kv_idx in index_tensor (for bwd kernel where S is transposed)
"""
# Index positions in the index_tensor tuple
# Forward: index_tensor[...][0] = q_idx, index_tensor[...][1] = kv_idx
# Backward (transposed): index_tensor[...][0] = kv_idx, index_tensor[...][1] = q_idx
if cutlass.const_expr(transpose_indices):
q_idx_pos = cutlass.const_expr(1)
kv_idx_pos = cutlass.const_expr(0)
else:
q_idx_pos = cutlass.const_expr(0)
kv_idx_pos = cutlass.const_expr(1)
n_vals = cutlass.const_expr(cute.size(score_tensor.shape))
score_vec = cute.make_rmem_tensor(vec_size, qk_acc_dtype)
kv_idx_vec = cute.make_rmem_tensor(vec_size, cutlass.Int32)
# SSA values for batch (constant across all elements)
batch_idx_ssa = utils.scalar_to_ssa(batch_idx, cutlass.Int32).broadcast_to((vec_size,))
# Handle q_idx based on whether it's constant
q_idx_vec = cute.make_rmem_tensor(vec_size, cutlass.Int32)
# For Pack-GQA with non-constant q_idx, we need per-element head indices
# since a thread may process multiple query head indices
if cutlass.const_expr(qhead_per_kvhead > 1 and constant_q_idx is None):
head_idx_vec = cute.make_rmem_tensor(vec_size, cutlass.Int32)
for i in cutlass.range(0, n_vals, vec_size, unroll_full=True):
for j in cutlass.range(vec_size, unroll_full=True):
score_vec[j] = score_tensor[i + j] * softmax_scale
# Extract head offset from packed q_idx for Pack-GQA
if cutlass.const_expr(qhead_per_kvhead > 1 and constant_q_idx is None):
q_idx_packed = index_tensor[i + j][q_idx_pos]
# Building up the logical q_head idx: final_q_head = kv_head * qhead_per_kvhead + (q_physical % qhead_per_kvhead)
q_idx_logical = q_idx_packed // qhead_per_kvhead
head_offset = q_idx_packed - q_idx_logical * qhead_per_kvhead
head_idx_vec[j] = head_idx * qhead_per_kvhead + head_offset
# If we will do loads we mod, in order to not read OOB
if cutlass.const_expr(aux_data.tensors is not None and fastdiv_mods is not None):
if cutlass.const_expr(constant_q_idx is None):
seqlen_q_divmod, seqlen_k_divmod = fastdiv_mods
q_idx_floored = floor_if_packed(
index_tensor[i + j][q_idx_pos], qhead_per_kvhead
)
_, q_idx_wrapped = divmod(q_idx_floored, seqlen_q_divmod)
q_idx_vec[j] = q_idx_wrapped
else:
_, seqlen_k_divmod = fastdiv_mods
_, kv_idx_wrapped = divmod(index_tensor[i + j][kv_idx_pos], seqlen_k_divmod)
kv_idx_vec[j] = kv_idx_wrapped
else:
# No bounds checking - direct indexing
if constant_q_idx is None:
q_idx_vec[j] = floor_if_packed(index_tensor[i + j][q_idx_pos], qhead_per_kvhead)
kv_idx_vec[j] = index_tensor[i + j][kv_idx_pos]
# Convert to SSA for score_mod call
score_ssa = score_vec.load()
kv_idx_ssa = kv_idx_vec.load()
if cutlass.const_expr(constant_q_idx is None):
q_idx_ssa = q_idx_vec.load()
else:
# NB we do not apply Pack-GQA division here, as constant_q_idx is assumed to already be logical
q_idx_const = constant_q_idx
q_idx_ssa = utils.scalar_to_ssa(q_idx_const, cutlass.Int32).broadcast_to((vec_size,))
# Compute head_idx_ssa: per-element for Pack-GQA with non-constant q_idx, constant otherwise
if cutlass.const_expr(qhead_per_kvhead > 1 and constant_q_idx is None):
head_idx_ssa = head_idx_vec.load()
else:
head_idx_ssa = utils.scalar_to_ssa(head_idx, cutlass.Int32).broadcast_to((vec_size,))
post_mod_scores = call_score_mod(
score_mod,
score_ssa,
batch_idx_ssa,
head_idx_ssa,
q_idx_ssa,
kv_idx_ssa,
seqlen_info,
aux_data,
)
# Write back modified scores
score_vec.store(post_mod_scores)
for j in cutlass.range(vec_size, unroll_full=True):
score_tensor[i + j] = score_vec[j]
@cute.jit
def apply_score_mod_bwd_inner(
grad_tensor,
score_tensor,
index_tensor,
score_mod_bwd: cutlass.Constexpr,
batch_idx,
head_idx,
softmax_scale,
vec_size: cutlass.Constexpr,
qk_acc_dtype: cutlass.Constexpr,
aux_data: AuxData,
fastdiv_mods,
seqlen_info,
constant_q_idx: cutlass.Constexpr,
qhead_per_kvhead: cutlass.Constexpr[int] = 1,
transpose_indices: cutlass.Constexpr[bool] = False,
):
"""Apply backward score modification (joint graph).
Args:
grad_tensor: in/out: dlogits rewritten in-place with d(scaled_scores)
score_tensor: pre-mod scores (unscaled QK tile), scaled by softmax_scale internally
index_tensor: Index positions (same as forward)
score_mod_bwd: The backward score modification function (joint graph)
batch_idx: Batch index
head_idx: Head index
softmax_scale: Scale to apply to score_tensor
vec_size: Vector size for processing elements
qk_acc_dtype: Data type for accumulator
aux_tensors: Optional aux_tensors for FlexAttention
aux_scalars: Optional runtime scalar captures for FlexAttention
fastdiv_mods: Tuple of (seqlen_q_divmod, seqlen_k_divmod) for wrapping
seqlen_info: Sequence length info
constant_q_idx: If provided, use this constant for all q_idx values
qhead_per_kvhead: Pack-GQA replication factor
transpose_indices: If True, swap q_idx/kv_idx in index_tensor
"""
# Index positions in the index_tensor tuple
# Forward: index_tensor[...][0] = q_idx, index_tensor[...][1] = kv_idx
# Backward (transposed): index_tensor[...][0] = kv_idx, index_tensor[...][1] = q_idx
if cutlass.const_expr(transpose_indices):
q_idx_pos = cutlass.const_expr(1)
kv_idx_pos = cutlass.const_expr(0)
else:
q_idx_pos = cutlass.const_expr(0)
kv_idx_pos = cutlass.const_expr(1)
n_vals = cutlass.const_expr(cute.size(grad_tensor.shape))
grad_vec = cute.make_rmem_tensor(vec_size, qk_acc_dtype)
score_vec = cute.make_rmem_tensor(vec_size, qk_acc_dtype)
kv_idx_vec = cute.make_rmem_tensor(vec_size, cutlass.Int32)
batch_idx_ssa = utils.scalar_to_ssa(batch_idx, cutlass.Int32).broadcast_to((vec_size,))
q_idx_vec = cute.make_rmem_tensor(vec_size, cutlass.Int32)
# For Pack-GQA with non-constant q_idx, we need per-element head indices
if cutlass.const_expr(qhead_per_kvhead > 1 and constant_q_idx is None):
head_idx_vec = cute.make_rmem_tensor(vec_size, cutlass.Int32)
for i in cutlass.range(0, n_vals, vec_size, unroll_full=True):
for j in cutlass.range(vec_size, unroll_full=True):
grad_vec[j] = grad_tensor[i + j]
# Scale score so joint graph sees same value as forward score_mod
score_vec[j] = score_tensor[i + j] * softmax_scale
if cutlass.const_expr(qhead_per_kvhead > 1 and constant_q_idx is None):
q_idx_packed = index_tensor[i + j][q_idx_pos]
q_idx_logical = q_idx_packed // qhead_per_kvhead
head_offset = q_idx_packed - q_idx_logical * qhead_per_kvhead
head_idx_vec[j] = head_idx * qhead_per_kvhead + head_offset
if cutlass.const_expr(aux_data.tensors is not None and fastdiv_mods is not None):
if cutlass.const_expr(constant_q_idx is None):
seqlen_q_divmod, seqlen_k_divmod = fastdiv_mods
q_idx_floored = floor_if_packed(
index_tensor[i + j][q_idx_pos], qhead_per_kvhead
)
_, q_idx_wrapped = divmod(q_idx_floored, seqlen_q_divmod)
q_idx_vec[j] = q_idx_wrapped
else:
_, seqlen_k_divmod = fastdiv_mods
_, kv_idx_wrapped = divmod(index_tensor[i + j][kv_idx_pos], seqlen_k_divmod)
kv_idx_vec[j] = kv_idx_wrapped
else:
# No bounds checking - direct indexing
if constant_q_idx is None:
q_idx_vec[j] = floor_if_packed(index_tensor[i + j][q_idx_pos], qhead_per_kvhead)
kv_idx_vec[j] = index_tensor[i + j][kv_idx_pos]
grad_ssa = grad_vec.load()
score_ssa = score_vec.load()
kv_idx_ssa = kv_idx_vec.load()
if cutlass.const_expr(constant_q_idx is None):
q_idx_ssa = q_idx_vec.load()
else:
q_idx_ssa = utils.scalar_to_ssa(constant_q_idx, cutlass.Int32).broadcast_to((vec_size,))
if cutlass.const_expr(qhead_per_kvhead > 1 and constant_q_idx is None):
head_idx_ssa = head_idx_vec.load()
else:
head_idx_ssa = utils.scalar_to_ssa(head_idx, cutlass.Int32).broadcast_to((vec_size,))
grad_out_ssa = call_score_mod_bwd(
score_mod_bwd,
grad_ssa,
score_ssa,
batch_idx_ssa,
head_idx_ssa,
q_idx_ssa,
kv_idx_ssa,
seqlen_info,
aux_data,
)
grad_vec.store(grad_out_ssa)
for j in cutlass.range(vec_size, unroll_full=True):
grad_tensor[i + j] = grad_vec[j]
+487
View File
@@ -0,0 +1,487 @@
import math
from contextlib import nullcontext
from functools import wraps
from typing import Optional
import torch
import torch.nn.functional as F
from einops import rearrange, repeat
from torch._guards import active_fake_mode
from torch._subclasses.fake_tensor import FakeTensorMode
class IndexFirstAxis(torch.autograd.Function):
@staticmethod
def forward(ctx, input, indices):
ctx.save_for_backward(indices)
assert input.ndim >= 2
ctx.first_axis_dim, other_shape = input.shape[0], input.shape[1:]
second_dim = other_shape.numel()
return torch.gather(
rearrange(input, "b ... -> b (...)"),
0,
repeat(indices, "z -> z d", d=second_dim),
).reshape(-1, *other_shape)
@staticmethod
def backward(ctx, grad_output):
(indices,) = ctx.saved_tensors
assert grad_output.ndim >= 2
other_shape = grad_output.shape[1:]
grad_output = rearrange(grad_output, "b ... -> b (...)")
grad_input = torch.zeros(
[ctx.first_axis_dim, grad_output.shape[1]],
device=grad_output.device,
dtype=grad_output.dtype,
)
grad_input.scatter_(0, repeat(indices, "z -> z d", d=grad_output.shape[1]), grad_output)
return grad_input.reshape(ctx.first_axis_dim, *other_shape), None
index_first_axis = IndexFirstAxis.apply
class IndexPutFirstAxis(torch.autograd.Function):
@staticmethod
def forward(ctx, values, indices, first_axis_dim):
ctx.save_for_backward(indices)
assert indices.ndim == 1
assert values.ndim >= 2
output = torch.zeros(
first_axis_dim, *values.shape[1:], device=values.device, dtype=values.dtype
)
output[indices] = values
return output
@staticmethod
def backward(ctx, grad_output):
(indices,) = ctx.saved_tensors
grad_values = grad_output[indices]
return grad_values, None, None
index_put_first_axis = IndexPutFirstAxis.apply
def unpad_input(hidden_states, attention_mask, unused_mask=None):
all_masks = (attention_mask + unused_mask) if unused_mask is not None else attention_mask
seqlens_in_batch = all_masks.sum(dim=-1, dtype=torch.int32)
used_seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
in_fake_mode = active_fake_mode() is not None
if not in_fake_mode:
indices = torch.nonzero(all_masks.flatten(), as_tuple=False).flatten()
max_seqlen_in_batch = seqlens_in_batch.max().item()
else:
# torch.nonzero and .item() are not supported in FakeTensorMode
batch_size, seqlen = attention_mask.shape
indices = torch.arange(batch_size * seqlen, device=hidden_states.device)
max_seqlen_in_batch = seqlen
cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))
return (
index_first_axis(rearrange(hidden_states, "b s ... -> (b s) ..."), indices),
indices,
cu_seqlens,
max_seqlen_in_batch,
used_seqlens_in_batch,
)
def pad_input(hidden_states, indices, batch, seqlen):
output = index_put_first_axis(hidden_states, indices, batch * seqlen)
return rearrange(output, "(b s) ... -> b s ...", b=batch)
def generate_random_padding_mask(
max_seqlen, batch_size, device, mode="random", zero_lengths=False, min_seqlen=None
):
assert mode in ["full", "random", "third"]
min_seqlen = min_seqlen if min_seqlen is not None else 0 if zero_lengths else 1
if mode == "full":
lengths = torch.full((batch_size, 1), max_seqlen, device=device, dtype=torch.int32)
elif mode == "random":
lengths = torch.randint(
max(min_seqlen, max_seqlen - 20),
max_seqlen + 1,
(batch_size, 1),
device=device,
)
else:
lengths = torch.randint(
max(min_seqlen, max_seqlen // 3),
max_seqlen + 1,
(batch_size, 1),
device=device,
)
if zero_lengths:
for i in range(batch_size):
if i % 5 == 0:
lengths[i] = 0
lengths[-1] = 0
padding_mask = (
repeat(torch.arange(max_seqlen, device=device), "s -> b s", b=batch_size) < lengths
)
return padding_mask
def generate_qkv(
q,
k,
v,
query_padding_mask=None,
key_padding_mask=None,
qv=None,
kvpacked=False,
qkvpacked=False,
query_unused_mask=None,
key_unused_mask=None,
):
assert not (kvpacked and qkvpacked)
batch_size, seqlen_q, nheads, d = q.shape
d_v = v.shape[-1]
_, seqlen_k, nheads_k, _ = k.shape
assert k.shape == (batch_size, seqlen_k, nheads_k, d)
assert v.shape == (batch_size, seqlen_k, nheads_k, d_v)
if query_unused_mask is not None or key_unused_mask is not None:
assert not kvpacked
assert not qkvpacked
if query_padding_mask is not None:
q_unpad, indices_q, cu_seqlens_q, max_seqlen_q, seqused_q = unpad_input(
q, query_padding_mask, query_unused_mask
)
output_pad_fn = lambda output_unpad: pad_input(
output_unpad, indices_q, batch_size, seqlen_q
)
qv_unpad = rearrange(qv, "b s ... -> (b s) ...")[indices_q] if qv is not None else None
else:
q_unpad = rearrange(q, "b s h d -> (b s) h d")
cu_seqlens_q = torch.arange(
0, (batch_size + 1) * seqlen_q, step=seqlen_q, dtype=torch.int32, device=q_unpad.device
)
seqused_q = None
max_seqlen_q = seqlen_q
output_pad_fn = lambda output_unpad: rearrange(
output_unpad, "(b s) h d -> b s h d", b=batch_size
)
qv_unpad = rearrange(qv, "b s ... -> (b s) ...") if qv is not None else None
if key_padding_mask is not None:
k_unpad, indices_k, cu_seqlens_k, max_seqlen_k, seqused_k = unpad_input(
k, key_padding_mask, key_unused_mask
)
v_unpad, *_ = unpad_input(v, key_padding_mask, key_unused_mask)
else:
k_unpad = rearrange(k, "b s h d -> (b s) h d")
v_unpad = rearrange(v, "b s h d -> (b s) h d")
cu_seqlens_k = torch.arange(
0, (batch_size + 1) * seqlen_k, step=seqlen_k, dtype=torch.int32, device=k_unpad.device
)
seqused_k = None
max_seqlen_k = seqlen_k
if qkvpacked:
assert (query_padding_mask == key_padding_mask).all()
assert nheads == nheads_k
qkv_unpad = torch.stack([q_unpad, k_unpad, v_unpad], dim=1)
qkv = torch.stack([q, k, v], dim=2)
if query_padding_mask is not None:
dqkv_pad_fn = lambda dqkv_unpad: pad_input(dqkv_unpad, indices_q, batch_size, seqlen_q)
else:
dqkv_pad_fn = lambda dqkv_unpad: rearrange(
dqkv_unpad, "(b s) t h d -> b s t h d", b=batch_size
)
return (
qkv_unpad.detach().requires_grad_(),
cu_seqlens_q,
max_seqlen_q,
qkv.detach().requires_grad_(),
output_pad_fn,
dqkv_pad_fn,
)
elif kvpacked:
kv_unpad = torch.stack([k_unpad, v_unpad], dim=1)
kv = torch.stack([k, v], dim=2)
dq_pad_fn = output_pad_fn
if key_padding_mask is not None:
dkv_pad_fn = lambda dkv_unpad: pad_input(dkv_unpad, indices_k, batch_size, seqlen_k)
else:
dkv_pad_fn = lambda dkv_unpad: rearrange(
dkv_unpad, "(b s) t h d -> b s t h d", b=batch_size
)
return (
q_unpad.detach().requires_grad_(),
kv_unpad.detach().requires_grad_(),
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
q.detach().requires_grad_(),
kv.detach().requires_grad_(),
output_pad_fn,
dq_pad_fn,
dkv_pad_fn,
)
else:
dq_pad_fn = output_pad_fn
if key_padding_mask is not None:
dk_pad_fn = lambda dk_unpad: pad_input(dk_unpad, indices_k, batch_size, seqlen_k)
else:
dk_pad_fn = lambda dk_unpad: rearrange(dk_unpad, "(b s) h d -> b s h d", b=batch_size)
return (
q_unpad.detach().requires_grad_(),
k_unpad.detach().requires_grad_(),
v_unpad.detach().requires_grad_(),
qv_unpad.detach() if qv is not None else None,
cu_seqlens_q,
cu_seqlens_k,
seqused_q,
seqused_k,
max_seqlen_q,
max_seqlen_k,
q.detach().requires_grad_(),
k.detach().requires_grad_(),
v.detach().requires_grad_(),
qv.detach() if qv is not None else None,
output_pad_fn,
dq_pad_fn,
dk_pad_fn,
)
def construct_local_mask(
seqlen_q,
seqlen_k,
window_size=(None, None),
sink_token_length=0,
query_padding_mask=None,
key_padding_mask=None,
key_leftpad=None,
device=None,
):
row_idx = rearrange(torch.arange(seqlen_q, device=device, dtype=torch.long), "s -> s 1")
col_idx = torch.arange(seqlen_k, device=device, dtype=torch.long)
if key_leftpad is not None:
key_leftpad = rearrange(key_leftpad, "b -> b 1 1 1")
col_idx = repeat(col_idx, "s -> b 1 1 s", b=key_leftpad.shape[0])
col_idx = torch.where(col_idx >= key_leftpad, col_idx - key_leftpad, 2**32)
sk = (
seqlen_k
if key_padding_mask is None
else rearrange(key_padding_mask.sum(-1), "b -> b 1 1 1")
)
sq = (
seqlen_q
if query_padding_mask is None
else rearrange(query_padding_mask.sum(-1), "b -> b 1 1 1")
)
if window_size[0] is None:
return col_idx > row_idx + sk - sq + window_size[1]
else:
sk = torch.full_like(col_idx, seqlen_k) if key_padding_mask is None else sk
if window_size[1] is None:
local_mask_left = col_idx > sk
else:
local_mask_left = col_idx > torch.minimum(row_idx + sk - sq + window_size[1], sk)
return torch.logical_or(
local_mask_left,
torch.logical_and(
col_idx < row_idx + sk - sq - window_size[0], col_idx >= sink_token_length
),
)
def construct_chunk_mask(
seqlen_q,
seqlen_k,
attention_chunk,
query_padding_mask=None,
key_padding_mask=None,
key_leftpad=None,
device=None,
):
row_idx = rearrange(torch.arange(seqlen_q, device=device, dtype=torch.long), "s -> s 1")
col_idx = torch.arange(seqlen_k, device=device, dtype=torch.long)
if key_leftpad is not None:
key_leftpad = rearrange(key_leftpad, "b -> b 1 1 1")
col_idx = repeat(col_idx, "s -> b 1 1 s", b=key_leftpad.shape[0])
col_idx = torch.where(col_idx >= key_leftpad, col_idx - key_leftpad, 2**32)
sk = (
seqlen_k
if key_padding_mask is None
else rearrange(key_padding_mask.sum(-1), "b -> b 1 1 1")
)
sq = (
seqlen_q
if query_padding_mask is None
else rearrange(query_padding_mask.sum(-1), "b -> b 1 1 1")
)
sk = torch.full_like(col_idx, seqlen_k) if key_padding_mask is None else sk
col_limit_left_chunk = row_idx + sk - sq - (row_idx + sk - sq) % attention_chunk
return torch.logical_or(
col_idx < col_limit_left_chunk, col_idx >= col_limit_left_chunk + attention_chunk
)
def attention_ref(
q,
k,
v,
query_padding_mask=None,
key_padding_mask=None,
key_leftpad=None,
attn_bias=None,
dropout_p=0.0,
dropout_mask=None,
causal=False,
qv=None,
q_descale=None,
k_descale=None,
v_descale=None,
window_size=(None, None),
attention_chunk=0,
sink_token_length=0,
learnable_sink: Optional[torch.Tensor] = None,
softcap=0.0,
upcast=True,
reorder_ops=False,
intermediate_dtype=None,
return_lse=False,
gather_kv_indices=None,
):
assert v is not None
has_qk = q is not None and k is not None
assert has_qk or qv is not None
if causal:
window_size = (window_size[0], 0)
dtype_og = v.dtype
q_shape = q.shape if q is not None else qv.shape
if upcast:
q, k, v, qv = [t.float() if t is not None else None for t in (q, k, v, qv)]
if q_descale is not None:
q_descale = repeat(q_descale, "b h -> b 1 (h g) 1", g=q_shape[2] // v.shape[2])
q, qv = [(t.float() * q_descale).to(t.dtype) if t is not None else None for t in (q, qv)]
if k_descale is not None:
k = (k.float() * rearrange(k_descale, "b h -> b 1 h 1")).to(dtype=k.dtype)
if v_descale is not None:
v = (v.float() * rearrange(v_descale, "b h -> b 1 h 1")).to(dtype=v.dtype)
seqlen_q, seqlen_k = q_shape[1], v.shape[1]
k, v = [
repeat(t, "b s h d -> b s (h g) d", g=q_shape[2] // t.shape[2]) if t is not None else None
for t in (k, v)
]
d = q_shape[-1] # == dv for qv
dv = v.shape[-1]
softmax_scale = 1.0 / math.sqrt(d if qv is None or q is None else d + dv)
if has_qk:
scores = torch.einsum(
"bthd,bshd->bhts",
q if reorder_ops else q * softmax_scale,
k * softmax_scale if reorder_ops else k,
)
if qv is not None:
qv_scores = torch.einsum(
"bthd,bshd->bhts",
qv if reorder_ops else qv * softmax_scale,
v * softmax_scale if reorder_ops else v,
)
scores = qv_scores if not has_qk else scores + qv_scores
if softcap > 0:
scores = torch.tanh(scores / softcap) * softcap
if key_padding_mask is not None:
scores.masked_fill_(rearrange(~key_padding_mask, "b s -> b 1 1 s"), float("-inf"))
local_mask = None
if window_size[0] is not None or window_size[1] is not None:
local_mask = construct_local_mask(
seqlen_q,
seqlen_k,
window_size,
sink_token_length,
query_padding_mask,
key_padding_mask,
key_leftpad=key_leftpad,
device=v.device,
)
if attention_chunk > 0:
chunk_mask = construct_chunk_mask(
seqlen_q,
seqlen_k,
attention_chunk,
query_padding_mask,
key_padding_mask,
key_leftpad=key_leftpad,
device=v.device,
)
local_mask = (
torch.logical_or(local_mask, chunk_mask) if local_mask is not None else chunk_mask
)
if gather_kv_indices is not None:
batch = q_shape[0]
topk_len = gather_kv_indices.shape[2]
if topk_len < seqlen_k:
topk_index_mask = torch.full(
(batch, seqlen_q, seqlen_k), False, device="cuda"
).scatter_(-1, gather_kv_indices, True)
scores.masked_fill_(rearrange(~topk_index_mask, "b t s -> b 1 t s"), float("-inf"))
if local_mask is not None:
scores.masked_fill_(local_mask, float("-inf"))
if attn_bias is not None:
scores = scores + attn_bias
# After all masks are applied, before softmax:
# scores shape: [b, h, t, s]
lse = torch.logsumexp(scores, dim=-1) # [b, h, t]
if learnable_sink is None:
attention = torch.softmax(scores, dim=-1).to(v.dtype)
else:
scores_fp32 = scores.to(torch.float32)
logits_max = torch.amax(scores_fp32, dim=-1, keepdim=True)
learnable_sink = rearrange(learnable_sink, "h -> h 1 1")
logits_or_sinks_max = torch.maximum(learnable_sink, logits_max)
unnormalized_scores = torch.exp(scores_fp32 - logits_or_sinks_max)
normalizer = unnormalized_scores.sum(dim=-1, keepdim=True) + torch.exp(
learnable_sink - logits_or_sinks_max
)
# LSE with sink: log(Z) = log(normalizer) + max
lse = (torch.log(normalizer.squeeze(-1)) + logits_or_sinks_max.squeeze(-1)).to(dtype_og)
attention = (unnormalized_scores / normalizer).to(v.dtype)
if query_padding_mask is not None:
attention = attention.masked_fill(rearrange(~query_padding_mask, "b s -> b 1 s 1"), 0.0)
if key_padding_mask is not None:
attention = attention.masked_fill(rearrange(~key_padding_mask, "b s -> b 1 1 s"), 0.0)
if local_mask is not None:
attention = attention.masked_fill(torch.all(local_mask, dim=-1, keepdim=True), 0.0)
dropout_scaling = 1.0 / (1 - dropout_p)
if dropout_mask is not None:
attention_drop = attention.masked_fill(~dropout_mask, 0.0)
else:
attention_drop = attention
if intermediate_dtype is not None:
attention_drop = attention_drop.to(intermediate_dtype).to(attention_drop.dtype)
output = torch.einsum("bhts,bshd->bthd", attention_drop, v * dropout_scaling)
if query_padding_mask is not None:
output.masked_fill_(rearrange(~query_padding_mask, "b s -> b s 1 1"), 0.0)
if return_lse:
return output.to(dtype_og), attention.to(dtype_og), lse.to(dtype_og)
return output.to(dtype=dtype_og), attention.to(dtype=dtype_og)
def maybe_fake_tensor_mode(fake: bool = True):
"""
One way to populate/pre-compile cache is to use torch fake tensor mode,
which does not allocate actual GPU tensors but retains tensor shape/dtype
metadata for cute.compile.
"""
def decorator(fn):
@wraps(fn)
def wrapper(*args, **kwargs):
with FakeTensorMode() if fake else nullcontext():
return fn(*args, **kwargs)
return wrapper
return decorator
def is_fake_mode() -> bool:
return active_fake_mode() is not None
File diff suppressed because it is too large Load Diff
+278
View File
@@ -0,0 +1,278 @@
from typing import Type, Optional
from dataclasses import dataclass
import operator
import cutlass
import cutlass.cute as cute
import cutlass.pipeline as pipeline
from cutlass.cute.nvgpu import cpasync
from cutlass import Int32, Uint32, const_expr, Boolean
from flash_attn.cute import utils
from flash_attn.cute.utils import warp_reduce
from quack.cute_dsl_utils import ParamsBase
import math
@dataclass
class CpasyncGatherKVManager(ParamsBase):
mIndexTopk: cute.Tensor
sBitmask: Optional[cute.Tensor]
cta_rank_in_cluster: Int32
thread_idx: Int32
warp_idx: Int32
topk_length: Int32
seqlen_k_limit: Int32
tile_n: Int32
num_threads: cutlass.Constexpr[Int32]
hdim: cutlass.Constexpr[Int32]
hdim_v: cutlass.Constexpr[Int32]
num_hdimv_splits: cutlass.Constexpr[Int32]
cta_group_size: cutlass.Constexpr[Int32]
gmem_threads_per_row: cutlass.Constexpr[Int32]
topk_indices_per_thread: Int32
async_copy_elems: Int32
gmem_tiled_copy_KV: cute.TiledCopy
gmem_thr_copy_KV: cute.TiledCopy
rTopk: cute.Tensor
rTopkHalf: cute.Tensor
# for bitmask
rTopk_NonInterleaved: cute.Tensor
pipeline_bitmask: Optional[pipeline.PipelineAsync]
cpasync_barrier: Optional[pipeline.NamedBarrier]
disable_bitmask: cutlass.Constexpr[Boolean]
@staticmethod
def create(
mIndexTopk: cute.Tensor,
cta_rank_in_cluster: Int32,
thread_idx: Int32,
warp_idx: Int32,
topk_length: Int32,
seqlen_k_limit: Int32,
tile_n: cutlass.Constexpr[Int32],
hdim: cutlass.Constexpr[Int32],
hdim_v: cutlass.Constexpr[Int32],
num_hdimv_splits: cutlass.Constexpr[Int32],
num_threads: cutlass.Constexpr[Int32],
dtype: Type[cutlass.Numeric],
cta_group_size: cutlass.Constexpr[Int32],
cpasync_barrier: Optional[pipeline.NamedBarrier] = None,
disable_bitmask: cutlass.Constexpr[Boolean] = False,
sBitmask: Optional[cute.Tensor] = None,
pipeline_bitmask: Optional[pipeline.PipelineAsync] = None,
):
assert tile_n % num_threads == 0
assert num_threads == 128
assert hdim % 64 == 0
assert (hdim_v // num_hdimv_splits // cta_group_size) % 64 == 0
assert num_threads % cute.arch.WARP_SIZE == 0
universal_copy_bits = 128
async_copy_elems = universal_copy_bits // dtype.width
dtype_bytes = dtype.width // 8
# assumes hdim is never part of transposed operand
gmem_k_block_size = math.gcd(
hdim,
hdim_v // num_hdimv_splits // cta_group_size,
128 // dtype_bytes,
)
assert gmem_k_block_size % async_copy_elems == 0
gmem_threads_per_row = gmem_k_block_size // async_copy_elems
assert cute.arch.WARP_SIZE % gmem_threads_per_row == 0
atom_async_copy = cute.make_copy_atom(
cpasync.CopyG2SOp(cache_mode=cpasync.LoadCacheMode.GLOBAL),
dtype,
num_bits_per_copy=universal_copy_bits,
)
thr_layout = cute.make_ordered_layout(
(num_threads // gmem_threads_per_row, gmem_threads_per_row),
order=(1, 0),
)
val_layout = cute.make_layout((1, async_copy_elems))
gmem_tiled_copy_KV = cute.make_tiled_copy_tv(atom_async_copy, thr_layout, val_layout)
gmem_thr_copy_KV = gmem_tiled_copy_KV.get_slice(thread_idx)
topk_indices_per_thread = tile_n // num_threads
rTopk = cute.make_rmem_tensor((topk_indices_per_thread,), Int32)
rTopkHalf = cute.make_rmem_tensor((topk_indices_per_thread,), Int32)
rTopk_NonInterleaved = cute.make_rmem_tensor((topk_indices_per_thread,), Int32)
return CpasyncGatherKVManager(
mIndexTopk,
sBitmask,
cta_rank_in_cluster,
thread_idx,
warp_idx,
topk_length,
seqlen_k_limit,
tile_n,
num_threads,
hdim,
hdim_v,
num_hdimv_splits,
cta_group_size,
gmem_threads_per_row,
topk_indices_per_thread,
async_copy_elems,
gmem_tiled_copy_KV,
gmem_thr_copy_KV,
rTopk,
rTopkHalf,
rTopk_NonInterleaved,
pipeline_bitmask,
cpasync_barrier,
disable_bitmask,
)
@cute.jit
def load_index_topk(
self,
n_block: Int32,
transpose: bool,
):
entries_per_thread = self.topk_indices_per_thread
rTopk = self.rTopk if const_expr(transpose) else self.rTopkHalf
for i in cutlass.range_constexpr(entries_per_thread):
row = (
i * self.num_threads
+ (self.thread_idx % self.gmem_threads_per_row)
* (self.num_threads // self.gmem_threads_per_row)
+ (self.thread_idx // self.gmem_threads_per_row)
)
# need this if not offset in load_X
# if const_expr(not transpose):
# row += self.cta_rank_in_cluster * (self.tile_n//self.cta_group_size)
# row = row % self.tile_n
row_idx = n_block * self.tile_n + row
rTopk[i] = self.mIndexTopk[row_idx]
if const_expr(not transpose and not self.disable_bitmask):
row_non_interleaved = i * self.num_threads + self.thread_idx
row_idx_non_interleaved = n_block * self.tile_n + row_non_interleaved
self.rTopk_NonInterleaved[0] = self.mIndexTopk[row_idx_non_interleaved]
@cute.jit
def compute_bitmask(
self,
producer_state_bitmask,
):
assert self.pipeline_bitmask is not None, "pipeline_bitmask not provided"
assert self.cpasync_barrier is not None, "cpasync barrier not provided"
lane_idx = cute.arch.lane_idx()
assert cute.size(self.rTopk_NonInterleaved) == 1
bitmask = Uint32(0)
# Step 1. Construct per-thread bitmask
topk_idx = self.rTopk_NonInterleaved[0]
is_valid = topk_idx >= 0 and topk_idx < self.seqlen_k_limit
if is_valid:
bitmask = Uint32(1 << lane_idx)
# Step 2. Warp shuffle bitwise OR = add since indices are exclusive.
bitmask = warp_reduce(bitmask, operator.add)
self.pipeline_bitmask.producer_acquire(producer_state_bitmask)
# store to smem and sync threads
if lane_idx == 0:
self.sBitmask[self.warp_idx, producer_state_bitmask.index] = bitmask
self.cpasync_barrier.arrive_and_wait()
self.pipeline_bitmask.producer_commit(producer_state_bitmask)
producer_state_bitmask.advance()
return producer_state_bitmask
@cute.jit
def compute_X_ptr(
self,
mX: cute.Tensor,
transpose: bool,
d_offset: int = 0,
):
entries_per_thread = self.topk_indices_per_thread
tPrXPtr = cute.make_rmem_tensor((entries_per_thread,), cutlass.Int64)
tPrRowValid = cute.make_rmem_tensor((entries_per_thread,), cutlass.Int32)
rTopk = self.rTopk if const_expr(transpose) else self.rTopkHalf
for i in cutlass.range_constexpr(entries_per_thread):
topk_idx = rTopk[i]
if const_expr(not self.disable_bitmask):
row_valid = topk_idx >= 0 and topk_idx < self.seqlen_k_limit
tPrRowValid[i] = row_valid
if const_expr(not transpose):
tPrXPtr[i] = utils.elem_pointer(mX, (topk_idx, d_offset)).toint()
else:
tPrXPtr[i] = utils.elem_pointer(mX, (d_offset, topk_idx)).toint()
return tPrXPtr, tPrRowValid
@cute.jit
def load_X(
self,
mX: cute.Tensor,
sX: cute.Tensor,
transpose: bool,
K_or_V: str,
d_offset: int = 0,
):
assert K_or_V in ("K", "V")
cta_tile_n = self.tile_n if const_expr(transpose) else self.tile_n // self.cta_group_size
head_dim = self.hdim if const_expr(K_or_V == "K") else self.hdim_v // self.num_hdimv_splits
if const_expr(transpose):
head_dim = head_dim // self.cta_group_size
order = (1, 0) if const_expr(transpose) else (0, 1)
sX_nd_layout = cute.make_ordered_layout((cta_tile_n, head_dim), order=order)
sX_nd = cute.composition(sX, sX_nd_layout)
cX = cute.make_identity_tensor((cta_tile_n, head_dim))
tXsX = self.gmem_thr_copy_KV.partition_D(sX_nd)
tXcX = self.gmem_thr_copy_KV.partition_S(cX)
tPrXPtr, tPrRowValid = self.compute_X_ptr(mX, transpose, d_offset)
if const_expr(not transpose):
offset = self.cta_rank_in_cluster * (self.gmem_threads_per_row // self.cta_group_size)
else:
offset = 0
for m in cutlass.range_constexpr(cute.size(tXsX, mode=[1])):
if const_expr(not self.disable_bitmask):
row_valid = utils.shuffle_sync(
tPrRowValid[m // self.gmem_threads_per_row],
(m + offset) % self.gmem_threads_per_row,
width=self.gmem_threads_per_row,
)
should_load = cute.make_fragment_like(tXsX[(0, None), m, 0], Boolean)
should_load.fill(Boolean(row_valid))
x_ptr_i64 = utils.shuffle_sync(
tPrXPtr[m // self.gmem_threads_per_row],
(m + offset) % self.gmem_threads_per_row,
width=self.gmem_threads_per_row,
)
x_gmem_ptr = cute.make_ptr(
mX.element_type, x_ptr_i64, cute.AddressSpace.gmem, assumed_align=16
)
mX_cur = cute.make_tensor(x_gmem_ptr, cute.make_layout((head_dim,)))
mX_cur_copy = cute.tiled_divide(mX_cur, (self.async_copy_elems,))
for k in cutlass.range_constexpr(cute.size(tXsX, mode=[2])):
ki = tXcX[0, 0, k][1] // self.async_copy_elems
mX_cur_copy_ki = mX_cur_copy[None, ki]
tXsX_k = tXsX[None, m, k]
mX_cur_copy_ki = cute.make_tensor(mX_cur_copy_ki.iterator, tXsX_k.layout)
cute.copy(
self.gmem_tiled_copy_KV,
mX_cur_copy_ki,
tXsX_k,
pred=should_load if const_expr(not self.disable_bitmask) else None,
)
+957
View File
@@ -0,0 +1,957 @@
# Copyright (c) 2025, Tri Dao.
import math
import hashlib
import inspect
import os
from functools import partial
from typing import Type, Callable, Optional, Tuple, overload, NamedTuple
import cutlass
import cutlass.cute as cute
from cutlass import Float32, Int32, const_expr
from cutlass.cute import FastDivmodDivisor
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import nvvm, llvm
from cutlass.cute.runtime import from_dlpack
import quack.activation
_MIXER_ATTRS = ("__vec_size__",)
class AuxData(NamedTuple):
tensors: tuple | list | None = None
scalars: tuple | None = None
# Obtained from sollya:
# fpminimax(exp(x * log(2.0)), 1, [|1,24...|],[0;1],relative);
POLY_EX2 = {
0: (1.0),
1: (
1.0,
0.922497093677520751953125,
),
2: (
1.0,
0.6657850742340087890625,
0.330107033252716064453125,
),
3: (
1.0,
0.695146143436431884765625,
0.227564394474029541015625,
0.077119089663028717041015625,
),
4: (
1.0,
0.693042695522308349609375,
0.2412912547588348388671875,
5.2225358784198760986328125e-2,
1.3434938155114650726318359375e-2,
),
5: (
1.0,
0.693151414394378662109375,
0.24016360938549041748046875,
5.5802188813686370849609375e-2,
9.01452265679836273193359375e-3,
1.86810153536498546600341796875e-3,
),
}
_fa_clc_enabled: bool = os.environ.get("FA_CLC", "0") == "1"
_fa_disable_2cta_enabled: bool = os.environ.get("FA_DISABLE_2CTA", "0") == "1"
def _is_cuda_12() -> bool:
"""Check if the CUDA toolkit version is 12.x.
2CTA forward non-causal has a codegen regression on CUDA 12 that causes
~18% slowdown compared to 1CTA. This is fixed in CUDA 13.x.
"""
try:
import torch
cuda_version = torch.version.cuda
if cuda_version is not None:
major = cuda_version.split(".")[0]
return int(major) == 12
except Exception:
pass
return False
_fa_disable_2cta_cuda12: bool = _is_cuda_12()
def _get_use_clc_scheduler_default() -> bool:
return _fa_clc_enabled
def _get_disable_2cta_default(is_fwd: bool = False) -> bool:
if is_fwd:
return _fa_disable_2cta_enabled or _fa_disable_2cta_cuda12
else:
return _fa_disable_2cta_enabled
def _compute_base_hash(func: Callable) -> str:
"""Compute hash from source code or bytecode and closure values."""
try:
data = inspect.getsource(func).encode()
except (OSError, TypeError):
if hasattr(func, "__code__") and func.__code__ is not None:
data = func.__code__.co_code
else:
data = repr(func).encode()
hasher = hashlib.sha256(data)
if hasattr(func, "__closure__") and func.__closure__ is not None:
for cell in func.__closure__:
hasher.update(repr(cell.cell_contents).encode())
return hasher.hexdigest()
def hash_callable(
func: Callable, mixer_attrs: Tuple[str] = _MIXER_ATTRS, set_cute_hash: bool = True
) -> str:
"""Hash a callable based on the source code or bytecode and closure values.
Fast-path: if the callable (or its __wrapped__ base) has a ``__cute_hash__``
attribute, that value is returned immediately as the base hash, then
metadata dunders are mixed in to produce the final dict-key hash.
set_cute_hash: whether or not to set func.__cute_hash__
"""
# Resolve base hash
if hasattr(func, "__cute_hash__"):
base_hash = func.__cute_hash__
else:
# Unwrap decorated functions (e.g., cute.jit wrappers).
base_func = getattr(func, "__wrapped__", func)
if hasattr(base_func, "__cute_hash__"):
base_hash = base_func.__cute_hash__
else:
base_hash = _compute_base_hash(base_func)
if set_cute_hash:
base_func.__cute_hash__ = base_hash
# Mix in mutable metadata dunders
mixer_values = tuple(getattr(func, attr, None) for attr in mixer_attrs)
if all(v is None for v in mixer_values):
return base_hash
hasher = hashlib.sha256(base_hash.encode())
for attr, val in zip(mixer_attrs, mixer_values):
hasher.update(f"{attr}={val!r}".encode())
return hasher.hexdigest()
def create_softcap_scoremod(softcap_val):
@cute.jit
def scoremod_premask_fn(
acc_S_SSA, batch_idx, head_idx, q_idx, kv_idx, seqlen_info, aux_tensors
):
scores = acc_S_SSA / softcap_val
return softcap_val * cute.math.tanh(scores, fastmath=True)
return scoremod_premask_fn
def create_softcap_scoremod_bwd(softcap_val):
@cute.jit
def scoremod_bwd_fn(
grad_out_SSA, score_SSA, batch_idx, head_idx, q_idx, kv_idx, seqlen_info, aux_tensors
):
scores = score_SSA / softcap_val
tanh_scores = cute.math.tanh(scores, fastmath=True)
return grad_out_SSA * (1.0 - tanh_scores * tanh_scores)
return scoremod_bwd_fn
LOG2_E = math.log2(math.e)
def compute_softmax_scale_log2(softmax_scale, score_mod):
"""Compute softmax_scale_log2 and adjusted softmax_scale based on whether score_mod is used.
When score_mod is None, fold the log2(e) factor into softmax_scale_log2 and set softmax_scale
to None. When score_mod is present, keep softmax_scale separate so it can be applied before
the score_mod, and set softmax_scale_log2 to just the change-of-base constant.
Returns (softmax_scale_log2, softmax_scale).
"""
if const_expr(score_mod is None):
return softmax_scale * LOG2_E, None
else:
return LOG2_E, softmax_scale
def compute_fastdiv_mods(mQ, mK, qhead_per_kvhead, pack_gqa, aux_tensors, mPageTable=None):
"""Compute FastDivmodDivisor pairs for aux_tensors index computation.
Returns a (seqlen_q_divmod, seqlen_k_divmod) tuple, or None if aux_tensors is None.
"""
if const_expr(aux_tensors is None):
return None
seqlen_q = cute.size(mQ.shape[0]) // (qhead_per_kvhead if const_expr(pack_gqa) else 1)
seqlen_k = (
cute.size(mK.shape[0])
if const_expr(mPageTable is None)
else mK.shape[0] * mPageTable.shape[1]
)
return (FastDivmodDivisor(seqlen_q), FastDivmodDivisor(seqlen_k))
def convert_from_dlpack(x, leading_dim, alignment=16, divisibility=1) -> cute.Tensor:
return (
from_dlpack(x, assumed_align=alignment)
.mark_layout_dynamic(leading_dim=leading_dim)
.mark_compact_shape_dynamic(
mode=leading_dim, stride_order=x.dim_order(), divisibility=divisibility
)
)
def convert_from_dlpack_compact_dynamic(
x,
*,
dynamic_modes: tuple[int, ...],
alignment: int = 16,
stride_order=None,
divisibility: int = 1,
enable_tvm_ffi: bool = False,
) -> cute.Tensor:
"""Convert via DLPack and mark selected compact dimensions as dynamic."""
if isinstance(dynamic_modes, int):
dynamic_modes = (dynamic_modes,)
if stride_order is None:
stride_order = x.dim_order()
t = (
from_dlpack(x, assumed_align=alignment, enable_tvm_ffi=True)
if enable_tvm_ffi
else from_dlpack(x, assumed_align=alignment)
)
for mode in dynamic_modes:
t = t.mark_compact_shape_dynamic(
mode=mode,
stride_order=stride_order,
divisibility=divisibility,
)
return t
def convert_from_dlpack_leading_static(
x, leading_dim, alignment=16, static_modes=None, stride_order=None
) -> cute.Tensor:
if stride_order is None:
stride_order = x.dim_order()
x_ = from_dlpack(x, assumed_align=alignment)
for i in range(x.ndim):
if i != leading_dim and (static_modes is None or i not in static_modes):
x_ = x_.mark_compact_shape_dynamic(mode=i, stride_order=stride_order)
return x_
def make_tiled_copy_A(
copy_atom: cute.CopyAtom, tiled_mma: cute.TiledMma, swapAB: cutlass.Constexpr[bool] = False
) -> cute.TiledCopy:
if const_expr(swapAB):
return cute.make_tiled_copy_B(copy_atom, tiled_mma)
else:
return cute.make_tiled_copy_A(copy_atom, tiled_mma)
def make_tiled_copy_B(
copy_atom: cute.CopyAtom, tiled_mma: cute.TiledMma, swapAB: cutlass.Constexpr[bool] = False
) -> cute.TiledCopy:
if const_expr(swapAB):
return cute.make_tiled_copy_A(copy_atom, tiled_mma)
else:
return cute.make_tiled_copy_B(copy_atom, tiled_mma)
def mma_make_fragment_A(
smem: cute.Tensor, thr_mma: cute.ThrMma, swapAB: cutlass.Constexpr[bool] = False
) -> cute.Tensor:
if const_expr(swapAB):
return mma_make_fragment_B(smem, thr_mma)
else:
return thr_mma.make_fragment_A(thr_mma.partition_A(smem))
def mma_make_fragment_B(
smem: cute.Tensor, thr_mma: cute.ThrMma, swapAB: cutlass.Constexpr[bool] = False
) -> cute.Tensor:
if const_expr(swapAB):
return mma_make_fragment_A(smem, thr_mma)
else:
return thr_mma.make_fragment_B(thr_mma.partition_B(smem))
def get_smem_store_atom(
arch: cutlass.Constexpr[int], element_type: Type[cute.Numeric], transpose: bool = False
) -> cute.CopyAtom:
if const_expr(arch < 90 or element_type.width != 16):
return cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
element_type,
num_bits_per_copy=2 * element_type.width,
)
else:
return cute.make_copy_atom(
cute.nvgpu.warp.StMatrix8x8x16bOp(transpose=transpose, num_matrices=4),
element_type,
)
@cute.jit
def warp_reduce(
val: cute.TensorSSA | cute.Numeric,
op: Callable,
width: cutlass.Constexpr[int] = cute.arch.WARP_SIZE,
) -> cute.TensorSSA | cute.Numeric:
if const_expr(isinstance(val, cute.TensorSSA)):
res = cute.make_rmem_tensor(val.shape, val.dtype)
res.store(val)
for i in cutlass.range_constexpr(cute.size(val.shape)):
res[i] = warp_reduce(res[i], op, width)
return res.load()
else:
for i in cutlass.range_constexpr(int(math.log2(width))):
val = op(val, cute.arch.shuffle_sync_bfly(val, offset=1 << i))
return val
@dsl_user_op
def smid(*, loc=None, ip=None) -> Int32:
return Int32(
llvm.inline_asm(
T.i32(),
[],
"mov.u32 $0, %smid;",
"=r",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def fmax(
a: float | Float32, b: float | Float32, c: float | Float32 | None = None, *, loc=None, ip=None
) -> Float32:
return Float32(
nvvm.fmax(
Float32(a).ir_value(loc=loc, ip=ip),
Float32(b).ir_value(loc=loc, ip=ip),
c=Float32(c).ir_value(loc=loc, ip=ip) if c is not None else None,
loc=loc,
ip=ip,
)
)
@cute.jit
def fmax_reduce(
x: cute.TensorSSA, init_val: float | Float32 | None = None, arch: cutlass.Constexpr[int] = 80
) -> Float32:
if const_expr(arch < 100 or cute.size(x.shape) % 8 != 0):
# if const_expr(init_val is None):
# init_val = -cutlass.Float32.if
# return x.reduce(cute.ReductionOp.MAX, init_val, 0)
res = cute.make_rmem_tensor(x.shape, Float32)
res.store(x)
# local_max = [res[0], res[1]]
# for i in cutlass.range_constexpr(2, cute.size(x.shape), 2):
# local_max[0] = fmax(local_max[0], res[i + 0])
# local_max[1] = fmax(local_max[1], res[i + 1])
# local_max[0] = fmax(local_max[0], local_max[1])
# return local_max[0] if const_expr(init_val is None) else fmax(local_max[0], init_val)
local_max = [res[0], res[1], res[2], res[3]]
for i in cutlass.range_constexpr(4, cute.size(x.shape), 4):
local_max[0] = fmax(local_max[0], res[i + 0])
local_max[1] = fmax(local_max[1], res[i + 1])
local_max[2] = fmax(local_max[2], res[i + 2])
local_max[3] = fmax(local_max[3], res[i + 3])
local_max[0] = fmax(local_max[0], local_max[1])
local_max[2] = fmax(local_max[2], local_max[3])
local_max[0] = fmax(local_max[0], local_max[2])
return local_max[0] if const_expr(init_val is None) else fmax(local_max[0], init_val)
else:
# [2025-06-15] x.reduce only seems to use 50% 3-input max and 50% 2-input max
# We instead force the 3-input max.
res = cute.make_rmem_tensor(x.shape, Float32)
res.store(x)
local_max_0 = (
fmax(init_val, res[0], res[1])
if const_expr(init_val is not None)
else fmax(res[0], res[1])
)
local_max = [
local_max_0,
fmax(res[2], res[3]),
fmax(res[4], res[5]),
fmax(res[6], res[7]),
]
for i in cutlass.range_constexpr(8, cute.size(x.shape), 8):
local_max[0] = fmax(local_max[0], res[i], res[i + 1])
local_max[1] = fmax(local_max[1], res[i + 2], res[i + 3])
local_max[2] = fmax(local_max[2], res[i + 4], res[i + 5])
local_max[3] = fmax(local_max[3], res[i + 6], res[i + 7])
local_max[0] = fmax(local_max[0], local_max[1])
return fmax(local_max[0], local_max[2], local_max[3])
@cute.jit
def fadd_reduce(
x: cute.TensorSSA, init_val: float | Float32 | None = None, arch: cutlass.Constexpr[int] = 80
) -> Float32:
if const_expr(arch < 100 or cute.size(x.shape) % 8 != 0):
if const_expr(init_val is None):
init_val = Float32.zero
return x.reduce(cute.ReductionOp.ADD, init_val, 0)
# res = cute.make_rmem_tensor(x.shape, Float32)
# res.store(x)
# local_sum = [res[0], res[1], res[2], res[3]]
# for i in cutlass.range_constexpr(4, cute.size(x.shape), 4):
# local_sum[0] += res[i + 0]
# local_sum[1] += res[i + 1]
# local_sum[2] += res[i + 2]
# local_sum[3] += res[i + 3]
# local_sum[0] += local_sum[1]
# local_sum[2] += local_sum[3]
# local_sum[0] += local_sum[2]
# return local_sum[0] if const_expr(init_val is None) else local_sum[0] + init_val
else:
res = cute.make_rmem_tensor(x.shape, Float32)
res.store(x)
local_sum_0 = (
cute.arch.add_packed_f32x2((init_val, 0.0), (res[0], res[1]))
# cute.arch.add_packed_f32x2((init_val / 2, init_val / 2), (res[0], res[1]))
if const_expr(init_val is not None)
else (res[0], res[1])
)
local_sum = [local_sum_0, (res[2], res[3]), (res[4], res[5]), (res[6], res[7])]
for i in cutlass.range_constexpr(8, cute.size(x.shape), 8):
local_sum[0] = cute.arch.add_packed_f32x2(local_sum[0], (res[i + 0], res[i + 1]))
local_sum[1] = cute.arch.add_packed_f32x2(local_sum[1], (res[i + 2], res[i + 3]))
local_sum[2] = cute.arch.add_packed_f32x2(local_sum[2], (res[i + 4], res[i + 5]))
local_sum[3] = cute.arch.add_packed_f32x2(local_sum[3], (res[i + 6], res[i + 7]))
local_sum[0] = cute.arch.add_packed_f32x2(local_sum[0], local_sum[1])
local_sum[2] = cute.arch.add_packed_f32x2(local_sum[2], local_sum[3])
local_sum[0] = cute.arch.add_packed_f32x2(local_sum[0], local_sum[2])
return local_sum[0][0] + local_sum[0][1]
@dsl_user_op
def atomic_add_fp32(a: float | Float32, gmem_ptr: cute.Pointer, *, loc=None, ip=None) -> None:
# gmem_ptr_i64 = gmem_ptr.toint(loc=loc, ip=ip).ir_value()
# # cache_hint = cutlass.Int64(0x12F0000000000000)
# llvm.inline_asm(
# None,
# [gmem_ptr_i64, Float32(a).ir_value(loc=loc, ip=ip)],
# # [gmem_ptr_i64, Float32(a).ir_value(loc=loc, ip=ip), cache_hint.ir_value()],
# "red.global.add.f32 [$0], $1;",
# # "red.global.add.L2::cache_hint.f32 [$0], $1, 0x12F0000000000000;",
# # "red.global.add.L2::cache_hint.f32 [$0], $1, $2;",
# "l,f",
# # "l,f,l",
# has_side_effects=True,
# is_align_stack=False,
# asm_dialect=llvm.AsmDialect.AD_ATT,
# )
nvvm.atomicrmw(
res=T.f32(), op=nvvm.AtomicOpKind.FADD, ptr=gmem_ptr.llvm_ptr, a=Float32(a).ir_value()
)
@dsl_user_op
def elem_pointer(x: cute.Tensor, coord: cute.Coord, *, loc=None, ip=None) -> cute.Pointer:
return x.iterator + cute.crd2idx(coord, x.layout, loc=loc, ip=ip)
@cute.jit
def predicate_k(tAcA: cute.Tensor, limit: cutlass.Int32) -> cute.Tensor:
# Only compute predicates for the "k" dimension. For the mn dimension, we will use "if"
tApA = cute.make_rmem_tensor(
cute.make_layout(
(cute.size(tAcA, mode=[0, 1]), cute.size(tAcA, mode=[1]), cute.size(tAcA, mode=[2])),
stride=(cute.size(tAcA, mode=[2]), 0, 1),
),
cutlass.Boolean,
)
for rest_v in cutlass.range_constexpr(tApA.shape[0]):
for rest_k in cutlass.range_constexpr(tApA.shape[2]):
tApA[rest_v, 0, rest_k] = cute.elem_less(tAcA[(0, rest_v), 0, rest_k][1], limit)
return tApA
def canonical_warp_group_idx(sync: bool = True) -> cutlass.Int32:
warp_group_idx = cute.arch.thread_idx()[0] // 128
if const_expr(sync):
warp_group_idx = cute.arch.make_warp_uniform(warp_group_idx)
return warp_group_idx
# @dsl_user_op
# def warp_vote_any_lt(a: float | Float32, b: float | Float32, *, loc=None, ip=None) -> cutlass.Boolean:
# mask = cutlass.Int32(-1)
# return cutlass.Boolean(
# llvm.inline_asm(
# T.i32(),
# [Float32(a).ir_value(loc=loc, ip=ip), Float32(b).ir_value(loc=loc, ip=ip), mask.ir_value(loc=loc, ip=ip)],
# ".pred p1, p2;\n"
# "setp.lt.f32 p1, $1, $2;\n"
# "vote.sync.any.pred p2, p1, $3;\n"
# "selp.u32 $0, 1, 0, p2;",
# # "selp.u32 $0, 1, 0, p1;",
# "=r,f,f,r",
# has_side_effects=False,
# is_align_stack=False,
# asm_dialect=llvm.AsmDialect.AD_ATT,
# )
# )
@cute.jit
def shuffle_sync(
value: cute.Numeric,
offset: cute.typing.Int,
width: cutlass.Constexpr[int] = cute.arch.WARP_SIZE,
) -> cute.Numeric:
assert value.width % 32 == 0, "value type must be a multiple of 32 bits"
# 1 -> 0b11111, 2 -> 0b11110, 4 -> 0b11100, 8 -> 0b11000, 16 -> 0b10000, 32 -> 0b00000
mask = cute.arch.WARP_SIZE - width
clamp = cute.arch.WARP_SIZE - 1
mask_and_clamp = mask << 8 | clamp
# important: need stride 1 and not 0 for recast_tensor to work
val = cute.make_rmem_tensor(cute.make_layout((1,), stride=(1,)), type(value))
val[0] = value
val_i32 = cute.recast_tensor(val, cutlass.Int32)
for i in cutlass.range_constexpr(cute.size(val_i32)):
val_i32[i] = cute.arch.shuffle_sync(val_i32[i], offset, mask_and_clamp=mask_and_clamp)
return val[0]
@dsl_user_op
def shl_u32(val: cutlass.Uint32, shift: cutlass.Uint32, *, loc=None, ip=None) -> cutlass.Uint32:
"""
Left-shift val by shift bits using PTX shl.b32 (sign-agnostic).
Named ``shl_u32`` (not ``shl_b32``) because python type annotations
distinguish signed/unsigned.
PTX semantics (§9.7.8.8): "Shift amounts greater than the register width N
are clamped to N." So ``shl.b32 d, a, 32`` is well-defined and yields 0.
This differs from C/C++ and LLVM IR, where shifting by >= the type width is
undefined behavior. CuTeDSL compiles through MLIR -> LLVM IR, so a plain
Python-level ``Uint32(x) << Uint32(n)`` inherits LLVM's UB: the optimizer
may treat the result as poison and eliminate dependent code. Inline PTX
bypasses the LLVM IR shift entirely — the instruction is emitted verbatim
into PTX where clamping makes it safe for all shift amounts.
"""
return cutlass.Uint32(
llvm.inline_asm(
T.i32(),
[
cutlass.Uint32(val).ir_value(loc=loc, ip=ip),
cutlass.Uint32(shift).ir_value(loc=loc, ip=ip),
],
"shl.b32 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def shr_u32(val: cutlass.Uint32, shift: cutlass.Uint32, *, loc=None, ip=None) -> cutlass.Uint32:
"""
Unsigned right-shift val by shift bits using PTX shr.u32 (zero-fills).
See ``shl_u32`` docstring for why inline PTX is used instead of plain
CuTeDSL shift operators (LLVM shift-by-type-width UB).
"""
return cutlass.Uint32(
llvm.inline_asm(
T.i32(),
[
cutlass.Uint32(val).ir_value(loc=loc, ip=ip),
cutlass.Uint32(shift).ir_value(loc=loc, ip=ip),
],
"shr.u32 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@cute.jit
def warp_prefix_sum(val: cutlass.Int32, lane: Optional[cutlass.Int32] = None) -> cutlass.Int32:
if const_expr(lane is None):
lane = cute.arch.lane_idx()
# if cute.arch.thread_idx()[0] >= 128 and cute.arch.thread_idx()[0] < 128 + 32 and cute.arch.block_idx()[0] == 0: cute.printf("tidx = %d, val = %d", cute.arch.thread_idx()[0] % 32, val)
for i in cutlass.range_constexpr(int(math.log2(cute.arch.WARP_SIZE))):
offset = 1 << i
# Very important that we set mask_and_clamp to 0
partial_sum = cute.arch.shuffle_sync_up(val, offset=offset, mask_and_clamp=0)
if lane >= offset:
val += partial_sum
# if cute.arch.thread_idx()[0] >= 128 and cute.arch.thread_idx()[0] < 128 + 32 and cute.arch.block_idx()[0] == 0: cute.printf("tidx = %d, partial_sum = %d, val = %d", cute.arch.thread_idx()[0] % 32, partial_sum, val)
return val
@dsl_user_op
def cvt_f16x2_f32(
a: float | Float32, b: float | Float32, to_dtype: Type, *, loc=None, ip=None
) -> cutlass.Int32:
assert to_dtype in [cutlass.BFloat16, cutlass.Float16], "to_dtype must be BFloat16 or Float16"
return cutlass.Int32(
llvm.inline_asm(
T.i32(),
[Float32(a).ir_value(loc=loc, ip=ip), Float32(b).ir_value(loc=loc, ip=ip)],
f"cvt.rn.{'bf16x2' if to_dtype is cutlass.BFloat16 else 'f16x2'}.f32 $0, $2, $1;",
"=r,f,f",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@overload
def cvt_f16(src: cute.Tensor, dst: cute.Tensor) -> None: ...
@overload
def cvt_f16(src: cute.Tensor, dtype: Type[cute.Numeric]) -> cute.Tensor: ...
@cute.jit
def cvt_f16(src: cute.Tensor, dst_or_dtype):
"""Convert Float32 tensor to Float16/BFloat16.
Args:
src: Source tensor with Float32 element type
dst_or_dtype: Either a destination tensor or a dtype (Float16/BFloat16)
Returns:
None if dst is a tensor, or a new tensor if dtype is provided
"""
if const_expr(isinstance(dst_or_dtype, type)):
# dtype variant: create new tensor and call the tensor variant
dtype = dst_or_dtype
dst = cute.make_rmem_tensor(src.shape, dtype)
cvt_f16(src, dst)
return dst
else:
# tensor variant: write to dst
dst = dst_or_dtype
assert cute.size(dst.shape) == cute.size(src.shape), "dst and src must have the same size"
assert cute.size(src.shape) % 2 == 0, "src must have an even number of elements"
assert dst.element_type in [cutlass.BFloat16, cutlass.Float16], (
"dst must be BFloat16 or Float16"
)
assert src.element_type is Float32, "src must be Float32"
dst_i32 = cute.recast_tensor(dst, cutlass.Int32)
assert cute.size(dst_i32.shape) * 2 == cute.size(src.shape)
for i in cutlass.range_constexpr(cute.size(dst_i32)):
dst_i32[i] = cvt_f16x2_f32(src[2 * i], src[2 * i + 1], dst.element_type)
@dsl_user_op
@cute.jit
def evaluate_polynomial(x: Float32, poly: Tuple[Float32, ...], *, loc=None, ip=None) -> Float32:
deg = len(poly) - 1
out = poly[deg]
for i in cutlass.range_constexpr(deg - 1, -1, -1):
out = out * x + poly[i]
return out
@dsl_user_op
@cute.jit
def evaluate_polynomial_2(
x: Float32, y: Float32, poly: Tuple[Float32, ...], *, loc=None, ip=None
) -> Tuple[Float32, Float32]:
deg = len(poly) - 1
out = (poly[deg], poly[deg])
for i in cutlass.range_constexpr(deg - 1, -1, -1):
out = cute.arch.fma_packed_f32x2(out, (x, y), (poly[i], poly[i]))
return out
@dsl_user_op
def add_round_down(x: float | Float32, y: float | Float32, *, loc=None, ip=None) -> Float32:
# There's probably a way to call llvm or nvvm to do this instead of ptx
return cutlass.Float32(
llvm.inline_asm(
T.f32(),
[Float32(x).ir_value(loc=loc, ip=ip), Float32(y).ir_value(loc=loc, ip=ip)],
"add.rm.ftz.f32 $0, $1, $2;",
"=f,f,f",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def combine_int_frac_ex2(x_rounded: Float32, frac_ex2: Float32, *, loc=None, ip=None) -> Float32:
return cutlass.Float32(
llvm.inline_asm(
T.f32(),
[
Float32(x_rounded).ir_value(loc=loc, ip=ip),
Float32(frac_ex2).ir_value(loc=loc, ip=ip),
],
"{\n\t"
".reg .s32 x_rounded_i, frac_ex_i, x_rounded_e, out_i;\n\t"
"mov.b32 x_rounded_i, $1;\n\t"
"mov.b32 frac_ex_i, $2;\n\t"
"shl.b32 x_rounded_e, x_rounded_i, 23;\n\t"
# add.u32 generates IMAD instruction and add.s32 generates LEA instruction
# IMAD uses the FMA pipeline and LEA uses the ALU pipeline, afaik
"add.s32 out_i, x_rounded_e, frac_ex_i;\n\t"
"mov.b32 $0, out_i;\n\t"
"}\n",
"=f,f,f",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def ex2_emulation(x: Float32, *, poly_degree: int = 3, loc=None, ip=None) -> Float32:
assert poly_degree in POLY_EX2, f"Polynomial degree {poly_degree} not supported"
# We assume x <= 127.0
fp32_round_int = float(2**23 + 2**22)
x_clamped = cute.arch.fmax(x, -127.0)
# We want to round down here, so that the fractional part is in [0, 1)
x_rounded = add_round_down(x_clamped, fp32_round_int, loc=loc, ip=ip)
# The integer floor of x is now in the last 8 bits of x_rounded
# We assume the next 2 ops round to nearest even. The rounding mode is important.
x_rounded_back = x_rounded - fp32_round_int
x_frac = x_clamped - x_rounded_back
x_frac_ex2 = evaluate_polynomial(x_frac, POLY_EX2[poly_degree], loc=loc, ip=ip)
return combine_int_frac_ex2(x_rounded, x_frac_ex2, loc=loc, ip=ip)
# TODO: check that the ex2_emulation_2 produces the same SASS as the ptx version
@dsl_user_op
def ex2_emulation_2(
x: Float32, y: Float32, *, poly_degree: int = 3, loc=None, ip=None
) -> Tuple[Float32, Float32]:
# We assume x <= 127.0 and y <= 127.0
fp32_round_int = float(2**23 + 2**22)
xy_clamped = (cute.arch.fmax(x, -127.0), cute.arch.fmax(y, -127.0))
# We want to round down here, so that the fractional part is in [0, 1)
xy_rounded = cute.arch.add_packed_f32x2(xy_clamped, (fp32_round_int, fp32_round_int), rnd="rm")
# The integer floor of x & y are now in the last 8 bits of xy_rounded
# We want the next 2 ops to round to nearest even. The rounding mode is important.
xy_rounded_back = quack.activation.sub_packed_f32x2(
xy_rounded, (fp32_round_int, fp32_round_int)
)
xy_frac = quack.activation.sub_packed_f32x2(xy_clamped, xy_rounded_back)
xy_frac_ex2 = evaluate_polynomial_2(*xy_frac, POLY_EX2[poly_degree], loc=loc, ip=ip)
x_out = combine_int_frac_ex2(xy_rounded[0], xy_frac_ex2[0], loc=loc, ip=ip)
y_out = combine_int_frac_ex2(xy_rounded[1], xy_frac_ex2[1], loc=loc, ip=ip)
return x_out, y_out
@dsl_user_op
def e2e_asm2(x: Float32, y: Float32, *, loc=None, ip=None) -> Tuple[Float32, Float32]:
out_f32x2 = llvm.inline_asm(
llvm.StructType.get_literal([T.f32(), T.f32()]),
[Float32(x).ir_value(loc=loc, ip=ip), Float32(y, loc=loc, ip=ip).ir_value()],
"{\n\t"
".reg .f32 f1, f2, f3, f4, f5, f6, f7;\n\t"
".reg .b64 l1, l2, l3, l4, l5, l6, l7, l8, l9, l10;\n\t"
".reg .s32 r1, r2, r3, r4, r5, r6, r7, r8;\n\t"
"max.ftz.f32 f1, $2, 0fC2FE0000;\n\t"
"max.ftz.f32 f2, $3, 0fC2FE0000;\n\t"
"mov.b64 l1, {f1, f2};\n\t"
"mov.f32 f3, 0f4B400000;\n\t"
"mov.b64 l2, {f3, f3};\n\t"
"add.rm.ftz.f32x2 l7, l1, l2;\n\t"
"sub.rn.ftz.f32x2 l8, l7, l2;\n\t"
"sub.rn.ftz.f32x2 l9, l1, l8;\n\t"
"mov.f32 f7, 0f3D9DF09D;\n\t"
"mov.b64 l6, {f7, f7};\n\t"
"mov.f32 f6, 0f3E6906A4;\n\t"
"mov.b64 l5, {f6, f6};\n\t"
"mov.f32 f5, 0f3F31F519;\n\t"
"mov.b64 l4, {f5, f5};\n\t"
"mov.f32 f4, 0f3F800000;\n\t"
"mov.b64 l3, {f4, f4};\n\t"
"fma.rn.ftz.f32x2 l10, l9, l6, l5;\n\t"
"fma.rn.ftz.f32x2 l10, l10, l9, l4;\n\t"
"fma.rn.ftz.f32x2 l10, l10, l9, l3;\n\t"
"mov.b64 {r1, r2}, l7;\n\t"
"mov.b64 {r3, r4}, l10;\n\t"
"shl.b32 r5, r1, 23;\n\t"
"add.s32 r7, r5, r3;\n\t"
"shl.b32 r6, r2, 23;\n\t"
"add.s32 r8, r6, r4;\n\t"
"mov.b32 $0, r7;\n\t"
"mov.b32 $1, r8;\n\t"
"}\n",
"=r,=r,f,f",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
out0 = Float32(llvm.extractvalue(T.f32(), out_f32x2, [0], loc=loc, ip=ip))
out1 = Float32(llvm.extractvalue(T.f32(), out_f32x2, [1], loc=loc, ip=ip))
return out0, out1
@dsl_user_op
def domain_offset_aligned(
coord: cute.Coord, tensor: cute.Tensor, *, loc=None, ip=None
) -> cute.Tensor:
assert isinstance(tensor.iterator, cute.Pointer)
# We assume that applying the offset does not change the pointer alignment
new_ptr = cute.make_ptr(
tensor.element_type,
elem_pointer(tensor, coord).toint(),
tensor.memspace,
assumed_align=tensor.iterator.alignment,
)
return cute.make_tensor(new_ptr, tensor.layout)
@dsl_user_op
def warp_reduction(
val: cute.Numeric, op: Callable, *, threads_in_group: int = 32, loc=None, ip=None
) -> cute.Numeric:
"""Warp-wide reduction helper for a custom binary op."""
offset = threads_in_group // 2
while offset > 0:
val = op(
val,
cute.arch.shuffle_sync_bfly(
val, offset=offset, mask=-1, mask_and_clamp=31, loc=loc, ip=ip
),
)
offset //= 2
return val
warp_reduction_max = partial(
warp_reduction, op=lambda x, y: fmax(x, y) if isinstance(x, Float32) else max(x, y)
)
warp_reduction_sum = partial(warp_reduction, op=lambda x, y: x + y) # noqa: FURB118
@dsl_user_op
def make_cotiled_copy(
atom: cute.CopyAtom, atom_layout_tv: cute.Layout, data_layout: cute.Layout, *, loc=None, ip=None
) -> cute.TiledCopy:
"""Compatibility wrapper for deprecated CuTeDSL `make_cotiled_copy`."""
assert cute.is_static(atom_layout_tv.type), "atom_layout_tv must be static"
assert cute.is_static(data_layout.type), "data_layout must be static"
inv_layout_ = cute.left_inverse(data_layout, loc=loc, ip=ip)
inv_data_layout = cute.make_layout(
(inv_layout_.shape, (1)), stride=(inv_layout_.stride, (0)), loc=loc, ip=ip
)
layout_tv_data = cute.composition(inv_data_layout, atom_layout_tv, loc=loc, ip=ip)
atom_layout_v_to_check = cute.coalesce(
cute.make_layout(atom_layout_tv.shape[1], stride=atom_layout_tv.stride[1], loc=loc, ip=ip),
loc=loc,
ip=ip,
)
data_layout_v_to_check = cute.coalesce(
cute.composition(
data_layout,
cute.make_layout(
layout_tv_data.shape[1], stride=layout_tv_data.stride[1], loc=loc, ip=ip
),
loc=loc,
ip=ip,
),
loc=loc,
ip=ip,
)
assert data_layout_v_to_check == atom_layout_v_to_check, (
"the memory pointed to by atom_layout_tv does not exist in the data_layout."
)
flat_data_shape = cute.product_each(data_layout.shape, loc=loc, ip=ip)
tiler = tuple(
cute.filter(
cute.composition(
cute.make_layout(
flat_data_shape,
stride=tuple(0 if j != i else 1 for j in range(cute.rank(flat_data_shape))),
loc=loc,
ip=ip,
),
layout_tv_data,
loc=loc,
ip=ip,
),
loc=loc,
ip=ip,
)
for i in range(cute.rank(flat_data_shape))
)
tile2data = cute.composition(
cute.make_layout(flat_data_shape, loc=loc, ip=ip), tiler, loc=loc, ip=ip
)
layout_tv = cute.composition(
cute.left_inverse(tile2data, loc=loc, ip=ip), layout_tv_data, loc=loc, ip=ip
)
return cute.make_tiled_copy(atom, layout_tv, tiler, loc=loc, ip=ip)
@cute.jit
def scalar_to_ssa(a: cute.Numeric, dtype) -> cute.TensorSSA:
"""Convert a scalar to a cute TensorSSA of shape (1,) and given dtype"""
vec = cute.make_rmem_tensor(1, dtype)
vec[0] = a
return vec.load()
def ssa_to_scalar(val):
"""Could inline but nice for reflecting the above api"""
return val[0]
@cute.jit
def get_batch_from_cu_tensor(idx: Int32, cu_tensor: cute.Tensor) -> Int32:
"""Binary search to determine batch from packed index in a cumulative tensor"""
batch_size = cute.size(cu_tensor) - 1
lo = Int32(0)
hi = batch_size
while lo < hi:
mid = (lo + hi) // 2
if cu_tensor[mid + 1] <= idx:
lo = mid + 1
else:
hi = mid
return lo
@@ -6,9 +6,13 @@ edge ("route A"), and requires no optional dependencies.
The FA4 CuTe block-sparse fastpath (intended for Blackwell sm_100+) is
*opt-in* via ``FASTVIDEO_VSA_CUTEDSL=1``. It routes to
:mod:`fastvideo_kernel.block_sparse_attn_cute_fwd`, which natively operates
on 128-token KV blocks (this wrapper expands the logical 256-block map /
sizes into that physical 128-block representation). The CuTe kernel
:mod:`fastvideo_kernel.block_sparse_attn_cute_fwd`. By default this wrapper
expands logical KV256 blocks to the historical physical KV128 tiles. Set
``FASTVIDEO_VSA_FA4_BLOCK_SHAPE=128x64`` for the fine-grained FA4 schedule.
The production ``64x64`` selection pairs adjacent Q64 children into physical
Q128 tiles: their KV lists are identical because this wrapper derives them
from one original Q256 metadata row. Direct Q64 callers with arbitrary masks
continue to use the native Q64/KV64 path. The CuTe kernel
(``flash_attn.cute`` with block-sparsity) is an optional dependency,
imported lazily only when this fastpath is selected.
@@ -30,7 +34,18 @@ from .block_sparse_attn import block_sparse_attn_triton, _force_triton
# FA4 CuTe build (``flash_attn.cute``) and make it a hard dependency of the
# default Triton path.
_KV_BLOCK_PHYS = 128 # FA4 CuTe BSA forward uses 128-token KV blocks.
_LOGICAL_BLOCK_SIZE = 256
_FA4_BLOCK_SHAPES = {
# The public name describes the original VSA mask. FA4 internally splits
# its logical KV256 edge into two physical KV128 tiles.
"256x256": (256, 128),
"128x64": (128, 64),
# This pairing is safe only at this original-Q256 wrapper boundary: all
# four Q64 children inherit exactly the same KV list from their parent.
# Keep arbitrary/distinct Q64 maps on block_sparse_attn_cute_fwd's native
# Q64/KV64 path, where they are not coalesced.
"64x64": (128, 64),
}
_KV_BLOCK_TRITON = 64 # Existing Triton path uses 64-token KV blocks.
@@ -49,28 +64,49 @@ def _resolve_backend() -> str:
return "triton"
def _expand_mask_and_sizes_256_to_128(
def _resolve_fa4_block_shape() -> Tuple[int, int]:
requested = os.environ.get("FASTVIDEO_VSA_FA4_BLOCK_SHAPE", "256x256").lower()
try:
return _FA4_BLOCK_SHAPES[requested]
except KeyError as exc:
choices = ", ".join(_FA4_BLOCK_SHAPES)
raise ValueError(
f"unsupported FASTVIDEO_VSA_FA4_BLOCK_SHAPE={requested!r}; choose one of {choices}"
) from exc
def _expand_mask_and_sizes_256_for_fa4(
logical_mask_256: torch.Tensor,
logical_kv_sizes_256: torch.Tensor,
q_block_size: int,
kv_block_size: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Expand a [B, H, Qb256, KVb256] map to [B, H, Qb256, KVb128].
"""Expand a logical Q256/KV256 map to one supported physical FA4 shape.
Each logical 256-token KV block splits into two physical 128-token
children. Each child inherits the logical mask edge; its valid-token
count is the logical count clamped into the child's window.
Every child inherits its parent edge. KV valid-token counts are clamped
into each child's window, preserving partial and empty tail blocks.
"""
expanded_mask = logical_mask_256.repeat_interleave(2, dim=3)
if _LOGICAL_BLOCK_SIZE % q_block_size or _LOGICAL_BLOCK_SIZE % kv_block_size:
raise ValueError(
f"FA4 QxKV blocks must divide {_LOGICAL_BLOCK_SIZE}, got "
f"{q_block_size}x{kv_block_size}"
)
q_factor = _LOGICAL_BLOCK_SIZE // q_block_size
kv_factor = _LOGICAL_BLOCK_SIZE // kv_block_size
expanded_mask = logical_mask_256.repeat_interleave(q_factor, dim=2)
expanded_mask = expanded_mask.repeat_interleave(kv_factor, dim=3)
sizes_i32 = logical_kv_sizes_256.to(torch.int32)
child0 = torch.clamp(sizes_i32, min=0, max=_KV_BLOCK_PHYS)
child1 = torch.clamp(sizes_i32 - _KV_BLOCK_PHYS, min=0, max=_KV_BLOCK_PHYS)
expanded_sizes = torch.empty(
(sizes_i32.numel() * 2, ),
offsets = torch.arange(
kv_factor,
dtype=torch.int32,
device=sizes_i32.device,
)
expanded_sizes[0::2] = child0
expanded_sizes[1::2] = child1
) * kv_block_size
expanded_sizes = torch.clamp(
sizes_i32[:, None] - offsets[None, :],
min=0,
max=kv_block_size,
).reshape(-1)
return expanded_mask, expanded_sizes
@@ -126,9 +162,15 @@ def block_sparse_attn_256(
if _resolve_backend() == "triton":
return _triton_via_route_a(q, k, v, logical_block_map_256, logical_variable_block_sizes_256)
mask_128, sizes_128 = _expand_mask_and_sizes_256_to_128(logical_block_map_256, logical_variable_block_sizes_256)
q_block_size, kv_block_size = _resolve_fa4_block_shape()
physical_mask, physical_sizes = _expand_mask_and_sizes_256_for_fa4(
logical_block_map_256,
logical_variable_block_sizes_256,
q_block_size,
kv_block_size,
)
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd
return block_sparse_attn_cute_fwd(q, k, v, mask_128, sizes_128)
return block_sparse_attn_cute_fwd(q, k, v, physical_mask, physical_sizes)
def block_sparse_attn_256_bshd(
@@ -156,6 +198,12 @@ def block_sparse_attn_256_bshd(
)
return out_bhsd.transpose(1, 2).contiguous(), aux
mask_128, sizes_128 = _expand_mask_and_sizes_256_to_128(logical_block_map_256, logical_variable_block_sizes_256)
q_block_size, kv_block_size = _resolve_fa4_block_shape()
physical_mask, physical_sizes = _expand_mask_and_sizes_256_for_fa4(
logical_block_map_256,
logical_variable_block_sizes_256,
q_block_size,
kv_block_size,
)
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
return block_sparse_attn_cute_fwd_bshd(q, k, v, mask_128, sizes_128)
return block_sparse_attn_cute_fwd_bshd(q, k, v, physical_mask, physical_sizes)
@@ -9,10 +9,11 @@ The BSHD variant is preferred from VSA-256 callers to avoid layout
round-trips on the hot path.
The FA4 CuTe block-sparse kernel (``flash_attn.cute`` with
``block_sparsity``) is an *optional* dependency: it is imported lazily and
only exercised when the VSA-256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default VSA-256 path is Triton and does
not require it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
``block_sparsity``) is an *optional* dependency: it is imported lazily. The
legacy VSA-256 route selects it with ``FASTVIDEO_VSA_CUTEDSL=1``; direct
fine-grained callers can use Q128/KV64 and, on SM100, Q64/KV64. The default
VSA-256 path is Triton and does not require FA4. The CuTe path also needs
``nvidia-cutlass-dsl`` and ``quack-kernels``.
"""
from __future__ import annotations
@@ -22,11 +23,11 @@ from typing import Tuple
import torch
_FA4_IMPORT_HINT = ("VSA-256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
_FA4_IMPORT_HINT = ("The CuTe block-sparse fastpath requires a FlashAttention-4 CuTe build that "
"provides `flash_attn.cute` with block-sparsity support (plus "
"`nvidia-cutlass-dsl` and `quack-kernels`). This is an optional "
"dependency; the default VSA-256 path is Triton. Install the FA4 CuTe "
"build and set FASTVIDEO_VSA_CUTEDSL=1 to enable the CuTe fastpath.")
"build; set FASTVIDEO_VSA_CUTEDSL=1 when selecting the VSA-256 route.")
def _load_fa4_cute():
@@ -44,7 +45,7 @@ def _load_fa4_cute():
return BlockSparseTensorsTorch, _flash_attn_fwd
# Q-side tile size; kv_block_size comes from the caller's VSA logical KV block.
# Default Q-side tile size; the Q64/KV64 specialization uses a 64-row tile.
_M_BLOCK_SIZE_DEFAULT = 128
@@ -72,6 +73,18 @@ def _choose_q_sparse_block_size(q_len: int, m_block_size: int = _M_BLOCK_SIZE_DE
return m_block_size
def _choose_m_block_size(q_block_size: int, kv_block_size: int) -> int:
"""Choose the physical FA4 Q tile without changing sparse-map semantics."""
if q_block_size == 64:
if kv_block_size != 64:
raise ValueError("the FA4 Q64 specialization requires KV64 blocks")
major, _minor = torch.cuda.get_device_capability()
if major != 10:
raise RuntimeError("the FA4 Q64/KV64 specialization currently requires SM100")
return 64
return _M_BLOCK_SIZE_DEFAULT
def _aggregate_q_block_map(
block_map: torch.Tensor,
q_sparse_block_size: int,
@@ -146,11 +159,18 @@ def _cute_forward(
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Internal: FA4 CuTe BSA fwd with BSHD inputs."""
BlockSparseTensorsTorch, _flash_attn_fwd = _load_fa4_cute()
q_sparse_candidate = _choose_q_sparse_block_size(q_bshd.shape[1])
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
m_block_size = _choose_m_block_size(q_block_size, kv_block_size)
if q_block_size == m_block_size:
# Q128 and Q64 maps may select different KV blocks in adjacent rows.
# Preserve their declared granularity instead of unioning them into the
# historical two-stage Q256 schedule.
q_sparse_block_size = q_block_size
else:
q_sparse_candidate = _choose_q_sparse_block_size(q_bshd.shape[1], m_block_size)
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
sparse_map = _aggregate_q_block_map(
block_map,
q_sparse_block_size=q_sparse_block_size,
@@ -177,7 +197,7 @@ def _cute_forward(
q_bshd,
k_bshd,
v_bshd,
tile_mn=(_M_BLOCK_SIZE_DEFAULT, kv_block_size),
tile_mn=(m_block_size, kv_block_size),
mask_mod=_build_vbs_mask_mod(kv_block_size),
block_sparse_tensors=sparse_tensors,
aux_tensors=[variable_block_sizes],
@@ -0,0 +1,526 @@
"""Correctness coverage for FastVideo's fine-grained FA4 sparse tiles."""
from __future__ import annotations
import math
import pytest
import torch
@pytest.fixture(autouse=True)
def _require_fa4_sm100(monkeypatch):
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
major, _minor = torch.cuda.get_device_capability()
if major not in (10, 11):
pytest.skip("fine-grained FA4 VSA requires SM100/SM110")
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FastVideo FA4 source package is not installed",
)
monkeypatch.setenv("FASTVIDEO_FA4_VSA_DUAL_STREAM", "1")
monkeypatch.setenv("FASTVIDEO_FA4_VSA_SP_DOUBLE_BUFFER", "1")
def _selected_block_map(
heads: int,
q_blocks: int,
kv_blocks: int,
selected_blocks: int,
device: torch.device,
) -> torch.Tensor:
block_map = torch.zeros(
1,
heads,
q_blocks,
kv_blocks,
dtype=torch.bool,
device=device,
)
for head in range(heads):
for q_block in range(q_blocks):
start = (3 * head + 5 * q_block) % kv_blocks
indices = [(start + 2 * offset) % kv_blocks for offset in range(selected_blocks)]
block_map[0, head, q_block, indices] = True
return block_map
def _ordered_sparse_tensors(
block_map: torch.Tensor,
q_block_size: int,
kv_block_size: int,
masked_blocks: int,
):
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
batch, heads, q_blocks, kv_blocks = block_map.shape
selected_blocks = int(block_map.sum(dim=-1).min().item())
assert bool((block_map.sum(dim=-1) == selected_blocks).all())
if not 0 <= masked_blocks <= selected_blocks:
raise ValueError("masked_blocks must be within the selected-block count")
selected = torch.arange(kv_blocks, dtype=torch.int32, device=block_map.device)
selected = selected.view(1, 1, 1, kv_blocks).expand(batch, heads, q_blocks, kv_blocks)
selected = selected.masked_select(block_map).view(batch, heads, q_blocks, selected_blocks)
selected = selected.sort(dim=-1).values
mask_count = masked_blocks
full_count = selected_blocks - masked_blocks
mask_idx = torch.zeros(
batch,
heads,
q_blocks,
max(mask_count, 1),
dtype=torch.int32,
device=block_map.device,
)
full_idx = torch.zeros(
batch,
heads,
q_blocks,
max(full_count, 1),
dtype=torch.int32,
device=block_map.device,
)
if mask_count:
mask_idx[..., :mask_count] = selected[..., :mask_count]
if full_count:
full_idx[..., :full_count] = selected[..., mask_count:]
return BlockSparseTensorsTorch(
mask_block_cnt=torch.full(
(batch, heads, q_blocks),
mask_count,
dtype=torch.int32,
device=block_map.device,
),
mask_block_idx=mask_idx,
full_block_cnt=torch.full(
(batch, heads, q_blocks),
full_count,
dtype=torch.int32,
device=block_map.device,
),
full_block_idx=full_idx,
block_size=(q_block_size, kv_block_size),
)
def _torch_reference(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
q_block_size: int,
kv_block_size: int,
variable_block_sizes: torch.Tensor | None = None,
) -> torch.Tensor:
token_mask = block_map.repeat_interleave(q_block_size, dim=2)
token_mask = token_mask.repeat_interleave(kv_block_size, dim=3)
if variable_block_sizes is not None:
kv_positions = torch.arange(k.shape[1], device=k.device)
kv_blocks = kv_positions // kv_block_size
kv_offsets = kv_positions % kv_block_size
valid_tokens = kv_offsets < variable_block_sizes[kv_blocks]
token_mask = token_mask & valid_tokens.view(1, 1, 1, -1)
q_heads = q.permute(0, 2, 1, 3).float()
k_heads = k.permute(0, 2, 1, 3).float()
v_heads = v.permute(0, 2, 1, 3).float()
scores = torch.matmul(q_heads, k_heads.transpose(-2, -1)) / math.sqrt(q.shape[-1])
scores.masked_fill_(~token_mask, float("-inf"))
out = torch.matmul(torch.softmax(scores, dim=-1), v_heads)
return out.permute(0, 2, 1, 3).contiguous()
@pytest.mark.cuda
@pytest.mark.parametrize(
(
"q_block_size",
"kv_block_size",
"selected_blocks",
"masked_blocks",
"sp_double_buffer",
),
[
pytest.param(128, 64, 1, 0, True, id="q128_kv64_one_stream_empty"),
pytest.param(128, 64, 2, 1, True, id="q128_kv64_even_mixed"),
pytest.param(128, 64, 3, 1, True, id="q128_kv64_odd_mixed"),
pytest.param(128, 64, 3, 0, True, id="q128_kv64_odd_full_only"),
pytest.param(128, 64, 3, 1, False, id="q128_kv64_legacy_optout"),
pytest.param(64, 64, 3, 1, True, id="q64_kv64_odd_mixed"),
],
)
def test_fa4_vsa_fine_grained_forward(
q_block_size: int,
kv_block_size: int,
selected_blocks: int,
masked_blocks: int,
sp_double_buffer: bool,
monkeypatch: pytest.MonkeyPatch,
) -> None:
from flash_attn.cute.interface import _flash_attn_fwd
monkeypatch.setenv(
"FASTVIDEO_FA4_VSA_SP_DOUBLE_BUFFER",
"1" if sp_double_buffer else "0",
)
torch.manual_seed(2026 + q_block_size + selected_blocks)
device = torch.device("cuda")
batch, heads, head_dim = 1, 8, 128
q_blocks, kv_blocks = 4, 7
q_len = q_blocks * q_block_size
kv_len = kv_blocks * kv_block_size
q = torch.randn(batch, q_len, heads, head_dim, dtype=torch.bfloat16, device=device)
k = torch.randn(batch, kv_len, heads, head_dim, dtype=torch.bfloat16, device=device)
v = torch.randn_like(k)
block_map = _selected_block_map(heads, q_blocks, kv_blocks, selected_blocks, device)
sparse_tensors = _ordered_sparse_tensors(
block_map,
q_block_size,
kv_block_size,
masked_blocks,
)
out = _flash_attn_fwd(
q,
k,
v,
tile_mn=(q_block_size, kv_block_size),
block_sparse_tensors=sparse_tensors,
causal=False,
return_lse=True,
)[0]
ref = _torch_reference(q, k, v, block_map, q_block_size, kv_block_size)
assert bool(torch.isfinite(out).all())
error = (out.float() - ref).abs()
avg_abs = float(error.mean())
max_rel = float(error.max() / (ref.abs().mean() + 1e-6))
print(
f"[fa4 {q_block_size}x{kv_block_size} selected={selected_blocks} "
f"masked={masked_blocks} spdb={sp_double_buffer}] "
f"avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}"
)
assert avg_abs < 1e-3
assert max_rel < 0.2
@pytest.mark.cuda
@pytest.mark.parametrize("q_block_size", [128, 64], ids=["q128_kv64", "q64_kv64"])
def test_fa4_vsa_dual_stream_persistent_phase_transitions(q_block_size: int) -> None:
"""Exercise even, odd, empty, and one-block tiles in one persistent launch."""
if q_block_size == 64 and torch.cuda.get_device_capability()[0] != 10:
pytest.skip("the Q64/KV64 specialization currently requires SM100")
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
from flash_attn.cute.interface import _flash_attn_fwd
torch.manual_seed(4554 + q_block_size)
device = torch.device("cuda")
kv_block_size = 64
counts = [2, 3, 0, 1, 4]
masked_counts = [1, 1, 0, 1, 2]
batch, heads, head_dim, kv_blocks = 1, 4, 128, 8
q_blocks = len(counts)
q = torch.randn(
batch,
q_blocks * q_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
k = torch.randn(
batch,
kv_blocks * kv_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
v = torch.randn_like(k)
mask_count = torch.empty(batch, heads, q_blocks, dtype=torch.int32, device=device)
full_count = torch.empty_like(mask_count)
mask_idx = torch.zeros(
batch,
heads,
q_blocks,
max(masked_counts),
dtype=torch.int32,
device=device,
)
full_idx = torch.zeros(
batch,
heads,
q_blocks,
max(count - masked for count, masked in zip(counts, masked_counts)),
dtype=torch.int32,
device=device,
)
block_map = torch.zeros(
batch,
heads,
q_blocks,
kv_blocks,
dtype=torch.bool,
device=device,
)
for head in range(heads):
for q_block, (count, masked) in enumerate(zip(counts, masked_counts)):
chosen = sorted({(head + 3 * q_block + 2 * offset) % kv_blocks for offset in range(count)})
assert len(chosen) == count
mask_count[0, head, q_block] = masked
full_count[0, head, q_block] = count - masked
if masked:
mask_idx[0, head, q_block, :masked] = torch.tensor(
chosen[:masked],
dtype=torch.int32,
device=device,
)
if count > masked:
full_idx[0, head, q_block, :count - masked] = torch.tensor(
chosen[masked:],
dtype=torch.int32,
device=device,
)
if count:
block_map[0, head, q_block, chosen] = True
sparse_tensors = BlockSparseTensorsTorch(
mask_block_cnt=mask_count,
mask_block_idx=mask_idx,
full_block_cnt=full_count,
full_block_idx=full_idx,
block_size=(q_block_size, kv_block_size),
)
out, lse = _flash_attn_fwd(
q,
k,
v,
tile_mn=(q_block_size, kv_block_size),
block_sparse_tensors=sparse_tensors,
causal=False,
return_lse=True,
)[:2]
ref = _torch_reference(q, k, v, block_map, q_block_size, kv_block_size)
nonempty_rows = torch.tensor(
[count > 0 for count in counts],
dtype=torch.bool,
device=device,
).repeat_interleave(q_block_size)
error = (out[:, nonempty_rows].float() - ref[:, nonempty_rows]).abs()
empty_out = out[:, ~nonempty_rows].float()
empty_lse = lse[:, :, ~nonempty_rows]
avg_abs = float(error.mean())
max_abs = float(error.max())
print(
f"[fa4 persistent {q_block_size}x{kv_block_size}] "
f"avg_abs={avg_abs:.6e}, max_abs={max_abs:.6e}"
)
assert avg_abs < 1e-3
assert max_abs < 1e-2
assert float(empty_out.abs().max()) == 0.0
assert bool(torch.isneginf(empty_lse).all())
@pytest.mark.cuda
def test_fa4_vsa_q128_kv64_bshd_wrapper_variable_blocks() -> None:
"""Cover partial and empty KV blocks through FastVideo's VBS mask mod."""
from fastvideo_kernel.block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
torch.manual_seed(5128)
device = torch.device("cuda")
batch, heads, head_dim = 1, 8, 128
q_block_size, kv_block_size = 128, 64
q_blocks, kv_blocks, selected_blocks = 4, 7, 3
q = torch.randn(
batch,
q_blocks * q_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
k = torch.randn(
batch,
kv_blocks * kv_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
v = torch.randn_like(k)
block_map = _selected_block_map(heads, q_blocks, kv_blocks, selected_blocks, device)
variable_block_sizes = torch.tensor(
[64, 23, 64, 7, 0, 51, 64],
dtype=torch.int32,
device=device,
)
out, lse = block_sparse_attn_cute_fwd_bshd(
q,
k,
v,
block_map,
variable_block_sizes,
)
ref = _torch_reference(
q,
k,
v,
block_map,
q_block_size,
kv_block_size,
variable_block_sizes,
)
assert bool(torch.isfinite(out).all())
assert bool(torch.isfinite(lse).all())
error = (out.float() - ref).abs()
avg_abs = float(error.mean())
max_rel = float(error.max() / (ref.abs().mean() + 1e-6))
print(f"[fa4 wrapper Q128/KV64 VBS] avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < 1e-3
assert max_rel < 0.2
@pytest.mark.cuda
@pytest.mark.parametrize(
("q_block_size", "kv_block_size"),
[
pytest.param(128, 64, id="q128_kv64"),
pytest.param(64, 64, id="q64_kv64"),
],
)
def test_fa4_vsa_fine_grained_bshd_wrapper(
q_block_size: int,
kv_block_size: int,
) -> None:
"""Exercise the FastVideo BSHD entrypoint used by the H3 backend."""
if q_block_size == 64 and torch.cuda.get_device_capability()[0] != 10:
pytest.skip("the Q64/KV64 specialization currently requires SM100")
from fastvideo_kernel.block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
torch.manual_seed(4048 + q_block_size)
device = torch.device("cuda")
batch, heads, head_dim = 1, 8, 128
q_blocks, kv_blocks, selected_blocks = 4, 7, 3
q = torch.randn(
batch,
q_blocks * q_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
k = torch.randn(
batch,
kv_blocks * kv_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
v = torch.randn_like(k)
block_map = _selected_block_map(heads, q_blocks, kv_blocks, selected_blocks, device)
variable_block_sizes = torch.full(
(kv_blocks,),
kv_block_size,
dtype=torch.int32,
device=device,
)
out, lse = block_sparse_attn_cute_fwd_bshd(
q,
k,
v,
block_map,
variable_block_sizes,
)
ref = _torch_reference(q, k, v, block_map, q_block_size, kv_block_size)
assert bool(torch.isfinite(out).all())
assert bool(torch.isfinite(lse).all())
error = (out.float() - ref).abs()
avg_abs = float(error.mean())
max_rel = float(error.max() / (ref.abs().mean() + 1e-6))
print(
f"[fa4 wrapper {q_block_size}x{kv_block_size}] "
f"avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}"
)
assert avg_abs < 1e-3
assert max_rel < 0.2
@pytest.mark.cuda
@pytest.mark.parametrize("block_shape", ["256x256", "128x64", "64x64"])
def test_vsa256_bshd_fa4_block_shape_routes(
block_shape: str,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Route one logical Q256/KV256 mask through every FA4 specialization."""
if block_shape == "64x64" and torch.cuda.get_device_capability()[0] != 10:
pytest.skip("the Q64/KV64 specialization currently requires SM100")
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.setenv("FASTVIDEO_VSA_FA4_BLOCK_SHAPE", block_shape)
torch.manual_seed(6256)
device = torch.device("cuda")
batch, heads, head_dim = 1, 8, 128
logical_block_size = 256
q_blocks, kv_blocks, selected_blocks = 4, 7, 3
q = torch.randn(
batch,
q_blocks * logical_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
k = torch.randn(
batch,
kv_blocks * logical_block_size,
heads,
head_dim,
dtype=torch.bfloat16,
device=device,
)
v = torch.randn_like(k)
block_map = _selected_block_map(heads, q_blocks, kv_blocks, selected_blocks, device)
variable_block_sizes = torch.tensor(
[256, 137, 255, 64, 201, 1, 192],
dtype=torch.int32,
device=device,
)
out, lse = block_sparse_attn_256_bshd(
q,
k,
v,
block_map,
variable_block_sizes,
)
ref = _torch_reference(
q,
k,
v,
block_map,
logical_block_size,
logical_block_size,
variable_block_sizes,
)
assert bool(torch.isfinite(out).all())
assert bool(torch.isfinite(lse).all())
error = (out.float() - ref).abs()
avg_abs = float(error.mean())
max_rel = float(error.max() / (ref.abs().mean() + 1e-6))
print(f"[VSA256 route {block_shape}] avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < 1e-3
assert max_rel < 0.2
+170
View File
@@ -0,0 +1,170 @@
# FA4 VSA block-64 handoff
## Status
This branch vendors FlashAttention-4's CuTe source as a first-class FastVideo
fork, adds the issue-#4554 benchmark, and implements forward-only SM100 sparse
attention paths for Q128/KV64 and Q64/KV64.
The requested single-Q-block/two-KV-workstream design is implemented. The two
workstreams split the canonical sparse traversal by ordinal parity, maintain
independent online-softmax state, and merge their FP32 output accumulators in
the correction epilogue. Q128/KV64 additionally uses four score/probability
slots to overlap successive blocks.
The measured result is not fully zero-loss:
- Native Q128/KV64 reaches about 90% of raw FA4 Q256/KV256 throughput.
- Native Q64/KV64 reaches about 55% of raw FA4 Q256/KV256 throughput.
- A newly added production-only VSA256 adapter reaches 80-86% of the public
Q256 wrapper. It is selected with
`FASTVIDEO_VSA_FA4_BLOCK_SHAPE=64x64`, but physically runs Q128/KV64 by
pairing adjacent Q64 children whose sparse lists are identical at the
original Q256 metadata boundary. It must not be presented as native
Q64/KV64 performance.
## Source and provenance
- Fork: `fastvideo-kernel/fa4/`
- Upstream base: Dao-AILab/flash-attention commit
`82d6441eec5d4dfec120153db2c0145ae855a083`
- Q64 primitives adapted from upstream commit
`526c18d25bcbc7fc7d6740ab3c7c84ed2d42cb0b`
- Full refresh instructions: `fastvideo-kernel/fa4/UPSTREAM.md`
- Root `pyproject.toml` resolves `flash-attn-4` to the editable local fork.
The fork retains the upstream `flash-attn-4` distribution name and
`flash_attn.cute` import namespace.
## Main implementation
- `fastvideo-kernel/fa4/interface.py`
- Strict compile-time gates for the dual-stream and double-buffer paths.
- `FASTVIDEO_FA4_VSA_DUAL_STREAM=0` retains the upstream-style single-stream
fallback.
- `FASTVIDEO_FA4_VSA_SP_DOUBLE_BUFFER=0` retains the earlier dual-stream
schedule for rollback and comparison.
- `fastvideo-kernel/fa4/block_sparse_utils.py`
- Canonical ordinal-parity sparse traversal shared by load, MMA, and
softmax code.
- Correct handling for masked-only, full-only, mixed, odd-length, and empty
streams.
- `fastvideo-kernel/fa4/flash_fwd_sm100.py`
- One Q block, two KV workstreams, four S/P slots for Q128/KV64, independent
phase tracking, and stable FP32 merge.
- Native M64 score/output layouts and row-pair softmax support.
- `fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_256.py`
- Public VSA256 shape selection via `FASTVIDEO_VSA_FA4_BLOCK_SHAPE`.
- The `64x64` option is the newly added coalescing adapter described above.
- `fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_cute_fwd.py`
- Direct Q128/KV64 and native Q64/KV64 BSHD entrypoints.
## Benchmark
Harness:
```bash
fastvideo-kernel/benchmarks/bench_vsa_blackwell.py
```
It derives from FlashInfer issue #4554 and immutable gist revision
`e15ac9066f23ef3690e33e1cc1fdac45b4b9099f`. The default `exact256` mask mode
uses identical selected token pairs across block shapes, checks every output
against FP32 masked SDPA, and reports sparse-aware algorithmic TFLOP/s plus MFU.
Four-GPU command:
```bash
PYTHONPATH="$PWD/fastvideo-kernel/fa4:$PWD/fastvideo-kernel/python" \
CUTE_DSL_ENABLE_TVM_FFI=1 \
FASTVIDEO_VSA_CUTEDSL=1 \
FASTVIDEO_FA4_VSA_DUAL_STREAM=1 \
FASTVIDEO_FA4_VSA_SP_DOUBLE_BUFFER=1 \
uv run --no-sync torchrun --standalone --nproc-per-node=4 \
fastvideo-kernel/benchmarks/bench_vsa_blackwell.py \
--seq_lens 32768 --sparsities dense 90 --mask_mode exact256 \
--arms cutedsl256 fa4_wrapper fa4 \
--block_shapes 256x256 128x64 64x64 \
--warmup 5 --rep 20 --out /tmp/fa4_final_tray.json
```
Environment used for the final run:
- 4x NVIDIA GB200, SM100
- PyTorch `2.12.0+cu130`
- CUDA 13.0, driver 580.159.04
- BF16, B=1, H=12, D=128, S=32768
- MFU denominator: 2.5 PFLOP/s dense BF16 per GPU, 10 PFLOP/s per tray
### Four-GPU aggregate results
| Raw schedule | Dense PFLOP/s | Dense MFU | 90% sparse PFLOP/s | 90% sparse MFU | Relative to raw Q256 |
|---|---:|---:|---:|---:|---:|
| Q256/KV256 | 6.353 | 63.53% | 5.852 | 58.52% | 100% |
| Q128/KV64 | 5.764 | 57.64% | 5.234 | 52.34% | 90.7% dense / 89.5% sparse |
| Native Q64/KV64 | 3.408 | 34.08% | 3.215 | 32.15% | 53.6% dense / 54.9% sparse |
| Public VSA256 route | Dense PFLOP/s | Dense MFU | 90% sparse PFLOP/s | 90% sparse MFU | Relative to public Q256 |
|---|---:|---:|---:|---:|---:|
| Default Q256 | 6.185 | 61.85% | 4.494 | 44.94% | 100% |
| Added `64x64` coalescing adapter | 5.317 | 53.17% | 3.584 | 35.84% | 86.0% dense / 79.7% sparse |
Fixed-base NCU showed that score/probability double buffering improves the
Q128/KV64 end-to-end wrapper by 8.74% dense and 4.35% at 90% sparsity over
the retained legacy dual-stream schedule.
## Validation
Final GB200 runs:
```bash
CUDA_VISIBLE_DEVICES=3 \
PYTHONPATH="$PWD/fastvideo-kernel/fa4:$PWD/fastvideo-kernel/python" \
CUTE_DSL_ENABLE_TVM_FFI=1 \
FASTVIDEO_VSA_CUTEDSL=1 \
FASTVIDEO_FA4_VSA_DUAL_STREAM=1 \
FASTVIDEO_FA4_VSA_SP_DOUBLE_BUFFER=1 \
uv run --no-sync pytest -vs \
fastvideo-kernel/tests/test_fa4_vsa_block_shapes.py
```
Result: 14 passed. Coverage includes selected counts 0-5, odd/even and
one-stream-empty cases, persistent phase transitions, masked/full/mixed
metadata, variable block sizes, native Q64, Q128, public VSA256 routes, and
the legacy double-buffer opt-out.
Existing compatibility tests:
```bash
uv run --no-sync pytest -vs \
fastvideo-kernel/tests/test_vsa256_forward.py \
fastvideo-kernel/tests/test_vsa256_forward_cross.py \
fastvideo-kernel/tests/test_vsa256_forward_vbs.py
```
Result: 4 passed.
Other completed gates:
- `pre-commit run --files ...`: passed for all scoped files.
- `python -m py_compile`: passed for the fork, wrappers, tests, and benchmark.
- FA4 sdist and wheel builds: passed.
- `twine check`: passed for both distributions.
- Isolated wheel import: resolved `flash_attn.cute` and the expected local
version.
- `git diff --check`: passed.
## Known limitations and follow-up
- These optimized specializations are forward-only.
- Native Q128/KV64 and Q64/KV64 do not meet the original zero-loss/30%-loss
targets; keep the native and coalesced-adapter numbers separate.
- The existing kernel CI lane uses H100. The SM100/SM110 tests skip there, so
the four-GB200 local run is currently the only hardware regression gate.
- The workspace's untracked `uv.lock` predates the local source selection and
was deliberately not overwritten. `uv lock --dry-run --offline` resolves
only the expected `flash-attn-4` source/version update.
- Rejected performance experiments include a second UMMA issuer, paired wait
batching, cooperative 2-CTA execution, cluster multicast, CTA-wide stats
barriers, and correction wait reordering. All either regressed or failed to
improve the fixed-base profiles and were not retained.
+5 -6
View File
@@ -122,9 +122,9 @@ fastvideo-kernel = [
# uv reads UV_TORCH_BACKEND / --torch-backend (not a pyproject setting). The
# supported backends (cu126, cu130) both provide torch 2.12.0.
imagebind = { git = "https://github.com/facebookresearch/ImageBind.git", rev = "53680b02d7e37b19b124fa37bae4b6c98c38f5be" }
# FA4 cute, pinned to a cutlass-4.5-compatible revision. torch.compile support
# comes from FastVideo's own custom_op wrappers.
flash-attn-4 = { git = "https://github.com/Dao-AILab/flash-attention.git", rev = "82d6441eec5d4dfec120153db2c0145ae855a083", subdirectory = "flash_attn/cute" }
# FastVideo's FA4 CuTe source fork. Its upstream commit and refresh procedure
# are recorded in fastvideo-kernel/fa4/UPSTREAM.md.
flash-attn-4 = { path = "fastvideo-kernel/fa4", editable = true }
[project.optional-dependencies]
@@ -189,9 +189,8 @@ streaming = [
dreamverse = [
"uvicorn[standard]>=0.41.0",
"cerebras-cloud-sdk",
# PyPI forbids direct URL deps in published metadata; pin upstream cute via
# [tool.uv.sources] above (same pattern as imagebind) so the wheel stays
# publishable while uv workspace installs still get it.
# Keep the published dependency name portable; [tool.uv.sources] above
# selects FastVideo's in-tree CuTe fork for workspace installs.
"flash-attn-4",
"flashinfer-python",
"openai>=1.40",