Ran into an issue while testing long sequences using FP16. Everything works great up to around 32k tokens, but as soon as I push past seq_len = 32769, the forward pass starts dumping NaNs in the attention output.
to reproduce:
import torch
from flash_qla import flash_qla_forward
Anything > 32768 causes issues in FP16
q = torch.randn(1, 32769, 16, 128, dtype=torch.float16, device="cuda")
k = torch.randn(1, 32769, 16, 128, dtype=torch.float16, device="cuda")
v = torch.randn(1, 32769, 16, 128, dtype=torch.float16, device="cuda")
out = flash_qla_forward(q, k, v)
print(torch.isnan(out).any()) # True
Switching over to bfloat16 completely resolves it, so it looks like a numerical overflow in the Triton kernel accumulator when processing long sequence blocks.
Tested on an H100 with CUDA 12.4 / PyTorch 2.4.
Ran into an issue while testing long sequences using FP16. Everything works great up to around 32k tokens, but as soon as I push past seq_len = 32769, the forward pass starts dumping NaNs in the attention output.
to reproduce:
import torch
from flash_qla import flash_qla_forward
Anything > 32768 causes issues in FP16
q = torch.randn(1, 32769, 16, 128, dtype=torch.float16, device="cuda")
k = torch.randn(1, 32769, 16, 128, dtype=torch.float16, device="cuda")
v = torch.randn(1, 32769, 16, 128, dtype=torch.float16, device="cuda")
out = flash_qla_forward(q, k, v)
print(torch.isnan(out).any()) # True
Switching over to bfloat16 completely resolves it, so it looks like a numerical overflow in the Triton kernel accumulator when processing long sequence blocks.
Tested on an H100 with CUDA 12.4 / PyTorch 2.4.