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:
+11
-5
@@ -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,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user