ArushBuilds commited on
Commit
61b358c
·
1 Parent(s): 4b04e69

Update kernel.py

Browse files
Files changed (1) hide show
  1. kernel.py +194 -1
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 turing_eligible:
 
 
 
 
 
 
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"