Skip to content

Commit 1a67359

Browse files
committed
[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.
1 parent c612890 commit 1a67359

1 file changed

Lines changed: 204 additions & 0 deletions

File tree

fastvideo-kernel/tests/test_vsa_varlen.py

Lines changed: 204 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
33
Reference: per-sequence calls to block_sparse_attn_from_indices.
44
Test: single-launch via block_sparse_attn_varlen.
5+
Tests cover both forward and backward (gradient) correctness.
56
"""
67

78
import torch
@@ -226,5 +227,208 @@ def test_without_q_vbs(self):
226227
assert max_rel < 0.02, f"max relative error {max_rel:.4e} exceeds threshold"
227228

228229

230+
def _run_varlen_backward_test(
231+
seq_configs: list,
232+
h: int = 8,
233+
d: int = 64,
234+
topk: int = 2,
235+
grad_rtol: float = 0.05,
236+
):
237+
"""Backward correctness: compare dQ/dK/dV from varlen vs per-sequence reference.
238+
239+
Both paths use the same underlying block_sparse_attn_from_indices kernel
240+
(which has registered autograd). The varlen wrapper's scatter/gather must
241+
correctly propagate gradients through PyTorch's in-place slice assignment.
242+
"""
243+
device = "cuda"
244+
num_seqs = len(seq_configs)
245+
246+
q_list = []
247+
k_list = []
248+
v_list = []
249+
block_masks = []
250+
vbs_list = []
251+
q_vbs_list = []
252+
non_pad_q_list = []
253+
non_pad_kv_list = []
254+
q_nblocks_list = []
255+
kv_nblocks_list = []
256+
q2k_idx_list = []
257+
q2k_num_list = []
258+
259+
cu_q = [0]
260+
cu_kv = [0]
261+
262+
for nq, nkv in seq_configs:
263+
vbs_kv = generate_variable_block_sizes(nkv, device=device)
264+
vbs_q = generate_variable_block_sizes(nq, device=device)
265+
sq = int(vbs_q.sum().item())
266+
skv = int(vbs_kv.sum().item())
267+
268+
q = generate_tensor((1, h, sq, d), torch.bfloat16, device)
269+
k = generate_tensor((1, h, skv, d), torch.bfloat16, device)
270+
v = generate_tensor((1, h, skv, d), torch.bfloat16, device)
271+
272+
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
273+
npq = get_non_pad_index(vbs_q, nq, BLOCK_M)
274+
npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M)
275+
276+
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
277+
278+
q_list.append(q)
279+
k_list.append(k)
280+
v_list.append(v)
281+
block_masks.append(mask)
282+
vbs_list.append(vbs_kv)
283+
q_vbs_list.append(vbs_q)
284+
non_pad_q_list.append(npq)
285+
non_pad_kv_list.append(npkv)
286+
q_nblocks_list.append(nq)
287+
kv_nblocks_list.append(nkv)
288+
q2k_idx_list.append(q2k_idx)
289+
q2k_num_list.append(q2k_num)
290+
291+
cu_q.append(cu_q[-1] + sq)
292+
cu_kv.append(cu_kv[-1] + skv)
293+
294+
# --- Reference: per-sequence backward ---
295+
ref_q_grads = []
296+
ref_k_grads = []
297+
ref_v_grads = []
298+
ref_outs = []
299+
for i in range(num_seqs):
300+
qi = q_list[i].detach().requires_grad_(True)
301+
ki = k_list[i].detach().requires_grad_(True)
302+
vi = v_list[i].detach().requires_grad_(True)
303+
304+
q_pad = vsa_pad(qi, non_pad_q_list[i], q_nblocks_list[i], BLOCK_M)
305+
k_pad = vsa_pad(ki, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
306+
v_pad = vsa_pad(vi, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
307+
308+
q2k_idx, q2k_num = _map_to_index(block_masks[i].unsqueeze(0))
309+
o_pad, _ = block_sparse_attn_from_indices(
310+
q_pad, k_pad, v_pad, q2k_idx, q2k_num, vbs_list[i],
311+
)
312+
o = o_pad[:, :, non_pad_q_list[i], :]
313+
o_flat = o.squeeze(0).transpose(0, 1)
314+
ref_outs.append(o_flat)
315+
316+
dO = torch.ones_like(o_flat)
317+
o_flat.backward(dO)
318+
319+
ref_q_grads.append(qi.grad.squeeze(0).transpose(0, 1))
320+
ref_k_grads.append(ki.grad.squeeze(0).transpose(0, 1))
321+
ref_v_grads.append(vi.grad.squeeze(0).transpose(0, 1))
322+
323+
ref_dq = torch.cat(ref_q_grads, dim=0)
324+
ref_dk = torch.cat(ref_k_grads, dim=0)
325+
ref_dv = torch.cat(ref_v_grads, dim=0)
326+
327+
# --- Varlen backward ---
328+
q_packed = torch.cat(
329+
[qi.squeeze(0).transpose(0, 1) for qi in q_list], dim=0,
330+
).detach().requires_grad_(True)
331+
k_packed = torch.cat(
332+
[ki.squeeze(0).transpose(0, 1) for ki in k_list], dim=0,
333+
).detach().requires_grad_(True)
334+
v_packed = torch.cat(
335+
[vi.squeeze(0).transpose(0, 1) for vi in v_list], dim=0,
336+
).detach().requires_grad_(True)
337+
338+
cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device)
339+
cu_seqlens_kv = torch.tensor(cu_kv, dtype=torch.int32, device=device)
340+
341+
varlen_out = block_sparse_attn_varlen(
342+
q_packed, k_packed, v_packed,
343+
cu_seqlens_q, cu_seqlens_kv,
344+
q2k_idx_list, q2k_num_list,
345+
vbs_list,
346+
q_variable_block_sizes_list=q_vbs_list,
347+
)
348+
349+
dO = torch.ones_like(varlen_out)
350+
varlen_out.backward(dO)
351+
352+
varlen_dq = q_packed.grad
353+
varlen_dk = k_packed.grad
354+
varlen_dv = v_packed.grad
355+
356+
for name, ref, actual in [
357+
("dQ", ref_dq, varlen_dq),
358+
("dK", ref_dk, varlen_dk),
359+
("dV", ref_dv, varlen_dv),
360+
]:
361+
assert actual is not None, f"{name}: gradient is None (autograd chain broken)"
362+
max_abs = (ref - actual).abs().max().item()
363+
mean_abs = ref.abs().mean().item()
364+
max_rel = max_abs / (mean_abs + 1e-8)
365+
print(f" {name}: max_abs={max_abs:.4e}, max_rel={max_rel:.4e}")
366+
assert max_rel < grad_rtol, (
367+
f"{name}: max relative error {max_rel:.4e} exceeds threshold {grad_rtol}"
368+
)
369+
370+
371+
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
372+
class TestVSAVarlenBackward:
373+
374+
def test_backward_equal_length(self):
375+
"""Backward: two sequences with same number of blocks."""
376+
_run_varlen_backward_test([(4, 4), (4, 4)], h=8, d=64)
377+
378+
def test_backward_different_lengths(self):
379+
"""Backward: three sequences with different block counts."""
380+
_run_varlen_backward_test([(2, 3), (5, 4), (3, 6)], h=8, d=64)
381+
382+
def test_backward_single_sequence(self):
383+
"""Backward: single sequence should match non-varlen gradient path."""
384+
_run_varlen_backward_test([(8, 8)], h=8, d=64)
385+
386+
def test_backward_many_heads(self):
387+
"""Backward: more heads to stress gradient routing."""
388+
_run_varlen_backward_test([(3, 4), (5, 3)], h=16, d=128)
389+
390+
def test_backward_asymmetric_q_kv(self):
391+
"""Backward: Q and KV have very different block counts."""
392+
_run_varlen_backward_test([(1, 8), (8, 1)], h=8, d=64, topk=1)
393+
394+
def test_backward_grad_nonzero(self):
395+
"""Smoke test: gradients are non-zero (autograd chain is connected)."""
396+
device = "cuda"
397+
h, d, topk = 4, 64, 2
398+
nq, nkv = 3, 4
399+
400+
vbs_kv = generate_variable_block_sizes(nkv, device=device)
401+
vbs_q = generate_variable_block_sizes(nq, device=device)
402+
sq = int(vbs_q.sum().item())
403+
skv = int(vbs_kv.sum().item())
404+
405+
q = torch.randn(sq, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
406+
k = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
407+
v = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
408+
409+
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
410+
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
411+
412+
cu_q = torch.tensor([0, sq], dtype=torch.int32, device=device)
413+
cu_kv = torch.tensor([0, skv], dtype=torch.int32, device=device)
414+
415+
out = block_sparse_attn_varlen(
416+
q, k, v,
417+
cu_q, cu_kv,
418+
[q2k_idx], [q2k_num], [vbs_kv],
419+
q_variable_block_sizes_list=[vbs_q],
420+
)
421+
422+
loss = out.sum()
423+
loss.backward()
424+
425+
assert q.grad is not None, "q.grad is None"
426+
assert k.grad is not None, "k.grad is None"
427+
assert v.grad is not None, "v.grad is None"
428+
assert q.grad.abs().sum().item() > 0, "q.grad is all zeros"
429+
assert k.grad.abs().sum().item() > 0, "k.grad is all zeros"
430+
assert v.grad.abs().sum().item() > 0, "v.grad is all zeros"
431+
432+
229433
if __name__ == "__main__":
230434
pytest.main([__file__, "-v", "-s"])

0 commit comments

Comments
 (0)