def sparse_sm89_attention(q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor,
vbs: torch.Tensor,
*,
int8_qk: bool = True) -> torch.Tensor:
"""Forward-only ``[B,H,S,128]`` BF16 attention on sm89, with 64-token tiles."""
if torch.is_grad_enabled():
raise ValueError("Sparse INT8/FP8 attention is inference-only")
if not q.is_cuda or torch.cuda.get_device_capability(q.device) != (8, 9):
raise ValueError("Sparse INT8/FP8 attention requires sm89 CUDA")
if q.dtype != torch.bfloat16 or q.shape[-1] != 128 or q.shape != k.shape or q.shape != v.shape:
raise ValueError("Sparse INT8/FP8 attention requires matching BF16 Q/K/V with head dimension 128")
b, h, length, dim = q.shape
if length != vbs.numel() * 64 or mask.shape != (b, h, length // 64, length // 64):
raise ValueError("Sparse INT8/FP8 attention requires a tile-64 mask and validity vector")
from fastvideo_kernel.triton_kernels.index import map_to_index
# The production INT8-QK/BF16-PV route reads BSHD-backed views directly.
# Quantized Q/K and the output remain contiguous BHSD. Other ablations
# retain their established layout and arithmetic.
if not int8_qk:
q, k, v = q.contiguous(), k.contiguous(), v.contiguous()
vbs = vbs.to(device=q.device, dtype=torch.int32).contiguous()
grid = (triton.cdiv(length, 16), b * h)
qi, ki = q, k
qs, ks = q, k # unused pointers in the BF16-QK ablation
if int8_qk:
# Tile pads are zero by contract; avoid a full FP32 copy for the reduction.
# Preserve the exact reduction used by the old contiguous adapter;
# its temporary copy dies before Q/K quantization and fine attention.
mean = k.contiguous().sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1)
qi = torch.empty(q.shape, device=q.device, dtype=torch.int8)
ki = torch.empty(k.shape, device=k.device, dtype=torch.int8)
qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32)
ks = torch.empty_like(qs)
_quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, h, *q.stride(), CENTER=False, ROWS=16, num_warps=4)
_quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, h, *k.stride(), CENTER=True, ROWS=16, num_warps=4)
index, count = map_to_index(mask.contiguous())
out = torch.empty(q.shape, device=q.device, dtype=q.dtype)
_sparse_int8_fp8[(length // 64, b * h)](qi,
ki,
v,
qs,
ks,
index,
count,
vbs,
out,
length,
dim,
h,
*v.stride(),
INT8_QK=int8_qk)
return out