[perf] Add FA4 CuTe backward support for VSA-256 (#1639)

Co-authored-by: Hyunsung Lee <hyunsungl@sizigistudios.com>
Co-authored-by: alexzms <3036648523@qq.com>
This commit is contained in:
Hyunsung Lee
2026-08-19 11:46:49 -07:00
committed by GitHub
co-authored by Hyunsung Lee alexzms
parent 8537dcd6de
commit 00338aa9ca
13 changed files with 929 additions and 140 deletions
+1 -1
View File
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
+40 -11
View File
@@ -17,7 +17,7 @@ Runtime-JIT kernels (no build step, ship in every wheel/image):
| Kernels | Where | Used when |
|---|---|---|
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| FA4 CuTe-DSL block-sparse forward/backward (VSA-128/256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
## What gets built where, and when
@@ -64,28 +64,49 @@ cd fastvideo-kernel
./build.sh --rocm
```
### Optional: FA4 CuTe block-sparse backend (VSA-256 fastpath)
### Optional: FA4 CuTe block-sparse backend (VSA-128/256 fastpath)
The VSA-256 fastpath (tile volume 256, on NVIDIA Blackwell / sm_100) routes to the
The VSA-128/256 fastpaths (tile volume 128 or 256, on NVIDIA Blackwell / sm_100) route 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 symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`,
`flash_attn.cute.interface._flash_attn_fwd`) are provided upstream by
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`
and the public/private forward-backward bridges in `flash_attn.cute.interface`) are provided upstream by
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention). Pin to
commit `940cd9680f3315f2f06b43ab5bea2c2cf2d96806`, the revision FastVideo pins as
commit `14c377950125c70b7a9dabf9c561fca53715ac7d`, 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.
have incompatible block-sparse forward/backward interfaces.
Install it under its distribution name so its own runtime stack resolves with it.
Do **not** pre-install `nvidia-cutlass-dsl` by hand: this revision pins
`nvidia-cutlass-dsl==4.6.0.dev0` exactly, and a hand-installed 4.5.x floor either
gets silently upgraded or, if something else holds it back, leaves the CuTe
kernels broken.
```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"
pip install torchvision
pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@14c377950125c70b7a9dabf9c561fca53715ac7d#subdirectory=flash_attn/cute"
```
The CuTe kernel JIT-compiles on first use. Verified on Blackwell (sm_100) against
`tests/test_vsa256_forward*.py`.
That resolves `nvidia-cutlass-dsl` to 4.6.0.dev0 and `quack-kernels` to 0.5.3, a
combination this revision works with. A mismatched CuTe DSL only surfaces when the
kernel JIT-compiles, so the error points at CuTe internals rather than at the
install:
| Error on first VSA-128/256 CuTe call | Cause |
|---|---|
| `TypeError: fmax() missing 1 required positional argument: 'b'` | `nvidia-cutlass-dsl` 4.5.x |
| `AttributeError: module 'cutlass.cute.core' has no attribute 'ThrMma'` | `quack-kernels` older than 0.5.1 |
| `ImportError: cannot import name 'alloc_reserved_mbarrier'` | `quack-kernels` 0.6.2 or newer |
An environment whose `flash_attn.cute` came from a prebuilt flash-attn wheel rather
than from this pin hits the first row; that is what the overlay step in
`docker/Dockerfile` works around.
The CuTe kernels JIT-compile on first use. Forward and backward are verified on
Blackwell (sm_100) against `tests/test_vsa128_*.py` and `tests/test_vsa256_*.py`.
## Usage
@@ -142,6 +163,14 @@ After building/installing `fastvideo-kernel`, run:
```bash
cd fastvideo-kernel
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
# VSA-256 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 256 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 20
# VSA-128 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 128 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 40
```
### TurboDiffusion Kernels
+48 -20
View File
@@ -2,8 +2,9 @@
"""
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
This script benchmarks the autograd-enabled wrapper:
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
This script benchmarks the autograd-enabled wrappers:
- 64-token TK/Triton: fastvideo_kernel.block_sparse_attn.block_sparse_attn
- 128/256-token Triton/CuTe: fastvideo_kernel.block_sparse_attn_256
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
"""
@@ -23,9 +24,6 @@ try:
except Exception as e: # pragma: no cover
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
BLOCK_M = 64
BLOCK_N = 64
def set_seed(seed: int = 42) -> None:
random.seed(seed)
@@ -41,7 +39,11 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--num_heads", type=int, default=12)
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
p.add_argument("--q_seq_lens",
type=int,
nargs="+",
default=[49152],
help="Q sequence lengths (must be divisible by --block_size)")
p.add_argument("--kv_seq_lens",
type=int,
nargs="+",
@@ -51,9 +53,13 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--rep", type=int, default=20)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
p.add_argument("--block_size", type=int, default=64, choices=[64, 128, 256])
p.add_argument("--force_triton",
action="store_true",
help="Force wrapper to use Triton path (if supported by shapes).")
p.add_argument("--use_cute",
action="store_true",
help="Use the optional FA4 CuTe forward/backward path (requires --block_size 128 or 256).")
return p.parse_args()
@@ -84,18 +90,38 @@ def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
def _configure_backend(args: argparse.Namespace) -> None:
if args.use_cute and args.block_size not in (128, 256):
raise ValueError("--use_cute requires --block_size 128 or 256")
if args.use_cute and args.force_triton:
raise ValueError("--use_cute and --force_triton are mutually exclusive")
if args.force_triton:
os.environ.pop("FASTVIDEO_VSA_CUTEDSL", None)
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
elif args.use_cute:
os.environ.pop("FASTVIDEO_VSA_TRITON", None)
os.environ.pop("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", None)
os.environ["FASTVIDEO_VSA_CUTEDSL"] = "1"
def main() -> None:
args = parse_arguments()
set_seed(args.seed)
_configure_backend(args)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
if args.force_triton:
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128, block_sparse_attn_256
bs, h, d = args.batch_size, args.num_heads, args.head_dim
block_size = args.block_size
attention = {
64: block_sparse_attn,
128: block_sparse_attn_128,
256: block_sparse_attn_256,
}[block_size]
kv_seq_lens = args.kv_seq_lens
if kv_seq_lens is None:
kv_seq_lens = args.q_seq_lens
@@ -105,20 +131,22 @@ def main() -> None:
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
print(f"device: {torch.cuda.get_device_name(0)}")
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
print(f"block_size={block_size}")
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
if args.force_triton:
if args.use_cute:
print("dispatch: FA4 CuTe")
elif args.force_triton:
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
else:
print("dispatch: SM90 if available, else Triton")
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
if q_len % block_size != 0 or kv_len % block_size != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by {block_size}")
continue
num_q_blocks = q_len // BLOCK_M
num_kv_blocks = kv_len // BLOCK_N
num_q_blocks = q_len // block_size
num_kv_blocks = kv_len // block_size
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
topk = min(topk, num_kv_blocks)
@@ -129,11 +157,11 @@ def main() -> None:
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
# Variable block sizes: default full blocks (64 tokens per KV block)
variable_block_sizes = torch.full((num_kv_blocks, ), BLOCK_N, dtype=torch.int32, device="cuda")
# Variable block sizes: default full logical blocks.
variable_block_sizes = torch.full((num_kv_blocks, ), block_size, dtype=torch.int32, device="cuda")
def _fwd():
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
return attention(q, k, v, block_map, variable_block_sizes)
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
@@ -142,7 +170,7 @@ def main() -> None:
q_ = q.detach().requires_grad_(True)
k_ = k.detach().requires_grad_(True)
v_ = v.detach().requires_grad_(True)
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
o_, _aux_ = attention(q_, k_, v_, block_map, variable_block_sizes)
og = torch.randn_like(o_)
loss = (o_ * og).sum()
@@ -156,7 +184,7 @@ def main() -> None:
rep=max(5, args.rep // 2),
)
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
flops = flops_sparse_attention(bs, h, d, q_len, topk, block_size)
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
# Rough backward multiplier (attention backward typically ~2-3x forward)
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
@@ -1,4 +1,4 @@
"""VSA-256 block-sparse attention wrapper.
"""VSA-128/256 block-sparse attention wrappers.
The default 256-block path is Triton: it expands the logical 256-block map
to the existing 64-block Triton kernel via a dense 4x4 expansion per logical
@@ -7,8 +7,8 @@ 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
on 128-token Q/KV blocks (the 256 wrapper expands its logical KV map and
sizes into that physical representation). The CuTe kernel
(``flash_attn.cute`` with block-sparsity) is an optional dependency,
imported lazily only when this fastpath is selected.
@@ -35,7 +35,7 @@ _KV_BLOCK_TRITON = 64 # Existing Triton path uses 64-token KV blocks.
def _resolve_backend() -> str:
"""Pick the backend for the 256-block VSA path.
"""Pick the backend for the 128/256-block VSA paths.
Default is Triton (no optional deps). The FA4 CuTe fastpath is opt-in
via ``FASTVIDEO_VSA_CUTEDSL=1`` and requires the optional FA4 CuTe
@@ -49,6 +49,26 @@ def _resolve_backend() -> str:
return "triton"
def _expand_mask_and_sizes_128_to_64(
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Expand a [B, H, Qb128, KVb128] map to 64-token Triton tiles."""
expanded_mask = logical_mask_128.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
sizes_i32 = logical_kv_sizes_128.to(torch.int32)
offsets = torch.tensor(
[0, _KV_BLOCK_TRITON],
dtype=torch.int32,
device=sizes_i32.device,
)
expanded_sizes = torch.clamp(
sizes_i32[:, None] - offsets[None, :],
min=0,
max=_KV_BLOCK_TRITON,
).reshape(-1)
return expanded_mask, expanded_sizes
def _expand_mask_and_sizes_256_to_128(
logical_mask_256: torch.Tensor,
logical_kv_sizes_256: torch.Tensor,
@@ -112,6 +132,63 @@ def _triton_via_route_a(
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def _triton_via_route_a_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
from .triton_kernels.index import map_to_index as triton_map_to_index
mask_64, sizes_64 = _expand_mask_and_sizes_128_to_64(logical_mask_128, logical_kv_sizes_128)
q2k_idx, q2k_num = triton_map_to_index(mask_64.to(torch.bool))
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def block_sparse_attn_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, H, S, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
return _triton_via_route_a_128(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd
return block_sparse_attn_cute_fwd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_128_bshd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, S, H, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
out_bhsd, aux = _triton_via_route_a_128(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
logical_block_map_128,
logical_variable_block_sizes_128,
)
return out_bhsd.transpose(1, 2).contiguous(), aux
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
return block_sparse_attn_cute_fwd_bshd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_256(
q: torch.Tensor,
k: torch.Tensor,
@@ -1,18 +1,18 @@
"""CuTe-DSL block-sparse attention forward kernel.
"""FA4 CuTe-DSL block-sparse attention adapter.
Thin wrapper around `flash_attn.cute.interface._flash_attn_fwd` that adapts
VSA's `(block_map, variable_block_sizes)` inputs into FA4's
`BlockSparseTensorsTorch` representation and the per-KV-block validity mask.
This module adapts VSA's ``(block_map, variable_block_sizes)`` inputs into
FA4's forward and backward ``BlockSparseTensorsTorch`` representations.
FA4's public ``flash_attn_func`` owns the forward/backward autograd bridge.
Both [B, H, S, D] (BHSD) and [B, S, H, D] (BSHD) entrypoints are provided.
The BSHD variant is preferred from VSA-256 callers to avoid layout
The BSHD variant is preferred from VSA-128/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``.
only exercised when the VSA-128/256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default path is Triton and does not require
it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
"""
from __future__ import annotations
@@ -22,13 +22,14 @@ from typing import Tuple
import torch
_FA4_IMPORT_HINT = ("VSA-256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
_FA4_IMPORT_HINT = ("VSA-128/256 CuTe 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 "
"dependency; the default path is Triton. Install the FA4 CuTe "
"build and set FASTVIDEO_VSA_CUTEDSL=1 to enable the CuTe fastpath.")
@functools.lru_cache(maxsize=1)
def _load_fa4_cute():
"""Lazily import the optional FA4 CuTe block-sparse symbols.
@@ -38,14 +39,39 @@ def _load_fa4_cute():
"""
try:
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
from flash_attn.cute.interface import _flash_attn_fwd
from flash_attn.cute.interface import (
_flash_attn_bwd,
_flash_attn_fwd,
flash_attn_func,
)
except ImportError as exc: # pragma: no cover - optional dependency
raise ImportError(_FA4_IMPORT_HINT) from exc
return BlockSparseTensorsTorch, _flash_attn_fwd
return BlockSparseTensorsTorch, flash_attn_func, _flash_attn_fwd, _flash_attn_bwd
# Q-side tile size; kv_block_size comes from the caller's VSA logical KV block.
_M_BLOCK_SIZE_DEFAULT = 128
# FA4's physical Q tile size; KV block size comes from the VSA caller.
_FA4_Q_BLOCK_SIZE = 128
class _SingleQStageLength(int):
"""Keep the real length while selecting FA4's one-stage Q128 path.
On sm_100 FA4 derives ``q_stage`` from ``max_seqlen_q > tile_m``. Its
kernel supports one 128-token Q stage, but the fixed-length public wrapper
does not expose that choice. VSA-128 must select it explicitly; otherwise
adjacent logical Q blocks are merged into a 256-token sparse block.
"""
def __mul__(self, other):
return type(self)(int(self) * int(other))
def __rmul__(self, other):
return type(self)(int(other) * int(self))
def __gt__(self, other):
if int(other) == _FA4_Q_BLOCK_SIZE:
return False
return int(self) > int(other)
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
@@ -64,12 +90,12 @@ def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
return triton_map_to_index(block_map)
def _choose_q_sparse_block_size(q_len: int, m_block_size: int = _M_BLOCK_SIZE_DEFAULT) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > m_block_size.
def _choose_q_sparse_block_size(q_len: int, q_tile_size: int = _FA4_Q_BLOCK_SIZE) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > q_tile_size.
major, _ = torch.cuda.get_device_capability()
if major >= 10 and q_len > m_block_size:
return 2 * m_block_size
return m_block_size
if major >= 10 and q_len > q_tile_size:
return 2 * q_tile_size
return q_tile_size
def _aggregate_q_block_map(
@@ -134,23 +160,35 @@ def _build_vbs_mask_mod(kv_block_size: int):
return _vbs_mask_mod
def _cute_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
def _build_sparse_tensors(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
q_len: int,
q_block_size: int,
kv_block_size: int,
) -> 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,
)
need_backward: bool,
force_q_sparse_block_size: int | None = None,
) -> Tuple[object, object | None]:
"""Build the Q-owned forward and KV-owned backward sparse metadata.
``need_backward`` is False on inference-only calls: the backward metadata
is a pair of dense ``[B, H, kv_blocks, q_blocks]`` int32 index tensors that
FA4 keeps alive on its autograd ctx until backward runs, so building it
when nothing requires grad is pure overhead (~80 MiB per call at Wan-14B
720p shape).
"""
BlockSparseTensorsTorch, _, _, _ = _load_fa4_cute()
if force_q_sparse_block_size is None:
q_sparse_candidate = _choose_q_sparse_block_size(q_len)
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
else:
q_sparse_block_size = force_q_sparse_block_size
if q_sparse_block_size < q_block_size or q_sparse_block_size % q_block_size != 0:
raise ValueError("force_q_sparse_block_size must be a positive multiple of q_block_size")
sparse_map = _aggregate_q_block_map(
block_map,
q_sparse_block_size=q_sparse_block_size,
@@ -158,35 +196,166 @@ def _cute_forward(
)
kv_full = (variable_block_sizes == kv_block_size).view(1, 1, 1, -1)
kv_partial = ((variable_block_sizes > 0) & (variable_block_sizes < kv_block_size)).view(1, 1, 1, -1)
full_map = sparse_map & kv_full
mask_map = sparse_map & kv_partial
full_block_idx, full_block_cnt = _map_to_index(full_map)
mask_block_idx, mask_block_cnt = _map_to_index(mask_map)
def from_maps(full_map: torch.Tensor, mask_map: torch.Tensor) -> object:
full_block_idx, full_block_cnt = _map_to_index(full_map.contiguous())
mask_block_idx, mask_block_cnt = _map_to_index(mask_map.contiguous())
return BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
)
sparse_tensors = BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
forward_sparse_tensors = from_maps(
sparse_map & kv_full,
sparse_map & kv_partial,
)
# _flash_attn_fwd returns (out, lse, p, row_max); keep the first two.
out, lse = _flash_attn_fwd(
if not need_backward:
return forward_sparse_tensors, None
# FA4 backward is KV-owned: for each physical KV tile, list the sparse
# query tiles that selected it. Full and partial KV tiles stay separate
# so the token-level validity mask only runs for padded tiles.
backward_sparse_tensors = from_maps(
(sparse_map & kv_full).transpose(2, 3),
(sparse_map & kv_partial).transpose(2, 3),
)
return forward_sparse_tensors, backward_sparse_tensors
def _cute_attention_q128_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
need_backward: bool,
) -> Tuple[torch.Tensor, torch.Tensor, object | None]:
"""Run FA4 with one physical Q stage per logical VSA-128 block."""
_, _, flash_attn_fwd, _ = _load_fa4_cute()
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=_FA4_Q_BLOCK_SIZE,
kv_block_size=_FA4_Q_BLOCK_SIZE,
need_backward=need_backward,
force_q_sparse_block_size=_FA4_Q_BLOCK_SIZE,
)
out, lse = flash_attn_fwd(
q_bshd,
k_bshd,
v_bshd,
tile_mn=(_M_BLOCK_SIZE_DEFAULT, kv_block_size),
mask_mod=_build_vbs_mask_mod(kv_block_size),
block_sparse_tensors=sparse_tensors,
tile_mn=(_FA4_Q_BLOCK_SIZE, _FA4_Q_BLOCK_SIZE),
max_seqlen_q=_SingleQStageLength(q_bshd.shape[1]),
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
block_sparse_tensors=forward_sparse_tensors,
aux_tensors=[variable_block_sizes],
causal=False,
return_lse=True,
)[:2]
return out, lse, backward_sparse_tensors
class _CuteAttentionQ128(torch.autograd.Function):
@staticmethod
def forward(ctx, q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes):
out, lse, backward_sparse_tensors = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=True,
)
ctx.save_for_backward(q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes)
ctx.backward_sparse_tensors = backward_sparse_tensors
ctx.mark_non_differentiable(lse)
ctx.set_materialize_grads(False)
return out, lse
@staticmethod
def backward(ctx, grad_out, grad_lse):
del grad_lse
q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes = ctx.saved_tensors
if grad_out is None:
grad_out = torch.zeros_like(out)
_, _, _, flash_attn_bwd = _load_fa4_cute()
dq, dk, dv = flash_attn_bwd(
q_bshd,
k_bshd,
v_bshd,
out,
grad_out.contiguous(),
lse,
softmax_scale=q_bshd.shape[-1]**-0.5,
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=ctx.backward_sparse_tensors,
)
return dq, dk, dv, None, None
def _cute_attention_q128(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
if need_backward:
return _CuteAttentionQ128.apply(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
out, lse, _ = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=False,
)
return out, lse
def _cute_attention(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Run FA4's autograd-enabled block-sparse attention with BSHD inputs."""
_, flash_attn_func, _, _ = _load_fa4_cute()
q_block_size = q_bshd.shape[1] // block_map.shape[2]
kv_block_size = k_bshd.shape[1] // block_map.shape[3]
if q_block_size == kv_block_size == _FA4_Q_BLOCK_SIZE:
return _cute_attention_q128(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=q_block_size,
kv_block_size=kv_block_size,
need_backward=need_backward,
)
return flash_attn_func(
q_bshd,
k_bshd,
v_bshd,
mask_mod=_build_vbs_mask_mod(kv_block_size),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=forward_sparse_tensors,
block_sparse_tensors_bwd=backward_sparse_tensors,
return_lse=True,
)
def block_sparse_attn_cute_fwd(
q: torch.Tensor,
k: torch.Tensor,
@@ -194,34 +363,25 @@ def block_sparse_attn_cute_fwd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""CuTe forward-only block-sparse attention with [B, H, S, D] inputs."""
"""Autograd-enabled CuTe block-sparse attention for [B, H, S, D]."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[2] // block_map.shape[2]
kv_block_size = k.shape[2] // block_map.shape[3]
q_bshd = q.transpose(1, 2).contiguous()
k_bshd = k.transpose(1, 2).contiguous()
v_bshd = v.transpose(1, 2).contiguous()
out_bshd, lse_bshd = _cute_forward(
out_bshd, lse = _cute_attention(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
out = out_bshd.transpose(1, 2).contiguous()
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
# FA4 already returns lse as [B, H, S], matching the Triton path's aux
# contract, so it needs no transpose. Detach before any further op: the
# value is informational and callers never backprop through it.
return out, lse.detach()
def block_sparse_attn_cute_fwd_bshd(
@@ -231,27 +391,16 @@ def block_sparse_attn_cute_fwd_bshd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""CuTe forward-only block-sparse attention with [B, S, H, D] inputs."""
"""Autograd-enabled CuTe block-sparse attention for [B, S, H, D]."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[1] // block_map.shape[2]
kv_block_size = k.shape[1] // block_map.shape[3]
out, lse_bshd = _cute_forward(
out, lse = _cute_attention(
q,
k,
v,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[2], q.shape[1]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
# lse is [B, H, S] regardless of the q/k/v layout; see above.
return out, lse.detach()
+24 -22
View File
@@ -2,6 +2,8 @@ import math
import torch
from .block_sparse_attn import block_sparse_attn
from .block_sparse_attn_256 import (
block_sparse_attn_128,
block_sparse_attn_128_bshd,
block_sparse_attn_256,
block_sparse_attn_256_bshd,
)
@@ -74,12 +76,13 @@ def video_sparse_attn(
Dispatches the sparse branch by ``block_elements = prod(block_size)``:
- 64 -> existing TK/Triton path (see ``block_sparse_attn_from_indices``).
- 128 -> Triton fallback or CuTe FA4 block-sparse attention.
- 256 -> CuTe FA4 block-sparse attention (see ``block_sparse_attn_256``).
Backend overrides:
- ``FASTVIDEO_VSA_TRITON=1`` forces Triton in either path.
- ``FASTVIDEO_VSA_TK=1`` prefers the sm_90 TK kernel in the 64-block path.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 256-block path.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 128/256-block paths.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
@@ -119,8 +122,9 @@ def video_sparse_attn(
# Sparse branch (fused Triton topk mask)
mask = fused_topk_mask(scores, topk)
if block_elements == 256:
out_s = block_sparse_attn_256(q, k, v, mask, variable_block_sizes)[0]
if block_elements in (128, 256):
attention = block_sparse_attn_128 if block_elements == 128 else block_sparse_attn_256
out_s = attention(q, k, v, mask, variable_block_sizes)[0]
else:
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
@@ -142,14 +146,14 @@ def video_sparse_attn_bshd(
"""VSA entrypoint for [B, S, H, D] tensors.
Avoids the BHSD<->BSHD round-trip that ``video_sparse_attn`` performs on
the CuTe 256-block path; the 64-block path still expects BHSD and is not
the CuTe 128/256-block paths; the 64-block path still expects BHSD and is not
supported here.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
if block_elements != 256:
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=256 "
if block_elements not in (128, 256):
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=128 or 256 "
f"(got {block_elements}); use video_sparse_attn for the 64-block path.")
batch, q_seq_len, heads, dim = q.shape
@@ -171,19 +175,15 @@ def video_sparse_attn_bshd(
raise ValueError(f"q_variable_block_sizes must have length q_num_blocks={q_num_blocks}, "
f"got {q_variable_block_sizes.numel()}")
# Compression branch (BSHD-native: mean over the 256-token axis).
token_idx = torch.arange(block_elements, device=q.device, dtype=torch.int32)
q_token_valid = (token_idx.view(1, -1) < q_variable_block_sizes.view(-1,
1)).view(1, q_num_blocks, block_elements, 1, 1)
kv_token_valid = (token_idx.view(1, -1) < variable_block_sizes.view(-1,
1)).view(1, kv_num_blocks, block_elements, 1, 1)
# Compression branch (BSHD-native: match fused_block_mean's semantics).
# Padding values are expected to be zero; gradients are broadcast across
# the full padded block, just like the BHSD fused common path.
q_c = q.view(batch, q_num_blocks, block_elements, heads, dim)
k_c = k.view(batch, kv_num_blocks, block_elements, heads, dim)
v_c = v.view(batch, kv_num_blocks, block_elements, heads, dim)
q_c = ((q_c.float() * q_token_valid).sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = ((k_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = ((v_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_c = (q_c.float().sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = (k_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = (v_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_ch = q_c.permute(0, 2, 1, 3).contiguous()
k_ch = k_c.permute(0, 2, 1, 3).contiguous()
v_ch = v_c.permute(0, 2, 1, 3).contiguous()
@@ -195,13 +195,15 @@ def video_sparse_attn_bshd(
# Sparse branch (fused Triton topk mask + CuTe BSHD).
mask = fused_topk_mask(scores, topk)
out_s, _ = block_sparse_attn_256_bshd(q, k, v, mask, variable_block_sizes)
attention = block_sparse_attn_128_bshd if block_elements == 128 else block_sparse_attn_256_bshd
out_s, _ = attention(q, k, v, mask, variable_block_sizes)
out = out_s
out_view = out.view(batch, q_num_blocks, block_elements, heads, dim)
# Out-of-place: ``out_s`` is the tensor FA4's autograd node saved for its
# backward, so mutating it in place invalidates the graph.
out_view = out_s.view(batch, q_num_blocks, block_elements, heads, dim)
if compress_attn_weight is not None:
gate_view = compress_attn_weight.view(batch, q_num_blocks, block_elements, heads, dim)
out_view.add_(out_c_blk.unsqueeze(2) * gate_view)
out = out_view + out_c_blk.unsqueeze(2) * gate_view
else:
out_view.add_(out_c_blk.unsqueeze(2))
return out
out = out_view + out_c_blk.unsqueeze(2)
return out.view(batch, q_seq_len, heads, dim)
@@ -0,0 +1,146 @@
"""VSA-128 CuTe/Triton forward and backward parity on Blackwell."""
from __future__ import annotations
import math
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 128
_BLOCK_SIZE_3D = (2, 8, 8)
def _select_backend(monkeypatch, backend: str) -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
if backend == "cute":
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
else:
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
def _dense_sparse_reference(q, k, v, block_map, variable_block_sizes):
token_mask = block_map.repeat_interleave(_BLOCK, dim=2).repeat_interleave(_BLOCK, dim=3)
kv_valid = torch.arange(_BLOCK, device=k.device) < variable_block_sizes[:, None]
token_mask = token_mask & kv_valid.reshape(1, 1, 1, -1)
logits = torch.matmul(q.float(), k.float().transpose(-2, -1)) / math.sqrt(q.shape[-1])
probabilities = torch.softmax(logits.masked_fill(~token_mask, float("-inf")), dim=-1)
return torch.matmul(probabilities, v.float()).to(q.dtype)
def _check(tag: str, expected: torch.Tensor, actual: torch.Tensor, avg_tol: float, rel_tol: float) -> None:
assert torch.isfinite(actual).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(expected, actual)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < avg_tol
assert max_rel < rel_tol
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
def test_vsa128_explicit_routes_forward_backward(backend: str, monkeypatch) -> None:
"""Adjacent Q128 blocks must keep independent routes instead of merging."""
_select_backend(monkeypatch, backend)
torch.manual_seed(53)
shape = (1, 1, 3 * _BLOCK, 128)
base = [torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3)]
grad_output = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
variable_block_sizes = torch.tensor([128, 91, 37], device="cuda", dtype=torch.int32)
block_map = torch.eye(3, device="cuda", dtype=torch.bool).view(1, 1, 3, 3)
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
actual, _ = block_sparse_attn_128(*actual_inputs, block_map, variable_block_sizes)
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
expected = _dense_sparse_reference(*reference_inputs, block_map, variable_block_sizes)
(expected * grad_output).sum().backward()
print(f"[vsa128-explicit-{backend}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_inputs, strict=True):
_check(name, reference.grad, candidate.grad, 2e-2, 0.5)
def _zero_kv_tail(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
valid = torch.arange(_BLOCK, device=x.device) < variable_block_sizes[:, None]
valid = valid.view(1, 1, -1, _BLOCK, 1).expand_as(x.view(1, x.shape[1], -1, _BLOCK, x.shape[-1]))
return x * valid.reshape_as(x).to(x.dtype)
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa128_wrapper_forward_backward(backend: str, layout: str, monkeypatch) -> None:
_select_backend(monkeypatch, backend)
torch.manual_seed(59)
batch, heads, dim = 1, 2, 128
q_blocks, kv_blocks, topk = 3, 4, 2
q_shape = (batch, heads, q_blocks * _BLOCK, dim)
kv_shape = (batch, heads, kv_blocks * _BLOCK, dim)
q_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
kv_sizes = torch.tensor([128, 91, 37, 128], device="cuda", dtype=torch.int32)
q_sizes = torch.full((q_blocks, ), _BLOCK, device="cuda", dtype=torch.int32)
k_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
v_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
gate_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16) * 0.1
grad_output = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
if layout == "bhsd":
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
actual_gate = gate_base.detach().clone().requires_grad_()
actual = video_sparse_attn(
*actual_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=actual_gate,
)
actual_grads = actual_inputs
else:
bshd_inputs = [tensor.transpose(1, 2).contiguous().detach().requires_grad_()
for tensor in (q_base, k_base, v_base)]
bshd_gate = gate_base.transpose(1, 2).contiguous().detach().requires_grad_()
actual = video_sparse_attn_bshd(
*bshd_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=bshd_gate,
).transpose(1, 2)
actual_grads = bshd_inputs
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
reference_gate = gate_base.detach().clone().requires_grad_()
expected = _torch_vsa256_reference(
*reference_inputs,
q_sizes,
kv_sizes,
topk,
compress_attn_weight=reference_gate,
)
(expected * grad_output).sum().backward()
print(f"[vsa128-wrapper-{backend}-{layout}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_grads, strict=True):
candidate_grad = candidate.grad if layout == "bhsd" else candidate.grad.transpose(1, 2)
_check(name, reference.grad, candidate_grad, 2e-2, 0.5)
actual_gate_grad = actual_gate.grad if layout == "bhsd" else bshd_gate.grad.transpose(1, 2)
_check("dgate", reference_gate.grad, actual_gate_grad, 1e-3, 0.2)
@@ -0,0 +1,224 @@
"""VSA-256 FA4 CuTe forward/backward parity for BHSD and BSHD APIs.
Covers the shapes the CuTe backward actually sees in production: the gated
compression branch (`compress_attn_weight`), partially filled Q tiles,
and q_len != kv_len. Also pins the inference fast path, which must skip the
KV-owned backward metadata without changing the forward result.
"""
from __future__ import annotations
from typing import Tuple
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 256
_BLOCK_SIZE_3D = (4, 8, 8) # prod == 256
# Measured on GB200 (sm_100) with bf16 inputs: grads land around 1e-4 avg_abs
# and <=0.11 max_rel across every case below, so these leave ~10x headroom
# without being loose enough to hide a real regression.
_OUT_TOL = (1e-3, 0.2)
_GRAD_TOL = (1e-3, 0.25)
@pytest.fixture(autouse=True)
def _require_cute_backend(monkeypatch):
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
def _zero_pad_tail(x: torch.Tensor, var: torch.Tensor) -> torch.Tensor:
"""Zero the padded tail of every 256-token tile of a [B, H, S, D] tensor.
VSA callers scatter into a zeroed tile buffer, so padded slots are zero;
both the kernel and the reference rely on that.
"""
bsz, heads, _, dim = x.shape
blocks = var.numel()
token_idx = torch.arange(_BLOCK, device=x.device, dtype=torch.int32)
valid = (token_idx.view(1, -1) < var.view(-1, 1)).view(1, 1, blocks, _BLOCK, 1)
valid = valid.expand(bsz, heads, blocks, _BLOCK, dim).reshape_as(x)
return x * valid.to(x.dtype)
def _make_inputs(
q_blocks: int,
kv_blocks: int,
kv_var: torch.Tensor,
q_var: torch.Tensor,
heads: int = 2,
dim: int = 128,
seed: int = 42,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch.manual_seed(seed)
device = torch.device("cuda")
dtype = torch.bfloat16
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=dtype)
k = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
v = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
grad_out = torch.randn_like(q)
return _zero_pad_tail(q, q_var), _zero_pad_tail(k, kv_var), _zero_pad_tail(v, kv_var), grad_out
def _check(tag: str, ref: torch.Tensor, got: torch.Tensor, tol: Tuple[float, float]) -> None:
assert torch.isfinite(got).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(ref, got)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < tol[0], f"{tag}: avg_abs {avg_abs:.3e} >= {tol[0]:.3e}"
assert max_rel < tol[1], f"{tag}: max_rel {max_rel:.3e} >= {tol[1]:.3e}"
def _run_bhsd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = video_sparse_attn(qg, kg, vg, kv_var, q_var, topk, block_size=_BLOCK_SIZE_3D, compress_attn_weight=gate)
return out, (qg, kg, vg)
def _run_bshd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.transpose(1, 2).contiguous().requires_grad_(True) for t in (q, k, v))
gate_bshd = None if gate is None else gate.transpose(1, 2).contiguous()
out = video_sparse_attn_bshd(qg,
kg,
vg,
kv_var,
q_var,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=gate_bshd)
return out.transpose(1, 2), (qg, kg, vg)
def _reference(q, k, v, q_var, kv_var, topk, gate=None):
qr, kr, vr = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = _torch_vsa256_reference(qr, kr, vr, q_var, kv_var, topk, compress_attn_weight=gate)
return out, (qr, kr, vr)
def _compare(tag, layout, q, k, v, kv_var, q_var, topk, grad_out, gate=None):
runner = _run_bhsd if layout == "bhsd" else _run_bshd
out, (qg, kg, vg) = runner(q, k, v, kv_var, q_var, topk, gate=gate)
(out * grad_out).sum().backward()
grads = [g.grad if g.grad.dim() == 4 and layout == "bhsd" else g.grad for g in (qg, kg, vg)]
if layout == "bshd":
grads = [g.transpose(1, 2) for g in grads]
out_ref, refs = _reference(q, k, v, q_var, kv_var, topk, gate=gate)
(out_ref * grad_out).sum().backward()
print(f"[{tag}-{layout}]")
_check("out", out_ref, out, _OUT_TOL)
for name, ref, got in zip(("dq", "dk", "dv"), refs, grads):
_check(name, ref.grad, got, _GRAD_TOL)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_forward_backward_vs_torch_ref(layout: str) -> None:
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var)
_compare("vsa256-cute", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_with_compress_gate(layout: str) -> None:
"""The gated compression branch is what Wan and MiniMax-H3 actually run.
It is also the branch that composes the sparse output with the compression
output, so it is the one that breaks if that composition mutates FA4's
saved output in place.
"""
kv_var = torch.tensor([256, 200, 256, 91], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var, seed=7)
gate = torch.randn_like(q) * 0.1
_compare("vsa256-cute-gated", layout, q, k, v, kv_var, q_var, 2, grad_out, gate=gate)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_partial_q_blocks(layout: str) -> None:
"""Q tiles that are not full: only the compression divisor depends on it,
but it is the one axis the existing coverage held constant."""
kv_var = torch.tensor([256, 128, 256], dtype=torch.int32, device="cuda")
q_var = torch.tensor([256, 61, 199], dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 3, kv_var, q_var, seed=11)
_compare("vsa256-cute-partial-q", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_cross_q_kv(layout: str) -> None:
"""q_len != kv_len: forward has coverage, backward did not."""
kv_var = torch.tensor([256, 143, 256, 256, 88], dtype=torch.int32, device="cuda")
q_var = torch.full((2, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(2, 5, kv_var, q_var, seed=13)
_compare("vsa256-cute-cross", layout, q, k, v, kv_var, q_var, 3, grad_out)
@pytest.mark.cuda
def test_vsa256_cute_inference_matches_training_forward() -> None:
"""The KV-owned backward metadata is only built when something requires
grad. Skipping it must not perturb the forward result."""
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, _ = _make_inputs(3, 4, kv_var, q_var, seed=5)
with torch.no_grad():
out_infer = video_sparse_attn_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
kv_var,
q_var,
2,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=None,
)
out_train, _ = _run_bshd(q, k, v, kv_var, q_var, 2)
torch.testing.assert_close(out_infer, out_train.transpose(1, 2).detach(), rtol=0, atol=0)
@pytest.mark.cuda
def test_vsa256_cute_lse_is_bhs() -> None:
"""The aux return is [B, H, S] on both entrypoints, matching the Triton
path's contract."""
from fastvideo_kernel.block_sparse_attn_256 import (block_sparse_attn_256, block_sparse_attn_256_bshd)
device = torch.device("cuda")
heads, dim, q_blocks, kv_blocks = 2, 128, 3, 4
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=torch.bfloat16)
k = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
v = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
vbs = torch.full((kv_blocks, ), _BLOCK, dtype=torch.int32, device=device)
mask = torch.zeros(1, heads, q_blocks, kv_blocks, dtype=torch.bool, device=device)
mask[..., :2] = True
_, lse_bhsd = block_sparse_attn_256(q, k, v, mask, vbs)
assert lse_bhsd.shape == (1, heads, sq), lse_bhsd.shape
_, lse_bshd = block_sparse_attn_256_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
mask,
vbs,
)
assert lse_bshd.shape == (1, heads, sq), lse_bshd.shape
@@ -22,6 +22,7 @@ def _torch_vsa256_reference(
q_var: torch.Tensor,
kv_var: torch.Tensor,
topk_logical: int,
compress_attn_weight: torch.Tensor | None = None,
) -> torch.Tensor:
bsz, heads, _sq, dim = q.shape
q_blocks = q_var.numel()
@@ -55,6 +56,8 @@ def _torch_vsa256_reference(
logits = logits.masked_fill(~token_mask, float("-inf"))
prob = torch.softmax(logits, dim=-1)
out_s = torch.matmul(prob, vf).to(q.dtype)
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
return out_c + out_s
+14
View File
@@ -23,6 +23,20 @@ from fastvideo_kernel.block_sparse_attn import (
from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen
@pytest.fixture(autouse=True)
def _seed_rng():
"""Pin the RNG so these cases do not depend on what ran before them.
Every tensor and every variable block size here comes from the global
torch RNG, and the gradient checks use a tight max_rel threshold. Without
a seed the inputs shift whenever an earlier test file draws a different
number of randoms, which surfaces as an unrelated-looking failure in
whichever case happens to land on unlucky data.
"""
torch.manual_seed(0)
torch.cuda.manual_seed_all(0)
def _reference_per_sequence(
q_list,
k_list,
@@ -319,8 +319,13 @@ class MiniMaxH3VSAImpl(AttentionImpl):
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes)
out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled) # [B, H, n_tiles, D]
out_c = out_c.permute(0, 2, 1, 3).to(out.dtype) # [B, n_tiles, H, D]
batch, _, heads, dim = out.shape
batch, seq_len, heads, dim = out.shape
n_tiles = attn_metadata.variable_block_sizes.numel()
out.view(batch, n_tiles, _TILE_ELEMS, heads,
dim).addcmul_(out_c.unsqueeze(2), gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim))
# Out-of-place: on the CuTe backend ``out`` is the tensor FA4's
# autograd node saved for its backward, so an in-place add here
# bumps its version counter and backward dies with "one of the
# variables needed for gradient computation has been modified".
out_tiled = out.view(batch, n_tiles, _TILE_ELEMS, heads, dim)
gate_tiled = gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim)
out = (out_tiled + out_c.unsqueeze(2) * gate_tiled).view(batch, seq_len, heads, dim)
return out
@@ -0,0 +1,110 @@
# SPDX-License-Identifier: Apache-2.0
"""GPU backward checks for the VSA-H3 backend.
The CuTe backend returns FA4's own output tensor, which FA4's autograd node
saved for its backward. Composing the compression branch onto it in place
therefore poisons the graph, and the failure only appears once the VSA-256
CuTe path has a backward at all. These tests pin the composition.
"""
import pytest
import torch
from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAImpl, MiniMaxH3VSAMetadataBuilder)
_SPEC = dict(raw_latent_shape=(16, 16, 24), patch_size=(1, 2, 2), prefix_segments=(64, 32, 16))
_HEADS = 2
_DIM = 128
def _build_meta(device, sparsity=0.5):
return MiniMaxH3VSAMetadataBuilder().build(
current_timestep=0,
raw_latent_shape=_SPEC["raw_latent_shape"],
patch_size=_SPEC["patch_size"],
VSA_sparsity=sparsity,
prefix_segments=_SPEC["prefix_segments"],
device=device,
)
def _select_backend(monkeypatch, backend):
if backend == "cute":
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
else:
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
def _forward_backward(impl, meta, gate_compress, device):
seq = meta.total_seq_length
torch.manual_seed(0)
q, k, v = (torch.randn(1, seq, _HEADS, _DIM, device=device, dtype=torch.bfloat16, requires_grad=True)
for _ in range(3))
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
gate = None
if gate_compress:
gate = torch.randn(1, tq.shape[1], _HEADS, _DIM, device=device, dtype=torch.bfloat16) * 0.1
out = impl.forward(tq, tk, tv, gate, meta)
out = impl.postprocess_output(out, meta)
out.float().pow(2).sum().backward()
return out, (q, k, v)
@pytest.mark.parametrize("backend", ["triton", "cute"])
@pytest.mark.parametrize("gate_compress", [False, True])
def test_h3_vsa_backward_runs(monkeypatch, backend: str, gate_compress: bool) -> None:
"""Regression: with the CuTe backend and a non-zero gate this used to die
with "one of the variables needed for gradient computation has been
modified by an inplace operation ... output 0 of FlashAttnFuncBackward".
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
_select_backend(monkeypatch, backend)
device = torch.device("cuda")
meta = _build_meta(device)
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
out, leaves = _forward_backward(impl, meta, gate_compress, device)
assert torch.isfinite(out).all().item()
for name, leaf in zip(("q", "k", "v"), leaves):
assert leaf.grad is not None, f"{name} received no gradient"
assert torch.isfinite(leaf.grad).all().item(), f"{name}.grad has non-finite values"
assert leaf.grad.abs().sum().item() > 0, f"{name}.grad is all zero"
@pytest.mark.parametrize("gate_compress", [False, True])
def test_h3_vsa_backward_cute_matches_triton(monkeypatch, gate_compress: bool) -> None:
"""CuTe and Triton take different routes to the same math; their gradients
should agree to bf16 tolerance."""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
device = torch.device("cuda")
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
grads = {}
for backend in ("triton", "cute"):
with monkeypatch.context() as m:
_select_backend(m, backend)
meta = _build_meta(device)
_, leaves = _forward_backward(impl, meta, gate_compress, device)
grads[backend] = [leaf.grad.detach().float() for leaf in leaves]
for name, ref, got in zip(("dq", "dk", "dv"), grads["triton"], grads["cute"]):
diff = (ref - got).abs()
avg_abs = diff.mean().item()
max_rel = (diff.max() / (ref.abs().mean() + 1e-6)).item()
print(f"[h3-vsa gate={gate_compress}] {name}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < 1e-2, f"{name}: avg_abs {avg_abs:.3e}"
assert max_rel < 0.5, f"{name}: max_rel {max_rel:.3e}"
+5 -3
View File
@@ -122,9 +122,11 @@ 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" }
# FA4 cute. This revision pins nvidia-cutlass-dsl==4.6.0.dev0; holding cutlass-dsl
# back at 4.5.x breaks the CuTe kernels at JIT time, not at import
# (see fastvideo-kernel/README.md).
# torch.compile support comes from FastVideo's own custom_op wrappers.
flash-attn-4 = { git = "https://github.com/Dao-AILab/flash-attention.git", rev = "14c377950125c70b7a9dabf9c561fca53715ac7d", subdirectory = "flash_attn/cute" }
[project.optional-dependencies]