[Bug] [FA4] [hdim 256] CUDA error 700 in backward pass when using zero-length q sequences with hdim=256 in varlen mode
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 条评论