[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:
co-authored by
Hyunsung Lee
alexzms
parent
8537dcd6de
commit
00338aa9ca
+1
-1
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user