Spaces:
Running
Running
Commit ·
61b358c
1
Parent(s): 4b04e69
Update kernel.py
Browse files
kernel.py
CHANGED
|
@@ -218,6 +218,133 @@ def _turing_flash_attention(q, k, v, causal, softmax_scale):
|
|
| 218 |
return out.transpose(1, 2) # back to (B,H,T,D) for this codebase's convention
|
| 219 |
|
| 220 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 221 |
# ---------------------------------------------------------------------------
|
| 222 |
# FlashAttention-2 (NVIDIA Ampere+)
|
| 223 |
# ---------------------------------------------------------------------------
|
|
@@ -512,6 +639,46 @@ def _build_attention_fn(device: torch.device):
|
|
| 512 |
Called once; result is cached in _ATTENTION_FN."""
|
| 513 |
backend = _backend_for(device)
|
| 514 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 515 |
if (
|
| 516 |
device.type == "cuda"
|
| 517 |
and _turing_flash_attn_func is not None
|
|
@@ -633,6 +800,26 @@ def print_attention_backend(device: "torch.device | str | None" = None) -> None:
|
|
| 633 |
)
|
| 634 |
lines.append(f" turing-flash not eligible ({reason})")
|
| 635 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 636 |
# flash_attn 2.x
|
| 637 |
if _flash_attn_available():
|
| 638 |
lines.append(f" flash_attn available (Ampere+ path, backend={backend!r})")
|
|
@@ -646,7 +833,13 @@ def print_attention_backend(device: "torch.device | str | None" = None) -> None:
|
|
| 646 |
lines.append(" xformers not installed")
|
| 647 |
|
| 648 |
# Summarise what will actually fire first.
|
| 649 |
-
if
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 650 |
active = "turing-flash (sm_75 + flash_attention_interface)"
|
| 651 |
elif backend == "flash":
|
| 652 |
active = "flash_attn 2.x"
|
|
|
|
| 218 |
return out.transpose(1, 2) # back to (B,H,T,D) for this codebase's convention
|
| 219 |
|
| 220 |
|
| 221 |
+
# ---------------------------------------------------------------------------
|
| 222 |
+
# FlexAttention (torch.nn.attention.flex_attention). Config-gated, not
|
| 223 |
+
# hardware-gated -- set use_flexattention=True in config.py to opt in.
|
| 224 |
+
# Unlike every other backend in this file, FlexAttention is a torch-level
|
| 225 |
+
# composition of ops, not a hand-fused CUDA kernel -- its speed comes
|
| 226 |
+
# entirely from torch.compile specializing the mask/score-mod into one
|
| 227 |
+
# fused kernel at trace time. Eager flex_attention runs at roughly
|
| 228 |
+
# reference-attention speed, so (unlike the @torch.compiler.disable
|
| 229 |
+
# backends elsewhere in this file) this one is deliberately compiled, once,
|
| 230 |
+
# and cached.
|
| 231 |
+
# ---------------------------------------------------------------------------
|
| 232 |
+
|
| 233 |
+
@functools.lru_cache(maxsize=1)
|
| 234 |
+
def _flex_attention_available() -> bool:
|
| 235 |
+
try:
|
| 236 |
+
from torch.nn.attention.flex_attention import flex_attention # noqa: F401
|
| 237 |
+
return True
|
| 238 |
+
except ImportError:
|
| 239 |
+
return False
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
_compiled_flex_attention = None
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
def _get_compiled_flex_attention():
|
| 246 |
+
global _compiled_flex_attention
|
| 247 |
+
if _compiled_flex_attention is None:
|
| 248 |
+
from torch.nn.attention.flex_attention import flex_attention
|
| 249 |
+
_compiled_flex_attention = torch.compile(flex_attention, dynamic=False)
|
| 250 |
+
return _compiled_flex_attention
|
| 251 |
+
|
| 252 |
+
@torch.compiler.disable
|
| 253 |
+
@functools.lru_cache(maxsize=32)
|
| 254 |
+
|
| 255 |
+
def _flex_block_mask(q_len: int, kv_len: int, causal: bool, window_size: "int | None", device_key: str):
|
| 256 |
+
"""Build (and cache) a FlexAttention BlockMask for one (shape, mask-kind,
|
| 257 |
+
device) combo. Constructing a block mask isn't free, and with a fixed
|
| 258 |
+
training CONTEXT this pays the cost once instead of every forward call."""
|
| 259 |
+
from torch.nn.attention.flex_attention import create_block_mask
|
| 260 |
+
|
| 261 |
+
if window_size is not None:
|
| 262 |
+
def mask_mod(b, h, q_idx, kv_idx):
|
| 263 |
+
return (kv_idx <= q_idx) & (q_idx - kv_idx < window_size)
|
| 264 |
+
elif causal:
|
| 265 |
+
def mask_mod(b, h, q_idx, kv_idx):
|
| 266 |
+
return kv_idx <= q_idx
|
| 267 |
+
else:
|
| 268 |
+
return None # full bidirectional attention -- no block mask needed
|
| 269 |
+
|
| 270 |
+
return create_block_mask(mask_mod, B=None, H=None, Q_LEN=q_len, KV_LEN=kv_len, device=device_key)
|
| 271 |
+
|
| 272 |
+
|
| 273 |
+
def _flex_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale):
|
| 274 |
+
|
| 275 |
+
if dropout_p > 0 and training:
|
| 276 |
+
raise _UnsupportedByBackend(
|
| 277 |
+
"FlexAttention path here doesn't wire up an attention-weight "
|
| 278 |
+
"dropout term -- falling back."
|
| 279 |
+
)
|
| 280 |
+
|
| 281 |
+
b, hq, t, d = q.shape
|
| 282 |
+
_, hkv, tk, _ = k.shape
|
| 283 |
+
|
| 284 |
+
block_mask = _flex_block_mask(t, tk, causal, window_size, str(q.device))
|
| 285 |
+
scale = softmax_scale if softmax_scale is not None else (1.0 / math.sqrt(d))
|
| 286 |
+
flex_fn = _get_compiled_flex_attention()
|
| 287 |
+
|
| 288 |
+
# FlexAttention takes (B, H, T, D) natively -- no transpose needed, unlike
|
| 289 |
+
# the flash_attn/xformers paths above which want (B, T, H, D).
|
| 290 |
+
return flex_fn(q, k, v, block_mask=block_mask, scale=scale, enable_gqa=(hkv != hq))
|
| 291 |
+
|
| 292 |
+
|
| 293 |
+
# ---------------------------------------------------------------------------
|
| 294 |
+
# Blackwell: FlashAttention-4 (CuTeDSL) with FlashAttention-2 fallback.
|
| 295 |
+
#
|
| 296 |
+
# FA4 (`pip install flash-attn-4`, imported as `flash_attn.cute`) targets
|
| 297 |
+
# Hopper and data-center Blackwell (sm_90 / sm_100 / sm_103 -- H100/B200/B300)
|
| 298 |
+
# via warp-specialized MMA instructions. It does NOT run on desktop Blackwell
|
| 299 |
+
# (sm_120, e.g. RTX 50-series): sm_120 uses the same register-to-register
|
| 300 |
+
# HMMA path NVIDIA GPUs have used since Volta, not the warp-specialized MMA
|
| 301 |
+
# FA4's kernel design requires, and this is a physical silicon difference
|
| 302 |
+
# (confirmed by multiple independent sm_120 FA4 build attempts as of early
|
| 303 |
+
# 2026), not a version-gating issue that a newer release fixes. So this
|
| 304 |
+
# block tries FA4 only on the capabilities it can actually run on, and
|
| 305 |
+
# every other Blackwell-family GPU (sm_120 included) falls through to the
|
| 306 |
+
# existing FlashAttention-2 path below.
|
| 307 |
+
# ---------------------------------------------------------------------------
|
| 308 |
+
|
| 309 |
+
_BLACKWELL_CAPABILITIES = {(9, 0), (10, 0), (10, 3), (12, 0)} # Hopper + all Blackwell variants
|
| 310 |
+
_FA4_CAPABILITIES = {(9, 0), (10, 0), (10, 3)} # NOT (12, 0) -- see note above
|
| 311 |
+
|
| 312 |
+
try:
|
| 313 |
+
from flash_attn.cute import flash_attn_func as _fa4_attn_func
|
| 314 |
+
except ImportError:
|
| 315 |
+
_fa4_attn_func = None
|
| 316 |
+
|
| 317 |
+
|
| 318 |
+
def _fa4_eligible(q: torch.Tensor, causal: bool, window_size: "int | None", dropout_p: float) -> bool:
|
| 319 |
+
if _fa4_attn_func is None:
|
| 320 |
+
return False
|
| 321 |
+
if q.device.type != "cuda":
|
| 322 |
+
return False
|
| 323 |
+
if torch.cuda.get_device_capability(q.device) not in _FA4_CAPABILITIES:
|
| 324 |
+
return False
|
| 325 |
+
if window_size is not None: # not confirmed supported by FA4's current (beta) public API
|
| 326 |
+
return False
|
| 327 |
+
if dropout_p > 0.0:
|
| 328 |
+
return False
|
| 329 |
+
return True
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
def _fa4_attention(q, k, v, causal, softmax_scale):
|
| 333 |
+
# FA4's public surface is still small/beta -- only `causal` is confirmed
|
| 334 |
+
# from the library's own usage example (flash_attn.cute docs show only
|
| 335 |
+
# `flash_attn_func(q, k, v, causal=True)`). Anything beyond that
|
| 336 |
+
# (window_size, dropout) is gated out by _fa4_eligible above rather than
|
| 337 |
+
# guessed at here. Assumes the same (B, T, H, D) layout as FA2/FA3; if
|
| 338 |
+
# that assumption is wrong for a given release the call raises and the
|
| 339 |
+
# dispatcher below falls back to FlashAttention-2 automatically.
|
| 340 |
+
qf = q.transpose(1, 2).contiguous()
|
| 341 |
+
kf = k.transpose(1, 2).contiguous()
|
| 342 |
+
vf = v.transpose(1, 2).contiguous()
|
| 343 |
+
kwargs = {"softmax_scale": softmax_scale} if softmax_scale is not None else {}
|
| 344 |
+
out = _fa4_attn_func(qf, kf, vf, causal=causal, **kwargs)
|
| 345 |
+
return out.transpose(1, 2)
|
| 346 |
+
|
| 347 |
+
|
| 348 |
# ---------------------------------------------------------------------------
|
| 349 |
# FlashAttention-2 (NVIDIA Ampere+)
|
| 350 |
# ---------------------------------------------------------------------------
|
|
|
|
| 639 |
Called once; result is cached in _ATTENTION_FN."""
|
| 640 |
backend = _backend_for(device)
|
| 641 |
|
| 642 |
+
if device.type == "cuda" and _kernel_config_attr("use_flexattention", False) and _flex_attention_available():
|
| 643 |
+
def _flex_dispatch(q, k, v, causal, window_size, dropout_p, training, softmax_scale):
|
| 644 |
+
try:
|
| 645 |
+
result = _flex_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale)
|
| 646 |
+
_record_backend("flex_success")
|
| 647 |
+
return result
|
| 648 |
+
except _UnsupportedByBackend as e:
|
| 649 |
+
_record_backend(f"flex_fallback:{type(e).__name__}")
|
| 650 |
+
_warn_once(f"flex_unsupported_{type(e).__name__}",
|
| 651 |
+
f"FlexAttention can't handle this call ({e!r}) -- falling back to xformers.")
|
| 652 |
+
except Exception as e: # noqa: BLE001
|
| 653 |
+
_record_backend(f"flex_fallback:{type(e).__name__}")
|
| 654 |
+
_warn_once(f"flex_fail_{type(e).__name__}",
|
| 655 |
+
f"FlexAttention failed ({e!r}) -- falling back to xformers.")
|
| 656 |
+
return _xformers_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale)
|
| 657 |
+
return "flex", _flex_dispatch
|
| 658 |
+
|
| 659 |
+
if device.type == "cuda" and torch.cuda.get_device_capability(device) in _BLACKWELL_CAPABILITIES:
|
| 660 |
+
def _blackwell_dispatch(q, k, v, causal, window_size, dropout_p, training, softmax_scale):
|
| 661 |
+
if _fa4_eligible(q, causal, window_size, dropout_p if training else 0.0):
|
| 662 |
+
try:
|
| 663 |
+
result = _fa4_attention(q, k, v, causal, softmax_scale)
|
| 664 |
+
_record_backend("fa4_success")
|
| 665 |
+
return result
|
| 666 |
+
except Exception as e: # noqa: BLE001
|
| 667 |
+
_record_backend(f"fa4_fallback:{type(e).__name__}")
|
| 668 |
+
_warn_once(f"fa4_fail_{type(e).__name__}",
|
| 669 |
+
f"FlashAttention-4 failed ({e!r}) -- falling back to FlashAttention-2.")
|
| 670 |
+
if _flash_attn_available():
|
| 671 |
+
try:
|
| 672 |
+
result = _flash_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale)
|
| 673 |
+
_record_backend("flash_success")
|
| 674 |
+
return result
|
| 675 |
+
except Exception as e: # noqa: BLE001
|
| 676 |
+
_record_backend(f"flash_fallback:{type(e).__name__}")
|
| 677 |
+
_warn_once(f"flash_fail_{type(e).__name__}",
|
| 678 |
+
f"FlashAttention-2 failed ({e!r}) -- falling back to xformers.")
|
| 679 |
+
return _xformers_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale)
|
| 680 |
+
return "blackwell", _blackwell_dispatch
|
| 681 |
+
|
| 682 |
if (
|
| 683 |
device.type == "cuda"
|
| 684 |
and _turing_flash_attn_func is not None
|
|
|
|
| 800 |
)
|
| 801 |
lines.append(f" turing-flash not eligible ({reason})")
|
| 802 |
|
| 803 |
+
# FlexAttention: config-gated (use_flexattention in config.py), not hardware-gated
|
| 804 |
+
flex_cfg_on = _kernel_config_attr("use_flexattention", False)
|
| 805 |
+
flex_active = flex_cfg_on and _flex_attention_available()
|
| 806 |
+
if flex_active:
|
| 807 |
+
lines.append(" flex-attention ACTIVE (use_flexattention=True in config.py)")
|
| 808 |
+
elif flex_cfg_on:
|
| 809 |
+
lines.append(" flex-attention requested but not importable (torch < 2.5?)")
|
| 810 |
+
else:
|
| 811 |
+
lines.append(" flex-attention off (set use_flexattention=True in config.py to enable)")
|
| 812 |
+
|
| 813 |
+
# Blackwell family (Hopper sm_90 + all Blackwell variants): FA4 with FA2 fallback
|
| 814 |
+
blackwell_eligible = (major, minor) in _BLACKWELL_CAPABILITIES
|
| 815 |
+
if blackwell_eligible:
|
| 816 |
+
if (major, minor) in _FA4_CAPABILITIES and _fa4_attn_func is not None:
|
| 817 |
+
lines.append(f" fa4 AVAILABLE (sm_{major}{minor})")
|
| 818 |
+
elif (major, minor) in _FA4_CAPABILITIES:
|
| 819 |
+
lines.append(f" fa4 not installed (pip install flash-attn-4)")
|
| 820 |
+
else:
|
| 821 |
+
lines.append(f" fa4 not eligible (sm_120 lacks FA4's required warp-specialized MMA -- uses FlashAttention-2 instead)")
|
| 822 |
+
|
| 823 |
# flash_attn 2.x
|
| 824 |
if _flash_attn_available():
|
| 825 |
lines.append(f" flash_attn available (Ampere+ path, backend={backend!r})")
|
|
|
|
| 833 |
lines.append(" xformers not installed")
|
| 834 |
|
| 835 |
# Summarise what will actually fire first.
|
| 836 |
+
if flex_active:
|
| 837 |
+
active = "flex-attention (config-enabled, torch.compile'd)"
|
| 838 |
+
elif blackwell_eligible and (major, minor) in _FA4_CAPABILITIES and _fa4_attn_func is not None:
|
| 839 |
+
active = f"fa4 (sm_{major}{minor})"
|
| 840 |
+
elif blackwell_eligible and _flash_attn_available():
|
| 841 |
+
active = f"flash_attn 2.x (sm_{major}{minor}, FA4 not eligible/installed)"
|
| 842 |
+
elif turing_eligible:
|
| 843 |
active = "turing-flash (sm_75 + flash_attention_interface)"
|
| 844 |
elif backend == "flash":
|
| 845 |
active = "flash_attn 2.x"
|