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