File size: 64,051 Bytes
5d96a1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b22e951
5d96a1f
 
 
 
 
 
b22e951
 
 
 
 
 
 
 
 
 
5d96a1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b22e951
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d96a1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b22e951
 
 
 
 
 
 
5d96a1f
b22e951
 
 
 
 
 
 
5d96a1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b22e951
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d96a1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b22e951
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d96a1f
 
 
 
 
 
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
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
1346
1347
1348
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375
1376
1377
1378
1379
1380
1381
1382
1383
1384
1385
1386
1387
1388
1389
1390
1391
1392
1393
1394
1395
1396
1397
1398
1399
1400
1401
1402
1403
1404
1405
1406
1407
1408
1409
1410
1411
1412
1413
1414
1415
1416
1417
1418
1419
1420
1421
1422
1423
1424
1425
1426
1427
"""

veylon_attention.py  β€”  Veylon Alpha 1

======================================

Block-tiled Native GQA Sliding Window Attention for JAX / Keras-3 / TPU v5e.



Why NOT lax.scan + dynamic_slice

---------------------------------

The previous implementation used:



    jax.lax.scan(step, None, (jnp.arange(S), q_scan))



with `dynamic_slice(k_pad, (0,0,t,0), (B,Hkv,W,D))` inside the body.



XLA/XLA-TPU has a known pathology with this pattern:

  - The scan body contains a dynamic gather (dynamic_slice on a traced int `t`).

  - XLA's while-loop lowering stages ALL window materializations into a single

    large buffer to enable pipeline prefetching.

  - On TPU this produces a hidden tensor  [S, B, Hkv, W, D]  which is exactly

    what you see in the OOM allocation log:  f32[1024,8,4,512,64].

  - jax.checkpoint does NOT protect against this because it is a compiler

    (HLO-level) allocation, not a JAX-level rematerialization artifact.



The fix: block-tiled computation with fully static tensor shapes

----------------------------------------------------------------

Instead of iterating over S individual tokens we iterate over

(S / BLK) blocks of queries.  For each block:



  - q_blk  : [B, Hq,  BLK,        D]   <- static shape

  - k_blk  : [B, Hkv, BLK + W - 1, D]  <- static shape

  - scores : [B, Hkv, G, BLK, BLK+W-1] <- static shape



XLA sees ONLY static shapes inside the map body.  There is no

gather-inside-loop pattern.  XLA can freely fuse, pipeline, and

tile these einsums onto TPU systolic arrays without hidden buffers.



lax.fori_loop with dynamic_update_slice

----------------------------------------

We use `jax.lax.fori_loop` (not `lax.map` or `lax.scan`) because:

  - `lax.map` returns [n_blocks, B, Hq, BLK, D] β€” XLA stages this entire

    stack in HBM before the final transpose+reshape.  On large batches or

    many blocks this causes the VRAM spike you see in the profiler.

  - `lax.scan` has the same problem (stacks carry outputs).

  - `fori_loop` carries a single pre-allocated [B, Hq, S_pad, D] output

    buffer and writes each block with dynamic_update_slice.  XLA sees ONE

    static-shape buffer (same footprint as the final output) throughout,

    eliminating the n_blocks-deep intermediate stack entirely.



Tensor size invariants

----------------------

NEVER created:

    [B, Hq,  S, S]               full attention score matrix

    [S, B, Hkv, W, D]            per-token window stack (old scan bug)

    [B, Hq, Hkv, S, D]           duplicated KV heads



Largest tensors inside map body (all STATIC shapes):

    k_blk, v_blk  : [B, Hkv, BLK+W-1, D]

    scores        : [B, Hkv, G, BLK, BLK+W-1]

    probs         : same



Memory complexity

-----------------

  k_pad, v_pad  : O(B Γ— Hkv Γ— (S + W) Γ— D)   linear in S, scales with Hkv

  q             : O(B Γ— Hq  Γ— S Γ— D)

  Per block     : O(B Γ— Hkv Γ— (BLK + W) Γ— D)  constant w.r.t. S

  Output        : O(B Γ— Hq  Γ— S Γ— D)

  TOTAL         : O(B Γ— (Hkv + Hq) Γ— S Γ— D)   strictly linear in S



TPU-specific notes

------------------

  * All shapes in map body are concrete at trace time β€” XLA never inserts

    shape-dependent conditionals or recompiles.

  * BF16 inputs are cast to FP32 before matmul/softmax, then cast back.

  * BLK should be a multiple of 128 on TPU v5e for optimal systolic array

    utilization (default 128, tunable via `block_size` parameter).

  * The precomputed mask is a static bool array passed as a closed-over

    constant β€” XLA fuses it into the einsum kernel at no extra memory cost.

  * jax.checkpoint on the map body protects backward-pass activations

    (one block at a time, not one token at a time β€” much coarser remat).

"""

from __future__ import annotations

import math
import os
from functools import partial
from typing import Optional

import jax
import jax.numpy as jnp

# ---------------------------------------------------------------------------
# Optional Pallas/Triton GPU kernel (FlashAttention-style, I/O-aware)
# ---------------------------------------------------------------------------
try:
    from jax.experimental import pallas as pl
    from jax.experimental.pallas import triton as plgpu
    _PALLAS_GPU_AVAILABLE = True
except Exception:
    _PALLAS_GPU_AVAILABLE = False


# ---------------------------------------------------------------------------
# Public tuning constants
# ---------------------------------------------------------------------------

# TPU v5e: 128-wide systolic arrays β†’ BLK=128 saturates MXU
TPU_BLOCK_SIZE: int = 128

# GPU: warp size=32, tensor cores tile 16Γ—16 or 8Γ—16 β†’ BLK=64 is safe default
# cuDNN FlashAttention internally tiles at 64 or 128 depending on head dim
GPU_BLOCK_SIZE: int = 64

# ---------------------------------------------------------------------------
# Backend detection
# ---------------------------------------------------------------------------

def _detect_backend() -> str:
    """

    Detect the current JAX backend.

    Returns 'tpu', 'gpu', or 'cpu'.

    """
    try:
        backend = jax.default_backend().lower()
        if 'tpu' in backend:
            return 'tpu'
        elif 'gpu' in backend or 'cuda' in backend:
            return 'gpu'
        return 'cpu'
    except Exception:
        return 'cpu'


# ---------------------------------------------------------------------------
# Pallas/Triton GPU kernel β€” I/O-aware FlashAttention-style GQA SWA
# ---------------------------------------------------------------------------
#
# This is a from-scratch FlashAttention-2 style kernel:
#   - Tiles Q into BLOCK_Q-sized blocks, K/V into BLOCK_K-sized blocks.
#   - Grid = (batch, kv_head, num_q_blocks). Each program instance owns one
#     Q block for one KV head (covering its G query-head siblings via GQA
#     broadcast inside the kernel β€” K/V are NEVER duplicated in HBM).
#   - Online softmax: running max `m`, running sum `l`, running weighted
#     accumulator `acc` are carried across the K-block loop via
#     jax.lax.fori_loop. The full [BLK_Q, S] or [BLK_Q, window] score matrix
#     is NEVER materialized β€” only one [BLK_Q, BLOCK_K] tile lives in VMEM
#     at a time. This is the actual "I/O-aware" property: HBM traffic is
#     O(S) reads of Q/K/V blocks, not O(S^2) score matrix writes.
#   - Sliding-window + causal masking is applied per K-block using the same
#     relative-offset trick as the existing XLA kernel (b-independent delta),
#     so only K-blocks that intersect [q_pos - W + 1, q_pos] are visited β€”
#     blocks fully outside the window are skipped via the loop bounds, not
#     just masked, which is where the real compute savings come from vs the
#     existing XLA block-tiled kernel (which still computes+masks every
#     block inside a fixed kv_len window).
#
# Custom VJP: backward recomputes scores per (Q-block, K-block) pair from
# saved Q, K, V, O, m, l (NOT saved scores/probs β€” that's the whole point,
# same trick as FlashAttention). This keeps backward memory O(S) instead of
# O(S * window).
# ---------------------------------------------------------------------------

_PALLAS_BLOCK_Q = 64
_PALLAS_BLOCK_K = 64


def _gpu_compute_capability() -> Optional[tuple]:
    """Returns (major, minor) compute capability of the current GPU, or None

    if it can't be determined. Used to gate the Pallas/Triton path, which

    JAX only supports on Ampere (SM 8.0) and newer β€” Turing (T4, SM 7.5) and

    older will FAIL_PRECONDITION at Triton compile time, not at import time,

    so we must check this explicitly before attempting the kernel."""
    try:
        dev = jax.devices('gpu')[0]
        # jaxlib exposes this via device_kind (e.g. "Tesla T4", "NVIDIA A100")
        # or via compute_capability on newer jaxlib versions.
        cc = getattr(dev, 'compute_capability', None)
        if cc is not None:
            major, minor = str(cc).split('.')[:2]
            return (int(major), int(minor))
        return None
    except Exception:
        return None


_GPU_COMPUTE_CAPABILITY = None  # lazily populated on first check


def _pallas_supported(D: int, dtype) -> bool:
    """Conservative gate: only use the Pallas path for configs we've reasoned

    through (head_dim multiple of 16 for tensor-core alignment, fp16/bf16/fp32,

    Ampere-or-newer GPU). Anything else falls back to the cuDNN/XLA path

    automatically."""
    global _GPU_COMPUTE_CAPABILITY
    if os.environ.get('VEYLON_DISABLE_PALLAS_ATTN', '0') == '1':
        return False
    if not _PALLAS_GPU_AVAILABLE:
        return False
    if D % 16 != 0:
        return False
    if dtype not in (jnp.float16, jnp.bfloat16, jnp.float32):
        return False
    if _GPU_COMPUTE_CAPABILITY is None:
        _GPU_COMPUTE_CAPABILITY = _gpu_compute_capability() or (0, 0)
    if _GPU_COMPUTE_CAPABILITY < (8, 0):
        # Triton (Pallas GPU backend) requires Ampere or newer. T4 (7.5),
        # V100 (7.0), P100 (6.0) all fail here β€” this is a hard hardware
        # limit, not a bug, so we skip Pallas entirely rather than let it
        # crash through a full Triton compile attempt.
        return False
    return True


def _fa_fwd_kernel(

    q_ref, k_ref, v_ref,       # inputs, VMEM-resident blocks

    o_ref, m_ref, l_ref,       # outputs

    *,

    window: int,

    block_q: int,

    block_k: int,

    seq_len: int,

    scale: float,

):
    """

    Pallas kernel body β€” one program instance handles ONE (batch, kv_head,

    q_block) triple, looping internally over the K-blocks that intersect

    the causal + sliding-window range for this Q block.



    Ref shapes (per-program, already sliced by BlockSpec / index_map):

      q_ref : [block_q, D]      (single query head's slice β€” see note below)

      k_ref : [seq_len, D]      (full K for this batch/kv_head; we slice

                                  inside the loop via pl.load with dynamic

                                  start so only ONE [block_k, D] tile is

                                  actually resident in VMEM at a time)

      v_ref : [seq_len, D]      (same as k_ref)

      o_ref : [block_q, D]      (output accumulator, written once at end)

      m_ref, l_ref : [block_q, 1]  (running softmax stats, scratch)

    """
    q_block_idx = pl.program_id(2)
    q_start = q_block_idx * block_q

    q = q_ref[...].astype(jnp.float32) * scale  # [block_q, D]

    m_i = jnp.full((block_q, 1), -jnp.inf, dtype=jnp.float32)
    l_i = jnp.zeros((block_q, 1), dtype=jnp.float32)
    acc = jnp.zeros_like(q)

    # Range of K-blocks that can possibly intersect this Q-block's
    # causal+window range. Query positions in this block span
    # [q_start, q_start + block_q - 1]. Each attends to
    # [q_pos - window + 1, q_pos]. So the union over the block spans
    # [q_start - window + 1, q_start + block_q - 1].
    k_lo = jnp.maximum(0, q_start - window + 1)
    k_hi = jnp.minimum(seq_len, q_start + block_q)  # exclusive, causal cap
    first_k_block = k_lo // block_k
    num_k_blocks = (k_hi - first_k_block * block_k + block_k - 1) // block_k
    num_k_blocks = jnp.maximum(num_k_blocks, 1)

    def body(i, carry):
        m_i, l_i, acc = carry
        k_start = (first_k_block + i) * block_k

        k_blk = pl.load(
            k_ref, (pl.dslice(k_start, block_k), slice(None))
        ).astype(jnp.float32)  # [block_k, D]
        v_blk = pl.load(
            v_ref, (pl.dslice(k_start, block_k), slice(None))
        ).astype(jnp.float32)  # [block_k, D]

        scores = jnp.dot(
            q, k_blk.T, preferred_element_type=jnp.float32
        )  # [block_q, block_k]

        q_pos = q_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 0)
        k_pos = k_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 1)
        causal_ok = k_pos <= q_pos
        window_ok = (q_pos - k_pos) < window
        bounds_ok = k_pos < seq_len
        mask = causal_ok & window_ok & bounds_ok
        scores = jnp.where(mask, scores, -jnp.inf)

        m_ij = jnp.max(scores, axis=-1, keepdims=True)           # [block_q, 1]
        m_new = jnp.maximum(m_i, m_ij)
        # Guard against all-masked rows (m_new stays -inf) -> exp(0)=1 issue
        m_new_safe = jnp.where(m_new == -jnp.inf, 0.0, m_new)

        p = jnp.exp(scores - m_new_safe)                          # [block_q, block_k]
        p = jnp.where(mask, p, 0.0)

        alpha = jnp.exp(jnp.where(m_i == -jnp.inf, m_new_safe, m_i) - m_new_safe)
        l_new = l_i * alpha + jnp.sum(p, axis=-1, keepdims=True)
        acc_new = acc * alpha + jnp.dot(p, v_blk, preferred_element_type=jnp.float32)

        return m_new, l_new, acc_new

    m_i, l_i, acc = jax.lax.fori_loop(0, num_k_blocks, body, (m_i, l_i, acc))

    l_safe = jnp.where(l_i > 0, l_i, 1.0)
    out = acc / l_safe

    o_ref[...] = out.astype(o_ref.dtype)
    m_ref[...] = m_i
    l_ref[...] = l_i


def _pallas_fwd_single_head(q, k, v, window, block_q, block_k, scale):
    """

    Runs the Pallas forward kernel for ONE query head against its KV head.

    q: [B, S, D]   k, v: [B, S, D]   (already the per-head slices)

    Returns: out [B, S, D], m [B, S, 1], l [B, S, 1]  (m, l saved for bwd)

    """
    B, S, D = map(int, q.shape)
    n_q_blocks = (S + block_q - 1) // block_q
    S_pad = n_q_blocks * block_q

    q_p = jnp.pad(q, ((0, 0), (0, S_pad - S), (0, 0)))
    # K/V padded on the right only; kernel bounds-checks k_pos < seq_len so
    # right-padding is safe (never read past the pad due to k_hi clamp), but
    # we still pad to a multiple of block_k so pl.load's static block shape
    # never reads out-of-bounds memory.
    n_k_blocks_total = (S + block_k - 1) // block_k
    S_pad_k = n_k_blocks_total * block_k
    k_p = jnp.pad(k, ((0, 0), (0, S_pad_k - S), (0, 0)))
    v_p = jnp.pad(v, ((0, 0), (0, S_pad_k - S), (0, 0)))

    kernel = partial(
        _fa_fwd_kernel,
        window=window, block_q=block_q, block_k=block_k,
        seq_len=S, scale=scale,
    )

    out, m, l = pl.pallas_call(
        kernel,
        grid=(B, 1, n_q_blocks),
        in_specs=[
            pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),
            pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)),
            pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)),
        ],
        out_specs=[
            pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),
            pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)),
            pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)),
        ],
        out_shape=[
            jax.ShapeDtypeStruct((B, S_pad, D), q.dtype),
            jax.ShapeDtypeStruct((B, S_pad, 1), jnp.float32),
            jax.ShapeDtypeStruct((B, S_pad, 1), jnp.float32),
        ],
    )(q_p, k_p, v_p)

    return out[:, :S, :], m[:, :S, :], l[:, :S, :]


def _fa_bwd_kernel(

    q_ref, k_ref, v_ref, o_ref, do_ref, m_ref, l_ref,

    dq_ref, dk_ref, dv_ref,

    *,

    window: int,

    block_q: int,

    block_k: int,

    seq_len: int,

    scale: float,

):
    """

    Backward kernel β€” one program per (batch, k_block). Recomputes scores

    for each intersecting Q-block on the fly (from saved Q, K, V, m, l) and

    accumulates dK/dV. dQ is accumulated via a separate pass below since it

    is indexed by q_block, not k_block (standard FlashAttention-2 backward

    split to avoid atomic adds across programs).

    """
    k_block_idx = pl.program_id(2)
    k_start = k_block_idx * block_k

    k_blk = k_ref[...].astype(jnp.float32)   # [block_k, D]
    v_blk = v_ref[...].astype(jnp.float32)   # [block_k, D]

    dk_acc = jnp.zeros_like(k_blk)
    dv_acc = jnp.zeros_like(v_blk)

    # Q-blocks that can intersect this K-block: q_pos >= k_pos (causal) and
    # q_pos - k_pos < window. q spans [k_start, seq_len-1] roughly, capped
    # by window on the upper side: q_pos < k_start + block_k + window - 1.
    q_lo = k_start
    q_hi = jnp.minimum(seq_len, k_start + block_k + window - 1)
    first_q_block = q_lo // block_q
    num_q_blocks = (q_hi - first_q_block * block_q + block_q - 1) // block_q
    num_q_blocks = jnp.maximum(num_q_blocks, 1)

    def body(i, carry):
        dk_acc, dv_acc = carry
        q_start = (first_q_block + i) * block_q

        q_blk = pl.load(q_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32) * scale
        do_blk = pl.load(do_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)
        m_blk = pl.load(m_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)
        l_blk = pl.load(l_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)
        o_blk = pl.load(o_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)

        scores = jnp.dot(q_blk, k_blk.T, preferred_element_type=jnp.float32)

        q_pos = q_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 0)
        k_pos = k_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 1)
        causal_ok = k_pos <= q_pos
        window_ok = (q_pos - k_pos) < window
        bounds_ok = k_pos < seq_len
        mask = causal_ok & window_ok & bounds_ok

        l_safe = jnp.where(l_blk > 0, l_blk, 1.0)
        p = jnp.where(mask, jnp.exp(scores - m_blk), 0.0) / l_safe  # [block_q, block_k]

        dv_acc = dv_acc + jnp.dot(p.T, do_blk, preferred_element_type=jnp.float32)

        dp = jnp.dot(do_blk, v_blk.T, preferred_element_type=jnp.float32)  # [block_q, block_k]
        Di = jnp.sum(do_blk * o_blk, axis=-1, keepdims=True)  # [block_q, 1]
        dscores = p * (dp - Di)
        dscores = jnp.where(mask, dscores, 0.0)

        dk_acc = dk_acc + jnp.dot(dscores.T, q_blk, preferred_element_type=jnp.float32) * scale

        return dk_acc, dv_acc

    dk_acc, dv_acc = jax.lax.fori_loop(0, num_q_blocks, body, (dk_acc, dv_acc))

    dk_ref[...] = dk_acc.astype(dk_ref.dtype)
    dv_ref[...] = dv_acc.astype(dv_ref.dtype)


def _fa_bwd_dq_kernel(

    q_ref, k_ref, v_ref, o_ref, do_ref, m_ref, l_ref,

    dq_ref,

    *,

    window: int,

    block_q: int,

    block_k: int,

    seq_len: int,

    scale: float,

):
    """Separate pass computing dQ, one program per (batch, q_block), looping

    over intersecting K-blocks. Kept separate from the dK/dV kernel because

    dQ is naturally indexed by q_block and dK/dV by k_block β€” fusing both

    into one kernel would need cross-program atomics, which Pallas/Triton

    doesn't support cleanly. Recomputation cost (~2x score matmuls total

    across both passes) is the standard FlashAttention-2 backward tradeoff."""
    q_block_idx = pl.program_id(2)
    q_start = q_block_idx * block_q

    q_blk = q_ref[...].astype(jnp.float32) * scale
    do_blk = do_ref[...].astype(jnp.float32)
    m_blk = m_ref[...].astype(jnp.float32)
    l_blk = l_ref[...].astype(jnp.float32)
    o_blk = o_ref[...].astype(jnp.float32)

    dq_acc = jnp.zeros_like(q_blk)

    k_lo = jnp.maximum(0, q_start - window + 1)
    k_hi = jnp.minimum(seq_len, q_start + block_q)
    first_k_block = k_lo // block_k
    num_k_blocks = (k_hi - first_k_block * block_k + block_k - 1) // block_k
    num_k_blocks = jnp.maximum(num_k_blocks, 1)

    def body(i, dq_acc):
        k_start = (first_k_block + i) * block_k
        k_blk = pl.load(k_ref, (pl.dslice(k_start, block_k), slice(None))).astype(jnp.float32)
        v_blk = pl.load(v_ref, (pl.dslice(k_start, block_k), slice(None))).astype(jnp.float32)

        scores = jnp.dot(q_blk, k_blk.T, preferred_element_type=jnp.float32)

        q_pos = q_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 0)
        k_pos = k_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 1)
        causal_ok = k_pos <= q_pos
        window_ok = (q_pos - k_pos) < window
        bounds_ok = k_pos < seq_len
        mask = causal_ok & window_ok & bounds_ok

        l_safe = jnp.where(l_blk > 0, l_blk, 1.0)
        p = jnp.where(mask, jnp.exp(scores - m_blk), 0.0) / l_safe

        dp = jnp.dot(do_blk, v_blk.T, preferred_element_type=jnp.float32)
        Di = jnp.sum(do_blk * o_blk, axis=-1, keepdims=True)
        dscores = p * (dp - Di)
        dscores = jnp.where(mask, dscores, 0.0)

        dq_acc = dq_acc + jnp.dot(dscores, k_blk, preferred_element_type=jnp.float32) * scale
        return dq_acc

    dq_acc = jax.lax.fori_loop(0, num_k_blocks, body, dq_acc)
    dq_ref[...] = dq_acc.astype(dq_ref.dtype)


def _pallas_bwd_single_head(q, k, v, o, do, m, l, window, block_q, block_k, scale):
    """Runs both backward kernels (dK/dV and dQ) for one query/KV head pair."""
    B, S, D = map(int, q.shape)
    n_q_blocks = (S + block_q - 1) // block_q
    n_k_blocks = (S + block_k - 1) // block_k
    S_pad_q = n_q_blocks * block_q
    S_pad_k = n_k_blocks * block_k

    pad_q = lambda x, fill=0.0: jnp.pad(x, ((0, 0), (0, S_pad_q - S), (0, 0)), constant_values=fill)
    pad_k = lambda x: jnp.pad(x, ((0, 0), (0, S_pad_k - S), (0, 0)))

    q_p, o_p, do_p = pad_q(q), pad_q(o), pad_q(do)
    m_p = jnp.pad(m, ((0, 0), (0, S_pad_q - S), (0, 0)), constant_values=jnp.inf)
    l_p = jnp.pad(l, ((0, 0), (0, S_pad_q - S), (0, 0)), constant_values=1.0)
    k_p, v_p = pad_k(k), pad_k(v)

    dkdv_kernel = partial(
        _fa_bwd_kernel, window=window, block_q=block_q, block_k=block_k,
        seq_len=S, scale=scale,
    )
    dk, dv = pl.pallas_call(
        dkdv_kernel,
        grid=(B, 1, n_k_blocks),
        in_specs=[
            pl.BlockSpec((None, S_pad_q, D), lambda b, h, j: (b, 0, 0)),   # q (full, sliced inside)
            pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)),    # k block
            pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)),    # v block
            pl.BlockSpec((None, S_pad_q, D), lambda b, h, j: (b, 0, 0)),   # o (full)
            pl.BlockSpec((None, S_pad_q, D), lambda b, h, j: (b, 0, 0)),   # do (full)
            pl.BlockSpec((None, S_pad_q, 1), lambda b, h, j: (b, 0, 0)),   # m (full)
            pl.BlockSpec((None, S_pad_q, 1), lambda b, h, j: (b, 0, 0)),   # l (full)
        ],
        out_specs=[
            pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)),
            pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)),
        ],
        out_shape=[
            jax.ShapeDtypeStruct((B, S_pad_k, D), k.dtype),
            jax.ShapeDtypeStruct((B, S_pad_k, D), v.dtype),
        ],
    )(q_p, k_p, v_p, o_p, do_p, m_p, l_p)

    dq_kernel = partial(
        _fa_bwd_dq_kernel, window=window, block_q=block_q, block_k=block_k,
        seq_len=S, scale=scale,
    )
    dq = pl.pallas_call(
        dq_kernel,
        grid=(B, 1, n_q_blocks),
        in_specs=[
            pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),   # q block
            pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)),   # k (full)
            pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)),   # v (full)
            pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),   # o block
            pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),   # do block
            pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)),   # m block
            pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)),   # l block
        ],
        out_specs=pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),
        out_shape=jax.ShapeDtypeStruct((B, S_pad_q, D), q.dtype),
    )(q_p, k_p, v_p, o_p, do_p, m_p, l_p)

    return dq[:, :S, :], dk[:, :S, :], dv[:, :S, :]


@partial(jax.custom_vjp, nondiff_argnums=(3, 4, 5, 6))
def _pallas_gqa_swa_head(q, k, v, window, block_q, block_k, scale):
    """Single (query-head, kv-head) FlashAttention call with custom VJP.

    q, k, v: [B, S, D] for ONE head pair (GQA broadcast handled by caller)."""
    out, _, _ = _pallas_fwd_single_head(q, k, v, window, block_q, block_k, scale)
    return out


def _pallas_gqa_swa_head_fwd(q, k, v, window, block_q, block_k, scale):
    out, m, l = _pallas_fwd_single_head(q, k, v, window, block_q, block_k, scale)
    return out, (q, k, v, out, m, l)


def _pallas_gqa_swa_head_bwd(window, block_q, block_k, scale, residuals, dout):
    q, k, v, out, m, l = residuals
    dq, dk, dv = _pallas_bwd_single_head(
        q, k, v, out, dout, m, l, window, block_q, block_k, scale
    )
    return dq, dk, dv


_pallas_gqa_swa_head.defvjp(_pallas_gqa_swa_head_fwd, _pallas_gqa_swa_head_bwd)


def _pallas_flash_gqa_swa(

    q: jnp.ndarray,

    k: jnp.ndarray,

    v: jnp.ndarray,

    window_size: int,

    block_q: int = _PALLAS_BLOCK_Q,

    block_k: int = _PALLAS_BLOCK_K,

) -> jnp.ndarray:
    """

    I/O-aware FlashAttention-style GQA SWA, entry point for the Pallas path.



    q: [B, Hq,  S, D]

    k: [B, Hkv, S, D]

    v: [B, Hkv, S, D]



    GQA is handled by vmapping the single-head kernel over KV heads, and

    within each KV head over its G query-head siblings β€” K/V are never

    physically duplicated; only the (small) grid iterates over G.

    """
    B, Hq, S, D = map(int, q.shape)
    _, Hkv, Sk, Dk = map(int, k.shape)
    if Hq % Hkv != 0:
        raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")
    G = Hq // Hkv
    scale = 1.0 / math.sqrt(float(D))

    # [B, Hkv, G, S, D]
    q_g = q.reshape(B, Hkv, G, S, D)

    # vmap over (Hkv, G): each call gets q[B,S,D] for one query head and the
    # matching k/v[B,S,D] for its KV head (broadcast across G, no copy of
    # the underlying K/V buffer beyond what vmap's batching rule does).
    def per_kv_head(q_kv, k_h, v_h):
        # q_kv: [G, B, S, D]  k_h, v_h: [B, S, D]
        fn = lambda qh: _pallas_gqa_swa_head(qh, k_h, v_h, window_size, block_q, block_k, scale)
        return jax.vmap(fn)(q_kv)  # [G, B, S, D]

    q_g_t = q_g.transpose(1, 2, 0, 3, 4)  # [Hkv, G, B, S, D]
    k_t = k.transpose(1, 0, 2, 3)          # [Hkv, B, S, D]
    v_t = v.transpose(1, 0, 2, 3)

    out = jax.vmap(per_kv_head)(q_g_t, k_t, v_t)  # [Hkv, G, B, S, D]
    out = out.transpose(2, 0, 1, 3, 4).reshape(B, Hq, S, D)  # [B, Hq, S, D]
    return out.astype(q.dtype)


# ---------------------------------------------------------------------------
# GPU path: JAX cuDNN FlashAttention via jax.nn.dot_product_attention
# ---------------------------------------------------------------------------

def _gpu_flash_gqa_swa(

    q: jnp.ndarray,

    k: jnp.ndarray,

    v: jnp.ndarray,

    window_size: int,

    use_remat: bool = True,

) -> jnp.ndarray:
    """

    GPU-optimized GQA Sliding Window Attention.



    cuDNN FlashAttention does NOT reliably support SWA masking across all

    JAX/cuDNN versions. Instead we use:



      - cuDNN for the raw QK^T matmul + softmax + V aggregation

        (via jax.nn.dot_product_attention without masking)

        only when window_size >= S (full attention β€” no masking needed).



      - XLA block-tiled path with GPU-friendly block_size=64 for SWA

        (window_size < S). This avoids the cuDNN engine config error

        while still running fast on CUDA via XLA's GPU backend.



    Both paths use BF16 compute and avoid materializing [S,S] matrices.

    """
    B,  Hq,  S,  D  = map(int, q.shape)
    _B, Hkv, Sk, Dk = map(int, k.shape)

    if Hq % Hkv != 0:
        raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")

    G     = Hq // Hkv
    scale = 1.0 / math.sqrt(float(D))

    # ── Full attention (no window mask): use cuDNN ────────────────────────────
    if window_size >= S:
        # [B, S, H, D] layout for cuDNN
        q_s = q.transpose(0, 2, 1, 3)
        k_s = k.transpose(0, 2, 1, 3)
        v_s = v.transpose(0, 2, 1, 3)

        def _full_attn(q_, k_, v_):
            return jax.nn.dot_product_attention(
                q_, k_, v_,
                scale=scale,
                is_causal=True,
                implementation='cudnn',
            )

        if use_remat:
            # jax.checkpoint recomputes _full_attn on the backward pass; the
            # function itself still runs exactly ONCE per forward pass.
            result = jax.checkpoint(_full_attn)(q_s, k_s, v_s)
        else:
            result = _full_attn(q_s, k_s, v_s)
        return result.transpose(0, 2, 1, 3).astype(q.dtype)

    # ── SWA path: XLA block-tiled kernel (GPU block_size=64) ─────────────────
    # This is the fast path for SWA on GPU.
    # XLA compiles this to efficient CUDA matmuls with BF16 tensor cores.
    # GPU_BLOCK_SIZE=64 matches CUDA warp/tensor-core tiling.
    return _block_gqa_swa(
        q=q,
        k=k,
        v=v,
        window_size=window_size,
        block_size=GPU_BLOCK_SIZE,
        use_remat=use_remat,
    )


# ---------------------------------------------------------------------------
# GPU decode path: single-token GQA SWA for inference
# ---------------------------------------------------------------------------

def _gpu_decode_swa(

    q: jnp.ndarray,

    k: jnp.ndarray,

    v: jnp.ndarray,

) -> jnp.ndarray:
    """

    GPU decode path for single-token generation (S=1).



    q: [B, Hq,  1, D]

    k: [B, Hkv, W, D]

    v: [B, Hkv, W, D]



    Returns: [B, Hq, 1, D]

    """
    B, Hq, S, D  = map(int, q.shape)
    _, Hkv, W, _ = map(int, k.shape)
    G = Hq // Hkv
    scale = 1.0 / math.sqrt(float(D))

    # [B, 1, Hq, D] and [B, W, Hkv, D] for cuDNN
    q_s = q.transpose(0, 2, 1, 3)
    k_s = k.transpose(0, 2, 1, 3)
    v_s = v.transpose(0, 2, 1, 3)

    try:
        # S=1 decode: all W tokens are past, no masking needed
        # cuDNN handles this as a batched GEMV β€” very fast
        out = jax.nn.dot_product_attention(
            q_s, k_s, v_s,
            scale=scale,
            is_causal=False,
            implementation='cudnn',
        )
        return out.transpose(0, 2, 1, 3).astype(q.dtype)
    except Exception:
        pass

    # Fallback: native GQA einsum (always works)
    q_g    = q.reshape(B, Hkv, G, 1, D)
    scores = (
        jnp.einsum('bngqd,bnkd->bngqk',
                   q_g.astype(jnp.float32),
                   k.astype(jnp.float32))
        * scale
    )
    probs  = jax.nn.softmax(scores, axis=-1)
    out_g  = jnp.einsum('bngqk,bnkd->bngqd', probs, v.astype(jnp.float32))
    return out_g.reshape(B, Hq, 1, D).astype(q.dtype)


# ---------------------------------------------------------------------------
# Utility helpers
# ---------------------------------------------------------------------------

def _next_power_of_2(x: int) -> int:
    x = int(x)
    if x <= 1:
        return 1
    return 1 << (x - 1).bit_length()


def apply_rope(

    x: jnp.ndarray,

    cos: jnp.ndarray,

    sin: jnp.ndarray,

    offset: int = 0,

) -> jnp.ndarray:
    """

    Apply Rotary Position Encoding to [..., S, D].

    cos/sin tables are pre-built for D//2 (half the head dim).

    Uses dynamic_slice β€” safe under jit and XLA outside of attention body.

    """
    if x.ndim < 2:
        raise ValueError(f"apply_rope: expected β‰₯2-D input, got shape {x.shape}")
    d = int(x.shape[-1])
    if d % 2 != 0:
        raise ValueError(f"apply_rope: head dim must be even, got {d}")
    half = d // 2
    seq_len = int(x.shape[-2])
    if offset + seq_len > int(cos.shape[0]):
        raise ValueError(
            f"apply_rope: RoPE table too small β€” "
            f"offset={offset}, seq_len={seq_len}, table_size={cos.shape[0]}"
        )
    cos_s = jax.lax.dynamic_slice(cos, (offset, 0), (seq_len, half)).astype(x.dtype)
    sin_s = jax.lax.dynamic_slice(sin, (offset, 0), (seq_len, half)).astype(x.dtype)
    while cos_s.ndim < x.ndim - 1:
        cos_s = cos_s[None]
        sin_s = sin_s[None]
    x1, x2 = x[..., :half], x[..., half:]
    return jnp.concatenate(
        [x1 * cos_s - x2 * sin_s, x1 * sin_s + x2 * cos_s],
        axis=-1,
    )


# ---------------------------------------------------------------------------
# Core kernel: block-tiled native GQA SWA
# ---------------------------------------------------------------------------

def _block_gqa_swa(

    q: jnp.ndarray,

    k: jnp.ndarray,

    v: jnp.ndarray,

    window_size: int,

    block_size: int = TPU_BLOCK_SIZE,

    use_remat: bool = True,

) -> jnp.ndarray:
    """

    Block-tiled Native GQA Sliding Window Attention.



    This function is the replacement for the scan+dynamic_slice approach.

    All tensor shapes inside the map body are STATIC β€” XLA never sees a

    gather-inside-loop pattern, so the hidden [S,B,H,W,D] buffer cannot form.



    Parameters

    ----------

    q           : [B, Hq,  S, D]  β€” any float dtype (BF16 in production)

    k           : [B, Hkv, S, D]  β€” same dtype

    v           : [B, Hkv, S, D]  β€” same dtype

    window_size : causal window W; token t attends to [max(0, t-W+1), t]

    block_size  : query block size BLK (tune to 128 for TPU v5e)

    use_remat   : wrap map body in jax.checkpoint (recommended for training)



    Returns

    -------

    [B, Hq, S, D] β€” same dtype as q

    """
    if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
        raise ValueError(
            f"_block_gqa_swa: expected 4-D inputs, "
            f"got q={q.shape}  k={k.shape}  v={v.shape}"
        )

    B,  Hq,  S,  D   = map(int, q.shape)
    _B, Hkv, Sk, Dk  = map(int, k.shape)

    if _B != B:
        raise ValueError(f"Batch mismatch: q B={B}, k B={_B}")
    if Dk != D:
        raise ValueError(f"Head-dim mismatch: q D={D}, k D={Dk}")
    if Sk != S:
        raise ValueError(f"Sequence-length mismatch: q S={S}, k S={Sk}")
    if Hq % Hkv != 0:
        raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")

    G   = Hq // Hkv                  # queries per KV head
    W   = int(window_size)
    BLK = int(block_size)
    if BLK <= 0:
        raise ValueError(f"block_size must be > 0, got {BLK}")

    scale = 1.0 / math.sqrt(float(D))

    # ── Pad sequence to a multiple of BLK ────────────────────────────────────
    n_blocks = (S + BLK - 1) // BLK
    S_pad    = n_blocks * BLK          # β‰₯ S, multiple of BLK

    # ── Pad q along sequence axis ─────────────────────────────────────────────
    # [B, Hq, S_pad, D]
    # The extra (S_pad - S) tokens are zero-padded and trimmed from output.
    q_pad = jnp.pad(q, ((0, 0), (0, 0), (0, S_pad - S), (0, 0)))

    # ── Pad k/v on the LEFT by (W-1) for causal alignment ────────────────────
    # [B, Hkv, S_pad + W - 1, D]
    #
    # After padding, for query block starting at global position b*BLK:
    #   kv slice starts at offset b*BLK in k_pad
    #   kv slice has static length  BLK + W - 1
    #   It covers original positions  [b*BLK - (W-1), b*BLK + BLK - 1]
    #   which, after clipping to β‰₯ 0, is exactly the causal window.
    kv_pad_len = S_pad + W - 1        # total padded kv length (static)
    lpad       = W - 1                # left zero-padding width

    k_pad = jnp.pad(k, ((0, 0), (0, 0), (lpad, S_pad - S), (0, 0)))  # [B,Hkv,kv_pad_len,D]
    v_pad = jnp.pad(v, ((0, 0), (0, 0), (lpad, S_pad - S), (0, 0)))

    # ── Static per-block mask ─────────────────────────────────────────────────
    # Compute a boolean mask of shape [BLK, BLK + W - 1].
    # Entry [q_local, k_local] is True (mask=attend) when:
    #   (a) k is not from left-padding   β†’ k_local >= W - 1 - (something)
    # We cannot compute the absolute positions statically because they depend
    # on block index b.  So we record the RELATIVE offsets and apply the
    # offset inside the map body using only static arithmetic on local indices.
    #
    # q_abs  = b*BLK + q_local          (q_local in 0..BLK-1)
    # k_abs  = b*BLK + k_local - lpad   (k_local in 0..BLK+W-2)
    #
    # Attend iff:
    #   k_abs >= 0                (not left padding)
    #   k_abs <= q_abs            (causal)
    #   q_abs - k_abs < W         (within window)
    #
    # q_abs - k_abs  = q_local - k_local + lpad   (b cancels out!)
    # k_abs >= 0     ↔ k_local >= lpad - b*BLK   (depends on b β†’ handle in body)
    # k_abs <= q_abs ↔ k_local - q_local <= lpad  (b cancels out! β†’ STATIC)
    #
    # Only the "k_abs >= 0" condition depends on b (first block only).
    # We handle it cheaply with a dynamic mask inside the body.
    # Everything else is b-independent and can be PRECOMPUTED once.

    kv_len = BLK + W - 1   # static length of kv slice per block

    q_local  = jnp.arange(BLK,    dtype=jnp.int32)   # [BLK]
    kv_local = jnp.arange(kv_len, dtype=jnp.int32)   # [kv_len]

    # Relative offset: delta[q, k] = q_local[q] - kv_local[k] + lpad
    #   = q_abs - k_abs    (b-independent)
    delta = q_local[:, None] - kv_local[None, :] + lpad  # [BLK, kv_len]

    # Static masks (b-independent)
    static_future_mask  = delta < 0          # k is in the future (causal)  β†’ mask
    static_window_mask  = delta >= W         # k is too far back (window)   β†’ mask
    static_base_mask    = static_future_mask | static_window_mask  # [BLK, kv_len]

    # ── Map body ──────────────────────────────────────────────────────────────

    def process_block(b: jnp.ndarray) -> jnp.ndarray:
        """

        b : scalar int32 block index in [0, n_blocks)



        Tensor shapes inside this function are ALL STATIC:

            q_blk   : [B, Hq,  BLK,    D]

            k_blk   : [B, Hkv, kv_len, D]

            v_blk   : [B, Hkv, kv_len, D]

            scores  : [B, Hkv, G, BLK, kv_len]

            probs   : same



        XLA has NO gather-inside-loop here.  The dynamic_slice start index

        `b * BLK` is a scalar multiply β€” XLA lowers this to a simple pointer

        offset, not a buffer materialisation.

        """
        blk_start = b * BLK  # scalar traced int32

        # Static-shape slices β€” the KEY difference from the scan approach.
        # XLA sees shapes (B,Hq,BLK,D) and (B,Hkv,kv_len,D) as compile-time
        # constants.  It cannot stage all blocks simultaneously because
        # lax.map gives it one block at a time with no output accumulation.
        q_blk = jax.lax.dynamic_slice(q_pad, (0, 0, blk_start,        0), (B, Hq,  BLK,    D))
        k_blk = jax.lax.dynamic_slice(k_pad, (0, 0, blk_start,        0), (B, Hkv, kv_len, D))
        v_blk = jax.lax.dynamic_slice(v_pad, (0, 0, blk_start,        0), (B, Hkv, kv_len, D))

        # ── Native GQA reshape (zero-copy view) ──────────────────────────
        # [B, Hq, BLK, D] β†’ [B, Hkv, G, BLK, D]
        q_g = q_blk.reshape(B, Hkv, G, BLK, D)

        # ── Dot-product scores ────────────────────────────────────────────
        # [B, Hkv, G, BLK, kv_len]   ← STATIC, FUSED by XLA
        scores = (
            jnp.einsum(
                "bngqd,bnkd->bngqk",
                q_g.astype(jnp.float32),
                k_blk.astype(jnp.float32),
            )
            * scale
        )

        # ── Masking ───────────────────────────────────────────────────────
        # (a) b-independent mask (precomputed, fused as constant)
        mask = static_base_mask  # [BLK, kv_len]

        # (b) Left-padding mask: k_abs < 0  ↔  kv_local < lpad - blk_start
        #     Only non-trivial for b=0 (the very first block).
        #     For all subsequent blocks, lpad - blk_start < 0, so no extra masking.
        #     We compute it dynamically but it is a single scalar comparison
        #     broadcast β€” XLA will constant-fold it for b > 0 at runtime.
        leftpad_cutoff = lpad - blk_start   # scalar int32 (may be negative)
        leftpad_mask   = kv_local[None, :] < leftpad_cutoff  # [1, kv_len]
        mask = mask | leftpad_mask  # [BLK, kv_len]

        scores = jnp.where(
            mask[None, None, None, :, :],   # [1,1,1,BLK,kv_len]
            jnp.full_like(scores, -1e30),
            scores,
        )

        # ── Softmax + value aggregation (all FP32) ────────────────────────
        probs  = jax.nn.softmax(scores, axis=-1)            # [B,Hkv,G,BLK,kv_len]
        out_g  = jnp.einsum(
            "bngqk,bnkd->bngqd",
            probs,
            v_blk.astype(jnp.float32),
        )  # [B, Hkv, G, BLK, D]

        # ── Reshape + cast back to input dtype ────────────────────────────
        return out_g.reshape(B, Hq, BLK, D).astype(q.dtype)  # [B, Hq, BLK, D]

    # ── Gradient checkpointing ────────────────────────────────────────────────
    map_fn = jax.checkpoint(process_block) if use_remat else process_block

    # ── fori_loop: write each block directly into pre-allocated output ────────
    # lax.map returns [n_blocks, B, Hq, BLK, D] β€” XLA stages the ENTIRE stack
    # in HBM before the transpose+reshape.  For large n_blocks / batch this
    # wastes memory.
    #
    # lax.fori_loop carries a single [B, Hq, S_pad, D] output buffer and uses
    # dynamic_update_slice to write each block in-place.  XLA sees one static-
    # shape buffer (same size as the final output) instead of n_blocks copies.
    out_init = jnp.zeros((B, Hq, S_pad, D), dtype=q.dtype)

    def _write_block(b, out_buf):
        blk_out = map_fn(b)                             # [B, Hq, BLK, D]
        return jax.lax.dynamic_update_slice(
            out_buf,
            blk_out,
            (0, 0, b * BLK, 0),
        )

    out = jax.lax.fori_loop(0, n_blocks, _write_block, out_init)
    # out: [B, Hq, S_pad, D] β€” trim padding
    return out[:, :, :S, :]


# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------

def flash_splash_attention(

    q: jnp.ndarray,

    k: jnp.ndarray,

    v: jnp.ndarray,

    window_size: int,

    backend: Optional[str] = None,

    use_gqa: bool = True,

    start_pos: int = 0,

    block_size: int = TPU_BLOCK_SIZE,

    use_remat: bool = True,

) -> jnp.ndarray:
    """

    Block-tiled GQA-native Sliding Window Attention β€” main entry point.



    Automatically dispatches to the optimal kernel for the current backend:



    β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”¬β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”

    β”‚ Backend  β”‚ Kernel                                                   β”‚

    β”œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”Όβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€

    β”‚ TPU v5e  β”‚ _block_gqa_swa β€” block-tiled lax.map, static shapes,    β”‚

    β”‚          β”‚ BF16 matmul, gradient checkpoint per block               β”‚

    β”œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”Όβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€

    β”‚ GPU/CUDA β”‚ _gpu_flash_gqa_swa β€” cuDNN FlashAttention v2/v3 via      β”‚

    β”‚          β”‚ jax.nn.dot_product_attention, native GQA, SWA mask       β”‚

    β”‚          β”‚ Falls back to JAX XLA SDPA then explicit einsum          β”‚

    β”œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”Όβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€

    β”‚ CPU      β”‚ _block_gqa_swa β€” same as TPU (block_size=64 for cache)  β”‚

    β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”΄β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜



    Parameters

    ----------

    q, k, v      : [B, Hq, S, D] / [B, Hkv, S, D]

    window_size  : causal window W

    backend      : override auto-detection ('tpu', 'gpu', 'cpu', or None)

    use_gqa      : API-compat flag; GQA is always native here

    start_pos    : KV-cache offset for decode mode (training path: leave at 0)

    block_size   : query block size BLK (128 for TPU, 64 for GPU)

    use_remat    : gradient checkpointing (recommended for training)



    Returns

    -------

    [B, Hq, S, D] β€” same dtype as q

    """
    _ = use_gqa
    _ = start_pos

    # ── Backend detection ─────────────────────────────────────────────────────
    active_backend = (backend or _detect_backend()).lower()
    if 'tpu' in active_backend:
        active_backend = 'tpu'
    elif 'gpu' in active_backend or 'cuda' in active_backend:
        active_backend = 'gpu'
    else:
        active_backend = 'cpu'

    # ── Dispatch ──────────────────────────────────────────────────────────────
    if active_backend == 'gpu':
        B, Hq, S, D = map(int, q.shape)
        if _pallas_supported(D, q.dtype):
            try:
                return _pallas_flash_gqa_swa(
                    q, k, v,
                    window_size=int(window_size),
                    block_q=min(_PALLAS_BLOCK_Q, S) if S < _PALLAS_BLOCK_Q else _PALLAS_BLOCK_Q,
                    block_k=min(_PALLAS_BLOCK_K, S) if S < _PALLAS_BLOCK_K else _PALLAS_BLOCK_K,
                )
            except Exception as e:
                # Any Pallas/Triton compile or runtime failure (unsupported
                # GPU arch, block size mismatch, etc.) falls back silently to
                # the proven cuDNN/XLA path below β€” training never crashes
                # because of this optimization. Set
                # VEYLON_DEBUG_PALLAS_ATTN=1 to see what actually failed.
                if os.environ.get('VEYLON_DEBUG_PALLAS_ATTN', '0') == '1':
                    print(f"[veylon_attention] Pallas path failed, falling back: "
                          f"{type(e).__name__}: {e}")
        # cuDNN fused attention only accepts fp16/bf16/fp8 β€” fp32 inputs must
        # go through the plain-XLA fallback further down in
        # _gpu_flash_gqa_swa rather than crashing on the cuDNN dtype check.
        if q.dtype not in (jnp.float16, jnp.bfloat16):
            return _block_gqa_swa(
                q=q, k=k, v=v,
                window_size=int(window_size),
                block_size=int(GPU_BLOCK_SIZE),
                use_remat=use_remat,
            )
        return _gpu_flash_gqa_swa(
            q=q,
            k=k,
            v=v,
            window_size=int(window_size),
            use_remat=use_remat,
        )
    else:
        # TPU and CPU both use block-tiled JAX kernel
        # GPU_BLOCK_SIZE is used for CPU (cache-friendly), TPU_BLOCK_SIZE for TPU
        blk = block_size if active_backend == 'tpu' else GPU_BLOCK_SIZE
        return _block_gqa_swa(
            q=q,
            k=k,
            v=v,
            window_size=int(window_size),
            block_size=int(blk),
            use_remat=use_remat,
        )
def decode_swa(

    q: jnp.ndarray,

    k: jnp.ndarray,

    v: jnp.ndarray,

) -> jnp.ndarray:
    """

    Decode-time SWA for one-token generation.



    Dispatches to the optimal kernel for the current backend:

      - GPU/CUDA: cuDNN SDPA (batched GEMV, extremely fast for S=1)

      - TPU/CPU:  native GQA einsum (same as before)



    q: [B, Hq,  1, D]

    k: [B, Hkv, W, D]  (sliding window cache)

    v: [B, Hkv, W, D]



    Returns: [B, Hq, 1, D]

    """
    if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
        raise ValueError(
            f"decode_swa: expected 4-D tensors, got q={q.shape}, k={k.shape}, v={v.shape}"
        )

    B, Hq, S, D = map(int, q.shape)
    _, Hkv, W, Dk = map(int, k.shape)

    if S != 1:
        raise ValueError(f"decode_swa expects one token, got S={S}")
    if D != Dk:
        raise ValueError(f"Head-dim mismatch: q D={D}, k D={Dk}")
    if Hq % Hkv != 0:
        raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")

    # ── Backend dispatch ──────────────────────────────────────────────────────
    active_backend = _detect_backend()

    if active_backend == 'gpu':
        return _gpu_decode_swa(q, k, v)

    # ── TPU/CPU: native GQA einsum (original implementation) ─────────────────
    G = Hq // Hkv
    scale = 1.0 / math.sqrt(float(D))

    q_g = q.reshape(B, Hkv, G, 1, D)

    scores = (
        jnp.einsum(
            "bngqd,bnkd->bngqk",
            q_g.astype(jnp.float32),
            k.astype(jnp.float32),
        )
        * scale
    )  # [B, Hkv, G, 1, W]

    probs = jax.nn.softmax(scores, axis=-1)

    out_g = jnp.einsum(
        "bngqk,bnkd->bngqd",
        probs,
        v.astype(jnp.float32),
    )

    return out_g.reshape(B, Hq, 1, D).astype(q.dtype)

def local_causal_attention(

    q: jnp.ndarray,

    k: jnp.ndarray,

    v: jnp.ndarray,

) -> jnp.ndarray:
    """

    Full causal attention with native GQA β€” no windowing.



    ⚠  Creates an O(S²) score tensor [B, Hkv, G, S, S].

       Use only for short sequences, unit tests, or reference baselines.



    q: [B, Hq, S, D]  |  k, v: [B, Hkv, S, D]

    returns: [B, Hq, S, D]

    """
    if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
        raise ValueError(
            f"local_causal_attention: expected 4-D tensors, "
            f"got q={q.shape}  k={k.shape}  v={v.shape}"
        )
    B, Hq, S, D = map(int, q.shape)
    _,  Hkv, K, _ = map(int, k.shape)
    if Hq % Hkv != 0:
        raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")
    G = Hq // Hkv
    scale = 1.0 / math.sqrt(float(D))
    q_g   = q.reshape(B, Hkv, G, S, D)
    scores = (
        jnp.einsum("bngsd,bnkd->bngsk", q_g.astype(jnp.float32), k.astype(jnp.float32))
        * scale
    )
    qi = jnp.arange(S, dtype=jnp.int32)[:, None]
    ki = jnp.arange(K, dtype=jnp.int32)[None, :]
    scores = jnp.where(ki > qi, -1e30, scores)
    probs  = jax.nn.softmax(scores, axis=-1)
    out_g  = jnp.einsum("bngsk,bnkd->bngsd", probs, v.astype(jnp.float32))
    return out_g.reshape(B, Hq, S, D).astype(q.dtype)


# ---------------------------------------------------------------------------
# Validation suite
# ---------------------------------------------------------------------------

if __name__ == "__main__":
    import sys

    PASS  = "\033[92mβœ“\033[0m"
    FAIL  = "\033[91mβœ—\033[0m"
    HDR   = "\033[1;94m"
    RST   = "\033[0m"
    failures = 0

    def section(t):  print(f"\n{HDR}{'─'*62}{RST}\n{HDR}  {t}{RST}\n{HDR}{'─'*62}{RST}")
    def ok(m):       print(f"  {PASS}  {m}")
    def fail(m):
        global failures; failures += 1
        print(f"  {FAIL}  {m}", file=sys.stderr)

    print(f"\n{HDR}{'═'*62}{RST}")
    print(f"{HDR}  Veylon Attention β€” Block-tiled GQA SWA β€” Validation{RST}")
    print(f"{HDR}{'═'*62}{RST}")

    # ── 1. Shape and dtype ───────────────────────────────────────────────────
    section("1 Β· Shape and dtype correctness")

    B, Hq, Hkv, S, D, W = 1, 8, 2, 128, 64, 32
    ks = jax.random.split(jax.random.PRNGKey(0), 3)
    q = jax.random.normal(ks[0], (B, Hq,  S, D), dtype=jnp.bfloat16)
    k = jax.random.normal(ks[1], (B, Hkv, S, D), dtype=jnp.bfloat16)
    v = jax.random.normal(ks[2], (B, Hkv, S, D), dtype=jnp.bfloat16)

    out = flash_splash_attention(q, k, v, window_size=W, block_size=32)

    if out.shape == (B, Hq, S, D):   ok(f"Output shape  : {out.shape}")
    else:                             fail(f"Shape wrong β€” expected {(B,Hq,S,D)}, got {out.shape}")
    if out.dtype == jnp.bfloat16:    ok(f"Output dtype  : {out.dtype}")
    else:                             fail(f"dtype wrong β€” expected bfloat16, got {out.dtype}")
    if not jnp.any(jnp.isnan(out)):  ok("No NaNs")
    else:                             fail("Output contains NaNs")

    # ── 2. Numerical agreement with reference SWA ────────────────────────────
    section("2 Β· Numerical agreement: block SWA β‰ˆ reference SWA (W=S)")

    B_, Hq_, Hkv_, S_, D_ = 1, 4, 2, 48, 16
    ks = jax.random.split(jax.random.PRNGKey(1), 3)
    qn = jax.random.normal(ks[0], (B_, Hq_,  S_, D_), dtype=jnp.float32)
    kn = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
    vn = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)

    for W_t in [4, 8, 16, S_]:
        out_blk  = flash_splash_attention(qn, kn, vn, window_size=W_t, block_size=16)
        out_ref  = local_causal_attention(qn, kn, vn) if W_t == S_ else None

        # Build reference SWA inline for each W_t
        G_ = Hq_ // Hkv_
        sc = 1.0 / math.sqrt(D_)
        qg = qn.reshape(B_, Hkv_, G_, S_, D_)
        sc_ref = jnp.einsum("bngsd,bnkd->bngsk", qg, kn) * sc
        qi = jnp.arange(S_)[:, None]; ki = jnp.arange(S_)[None, :]
        sc_ref = jnp.where((ki > qi) | (qi - ki >= W_t), -1e30, sc_ref)
        pr_ref = jax.nn.softmax(sc_ref, axis=-1)
        out_swa_ref = jnp.einsum("bngsk,bnkd->bngsd", pr_ref, vn).reshape(B_, Hq_, S_, D_)

        err = float(jnp.max(jnp.abs(out_blk - out_swa_ref)))
        if err < 1e-4:  ok(f"W={W_t:3d}  max|block - ref_swa| = {err:.2e}")
        else:           fail(f"W={W_t:3d}  MISMATCH: max err = {err:.2e}")

    # ── 3. Memory scaling ────────────────────────────────────────────────────
    section("3 Β· Memory scaling  (2Γ— S β†’ ~2Γ— memory, not 4Γ—)")
    print("  Shape + completion check at S = 256, 512, 1024, 2048")
    for S_t in [256, 512, 1024, 2048]:
        ks = jax.random.split(jax.random.PRNGKey(S_t), 3)
        qt = jax.random.normal(ks[0], (1, 8, S_t, 64), dtype=jnp.bfloat16)
        kt = jax.random.normal(ks[1], (1, 2, S_t, 64), dtype=jnp.bfloat16)
        vt = jax.random.normal(ks[2], (1, 2, S_t, 64), dtype=jnp.bfloat16)
        ot = flash_splash_attention(qt, kt, vt, window_size=128)
        if ot.shape == (1, 8, S_t, 64): ok(f"S={S_t:5d}  β†’  {ot.shape}")
        else:                            fail(f"S={S_t}  wrong shape {ot.shape}")

    # ── 4. Native GQA isolation ──────────────────────────────────────────────
    section("4 Β· Native GQA head isolation (no KV duplication)")
    B_, Hq_, Hkv_, S_, D_, W_ = 1, 4, 2, 32, 16, 16
    G_ = Hq_ // Hkv_
    ks  = jax.random.split(jax.random.PRNGKey(7), 3)
    qg  = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32)
    kgq = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
    vgq = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)
    base  = flash_splash_attention(qg, kgq,                               vgq,                               window_size=W_, block_size=16)
    zero0 = flash_splash_attention(qg, kgq.at[:,0].set(0.), vgq.at[:,0].set(0.), window_size=W_, block_size=16)
    changed   = not jnp.allclose(base[:, :G_],  zero0[:, :G_],  atol=1e-4)
    unchanged = jnp.allclose(    base[:, G_:],  zero0[:, G_:],  atol=1e-4)
    ok(f"Q heads 0..{G_-1} changed when KV head 0 zeroed")   if changed   else fail("GQA dependency broken")
    ok(f"Q heads {G_}..{Hq_-1} unchanged (correct isolation)") if unchanged else fail("Cross-group contamination")

    # ── 5. Causality ─────────────────────────────────────────────────────────
    section("5 Β· Causality")
    B_, Hq_, Hkv_, S_, D_, W_ = 1, 4, 2, 24, 8, 8
    pivot = S_ // 2
    ks = jax.random.split(jax.random.PRNGKey(42), 3)
    qc = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32)
    kc = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
    vc = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)
    out_base = flash_splash_attention(qc, kc, vc, window_size=W_, block_size=8)
    out_corr = flash_splash_attention(qc,
                                      kc.at[:,:,pivot:].set(999.),
                                      vc.at[:,:,pivot:].set(999.),
                                      window_size=W_, block_size=8)
    if jnp.allclose(out_base[:,:,:pivot], out_corr[:,:,:pivot], atol=1e-5):
        ok(f"Tokens 0..{pivot-1} unaffected by corruption of tokens {pivot}+")
    else:
        fail("Causality violated β€” future tokens leaked into past outputs")

    # ── 6. Window boundary ───────────────────────────────────────────────────
    section("6 Β· Window boundary  (no leakage past W tokens)")
    B_, Hq_, Hkv_, S_, D_, W_ = 1, 2, 1, 24, 8, 4
    ks = jax.random.split(jax.random.PRNGKey(11), 3)
    qw = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32)
    kw = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
    vw = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)
    out_w  = flash_splash_attention(qw, kw, vw, window_size=W_, block_size=4)
    out_wm = flash_splash_attention(qw,
                                    kw.at[:,:,0].set(999.),
                                    vw.at[:,:,0].set(999.),
                                    window_size=W_, block_size=4)
    pivot_w = W_  # first token where token 0 is outside the window
    if jnp.allclose(out_w[:,:,pivot_w:], out_wm[:,:,pivot_w:], atol=1e-5):
        ok(f"Tokens {pivot_w}+ unaffected by modifying token 0  (W={W_} boundary)")
    else:
        fail(f"Window boundary violated β€” token 0 leaked into token {pivot_w}+")

    # ── 7. Non-power-of-2 sequence length ────────────────────────────────────
    section("7 Β· Non-power-of-2 sequence lengths  (S=100, 200, 500)")
    for S_t in [100, 200, 500]:
        ks = jax.random.split(jax.random.PRNGKey(S_t+1), 3)
        qt = jax.random.normal(ks[0], (1, 4, S_t, 16), dtype=jnp.float32)
        kt = jax.random.normal(ks[1], (1, 2, S_t, 16), dtype=jnp.float32)
        vt = jax.random.normal(ks[2], (1, 2, S_t, 16), dtype=jnp.float32)
        ot = flash_splash_attention(qt, kt, vt, window_size=32, block_size=32)
        if ot.shape == (1, 4, S_t, 16): ok(f"S={S_t}  β†’  {ot.shape}")
        else:                            fail(f"S={S_t}  wrong shape {ot.shape}")

    # ── 8. Pallas GPU kernel (forward correctness + gradient check) ─────────
    section("8 Β· Pallas/Triton FlashAttention kernel (GPU only)")
    if not _PALLAS_GPU_AVAILABLE:
        print("  (skipped β€” Pallas not importable in this environment)")
    elif _detect_backend() != 'gpu':
        print("  (skipped β€” no GPU backend detected)")
    elif not _pallas_supported(16, jnp.float16):
        cc = _gpu_compute_capability()
        if cc is not None and cc < (8, 0):
            print(f"  (skipped β€” GPU compute capability {cc[0]}.{cc[1]} < 8.0; "
                  f"Triton/Pallas requires Ampere or newer. cuDNN path handles "
                  f"FlashAttention on this GPU instead.)")
        else:
            print("  (skipped β€” Pallas gated off for this config; "
                  "set VEYLON_DEBUG_PALLAS_ATTN=1 for details)")
    else:
        B_, Hq_, Hkv_, S_, D_, W_ = 1, 4, 2, 130, 16, 24
        ks = jax.random.split(jax.random.PRNGKey(99), 3)
        qp = jax.random.normal(ks[0], (B_, Hq_,  S_, D_), dtype=jnp.float32) * 0.1
        kp = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32) * 0.1
        vp = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32) * 0.1

        try:
            out_pallas = _pallas_flash_gqa_swa(qp, kp, vp, window_size=W_, block_q=32, block_k=32)

            # Reference via existing XLA block-tiled kernel
            out_ref = _block_gqa_swa(qp, kp, vp, window_size=W_, block_size=32, use_remat=False)

            err = float(jnp.max(jnp.abs(out_pallas - out_ref)))
            if err < 1e-3:
                ok(f"Forward matches XLA reference: max err = {err:.2e}")
            else:
                fail(f"Forward MISMATCH vs XLA reference: max err = {err:.2e}")

            # Gradient check: compare d(sum(out))/d(q,k,v) against XLA reference
            def loss_pallas(q, k, v):
                return jnp.sum(_pallas_flash_gqa_swa(q, k, v, window_size=W_, block_q=32, block_k=32))

            def loss_ref(q, k, v):
                return jnp.sum(_block_gqa_swa(q, k, v, window_size=W_, block_size=32, use_remat=False))

            gp = jax.grad(loss_pallas, argnums=(0, 1, 2))(qp, kp, vp)
            gr = jax.grad(loss_ref, argnums=(0, 1, 2))(qp, kp, vp)

            names = ['dQ', 'dK', 'dV']
            for name, gp_i, gr_i in zip(names, gp, gr):
                gerr = float(jnp.max(jnp.abs(gp_i - gr_i)))
                if gerr < 1e-2:
                    ok(f"{name} matches XLA autodiff: max err = {gerr:.2e}")
                else:
                    fail(f"{name} MISMATCH vs XLA autodiff: max err = {gerr:.2e}")

        except Exception as e:
            fail(f"Pallas kernel raised an exception: {type(e).__name__}: {e}")

    # ── Summary ──────────────────────────────────────────────────────────────
    print(f"\n{HDR}{'═'*62}{RST}")
    if failures == 0: print(f"  {PASS}  All tests passed.")
    else:             print(f"  {FAIL}  {failures} test(s) failed.", file=sys.stderr)
    print(f"{HDR}{'═'*62}{RST}\n")
    sys.exit(failures)