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