Spaces:
Sleeping
Sleeping
File size: 35,673 Bytes
c0704f7 6cd4a29 c0704f7 a27f593 c0704f7 61b358c c0704f7 a27f593 c0704f7 a27f593 c0704f7 6cd4a29 c0704f7 61b358c 6cd4a29 a27f593 6cd4a29 c0704f7 6cd4a29 c0704f7 a27f593 c0704f7 61b358c c0704f7 6cd4a29 c0704f7 a27f593 61b358c 6cd4a29 c0704f7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 | 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)) |