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))