ITADN

[Bug] [FA4] [hdim 256] CUDA error 700 in backward pass when using zero-length q sequences with hdim=256 in varlen mode

#2562Openumiswing 创建于 2026-05-13
U
umiswingcommented
Hi! I'm trying to do some padding in varlen with the trick that setting corresponding seqlen_q to 0, e.g, setting cu_seqlens_q=[0, **0**, 1208, 2048, **2048**], cu_seqlens_k=[0, 2538, 4280, 5120, 8192]. Such packing can pass with small hdim like 128, but with hdim 256, which will dispatch to a dedicated kernel in fa4, will fail and raise CUDA error 700 in backward. Any way to fix it? I'm using `flash-attn-4 4.0.0b13`, and run the following test with `PYTORCH_NO_CUDA_MEMORY_CACHING=1` can reproduce the CUDA 700 error. @wangsiyu @Johnsonms ```python import math import torch from flash_attn.cute import flash_attn_varlen_func def main(): dtype = torch.bfloat16 device = "cuda" total_q = 2048 total_k = 8192 nheads = 4 nheads_kv = 1 d = 256 torch.manual_seed(42) q = torch.randn(total_q, nheads, d, dtype=dtype, device=device) k = torch.randn(total_k, nheads_kv, d, dtype=dtype, device=device) v = torch.randn(total_k, nheads_kv, d, dtype=dtype, device=device) cu_seqlens_q = torch.tensor([0, 0, 1208, 2048, 2048], dtype=torch.int32, device=device) cu_seqlens_k = torch.tensor([0, 2538, 4280, 5120, 8192], dtype=torch.int32, device=device) max_seqlen_q = 1208 max_seqlen_k = 3072 softmax_scale = 1.0 / math.sqrt(d) print("Running FA4 varlen with:") print(f" q: ({total_q}, {nheads}, {d})") print(f" k: ({total_k}, {nheads_kv}, {d})") print(f" v: ({total_k}, {nheads_kv}, {d})") print(f" cu_seqlens_q: {cu_seqlens_q.tolist()}") print(f" cu_seqlens_k: {cu_seqlens_k.tolist()}") print(f" max_seqlen_q: {max_seqlen_q}, max_seqlen_k: {max_seqlen_k}") print(f" causal: True") print() # Enable grad for backward q.requires_grad_(True) k.requires_grad_(True) v.requires_grad_(True) try: out, _ = flash_attn_varlen_func( q=q, k=k, v=v, cu_seqlens_q=cu_seqlens_q, cu_seqlens_k=cu_seqlens_k, max_seqlen_q=max_seqlen_q, max_seqlen_k=max_seqlen_k, softmax_scale=softmax_scale, causal=True, ) torch.cuda.synchronize() has_nan = torch.isnan(out).any().item() if has_nan: print("FAIL (forward): NaN detected in output") else: print("PASS (forward): No CUDA error, no NaN") print(f" out shape: {out.shape}, max abs: {out.abs().max().item():.4f}") except RuntimeError as e: if "illegal memory access" in str(e) or "error 700" in str(e).lower(): print(f"FAIL (forward): CUDA error 700 (illegal memory access)") print(f" {e}") else: print(f"FAIL (forward): RuntimeError: {e}") return # Backward pass print("\nRunning backward pass...") try: g = torch.randn_like(out) out.backward(g) torch.cuda.synchronize() dq_nan = torch.isnan(q.grad).any().item() dk_nan = torch.isnan(k.grad).any().item() dv_nan = torch.isnan(v.grad).any().item() if dq_nan or dk_nan or dv_nan: print(f"FAIL (backward): NaN in gradients - dQ:{dq_nan}, dK:{dk_nan}, dV:{dv_nan}") else: print("PASS (backward): No CUDA error, no NaN in gradients") print(f" dQ max abs: {q.grad.abs().max().item():.4f}") print(f" dK max abs: {k.grad.abs().max().item():.4f}") print(f" dV max abs: {v.grad.abs().max().item():.4f}") except RuntimeError as e: if "illegal memory access" in str(e) or "error 700" in str(e).lower(): print(f"FAIL (backward): CUDA error 700 (illegal memory access)") print(f" {e}") else: print(f"FAIL (backward): RuntimeError: {e}") if __name__ == "__main__": main() ```
5 条评论