Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a24b8fc02a |
+56
-16
@@ -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())
|
||||
@@ -0,0 +1,4 @@
|
||||
[flake8]
|
||||
max-line-length = 100
|
||||
# W503: line break before binary operator
|
||||
ignore = E731, E741, F841, W503
|
||||
@@ -0,0 +1,8 @@
|
||||
Tri Dao
|
||||
Jay Shah
|
||||
Ted Zadouri
|
||||
Markus Hoehnerbach
|
||||
Vijay Thakkar
|
||||
Timmy Liu
|
||||
Driss Guessous
|
||||
Reuben Stern
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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/
|
||||
```
|
||||
@@ -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.
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
|
||||
@@ -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 = {}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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 ---")
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
]
|
||||
@@ -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
@@ -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)
|
||||
@@ -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]
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user