From 2b82fbf29d78d9038c8fd17c7f3074513b7b010f Mon Sep 17 00:00:00 2001 From: freemty Date: Mon, 11 May 2026 21:46:44 +0800 Subject: [PATCH 1/3] [kernel] Add varlen support for block-sparse attention (#917) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add `block_sparse_attn_varlen` — a sequence packing wrapper that enables variable-length block-sparse attention in a single kernel launch. - Pack multiple variable-length sequences with block-aligned padding - Rebase per-sequence sparse indices (q2k_idx) to global offsets - Unpack output back to per-sequence layout - Support both Q and KV variable block sizes - Zero kernel modifications — delegates to existing block_sparse_attn_from_indices Tested on RTX 5880 Ada (Triton backend): - 9 forward correctness tests: all pass (error = 0.0) - Backward gradient verification: dQ/dK/dV match reference exactly - Covers: equal/unequal lengths, single sequence, many sequences, asymmetric Q/KV, dense attention, minimal blocks, default path Closes #917 --- .../python/fastvideo_kernel/__init__.py | 5 + .../block_sparse_attn_varlen.py | 207 ++++++++++++++++ fastvideo-kernel/tests/test_vsa_varlen.py | 230 ++++++++++++++++++ 3 files changed, 442 insertions(+) create mode 100644 fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py create mode 100644 fastvideo-kernel/tests/test_vsa_varlen.py diff --git a/fastvideo-kernel/python/fastvideo_kernel/__init__.py b/fastvideo-kernel/python/fastvideo_kernel/__init__.py index 0e6b8da6c8..e8b2e7671e 100644 --- a/fastvideo-kernel/python/fastvideo_kernel/__init__.py +++ b/fastvideo-kernel/python/fastvideo_kernel/__init__.py @@ -24,11 +24,16 @@ int8_quant, ) +from fastvideo_kernel.block_sparse_attn_varlen import ( + block_sparse_attn_varlen, +) + __all__ = [ "sliding_tile_attention", "video_sparse_attn", "block_sparse_attn", "block_sparse_attn_from_indices", + "block_sparse_attn_varlen", "moba_attn_varlen", "process_moba_input", "process_moba_output", diff --git a/fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py b/fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py new file mode 100644 index 0000000000..81101f9b69 --- /dev/null +++ b/fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py @@ -0,0 +1,207 @@ +"""Variable-length block-sparse attention via sequence packing. + +Packs multiple variable-length sequences into a single [1, H, T_total, D] +tensor and delegates to the existing block_sparse_attn_from_indices kernel +in a single launch. No kernel modifications required. +""" + +from __future__ import annotations + +from typing import Sequence + +import torch + +from .block_sparse_attn import block_sparse_attn_from_indices + +BLOCK_SIZE = 64 + + +def _scatter_to_padded( + src: torch.Tensor, + block_sizes: torch.Tensor, + block_size: int, + dst: torch.Tensor, + dst_offset: int, + src_start: int, + src_end: int, +) -> None: + """Copy tokens from a flat source into block-aligned positions in dst. + + Each block occupies exactly `block_size` slots in dst. The first + `block_sizes[b]` slots of block *b* receive real tokens; the remainder + stays zero (padding the kernel expects). + + src: [total_tokens, H, D] + dst: [1, H, total_padded, D] + block_sizes: [num_blocks] int32, actual token count per block. + """ + src_pos = src_start + dst_pos = dst_offset + for b in range(block_sizes.numel()): + actual = int(block_sizes[b].item()) + actual = min(actual, src_end - src_pos) + if actual > 0: + dst[:, :, dst_pos:dst_pos + actual, :] = ( + src[src_pos:src_pos + actual].transpose(0, 1).unsqueeze(0) + ) + src_pos += actual + dst_pos += block_size + + +def _gather_from_padded( + src: torch.Tensor, + block_sizes: torch.Tensor, + block_size: int, + dst: torch.Tensor, + src_offset: int, + dst_start: int, + dst_end: int, +) -> None: + """Inverse of _scatter_to_padded: extract real tokens from padded blocks. + + src: [1, H, total_padded, D] + dst: [total_tokens, H, D] + """ + src_pos = src_offset + dst_pos = dst_start + for b in range(block_sizes.numel()): + actual = int(block_sizes[b].item()) + actual = min(actual, dst_end - dst_pos) + if actual > 0: + dst[dst_pos:dst_pos + actual] = ( + src[0, :, src_pos:src_pos + actual, :].transpose(0, 1) + ) + dst_pos += actual + src_pos += block_size + + +def block_sparse_attn_varlen( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + cu_seqlens_q: torch.Tensor, + cu_seqlens_kv: torch.Tensor, + q2k_idx_list: Sequence[torch.Tensor], + q2k_num_list: Sequence[torch.Tensor], + variable_block_sizes_list: Sequence[torch.Tensor], + q_variable_block_sizes_list: Sequence[torch.Tensor] | None = None, + block_size: int = BLOCK_SIZE, +) -> torch.Tensor: + """Block-sparse attention over packed variable-length sequences. + + Args: + q: [total_q_tokens, H, D] packed query tensor. + k: [total_kv_tokens, H, D] packed key tensor. + v: [total_kv_tokens, H, D] packed value tensor. + cu_seqlens_q: [N+1] int32, cumulative Q token offsets. + cu_seqlens_kv: [N+1] int32, cumulative KV token offsets. + q2k_idx_list: Per-sequence q2k_idx tensors, each [1, H, Nq_i, Mk]. + q2k_num_list: Per-sequence q2k_num tensors, each [1, H, Nq_i]. + variable_block_sizes_list: Per-sequence KV block sizes, each [Nkv_i]. + q_variable_block_sizes_list: Per-sequence Q block sizes, each [Nq_i]. + If None, each Q block is assumed to be exactly `block_size` tokens. + block_size: Attention block size (default 64). + + Returns: + out: [total_q_tokens, H, D] packed output tensor. + """ + device = q.device + dtype = q.dtype + num_heads = q.shape[1] + head_dim = q.shape[2] + num_seqs = cu_seqlens_q.shape[0] - 1 + + cu_q = cu_seqlens_q.cpu().tolist() + cu_kv = cu_seqlens_kv.cpu().tolist() + + padded_q_lens = [] + padded_kv_lens = [] + q_block_offsets = [0] + kv_block_offsets = [0] + q_vbs_resolved = [] + + for i in range(num_seqs): + n_q_blocks = q2k_num_list[i].shape[-1] + n_kv_blocks = variable_block_sizes_list[i].numel() + padded_q_lens.append(n_q_blocks * block_size) + padded_kv_lens.append(n_kv_blocks * block_size) + q_block_offsets.append(q_block_offsets[-1] + n_q_blocks) + kv_block_offsets.append(kv_block_offsets[-1] + n_kv_blocks) + + if q_variable_block_sizes_list is not None: + q_vbs_resolved.append(q_variable_block_sizes_list[i]) + else: + q_vbs_resolved.append( + torch.full((n_q_blocks,), block_size, dtype=torch.int32, device=device) + ) + + total_padded_q = sum(padded_q_lens) + total_padded_kv = sum(padded_kv_lens) + + q_packed = torch.zeros(1, num_heads, total_padded_q, head_dim, device=device, dtype=dtype) + k_packed = torch.zeros(1, num_heads, total_padded_kv, head_dim, device=device, dtype=dtype) + v_packed = torch.zeros(1, num_heads, total_padded_kv, head_dim, device=device, dtype=dtype) + + q_offset = 0 + kv_offset = 0 + for i in range(num_seqs): + _scatter_to_padded( + q, q_vbs_resolved[i], block_size, + q_packed, q_offset, cu_q[i], cu_q[i + 1], + ) + _scatter_to_padded( + k, variable_block_sizes_list[i], block_size, + k_packed, kv_offset, cu_kv[i], cu_kv[i + 1], + ) + _scatter_to_padded( + v, variable_block_sizes_list[i], block_size, + v_packed, kv_offset, cu_kv[i], cu_kv[i + 1], + ) + q_offset += padded_q_lens[i] + kv_offset += padded_kv_lens[i] + + total_q_blocks = q_block_offsets[-1] + max_kv_per_q = max(t.shape[-1] for t in q2k_idx_list) + + global_q2k_idx = torch.zeros( + 1, num_heads, total_q_blocks, max_kv_per_q, + dtype=torch.int32, device=device, + ) + global_q2k_num = torch.zeros( + 1, num_heads, total_q_blocks, + dtype=torch.int32, device=device, + ) + global_vbs_parts = [] + + for i in range(num_seqs): + qb_start = q_block_offsets[i] + qb_end = q_block_offsets[i + 1] + n_q_blocks = qb_end - qb_start + kv_offset_blocks = kv_block_offsets[i] + + idx = q2k_idx_list[i] + num = q2k_num_list[i] + vbs = variable_block_sizes_list[i] + + mk = idx.shape[-1] + global_q2k_idx[:, :, qb_start:qb_end, :mk] = idx[:, :, :n_q_blocks, :] + kv_offset_blocks + global_q2k_num[:, :, qb_start:qb_end] = num[:, :, :n_q_blocks] + global_vbs_parts.append(vbs) + + global_vbs = torch.cat(global_vbs_parts, dim=0).to(torch.int32).contiguous() + + out_packed, _ = block_sparse_attn_from_indices( + q_packed, k_packed, v_packed, + global_q2k_idx, global_q2k_num, global_vbs, + ) + + out = torch.zeros(cu_q[-1], num_heads, head_dim, device=device, dtype=dtype) + q_offset = 0 + for i in range(num_seqs): + _gather_from_padded( + out_packed, q_vbs_resolved[i], block_size, + out, q_offset, cu_q[i], cu_q[i + 1], + ) + q_offset += padded_q_lens[i] + + return out diff --git a/fastvideo-kernel/tests/test_vsa_varlen.py b/fastvideo-kernel/tests/test_vsa_varlen.py new file mode 100644 index 0000000000..b2c009d9bc --- /dev/null +++ b/fastvideo-kernel/tests/test_vsa_varlen.py @@ -0,0 +1,230 @@ +"""Correctness tests for variable-length block-sparse attention. + +Reference: per-sequence calls to block_sparse_attn_from_indices. +Test: single-launch via block_sparse_attn_varlen. +""" + +import torch +import pytest + +from .test_vsa import ( + BLOCK_M, + generate_variable_block_sizes, + get_non_pad_index, + vsa_pad, + generate_tensor, +) +from .utils import generate_block_sparse_mask_for_function +from fastvideo_kernel.block_sparse_attn import ( + block_sparse_attn_from_indices, + _map_to_index, +) +from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen + + +def _reference_per_sequence( + q_list, k_list, v_list, + block_masks, vbs_list, + non_pad_q_list, non_pad_kv_list, + q_nblocks_list, kv_nblocks_list, +): + """Run per-sequence block_sparse_attn and concat outputs.""" + outs = [] + for i in range(len(q_list)): + q_pad = vsa_pad(q_list[i], non_pad_q_list[i], q_nblocks_list[i], BLOCK_M) + k_pad = vsa_pad(k_list[i], non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M) + v_pad = vsa_pad(v_list[i], non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M) + + q2k_idx, q2k_num = _map_to_index(block_masks[i].unsqueeze(0)) + o_pad, _ = block_sparse_attn_from_indices( + q_pad, k_pad, v_pad, q2k_idx, q2k_num, vbs_list[i], + ) + o = o_pad[:, :, non_pad_q_list[i], :] + outs.append(o.squeeze(0).transpose(0, 1)) + return torch.cat(outs, dim=0) + + +def _run_varlen_test( + seq_configs: list, + h: int = 8, + d: int = 64, + topk: int = 2, + atol: float = 0.05, + rtol: float = 0.02, +): + """Core test: compare varlen vs per-sequence reference. + + seq_configs: list of (num_q_blocks, num_kv_blocks) per sequence. + """ + device = "cuda" + num_seqs = len(seq_configs) + + q_list = [] + k_list = [] + v_list = [] + block_masks = [] + vbs_list = [] + q_vbs_list = [] + non_pad_q_list = [] + non_pad_kv_list = [] + q_nblocks_list = [] + kv_nblocks_list = [] + q2k_idx_list = [] + q2k_num_list = [] + q_vbs_for_varlen = [] + + cu_q = [0] + cu_kv = [0] + + for nq, nkv in seq_configs: + vbs_kv = generate_variable_block_sizes(nkv, device=device) + vbs_q = generate_variable_block_sizes(nq, device=device) + sq = int(vbs_q.sum().item()) + skv = int(vbs_kv.sum().item()) + + q = generate_tensor((1, h, sq, d), torch.bfloat16, device) + k = generate_tensor((1, h, skv, d), torch.bfloat16, device) + v = generate_tensor((1, h, skv, d), torch.bfloat16, device) + + mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device) + npq = get_non_pad_index(vbs_q, nq, BLOCK_M) + npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M) + + q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0)) + + q_list.append(q) + k_list.append(k) + v_list.append(v) + block_masks.append(mask) + vbs_list.append(vbs_kv) + q_vbs_list.append(vbs_q) + non_pad_q_list.append(npq) + non_pad_kv_list.append(npkv) + q_nblocks_list.append(nq) + kv_nblocks_list.append(nkv) + q2k_idx_list.append(q2k_idx) + q2k_num_list.append(q2k_num) + q_vbs_for_varlen.append(vbs_q) + + cu_q.append(cu_q[-1] + sq) + cu_kv.append(cu_kv[-1] + skv) + + ref_out = _reference_per_sequence( + q_list, k_list, v_list, + block_masks, vbs_list, + non_pad_q_list, non_pad_kv_list, + q_nblocks_list, kv_nblocks_list, + ) + + q_packed = torch.cat( + [qi.squeeze(0).transpose(0, 1) for qi in q_list], dim=0, + ) + k_packed = torch.cat( + [ki.squeeze(0).transpose(0, 1) for ki in k_list], dim=0, + ) + v_packed = torch.cat( + [vi.squeeze(0).transpose(0, 1) for vi in v_list], dim=0, + ) + + cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device) + cu_seqlens_kv = torch.tensor(cu_kv, dtype=torch.int32, device=device) + + varlen_out = block_sparse_attn_varlen( + q_packed, k_packed, v_packed, + cu_seqlens_q, cu_seqlens_kv, + q2k_idx_list, q2k_num_list, + vbs_list, + q_variable_block_sizes_list=q_vbs_for_varlen, + ) + + max_abs = (ref_out - varlen_out).abs().max().item() + mean_abs = ref_out.abs().mean().item() + max_rel = max_abs / (mean_abs + 1e-8) + + print(f" seqs={[c for c in seq_configs]}, max_abs={max_abs:.4e}, max_rel={max_rel:.4e}") + assert max_rel < rtol, f"max relative error {max_rel:.4e} exceeds threshold {rtol}" + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +class TestVSAVarlen: + + def test_equal_length(self): + """Two sequences with same number of blocks.""" + _run_varlen_test([(4, 4), (4, 4)], h=8, d=64) + + def test_different_lengths(self): + """Three sequences with different block counts.""" + _run_varlen_test([(2, 3), (5, 4), (3, 6)], h=8, d=64) + + def test_single_sequence(self): + """Degenerate case: single sequence should match non-varlen path.""" + _run_varlen_test([(8, 8)], h=8, d=64) + + def test_many_heads(self): + """More heads to stress the packing logic.""" + _run_varlen_test([(3, 4), (5, 3)], h=16, d=128) + + def test_many_sequences(self): + """Stress test: 8 sequences with varying block counts.""" + configs = [(i + 2, i + 3) for i in range(8)] + _run_varlen_test(configs, h=8, d=64) + + def test_topk_equals_num_blocks(self): + """Edge: topk covers all KV blocks (dense attention).""" + _run_varlen_test([(3, 3), (4, 4)], h=8, d=64, topk=8) + + def test_single_block_per_sequence(self): + """Minimal: each sequence has exactly 1 Q block and 1 KV block.""" + _run_varlen_test([(1, 1), (1, 1), (1, 1)], h=8, d=64, topk=1) + + def test_asymmetric_q_kv(self): + """Q and KV have very different block counts.""" + _run_varlen_test([(1, 8), (8, 1)], h=8, d=64, topk=1) + + def test_without_q_vbs(self): + """Test the default path where q_variable_block_sizes_list is None. + + Uses full block_size=64 for Q blocks so the None path is valid. + """ + device = "cuda" + h, d, topk = 4, 64, 2 + nq, nkv = 3, 4 + + vbs_kv = generate_variable_block_sizes(nkv, device=device) + sq = nq * BLOCK_M + skv = int(vbs_kv.sum().item()) + + q = generate_tensor((1, h, sq, d), torch.bfloat16, device) + k = generate_tensor((1, h, skv, d), torch.bfloat16, device) + v = generate_tensor((1, h, skv, d), torch.bfloat16, device) + + mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device) + npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M) + q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0)) + + k_pad = vsa_pad(k, npkv, nkv, BLOCK_M) + v_pad = vsa_pad(v, npkv, nkv, BLOCK_M) + ref_out, _ = block_sparse_attn_from_indices(q, k_pad, v_pad, q2k_idx, q2k_num, vbs_kv) + ref_flat = ref_out.squeeze(0).transpose(0, 1) + + q_flat = q.squeeze(0).transpose(0, 1) + k_flat = k.squeeze(0).transpose(0, 1) + v_flat = v.squeeze(0).transpose(0, 1) + cu_q = torch.tensor([0, sq], dtype=torch.int32, device=device) + cu_kv = torch.tensor([0, skv], dtype=torch.int32, device=device) + + varlen_out = block_sparse_attn_varlen( + q_flat, k_flat, v_flat, + cu_q, cu_kv, + [q2k_idx], [q2k_num], [vbs_kv], + ) + + max_abs = (ref_flat - varlen_out).abs().max().item() + mean_abs = ref_flat.abs().mean().item() + max_rel = max_abs / (mean_abs + 1e-8) + print(f" without_q_vbs: max_abs={max_abs:.4e}, max_rel={max_rel:.4e}") + assert max_rel < 0.02, f"max relative error {max_rel:.4e} exceeds threshold" + + +if __name__ == "__main__": + pytest.main([__file__, "-v", "-s"]) From c612890acc476f32d9181c608d8f5ffa1b984923 Mon Sep 17 00:00:00 2001 From: freemty Date: Tue, 12 May 2026 13:40:44 +0800 Subject: [PATCH 2/3] [kernel] Avoid per-iteration CPU-GPU sync in pack/unpack loops - Use block_sizes.cpu().tolist() before looping instead of .item() per block - Create default Q block size tensor on CPU (only used for loop iteration) Addresses review feedback from gemini-code-assist. --- .../fastvideo_kernel/block_sparse_attn_varlen.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py b/fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py index 81101f9b69..60197ec365 100644 --- a/fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py +++ b/fastvideo-kernel/python/fastvideo_kernel/block_sparse_attn_varlen.py @@ -37,8 +37,8 @@ def _scatter_to_padded( """ src_pos = src_start dst_pos = dst_offset - for b in range(block_sizes.numel()): - actual = int(block_sizes[b].item()) + sizes = block_sizes.cpu().tolist() + for actual in sizes: actual = min(actual, src_end - src_pos) if actual > 0: dst[:, :, dst_pos:dst_pos + actual, :] = ( @@ -64,8 +64,8 @@ def _gather_from_padded( """ src_pos = src_offset dst_pos = dst_start - for b in range(block_sizes.numel()): - actual = int(block_sizes[b].item()) + sizes = block_sizes.cpu().tolist() + for actual in sizes: actual = min(actual, dst_end - dst_pos) if actual > 0: dst[dst_pos:dst_pos + actual] = ( @@ -132,7 +132,7 @@ def block_sparse_attn_varlen( q_vbs_resolved.append(q_variable_block_sizes_list[i]) else: q_vbs_resolved.append( - torch.full((n_q_blocks,), block_size, dtype=torch.int32, device=device) + torch.full((n_q_blocks,), block_size, dtype=torch.int32) ) total_padded_q = sum(padded_q_lens) From 1a67359750e2e4dacc3bc96b3ec560dbd69be778 Mon Sep 17 00:00:00 2001 From: freemty Date: Wed, 27 May 2026 17:17:05 +0800 Subject: [PATCH 3/3] [kernel] Add backward correctness tests for block_sparse_attn_varlen Addresses reviewer feedback requesting gradient tests. Adds TestVSAVarlenBackward class with 6 tests that verify dQ/dK/dV from the varlen wrapper match per-sequence reference gradients, confirming the implicit autograd path through scatter/gather slice assignment is correct. --- fastvideo-kernel/tests/test_vsa_varlen.py | 204 ++++++++++++++++++++++ 1 file changed, 204 insertions(+) diff --git a/fastvideo-kernel/tests/test_vsa_varlen.py b/fastvideo-kernel/tests/test_vsa_varlen.py index b2c009d9bc..31231dc294 100644 --- a/fastvideo-kernel/tests/test_vsa_varlen.py +++ b/fastvideo-kernel/tests/test_vsa_varlen.py @@ -2,6 +2,7 @@ Reference: per-sequence calls to block_sparse_attn_from_indices. Test: single-launch via block_sparse_attn_varlen. +Tests cover both forward and backward (gradient) correctness. """ import torch @@ -226,5 +227,208 @@ def test_without_q_vbs(self): assert max_rel < 0.02, f"max relative error {max_rel:.4e} exceeds threshold" +def _run_varlen_backward_test( + seq_configs: list, + h: int = 8, + d: int = 64, + topk: int = 2, + grad_rtol: float = 0.05, +): + """Backward correctness: compare dQ/dK/dV from varlen vs per-sequence reference. + + Both paths use the same underlying block_sparse_attn_from_indices kernel + (which has registered autograd). The varlen wrapper's scatter/gather must + correctly propagate gradients through PyTorch's in-place slice assignment. + """ + device = "cuda" + num_seqs = len(seq_configs) + + q_list = [] + k_list = [] + v_list = [] + block_masks = [] + vbs_list = [] + q_vbs_list = [] + non_pad_q_list = [] + non_pad_kv_list = [] + q_nblocks_list = [] + kv_nblocks_list = [] + q2k_idx_list = [] + q2k_num_list = [] + + cu_q = [0] + cu_kv = [0] + + for nq, nkv in seq_configs: + vbs_kv = generate_variable_block_sizes(nkv, device=device) + vbs_q = generate_variable_block_sizes(nq, device=device) + sq = int(vbs_q.sum().item()) + skv = int(vbs_kv.sum().item()) + + q = generate_tensor((1, h, sq, d), torch.bfloat16, device) + k = generate_tensor((1, h, skv, d), torch.bfloat16, device) + v = generate_tensor((1, h, skv, d), torch.bfloat16, device) + + mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device) + npq = get_non_pad_index(vbs_q, nq, BLOCK_M) + npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M) + + q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0)) + + q_list.append(q) + k_list.append(k) + v_list.append(v) + block_masks.append(mask) + vbs_list.append(vbs_kv) + q_vbs_list.append(vbs_q) + non_pad_q_list.append(npq) + non_pad_kv_list.append(npkv) + q_nblocks_list.append(nq) + kv_nblocks_list.append(nkv) + q2k_idx_list.append(q2k_idx) + q2k_num_list.append(q2k_num) + + cu_q.append(cu_q[-1] + sq) + cu_kv.append(cu_kv[-1] + skv) + + # --- Reference: per-sequence backward --- + ref_q_grads = [] + ref_k_grads = [] + ref_v_grads = [] + ref_outs = [] + for i in range(num_seqs): + qi = q_list[i].detach().requires_grad_(True) + ki = k_list[i].detach().requires_grad_(True) + vi = v_list[i].detach().requires_grad_(True) + + q_pad = vsa_pad(qi, non_pad_q_list[i], q_nblocks_list[i], BLOCK_M) + k_pad = vsa_pad(ki, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M) + v_pad = vsa_pad(vi, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M) + + q2k_idx, q2k_num = _map_to_index(block_masks[i].unsqueeze(0)) + o_pad, _ = block_sparse_attn_from_indices( + q_pad, k_pad, v_pad, q2k_idx, q2k_num, vbs_list[i], + ) + o = o_pad[:, :, non_pad_q_list[i], :] + o_flat = o.squeeze(0).transpose(0, 1) + ref_outs.append(o_flat) + + dO = torch.ones_like(o_flat) + o_flat.backward(dO) + + ref_q_grads.append(qi.grad.squeeze(0).transpose(0, 1)) + ref_k_grads.append(ki.grad.squeeze(0).transpose(0, 1)) + ref_v_grads.append(vi.grad.squeeze(0).transpose(0, 1)) + + ref_dq = torch.cat(ref_q_grads, dim=0) + ref_dk = torch.cat(ref_k_grads, dim=0) + ref_dv = torch.cat(ref_v_grads, dim=0) + + # --- Varlen backward --- + q_packed = torch.cat( + [qi.squeeze(0).transpose(0, 1) for qi in q_list], dim=0, + ).detach().requires_grad_(True) + k_packed = torch.cat( + [ki.squeeze(0).transpose(0, 1) for ki in k_list], dim=0, + ).detach().requires_grad_(True) + v_packed = torch.cat( + [vi.squeeze(0).transpose(0, 1) for vi in v_list], dim=0, + ).detach().requires_grad_(True) + + cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device) + cu_seqlens_kv = torch.tensor(cu_kv, dtype=torch.int32, device=device) + + varlen_out = block_sparse_attn_varlen( + q_packed, k_packed, v_packed, + cu_seqlens_q, cu_seqlens_kv, + q2k_idx_list, q2k_num_list, + vbs_list, + q_variable_block_sizes_list=q_vbs_list, + ) + + dO = torch.ones_like(varlen_out) + varlen_out.backward(dO) + + varlen_dq = q_packed.grad + varlen_dk = k_packed.grad + varlen_dv = v_packed.grad + + for name, ref, actual in [ + ("dQ", ref_dq, varlen_dq), + ("dK", ref_dk, varlen_dk), + ("dV", ref_dv, varlen_dv), + ]: + assert actual is not None, f"{name}: gradient is None (autograd chain broken)" + max_abs = (ref - actual).abs().max().item() + mean_abs = ref.abs().mean().item() + max_rel = max_abs / (mean_abs + 1e-8) + print(f" {name}: max_abs={max_abs:.4e}, max_rel={max_rel:.4e}") + assert max_rel < grad_rtol, ( + f"{name}: max relative error {max_rel:.4e} exceeds threshold {grad_rtol}" + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +class TestVSAVarlenBackward: + + def test_backward_equal_length(self): + """Backward: two sequences with same number of blocks.""" + _run_varlen_backward_test([(4, 4), (4, 4)], h=8, d=64) + + def test_backward_different_lengths(self): + """Backward: three sequences with different block counts.""" + _run_varlen_backward_test([(2, 3), (5, 4), (3, 6)], h=8, d=64) + + def test_backward_single_sequence(self): + """Backward: single sequence should match non-varlen gradient path.""" + _run_varlen_backward_test([(8, 8)], h=8, d=64) + + def test_backward_many_heads(self): + """Backward: more heads to stress gradient routing.""" + _run_varlen_backward_test([(3, 4), (5, 3)], h=16, d=128) + + def test_backward_asymmetric_q_kv(self): + """Backward: Q and KV have very different block counts.""" + _run_varlen_backward_test([(1, 8), (8, 1)], h=8, d=64, topk=1) + + def test_backward_grad_nonzero(self): + """Smoke test: gradients are non-zero (autograd chain is connected).""" + device = "cuda" + h, d, topk = 4, 64, 2 + nq, nkv = 3, 4 + + vbs_kv = generate_variable_block_sizes(nkv, device=device) + vbs_q = generate_variable_block_sizes(nq, device=device) + sq = int(vbs_q.sum().item()) + skv = int(vbs_kv.sum().item()) + + q = torch.randn(sq, h, d, device=device, dtype=torch.bfloat16, requires_grad=True) + k = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True) + v = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True) + + mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device) + q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0)) + + cu_q = torch.tensor([0, sq], dtype=torch.int32, device=device) + cu_kv = torch.tensor([0, skv], dtype=torch.int32, device=device) + + out = block_sparse_attn_varlen( + q, k, v, + cu_q, cu_kv, + [q2k_idx], [q2k_num], [vbs_kv], + q_variable_block_sizes_list=[vbs_q], + ) + + loss = out.sum() + loss.backward() + + assert q.grad is not None, "q.grad is None" + assert k.grad is not None, "k.grad is None" + assert v.grad is not None, "v.grad is None" + assert q.grad.abs().sum().item() > 0, "q.grad is all zeros" + assert k.grad.abs().sum().item() > 0, "k.grad is all zeros" + assert v.grad.abs().sum().item() > 0, "v.grad is all zeros" + + if __name__ == "__main__": pytest.main([__file__, "-v", "-s"])