Spaces:
Sleeping
Sleeping
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) |