Pragya / kernel.py
ArushBuilds's picture
Update kernel.py
61b358c
Raw History Blame
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] = {}
@torch.compiler.allow_in_graph
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()
@torch.compiler.allow_in_graph
@functools.lru_cache(maxsize=1)
def _flash_attn_available() -> bool:
try:
import flash_attn # noqa: F401
return True
except ImportError:
return False
@torch.compiler.allow_in_graph
@functools.lru_cache(maxsize=1)
def _xformers_available() -> bool:
try:
import xformers.ops # noqa: F401
return True
except ImportError:
return False
@torch.compiler.allow_in_graph
@functools.lru_cache(maxsize=1)
def _tpu_kernel_available() -> bool:
try:
from torch_xla.experimental import custom_kernel # noqa: F401
return True
except ImportError:
return False
@torch.compiler.allow_in_graph
@functools.lru_cache(maxsize=None)
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)
@torch.compiler.disable
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.
# ---------------------------------------------------------------------------
@functools.lru_cache(maxsize=1)
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
@torch.compiler.disable
@functools.lru_cache(maxsize=32)
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)
# ---------------------------------------------------------------------------
@torch.compiler.disable
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)
# ---------------------------------------------------------------------------
@functools.lru_cache(maxsize=1)
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):
@staticmethod
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)
@staticmethod
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)
@torch.compiler.disable
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
@torch.compiler.disable
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))