Spaces:
Sleeping
Sleeping
Download veylon_attention.py from ArushBuilds/Pragya: direct link, hf CLI and curl.
- Browser
- Download file 29.7 kB
-
https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/fd6ac4ca5059d8d46aed9c7b6f6ac67273ecb7f0/veylon_attention.py
- Command line
-
hf download hf://spaces/ArushBuilds/Pragya@fd6ac4ca5059d8d46aed9c7b6f6ac67273ecb7f0/veylon_attention.py
-
curl -L -o veylon_attention.py https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/fd6ac4ca5059d8d46aed9c7b6f6ac67273ecb7f0/veylon_attention.py
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) | |