Pragya / veylon_attention.py
Arush kumar
Update veylon_attention.py
b22e951
Raw History Blame
64.1 kB
"""
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)