feat: tensor-core mma prefill with build-time dispatch

- add register-resident flash-attention kernel using mma.sync.m16n8k16
- dispatch mma vs scalar at build time: pre-sm_80 defines
  -DASTRAI_NO_MMA, else defaults to mma
- scalar path vectorized with float4 smem loads (ld8)
This commit is contained in:
2026-07-06 20:33:24 +08:00
parent cc36530c73
commit ddc4bd1cf6
5 changed files with 291 additions and 21 deletions
+11 -5
View File
@@ -1,14 +1,20 @@
from pathlib import Path
def _arch_flag():
def _arch_flags() -> list[str]:
import torch
if torch.cuda.is_available():
cap = torch.cuda.get_device_capability()
ver = f"{cap[0]}{cap[1]}"
return f"-gencode=arch=compute_{ver},code=sm_{ver}"
return "-gencode=arch=compute_80,code=sm_80"
else:
cap = (8, 0)
ver = f"{cap[0]}{cap[1]}"
flags = [f"-gencode=arch=compute_{ver},code=sm_{ver}"]
# tensor-core mma path (mma.sync.m16n8k16.bf16) requires sm_80+; decide the
# kernel dispatch at build time via this define rather than at runtime.
if cap[0] < 8:
flags.append("-DASTRAI_NO_MMA")
return flags
_kernels_dir = Path("csrc/kernels")
@@ -20,7 +26,7 @@ def register(name: str, sources: list[str] | None = None, **kwargs):
sources = [str(_kernels_dir / f"{name}.cu")]
REGISTRY[name] = {
"sources": sources,
"nvcc_flags": ["-O3", "--expt-relaxed-constexpr", _arch_flag()],
"nvcc_flags": ["-O3", "--expt-relaxed-constexpr", *_arch_flags()],
"extra_link_args": kwargs.pop("extra_link_args", []),
**kwargs,
}