Spaces:
Sleeping
Sleeping
Download kernel.py from ArushBuilds/Pragya: direct link, hf CLI and curl.
- Browser
- Download file 35.7 kB
-
https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/61b358c953587a838c340fe97c699e643d70ce84/kernel.py
- Command line
-
hf download hf://spaces/ArushBuilds/Pragya@61b358c953587a838c340fe97c699e643d70ce84/kernel.py
-
curl -L -o kernel.py https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/61b358c953587a838c340fe97c699e643d70ce84/kernel.py
35.7 kB
| import functools | |
| import math | |
| import warnings | |
| import torch | |
| import torch.utils.checkpoint | |
| from torch.nn import functional as F | |
| try: | |
| import config as _alpha_config | |
| except ImportError: | |
| _alpha_config = None | |
| def _kernel_config_attr(name: str, default): | |
| if _alpha_config is None: | |
| return default | |
| if not hasattr(_alpha_config, name): | |
| warnings.warn( | |
| f"[kernel.py] config module has no attribute '{name}' -- falling back to " | |
| f"default {default!r}. If this is unexpected, check for a casing mismatch " | |
| f"or rename in your config.py.", | |
| stacklevel=3, | |
| ) | |
| return default | |
| return getattr(_alpha_config, name) | |
| class _UnsupportedByBackend(Exception): | |
| """Raised internally to signal 'this backend can't safely do this' -> fall back.""" | |
| _warned: set[str] = set() | |
| def _warn_once(key: str, msg: str) -> None: | |
| if key not in _warned: | |
| _warned.add(key) | |
| warnings.warn(f"[kernel.py] {msg}", stacklevel=3) | |
| _backend_stats: dict[str, int] = {} | |
| def _record_backend(name: str) -> None: | |
| _backend_stats[name] = _backend_stats.get(name, 0) + 1 | |
| def get_backend_stats() -> dict[str, int]: | |
| return dict(_backend_stats) | |
| def reset_backend_stats() -> None: | |
| _backend_stats.clear() | |
| def _flash_attn_available() -> bool: | |
| try: | |
| import flash_attn # noqa: F401 | |
| return True | |
| except ImportError: | |
| return False | |
| def _xformers_available() -> bool: | |
| try: | |
| import xformers.ops # noqa: F401 | |
| return True | |
| except ImportError: | |
| return False | |
| def _tpu_kernel_available() -> bool: | |
| try: | |
| from torch_xla.experimental import custom_kernel # noqa: F401 | |
| return True | |
| except ImportError: | |
| return False | |
| def _detect_backend(device_key: str) -> str: | |
| if device_key.startswith("xla"): | |
| if _tpu_kernel_available(): | |
| return "tpu" | |
| _warn_once("tpu_missing", "torch_xla not importable on an XLA device -- using reference attention.") | |
| return "reference" | |
| if device_key.startswith("cuda"): | |
| idx = int(device_key.split(":")[1]) if ":" in device_key else torch.cuda.current_device() | |
| major, _minor = torch.cuda.get_device_capability(idx) | |
| if major >= 8 and _flash_attn_available(): | |
| return "flash" | |
| if _xformers_available(): | |
| return "xformers" | |
| if major >= 8: | |
| _warn_once("flash_missing", "Ampere+ GPU detected but `flash_attn` isn't installed, and neither is `xformers` -- using reference attention (slow).") | |
| else: | |
| _warn_once("xformers_missing", "Pre-Ampere GPU (e.g. T4) detected and `xformers` isn't installed -- using reference attention (slow, high memory).") | |
| return "reference" | |
| return "reference" | |
| def _backend_for(device: torch.device) -> str: | |
| if device.type == "cuda": | |
| return _detect_backend(f"cuda:{device.index if device.index is not None else torch.cuda.current_device()}") | |
| return _detect_backend(device.type) | |
| def _attention_block(q_block, k_block, v_block, q_start, k_start, causal, window_size, dropout_p, training, scale): | |
| """One tile's worth of causal/windowed attention. Kept as a standalone | |
| function (not a closure) so torch.utils.checkpoint can call it directly.""" | |
| qb, kb = q_block.shape[2], k_block.shape[2] | |
| scores = torch.matmul(q_block, k_block.transpose(-2, -1)) * scale | |
| if causal or window_size is not None: | |
| q_idx = torch.arange(q_start, q_start + qb, device=q_block.device).view(qb, 1) | |
| k_idx = torch.arange(k_start, k_start + kb, device=q_block.device).view(1, kb) | |
| allowed = (k_idx <= q_idx) if causal else torch.ones(qb, kb, dtype=torch.bool, device=q_block.device) | |
| if window_size is not None: | |
| allowed = allowed & (q_idx - k_idx < window_size) | |
| mask = torch.where(allowed, torch.zeros(1, device=q_block.device), torch.full((1,), -1e4, device=q_block.device)) | |
| scores = scores + mask.to(scores.dtype) | |
| weights = torch.softmax(scores.float(), dim=-1).to(scores.dtype) | |
| if training and dropout_p > 0: | |
| weights = torch.nn.functional.dropout(weights, p=dropout_p) | |
| return torch.matmul(weights, v_block) | |
| def _reference_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale, query_block_size=128): | |
| b, hq, t, d = q.shape | |
| _, hkv, tk, _ = k.shape | |
| if hkv != hq: | |
| reps = hq // hkv | |
| k = k.repeat_interleave(reps, dim=1) | |
| v = v.repeat_interleave(reps, dim=1) | |
| scale = softmax_scale if softmax_scale is not None else (1.0 / math.sqrt(d)) | |
| outputs = [] | |
| for q_start in range(0, t, query_block_size): | |
| q_end = min(q_start + query_block_size, t) | |
| q_block = q[:, :, q_start:q_end, :] | |
| k_lo = max(0, q_start - window_size + 1) if window_size is not None else 0 | |
| k_hi = min(q_end, tk) if causal else tk | |
| k_block = k[:, :, k_lo:k_hi, :] | |
| v_block = v[:, :, k_lo:k_hi, :] | |
| if q_block.device.type == "xla": | |
| out_block = _attention_block( | |
| q_block, k_block, v_block, q_start, k_lo, | |
| causal, window_size, dropout_p, training, scale, | |
| ) | |
| else: | |
| out_block = torch.utils.checkpoint.checkpoint( | |
| _attention_block, q_block, k_block, v_block, q_start, k_lo, | |
| causal, window_size, dropout_p, training, scale, | |
| use_reentrant=False, | |
| ) | |
| outputs.append(out_block) | |
| return torch.cat(outputs, dim=2) | |
| try: | |
| from flash_attention_interface import flash_attn_func as _turing_flash_attn_func | |
| except ImportError: | |
| _turing_flash_attn_func = None | |
| _TURING_SUPPORTED_HEAD_DIMS = (64, 128) | |
| def _turing_flash_eligible(q: torch.Tensor, k: torch.Tensor, causal: bool, | |
| window_size: int | None, dropout_p: float) -> bool: | |
| if _turing_flash_attn_func is None: | |
| return False | |
| if q.device.type != "cuda": | |
| return False | |
| major, minor = torch.cuda.get_device_capability(q.device) | |
| if (major, minor) != (7, 5): | |
| return False | |
| if q.shape[-1] not in _TURING_SUPPORTED_HEAD_DIMS: | |
| return False | |
| if window_size is not None: | |
| return False | |
| if dropout_p > 0.0: | |
| return False | |
| return True | |
| def _turing_flash_attention(q, k, v, causal, softmax_scale): | |
| q_ = q.transpose(1, 2) # (B,H,T,D) -> (B,T,H,D), this repo's expected layout | |
| k_ = k.transpose(1, 2) | |
| v_ = v.transpose(1, 2) | |
| out = _turing_flash_attn_func(q_, k_, v_, softmax_scale=softmax_scale, causal=causal) | |
| return out.transpose(1, 2) # back to (B,H,T,D) for this codebase's convention | |
| # --------------------------------------------------------------------------- | |
| # FlexAttention (torch.nn.attention.flex_attention). Config-gated, not | |
| # hardware-gated -- set use_flexattention=True in config.py to opt in. | |
| # Unlike every other backend in this file, FlexAttention is a torch-level | |
| # composition of ops, not a hand-fused CUDA kernel -- its speed comes | |
| # entirely from torch.compile specializing the mask/score-mod into one | |
| # fused kernel at trace time. Eager flex_attention runs at roughly | |
| # reference-attention speed, so (unlike the @torch.compiler.disable | |
| # backends elsewhere in this file) this one is deliberately compiled, once, | |
| # and cached. | |
| # --------------------------------------------------------------------------- | |
| def _flex_attention_available() -> bool: | |
| try: | |
| from torch.nn.attention.flex_attention import flex_attention # noqa: F401 | |
| return True | |
| except ImportError: | |
| return False | |
| _compiled_flex_attention = None | |
| def _get_compiled_flex_attention(): | |
| global _compiled_flex_attention | |
| if _compiled_flex_attention is None: | |
| from torch.nn.attention.flex_attention import flex_attention | |
| _compiled_flex_attention = torch.compile(flex_attention, dynamic=False) | |
| return _compiled_flex_attention | |
| def _flex_block_mask(q_len: int, kv_len: int, causal: bool, window_size: "int | None", device_key: str): | |
| """Build (and cache) a FlexAttention BlockMask for one (shape, mask-kind, | |
| device) combo. Constructing a block mask isn't free, and with a fixed | |
| training CONTEXT this pays the cost once instead of every forward call.""" | |
| from torch.nn.attention.flex_attention import create_block_mask | |
| if window_size is not None: | |
| def mask_mod(b, h, q_idx, kv_idx): | |
| return (kv_idx <= q_idx) & (q_idx - kv_idx < window_size) | |
| elif causal: | |
| def mask_mod(b, h, q_idx, kv_idx): | |
| return kv_idx <= q_idx | |
| else: | |
| return None # full bidirectional attention -- no block mask needed | |
| return create_block_mask(mask_mod, B=None, H=None, Q_LEN=q_len, KV_LEN=kv_len, device=device_key) | |
| def _flex_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| if dropout_p > 0 and training: | |
| raise _UnsupportedByBackend( | |
| "FlexAttention path here doesn't wire up an attention-weight " | |
| "dropout term -- falling back." | |
| ) | |
| b, hq, t, d = q.shape | |
| _, hkv, tk, _ = k.shape | |
| block_mask = _flex_block_mask(t, tk, causal, window_size, str(q.device)) | |
| scale = softmax_scale if softmax_scale is not None else (1.0 / math.sqrt(d)) | |
| flex_fn = _get_compiled_flex_attention() | |
| # FlexAttention takes (B, H, T, D) natively -- no transpose needed, unlike | |
| # the flash_attn/xformers paths above which want (B, T, H, D). | |
| return flex_fn(q, k, v, block_mask=block_mask, scale=scale, enable_gqa=(hkv != hq)) | |
| # --------------------------------------------------------------------------- | |
| # Blackwell: FlashAttention-4 (CuTeDSL) with FlashAttention-2 fallback. | |
| # | |
| # FA4 (`pip install flash-attn-4`, imported as `flash_attn.cute`) targets | |
| # Hopper and data-center Blackwell (sm_90 / sm_100 / sm_103 -- H100/B200/B300) | |
| # via warp-specialized MMA instructions. It does NOT run on desktop Blackwell | |
| # (sm_120, e.g. RTX 50-series): sm_120 uses the same register-to-register | |
| # HMMA path NVIDIA GPUs have used since Volta, not the warp-specialized MMA | |
| # FA4's kernel design requires, and this is a physical silicon difference | |
| # (confirmed by multiple independent sm_120 FA4 build attempts as of early | |
| # 2026), not a version-gating issue that a newer release fixes. So this | |
| # block tries FA4 only on the capabilities it can actually run on, and | |
| # every other Blackwell-family GPU (sm_120 included) falls through to the | |
| # existing FlashAttention-2 path below. | |
| # --------------------------------------------------------------------------- | |
| _BLACKWELL_CAPABILITIES = {(9, 0), (10, 0), (10, 3), (12, 0)} # Hopper + all Blackwell variants | |
| _FA4_CAPABILITIES = {(9, 0), (10, 0), (10, 3)} # NOT (12, 0) -- see note above | |
| try: | |
| from flash_attn.cute import flash_attn_func as _fa4_attn_func | |
| except ImportError: | |
| _fa4_attn_func = None | |
| def _fa4_eligible(q: torch.Tensor, causal: bool, window_size: "int | None", dropout_p: float) -> bool: | |
| if _fa4_attn_func is None: | |
| return False | |
| if q.device.type != "cuda": | |
| return False | |
| if torch.cuda.get_device_capability(q.device) not in _FA4_CAPABILITIES: | |
| return False | |
| if window_size is not None: # not confirmed supported by FA4's current (beta) public API | |
| return False | |
| if dropout_p > 0.0: | |
| return False | |
| return True | |
| def _fa4_attention(q, k, v, causal, softmax_scale): | |
| # FA4's public surface is still small/beta -- only `causal` is confirmed | |
| # from the library's own usage example (flash_attn.cute docs show only | |
| # `flash_attn_func(q, k, v, causal=True)`). Anything beyond that | |
| # (window_size, dropout) is gated out by _fa4_eligible above rather than | |
| # guessed at here. Assumes the same (B, T, H, D) layout as FA2/FA3; if | |
| # that assumption is wrong for a given release the call raises and the | |
| # dispatcher below falls back to FlashAttention-2 automatically. | |
| qf = q.transpose(1, 2).contiguous() | |
| kf = k.transpose(1, 2).contiguous() | |
| vf = v.transpose(1, 2).contiguous() | |
| kwargs = {"softmax_scale": softmax_scale} if softmax_scale is not None else {} | |
| out = _fa4_attn_func(qf, kf, vf, causal=causal, **kwargs) | |
| return out.transpose(1, 2) | |
| # --------------------------------------------------------------------------- | |
| # FlashAttention-2 (NVIDIA Ampere+) | |
| # --------------------------------------------------------------------------- | |
| def _flash_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| from flash_attn import flash_attn_func | |
| # flash_attn wants (B, T, H, D); we work in (B, H, T, D) throughout. | |
| qf = q.transpose(1, 2).contiguous() | |
| kf = k.transpose(1, 2).contiguous() | |
| vf = v.transpose(1, 2).contiguous() | |
| # window_size=(-1,-1) means "unbounded" (plain causal/full) per flash-attn's | |
| # own convention; (W-1, 0) means "attend to self + W-1 previous tokens". | |
| ws = (window_size - 1, 0) if window_size is not None else (-1, -1) | |
| out = flash_attn_func( | |
| qf, kf, vf, | |
| dropout_p=dropout_p if training else 0.0, | |
| softmax_scale=softmax_scale, | |
| causal=causal, | |
| window_size=ws, | |
| ) # GQA/MQA handled natively by flash_attn_func (kf/vf may have fewer heads than qf) | |
| return out.transpose(1, 2) # back to (B, H, T, D) | |
| # --------------------------------------------------------------------------- | |
| # xFormers (NVIDIA pre-Ampere, e.g. T4) | |
| # --------------------------------------------------------------------------- | |
| def _xformers_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| import xformers.ops as xops | |
| from xformers.ops.fmha.attn_bias import LowerTriangularMask, BlockDiagonalCausalMask | |
| qf = q.transpose(1, 2).contiguous() # (B, T, H, D) | |
| kf = k.transpose(1, 2).contiguous() | |
| vf = v.transpose(1, 2).contiguous() | |
| b, t, hq, d = qf.shape | |
| _, tk, hkv, _ = kf.shape | |
| if hkv != hq: | |
| reps = hq // hkv | |
| kf = kf.repeat_interleave(reps, dim=2) | |
| vf = vf.repeat_interleave(reps, dim=2) | |
| if window_size is not None: | |
| if not causal: | |
| raise _UnsupportedByBackend("non-causal windowed attention not implemented for the xFormers path") | |
| bias = BlockDiagonalCausalMask.from_seqlens( | |
| q_seqlen=[t] * b, kv_seqlen=[tk] * b, | |
| ).make_local_attention(window_size) | |
| qf_packed = qf.reshape(1, b * t, hq, d) | |
| kf_packed = kf.reshape(1, b * tk, hq, d) | |
| vf_packed = vf.reshape(1, b * tk, hq, d) | |
| out = xops.memory_efficient_attention( | |
| qf_packed, kf_packed, vf_packed, attn_bias=bias, | |
| p=dropout_p if training else 0.0, | |
| scale=softmax_scale, | |
| ) | |
| out = out.reshape(b, t, hq, d) | |
| return out.transpose(1, 2) | |
| bias = LowerTriangularMask() if causal else None | |
| out = xops.memory_efficient_attention( | |
| qf, kf, vf, attn_bias=bias, | |
| p=dropout_p if training else 0.0, | |
| scale=softmax_scale, | |
| ) | |
| return out.transpose(1, 2) | |
| # --------------------------------------------------------------------------- | |
| # TPU (torch_xla) | |
| # --------------------------------------------------------------------------- | |
| def _splash_attention_available() -> bool: | |
| try: | |
| from jax.experimental.pallas.ops.tpu.splash_attention import ( # noqa: F401 | |
| splash_attention_kernel, | |
| splash_attention_mask, | |
| ) | |
| import torch_xla.core.xla_builder # noqa: F401 | |
| return True | |
| except ImportError: | |
| return False | |
| _splash_disabled_reason: str | None = None | |
| def reset_splash_circuit_breaker() -> None: | |
| global _splash_disabled_reason | |
| _splash_disabled_reason = None | |
| def get_splash_disabled_reason() -> str | None: | |
| return _splash_disabled_reason | |
| class _SplashAttentionFn(torch.autograd.Function): | |
| def forward(ctx, q, k, v, num_heads, q_len, k_len, scale): | |
| import torch_xla | |
| import torch_xla.core.xla_builder as xb | |
| orig_dtype = q.dtype | |
| q_b, k_b, v_b = q.to(torch.bfloat16), k.to(torch.bfloat16), v.to(torch.bfloat16) | |
| def fwd_jax(qj, kj, vj): | |
| from jax.experimental.pallas.ops.tpu.splash_attention import ( | |
| splash_attention_kernel, splash_attention_mask, | |
| ) | |
| import jax | |
| mask = splash_attention_mask.MultiHeadMask( | |
| masks=[splash_attention_mask.CausalMask(shape=(q_len, k_len)) for _ in range(num_heads)] | |
| ) | |
| kernel_fn = splash_attention_kernel.make_splash_mha(mask=mask, head_shards=1, q_seq_shards=1) | |
| return jax.vmap(kernel_fn)(q=qj * scale, k=kj, v=vj) | |
| out = xb.call_jax(fwd_jax, (q_b, k_b, v_b), {}, "arya_splash_attention_fwd") | |
| torch_xla.sync(reset_scope=False) | |
| ctx.save_for_backward(q_b, k_b, v_b) | |
| ctx.num_heads, ctx.q_len, ctx.k_len, ctx.scale, ctx.orig_dtype = num_heads, q_len, k_len, scale, orig_dtype | |
| return out.to(orig_dtype) | |
| def backward(ctx, grad_output): | |
| import torch_xla | |
| import torch_xla.core.xla_builder as xb | |
| q_b, k_b, v_b = ctx.saved_tensors | |
| num_heads, q_len, k_len, scale = ctx.num_heads, ctx.q_len, ctx.k_len, ctx.scale | |
| grad_output_b = grad_output.to(torch.bfloat16) | |
| def bwd_jax(qj, kj, vj, gj): | |
| from jax.experimental.pallas.ops.tpu.splash_attention import ( | |
| splash_attention_kernel, splash_attention_mask, | |
| ) | |
| import jax | |
| def raw(qj_, kj_, vj_): | |
| mask = splash_attention_mask.MultiHeadMask( | |
| masks=[splash_attention_mask.CausalMask(shape=(q_len, k_len)) for _ in range(num_heads)] | |
| ) | |
| kernel_fn = splash_attention_kernel.make_splash_mha(mask=mask, head_shards=1, q_seq_shards=1) | |
| return jax.vmap(kernel_fn)(q=qj_ * scale, k=kj_, v=vj_) | |
| _, vjp_fn = jax.vjp(raw, qj, kj, vj) | |
| return vjp_fn(gj) # (dq, dk, dv) -- call_jax supports PyTree returns | |
| dq, dk, dv = xb.call_jax(bwd_jax, (q_b, k_b, v_b, grad_output_b), {}, "arya_splash_attention_bwd") | |
| torch_xla.sync(reset_scope=False) # same reasoning as forward() -- must fail here, not later, elsewhere | |
| orig_dtype = ctx.orig_dtype | |
| return dq.to(orig_dtype), dk.to(orig_dtype), dv.to(orig_dtype), None, None, None, None | |
| def _splash_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| if not causal: | |
| raise _UnsupportedByBackend("Splash Attention wiring here only supports causal=True -- falling back.") | |
| if window_size is not None: | |
| raise _UnsupportedByBackend("Splash Attention wiring here doesn't support window_size -- falling back.") | |
| if dropout_p > 0 and training: | |
| raise _UnsupportedByBackend("Splash Attention wiring here doesn't support dropout -- falling back.") | |
| q_, k_, v_ = q, k, v | |
| hq, hkv = q.shape[1], k.shape[1] | |
| if hkv != hq: | |
| reps = hq // hkv | |
| k_ = k.repeat_interleave(reps, dim=1) | |
| v_ = v.repeat_interleave(reps, dim=1) | |
| num_heads = q_.shape[1] | |
| q_len, k_len = q_.shape[2], k_.shape[2] | |
| scale = softmax_scale if softmax_scale is not None else (1.0 / math.sqrt(q_.shape[-1])) | |
| return _SplashAttentionFn.apply(q_, k_, v_, num_heads, q_len, k_len, scale) | |
| def _tpu_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| from torch_xla.experimental.custom_kernel import flash_attention as xla_flash_attention | |
| global _splash_disabled_reason | |
| if _splash_disabled_reason is not None: | |
| _record_backend("splash_circuit_broken") | |
| elif _splash_attention_available(): | |
| try: | |
| result = _splash_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| _record_backend("splash_success") | |
| return result | |
| except _UnsupportedByBackend as e: | |
| _record_backend(f"splash_fallback:{type(e).__name__}") | |
| _warn_once( | |
| f"splash_fail_{type(e).__name__}", | |
| f"Splash Attention call failed ({e!r}) -- falling back to flash_attention.", | |
| ) | |
| except Exception as e: # noqa: BLE001 -- a REAL execution failure (post-sync) -- trip the breaker | |
| _record_backend(f"splash_fallback:{type(e).__name__}") | |
| _splash_disabled_reason = repr(e) | |
| _warn_once( | |
| f"splash_fail_{type(e).__name__}", | |
| f"Splash Attention call failed with a real execution error ({e!r}) -- " | |
| f"this usually means a jax/jaxlib/libtpu version mismatch in this " | |
| f"environment (e.g. 'Unsupported version: expected <= N but got M' is " | |
| f"Mosaic IR version skew between jax and libtpu -- fix by aligning " | |
| f"`pip install -U \"jax[tpu]\" jaxlib` as a matched pair). Disabling " | |
| f"Splash Attention for the rest of this run and falling back to " | |
| f"flash_attention -- call kernel.reset_splash_circuit_breaker() to " | |
| f"retry after fixing the environment.", | |
| ) | |
| else: | |
| _record_backend("splash_unavailable") | |
| if window_size is not None: | |
| raise _UnsupportedByBackend( | |
| "torch_xla's built-in flash_attention wrapper doesn't expose a " | |
| "sliding-window argument -- falling back to reference." | |
| ) | |
| if dropout_p > 0 and training: | |
| raise _UnsupportedByBackend( | |
| "torch_xla's built-in flash_attention wrapper doesn't take a " | |
| "dropout argument -- falling back to reference." | |
| ) | |
| q_, k_, v_ = q, k, v | |
| hq, hkv = q.shape[1], k.shape[1] | |
| if hkv != hq: | |
| reps = hq // hkv | |
| k_ = k.repeat_interleave(reps, dim=1) | |
| v_ = v.repeat_interleave(reps, dim=1) | |
| orig_dtype = q_.dtype | |
| q_b = q_.to(torch.bfloat16) | |
| k_b = k_.to(torch.bfloat16) | |
| v_b = v_.to(torch.bfloat16) | |
| result = xla_flash_attention(q_b, k_b, v_b, causal=causal) | |
| result = result.to(orig_dtype) | |
| _record_backend("tpu_flash_success") | |
| return result | |
| def fused_block_sparse_attention( | |
| q_blocks: torch.Tensor, | |
| k_sel: torch.Tensor, | |
| v_sel: torch.Tensor, | |
| bias: torch.Tensor, | |
| dropout_p: float = 0.0, | |
| training: bool = True, | |
| softmax_scale: float | None = None, | |
| ) -> torch.Tensor | None: | |
| if q_blocks.device.type == "cuda" and _xformers_available(): | |
| try: | |
| return _xformers_block_sparse_attention(q_blocks, k_sel, v_sel, bias, dropout_p, training, softmax_scale) | |
| except Exception as e: # noqa: BLE001 -- must never crash training, just fall back | |
| _warn_once( | |
| f"block_sparse_xformers_fail_{type(e).__name__}", | |
| f"xFormers block-sparse attention failed ({e!r}) -- falling back to reference.", | |
| ) | |
| return None | |
| _ATTENTION_FN: "callable | None" = None | |
| def _build_attention_fn(device: torch.device): | |
| """Probe backends in priority order and return a direct callable. | |
| Called once; result is cached in _ATTENTION_FN.""" | |
| backend = _backend_for(device) | |
| if device.type == "cuda" and _kernel_config_attr("use_flexattention", False) and _flex_attention_available(): | |
| def _flex_dispatch(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| try: | |
| result = _flex_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| _record_backend("flex_success") | |
| return result | |
| except _UnsupportedByBackend as e: | |
| _record_backend(f"flex_fallback:{type(e).__name__}") | |
| _warn_once(f"flex_unsupported_{type(e).__name__}", | |
| f"FlexAttention can't handle this call ({e!r}) -- falling back to xformers.") | |
| except Exception as e: # noqa: BLE001 | |
| _record_backend(f"flex_fallback:{type(e).__name__}") | |
| _warn_once(f"flex_fail_{type(e).__name__}", | |
| f"FlexAttention failed ({e!r}) -- falling back to xformers.") | |
| return _xformers_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| return "flex", _flex_dispatch | |
| if device.type == "cuda" and torch.cuda.get_device_capability(device) in _BLACKWELL_CAPABILITIES: | |
| def _blackwell_dispatch(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| if _fa4_eligible(q, causal, window_size, dropout_p if training else 0.0): | |
| try: | |
| result = _fa4_attention(q, k, v, causal, softmax_scale) | |
| _record_backend("fa4_success") | |
| return result | |
| except Exception as e: # noqa: BLE001 | |
| _record_backend(f"fa4_fallback:{type(e).__name__}") | |
| _warn_once(f"fa4_fail_{type(e).__name__}", | |
| f"FlashAttention-4 failed ({e!r}) -- falling back to FlashAttention-2.") | |
| if _flash_attn_available(): | |
| try: | |
| result = _flash_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| _record_backend("flash_success") | |
| return result | |
| except Exception as e: # noqa: BLE001 | |
| _record_backend(f"flash_fallback:{type(e).__name__}") | |
| _warn_once(f"flash_fail_{type(e).__name__}", | |
| f"FlashAttention-2 failed ({e!r}) -- falling back to xformers.") | |
| return _xformers_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| return "blackwell", _blackwell_dispatch | |
| if ( | |
| device.type == "cuda" | |
| and _turing_flash_attn_func is not None | |
| and torch.cuda.get_device_capability(device) == (7, 5) | |
| ): | |
| def _turing_dispatch(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| if _turing_flash_eligible(q, causal, window_size, dropout_p if training else 0.0): | |
| try: | |
| result = _turing_flash_attention(q, k, v, causal, softmax_scale) | |
| _record_backend("turing_flash_success") | |
| return result | |
| except Exception as e: # noqa: BLE001 | |
| _record_backend(f"turing_flash_fallback:{type(e).__name__}") | |
| _warn_once(f"turing_flash_fail_{type(e).__name__}", | |
| f"flash-attention-turing failed ({e!r}) -- falling back to xformers.") | |
| return _xformers_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| return "turing_flash", _turing_dispatch | |
| if backend == "flash": | |
| def _flash_dispatch(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| try: | |
| result = _flash_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| _record_backend("flash_success") | |
| return result | |
| except Exception as e: # noqa: BLE001 | |
| _record_backend(f"flash_fallback:{type(e).__name__}") | |
| _warn_once(f"flash_fail_{type(e).__name__}", f"flash_attn failed ({e!r}) -- falling back to xformers.") | |
| return _xformers_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| return "flash", _flash_dispatch | |
| if backend == "xformers": | |
| # Direct call — hits @torch.compiler.disable immediately, no wasted tracing. | |
| return "xformers", _xformers_attention | |
| if backend == "tpu": | |
| def _tpu_dispatch(q, k, v, causal, window_size, dropout_p, training, softmax_scale): | |
| try: | |
| result = _tpu_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| return result | |
| except Exception as e: # noqa: BLE001 | |
| _record_backend(f"tpu_fallback:{type(e).__name__}") | |
| _warn_once(f"tpu_fail_{type(e).__name__}", f"torch_xla flash_attention failed ({e!r}) -- falling back.") | |
| return _reference_attention(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| return "tpu", _tpu_dispatch | |
| return "reference", _reference_attention | |
| def warmup_attention_backend(device: "torch.device | str | None" = None) -> str: | |
| """Resolve and cache the attention backend for *device*. | |
| Call this once after model construction (before torch.compile) so that | |
| fused_attention() contains a single unconditional dispatch with no | |
| branching inside the compiled graph. | |
| Returns the resolved backend name string (e.g. 'xformers', 'flash'). | |
| """ | |
| global _ATTENTION_FN | |
| if isinstance(device, str): | |
| device = torch.device(device) | |
| if device is None: | |
| if torch.cuda.is_available(): | |
| device = torch.device("cuda", torch.cuda.current_device()) | |
| else: | |
| device = torch.device("cpu") | |
| name, fn = _build_attention_fn(device) | |
| _ATTENTION_FN = fn | |
| return name | |
| def fused_attention( | |
| q: torch.Tensor, | |
| k: torch.Tensor, | |
| v: torch.Tensor, | |
| causal: bool = True, | |
| window_size: int | None = None, | |
| dropout_p: float = 0.0, | |
| training: bool = True, | |
| softmax_scale: float | None = None, | |
| ) -> torch.Tensor: | |
| global _ATTENTION_FN | |
| if _ATTENTION_FN is None: | |
| # First-call lazy init (warmup_attention_backend() not called yet). | |
| warmup_attention_backend(q.device) | |
| return _ATTENTION_FN(q, k, v, causal, window_size, dropout_p, training, softmax_scale) | |
| def print_attention_backend(device: "torch.device | str | None" = None) -> None: | |
| if device is None: | |
| if torch.cuda.is_available(): | |
| device = torch.device("cuda", torch.cuda.current_device()) | |
| else: | |
| device = torch.device("cpu") | |
| elif isinstance(device, str): | |
| device = torch.device(device) | |
| dev_str = str(device) | |
| backend = _backend_for(device) | |
| lines = [f"[kernel] attention backend device={dev_str}"] | |
| if device.type == "cuda": | |
| major, minor = torch.cuda.get_device_capability(device) | |
| gpu_name = torch.cuda.get_device_name(device) | |
| lines.append(f" gpu {gpu_name} (sm_{major}{minor})") | |
| # Turing flash: only sm_75 + flash_attention_interface installed | |
| turing_eligible = ( | |
| (major, minor) == (7, 5) | |
| and _turing_flash_attn_func is not None | |
| ) | |
| if turing_eligible: | |
| lines.append(" turing-flash AVAILABLE (sm_75 + flash_attention_interface)") | |
| else: | |
| reason = ( | |
| "sm != 7.5" if (major, minor) != (7, 5) | |
| else "flash_attention_interface not installed" | |
| ) | |
| lines.append(f" turing-flash not eligible ({reason})") | |
| # FlexAttention: config-gated (use_flexattention in config.py), not hardware-gated | |
| flex_cfg_on = _kernel_config_attr("use_flexattention", False) | |
| flex_active = flex_cfg_on and _flex_attention_available() | |
| if flex_active: | |
| lines.append(" flex-attention ACTIVE (use_flexattention=True in config.py)") | |
| elif flex_cfg_on: | |
| lines.append(" flex-attention requested but not importable (torch < 2.5?)") | |
| else: | |
| lines.append(" flex-attention off (set use_flexattention=True in config.py to enable)") | |
| # Blackwell family (Hopper sm_90 + all Blackwell variants): FA4 with FA2 fallback | |
| blackwell_eligible = (major, minor) in _BLACKWELL_CAPABILITIES | |
| if blackwell_eligible: | |
| if (major, minor) in _FA4_CAPABILITIES and _fa4_attn_func is not None: | |
| lines.append(f" fa4 AVAILABLE (sm_{major}{minor})") | |
| elif (major, minor) in _FA4_CAPABILITIES: | |
| lines.append(f" fa4 not installed (pip install flash-attn-4)") | |
| else: | |
| lines.append(f" fa4 not eligible (sm_120 lacks FA4's required warp-specialized MMA -- uses FlashAttention-2 instead)") | |
| # flash_attn 2.x | |
| if _flash_attn_available(): | |
| lines.append(f" flash_attn available (Ampere+ path, backend={backend!r})") | |
| else: | |
| lines.append(" flash_attn not installed") | |
| # xformers | |
| if _xformers_available(): | |
| lines.append(" xformers available") | |
| else: | |
| lines.append(" xformers not installed") | |
| # Summarise what will actually fire first. | |
| if flex_active: | |
| active = "flex-attention (config-enabled, torch.compile'd)" | |
| elif blackwell_eligible and (major, minor) in _FA4_CAPABILITIES and _fa4_attn_func is not None: | |
| active = f"fa4 (sm_{major}{minor})" | |
| elif blackwell_eligible and _flash_attn_available(): | |
| active = f"flash_attn 2.x (sm_{major}{minor}, FA4 not eligible/installed)" | |
| elif turing_eligible: | |
| active = "turing-flash (sm_75 + flash_attention_interface)" | |
| elif backend == "flash": | |
| active = "flash_attn 2.x" | |
| elif backend == "xformers": | |
| active = "xformers memory_efficient_attention" | |
| else: | |
| active = "reference (tiled matmul — slow; install flash_attn or xformers)" | |
| lines.append(f" >>> active <<< {active}") | |
| elif device.type == "xla": | |
| splash = _splash_attention_available() | |
| lines.append(f" splash-attn {'available' if splash else 'not available'}") | |
| lines.append(f" torch_xla flash {'available' if _tpu_kernel_available() else 'not available'}") | |
| active = "splash-attention" if splash else ("torch_xla flash_attention" if _tpu_kernel_available() else "reference") | |
| lines.append(f" >>> active <<< {active}") | |
| else: | |
| lines.append(" >>> active <<< reference (CPU — tiled matmul)") | |
| print("\n".join(lines)) |