Pragya / veylon_attention.py
Arush kumar
Upload 14 files
54ad1e5
Raw History Blame
29.7 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.map vs lax.scan
--------------------
We use `jax.lax.map` (not `lax.scan`) because:
- Each block is fully independent β€” no carry state is needed.
- `lax.map` lowers to a while_loop with no accumulation buffer.
- `lax.scan` always allocates an output stacked along the scan axis;
lax.map lets XLA write directly into the preallocated output slice.
- This eliminates the [n_blocks, B, Hq, BLK, D] intermediate stack.
(We still do one reshape at the end, but that is a view, not a copy.)
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
from functools import partial
from typing import Optional
import jax
import jax.numpy as jnp
# ---------------------------------------------------------------------------
# Public tuning constant
# ---------------------------------------------------------------------------
# Default block size for TPU v5e. Must be a power of 2 and β‰₯ 1.
# Larger blocks β†’ fewer kernel launches, better systolic utilisation.
# Smaller blocks β†’ lower peak HBM per block (useful when W is huge).
TPU_BLOCK_SIZE: int = 128
# ---------------------------------------------------------------------------
# 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 ────────────────────────────────────────────────
# With use_remat=True: JAX recomputes each block's forward during backward.
# Granularity is per-block (not per-token), so remat overhead is manageable.
map_fn = jax.checkpoint(process_block) if use_remat else process_block
# ── lax.map over block indices ────────────────────────────────────────────
# Unlike lax.scan, lax.map does NOT accumulate a carry β€” it writes each
# block's output directly. No hidden [n_blocks, B, Hq, BLK, D] stack.
out_blocks = jax.lax.map(map_fn, jnp.arange(n_blocks, dtype=jnp.int32))
# out_blocks: [n_blocks, B, Hq, BLK, D]
# ── Assemble output ───────────────────────────────────────────────────────
# Transpose to [B, Hq, n_blocks, BLK, D] then reshape to [B, Hq, S_pad, D].
# The final slice trims padding tokens. All zero-copy views on TPU.
out = out_blocks.transpose(1, 2, 0, 3, 4).reshape(B, Hq, S_pad, D)
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.
Drop-in replacement for the previous flash_splash_attention.
GQA is always native (no KV expansion). `use_gqa` is accepted for
API compatibility only.
Parameters
----------
q, k, v : [B, Hq, S, D] / [B, Hkv, S, D]
window_size : causal window W
backend : reserved for future Pallas/Flash kernel dispatch (no-op)
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 recommended for TPU v5e)
use_remat : gradient checkpointing in map body (recommended for training)
Returns
-------
[B, Hq, S, D] β€” same dtype as q
"""
_ = backend
_ = use_gqa
_ = start_pos # TODO: shift q/kv slicing for decode-mode KV cache
return _block_gqa_swa(
q=q,
k=k,
v=v,
window_size=int(window_size),
block_size=int(block_size),
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.
q: [B, Hq, 1, D]
k: [B, Hkv, W, D]
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}")
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}")
# ── 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)