Spaces:
Sleeping
Sleeping
Arush kumar commited on
Commit Β·
b22e951
1
Parent(s): 3554d5b
Update veylon_attention.py
Browse files- veylon_attention.py +600 -17
veylon_attention.py
CHANGED
|
@@ -81,12 +81,23 @@ TPU-specific notes
|
|
| 81 |
from __future__ import annotations
|
| 82 |
|
| 83 |
import math
|
|
|
|
| 84 |
from functools import partial
|
| 85 |
from typing import Optional
|
| 86 |
|
| 87 |
import jax
|
| 88 |
import jax.numpy as jnp
|
| 89 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 90 |
|
| 91 |
# ---------------------------------------------------------------------------
|
| 92 |
# Public tuning constants
|
|
@@ -119,6 +130,498 @@ def _detect_backend() -> str:
|
|
| 119 |
return 'cpu'
|
| 120 |
|
| 121 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 122 |
# ---------------------------------------------------------------------------
|
| 123 |
# GPU path: JAX cuDNN FlashAttention via jax.nn.dot_product_attention
|
| 124 |
# ---------------------------------------------------------------------------
|
|
@@ -162,24 +665,21 @@ def _gpu_flash_gqa_swa(
|
|
| 162 |
k_s = k.transpose(0, 2, 1, 3)
|
| 163 |
v_s = v.transpose(0, 2, 1, 3)
|
| 164 |
|
| 165 |
-
def _full_attn(
|
| 166 |
-
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
)
|
| 173 |
-
return out
|
| 174 |
-
except Exception:
|
| 175 |
-
# cuDNN unavailable: fall through to XLA path below
|
| 176 |
-
return None
|
| 177 |
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
|
|
|
|
|
|
| 183 |
|
| 184 |
# ββ SWA path: XLA block-tiled kernel (GPU block_size=64) βββββββββββββββββ
|
| 185 |
# This is the fast path for SWA on GPU.
|
|
@@ -573,6 +1073,34 @@ def flash_splash_attention(
|
|
| 573 |
|
| 574 |
# ββ Dispatch ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 575 |
if active_backend == 'gpu':
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 576 |
return _gpu_flash_gqa_swa(
|
| 577 |
q=q,
|
| 578 |
k=k,
|
|
@@ -836,6 +1364,61 @@ if __name__ == "__main__":
|
|
| 836 |
if ot.shape == (1, 4, S_t, 16): ok(f"S={S_t} β {ot.shape}")
|
| 837 |
else: fail(f"S={S_t} wrong shape {ot.shape}")
|
| 838 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 839 |
# ββ Summary ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 840 |
print(f"\n{HDR}{'β'*62}{RST}")
|
| 841 |
if failures == 0: print(f" {PASS} All tests passed.")
|
|
|
|
| 81 |
from __future__ import annotations
|
| 82 |
|
| 83 |
import math
|
| 84 |
+
import os
|
| 85 |
from functools import partial
|
| 86 |
from typing import Optional
|
| 87 |
|
| 88 |
import jax
|
| 89 |
import jax.numpy as jnp
|
| 90 |
|
| 91 |
+
# ---------------------------------------------------------------------------
|
| 92 |
+
# Optional Pallas/Triton GPU kernel (FlashAttention-style, I/O-aware)
|
| 93 |
+
# ---------------------------------------------------------------------------
|
| 94 |
+
try:
|
| 95 |
+
from jax.experimental import pallas as pl
|
| 96 |
+
from jax.experimental.pallas import triton as plgpu
|
| 97 |
+
_PALLAS_GPU_AVAILABLE = True
|
| 98 |
+
except Exception:
|
| 99 |
+
_PALLAS_GPU_AVAILABLE = False
|
| 100 |
+
|
| 101 |
|
| 102 |
# ---------------------------------------------------------------------------
|
| 103 |
# Public tuning constants
|
|
|
|
| 130 |
return 'cpu'
|
| 131 |
|
| 132 |
|
| 133 |
+
# ---------------------------------------------------------------------------
|
| 134 |
+
# Pallas/Triton GPU kernel β I/O-aware FlashAttention-style GQA SWA
|
| 135 |
+
# ---------------------------------------------------------------------------
|
| 136 |
+
#
|
| 137 |
+
# This is a from-scratch FlashAttention-2 style kernel:
|
| 138 |
+
# - Tiles Q into BLOCK_Q-sized blocks, K/V into BLOCK_K-sized blocks.
|
| 139 |
+
# - Grid = (batch, kv_head, num_q_blocks). Each program instance owns one
|
| 140 |
+
# Q block for one KV head (covering its G query-head siblings via GQA
|
| 141 |
+
# broadcast inside the kernel β K/V are NEVER duplicated in HBM).
|
| 142 |
+
# - Online softmax: running max `m`, running sum `l`, running weighted
|
| 143 |
+
# accumulator `acc` are carried across the K-block loop via
|
| 144 |
+
# jax.lax.fori_loop. The full [BLK_Q, S] or [BLK_Q, window] score matrix
|
| 145 |
+
# is NEVER materialized β only one [BLK_Q, BLOCK_K] tile lives in VMEM
|
| 146 |
+
# at a time. This is the actual "I/O-aware" property: HBM traffic is
|
| 147 |
+
# O(S) reads of Q/K/V blocks, not O(S^2) score matrix writes.
|
| 148 |
+
# - Sliding-window + causal masking is applied per K-block using the same
|
| 149 |
+
# relative-offset trick as the existing XLA kernel (b-independent delta),
|
| 150 |
+
# so only K-blocks that intersect [q_pos - W + 1, q_pos] are visited β
|
| 151 |
+
# blocks fully outside the window are skipped via the loop bounds, not
|
| 152 |
+
# just masked, which is where the real compute savings come from vs the
|
| 153 |
+
# existing XLA block-tiled kernel (which still computes+masks every
|
| 154 |
+
# block inside a fixed kv_len window).
|
| 155 |
+
#
|
| 156 |
+
# Custom VJP: backward recomputes scores per (Q-block, K-block) pair from
|
| 157 |
+
# saved Q, K, V, O, m, l (NOT saved scores/probs β that's the whole point,
|
| 158 |
+
# same trick as FlashAttention). This keeps backward memory O(S) instead of
|
| 159 |
+
# O(S * window).
|
| 160 |
+
# ---------------------------------------------------------------------------
|
| 161 |
+
|
| 162 |
+
_PALLAS_BLOCK_Q = 64
|
| 163 |
+
_PALLAS_BLOCK_K = 64
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
def _gpu_compute_capability() -> Optional[tuple]:
|
| 167 |
+
"""Returns (major, minor) compute capability of the current GPU, or None
|
| 168 |
+
if it can't be determined. Used to gate the Pallas/Triton path, which
|
| 169 |
+
JAX only supports on Ampere (SM 8.0) and newer β Turing (T4, SM 7.5) and
|
| 170 |
+
older will FAIL_PRECONDITION at Triton compile time, not at import time,
|
| 171 |
+
so we must check this explicitly before attempting the kernel."""
|
| 172 |
+
try:
|
| 173 |
+
dev = jax.devices('gpu')[0]
|
| 174 |
+
# jaxlib exposes this via device_kind (e.g. "Tesla T4", "NVIDIA A100")
|
| 175 |
+
# or via compute_capability on newer jaxlib versions.
|
| 176 |
+
cc = getattr(dev, 'compute_capability', None)
|
| 177 |
+
if cc is not None:
|
| 178 |
+
major, minor = str(cc).split('.')[:2]
|
| 179 |
+
return (int(major), int(minor))
|
| 180 |
+
return None
|
| 181 |
+
except Exception:
|
| 182 |
+
return None
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
_GPU_COMPUTE_CAPABILITY = None # lazily populated on first check
|
| 186 |
+
|
| 187 |
+
|
| 188 |
+
def _pallas_supported(D: int, dtype) -> bool:
|
| 189 |
+
"""Conservative gate: only use the Pallas path for configs we've reasoned
|
| 190 |
+
through (head_dim multiple of 16 for tensor-core alignment, fp16/bf16/fp32,
|
| 191 |
+
Ampere-or-newer GPU). Anything else falls back to the cuDNN/XLA path
|
| 192 |
+
automatically."""
|
| 193 |
+
global _GPU_COMPUTE_CAPABILITY
|
| 194 |
+
if os.environ.get('VEYLON_DISABLE_PALLAS_ATTN', '0') == '1':
|
| 195 |
+
return False
|
| 196 |
+
if not _PALLAS_GPU_AVAILABLE:
|
| 197 |
+
return False
|
| 198 |
+
if D % 16 != 0:
|
| 199 |
+
return False
|
| 200 |
+
if dtype not in (jnp.float16, jnp.bfloat16, jnp.float32):
|
| 201 |
+
return False
|
| 202 |
+
if _GPU_COMPUTE_CAPABILITY is None:
|
| 203 |
+
_GPU_COMPUTE_CAPABILITY = _gpu_compute_capability() or (0, 0)
|
| 204 |
+
if _GPU_COMPUTE_CAPABILITY < (8, 0):
|
| 205 |
+
# Triton (Pallas GPU backend) requires Ampere or newer. T4 (7.5),
|
| 206 |
+
# V100 (7.0), P100 (6.0) all fail here β this is a hard hardware
|
| 207 |
+
# limit, not a bug, so we skip Pallas entirely rather than let it
|
| 208 |
+
# crash through a full Triton compile attempt.
|
| 209 |
+
return False
|
| 210 |
+
return True
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
def _fa_fwd_kernel(
|
| 214 |
+
q_ref, k_ref, v_ref, # inputs, VMEM-resident blocks
|
| 215 |
+
o_ref, m_ref, l_ref, # outputs
|
| 216 |
+
*,
|
| 217 |
+
window: int,
|
| 218 |
+
block_q: int,
|
| 219 |
+
block_k: int,
|
| 220 |
+
seq_len: int,
|
| 221 |
+
scale: float,
|
| 222 |
+
):
|
| 223 |
+
"""
|
| 224 |
+
Pallas kernel body β one program instance handles ONE (batch, kv_head,
|
| 225 |
+
q_block) triple, looping internally over the K-blocks that intersect
|
| 226 |
+
the causal + sliding-window range for this Q block.
|
| 227 |
+
|
| 228 |
+
Ref shapes (per-program, already sliced by BlockSpec / index_map):
|
| 229 |
+
q_ref : [block_q, D] (single query head's slice β see note below)
|
| 230 |
+
k_ref : [seq_len, D] (full K for this batch/kv_head; we slice
|
| 231 |
+
inside the loop via pl.load with dynamic
|
| 232 |
+
start so only ONE [block_k, D] tile is
|
| 233 |
+
actually resident in VMEM at a time)
|
| 234 |
+
v_ref : [seq_len, D] (same as k_ref)
|
| 235 |
+
o_ref : [block_q, D] (output accumulator, written once at end)
|
| 236 |
+
m_ref, l_ref : [block_q, 1] (running softmax stats, scratch)
|
| 237 |
+
"""
|
| 238 |
+
q_block_idx = pl.program_id(2)
|
| 239 |
+
q_start = q_block_idx * block_q
|
| 240 |
+
|
| 241 |
+
q = q_ref[...].astype(jnp.float32) * scale # [block_q, D]
|
| 242 |
+
|
| 243 |
+
m_i = jnp.full((block_q, 1), -jnp.inf, dtype=jnp.float32)
|
| 244 |
+
l_i = jnp.zeros((block_q, 1), dtype=jnp.float32)
|
| 245 |
+
acc = jnp.zeros_like(q)
|
| 246 |
+
|
| 247 |
+
# Range of K-blocks that can possibly intersect this Q-block's
|
| 248 |
+
# causal+window range. Query positions in this block span
|
| 249 |
+
# [q_start, q_start + block_q - 1]. Each attends to
|
| 250 |
+
# [q_pos - window + 1, q_pos]. So the union over the block spans
|
| 251 |
+
# [q_start - window + 1, q_start + block_q - 1].
|
| 252 |
+
k_lo = jnp.maximum(0, q_start - window + 1)
|
| 253 |
+
k_hi = jnp.minimum(seq_len, q_start + block_q) # exclusive, causal cap
|
| 254 |
+
first_k_block = k_lo // block_k
|
| 255 |
+
num_k_blocks = (k_hi - first_k_block * block_k + block_k - 1) // block_k
|
| 256 |
+
num_k_blocks = jnp.maximum(num_k_blocks, 1)
|
| 257 |
+
|
| 258 |
+
def body(i, carry):
|
| 259 |
+
m_i, l_i, acc = carry
|
| 260 |
+
k_start = (first_k_block + i) * block_k
|
| 261 |
+
|
| 262 |
+
k_blk = pl.load(
|
| 263 |
+
k_ref, (pl.dslice(k_start, block_k), slice(None))
|
| 264 |
+
).astype(jnp.float32) # [block_k, D]
|
| 265 |
+
v_blk = pl.load(
|
| 266 |
+
v_ref, (pl.dslice(k_start, block_k), slice(None))
|
| 267 |
+
).astype(jnp.float32) # [block_k, D]
|
| 268 |
+
|
| 269 |
+
scores = jnp.dot(
|
| 270 |
+
q, k_blk.T, preferred_element_type=jnp.float32
|
| 271 |
+
) # [block_q, block_k]
|
| 272 |
+
|
| 273 |
+
q_pos = q_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 0)
|
| 274 |
+
k_pos = k_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 1)
|
| 275 |
+
causal_ok = k_pos <= q_pos
|
| 276 |
+
window_ok = (q_pos - k_pos) < window
|
| 277 |
+
bounds_ok = k_pos < seq_len
|
| 278 |
+
mask = causal_ok & window_ok & bounds_ok
|
| 279 |
+
scores = jnp.where(mask, scores, -jnp.inf)
|
| 280 |
+
|
| 281 |
+
m_ij = jnp.max(scores, axis=-1, keepdims=True) # [block_q, 1]
|
| 282 |
+
m_new = jnp.maximum(m_i, m_ij)
|
| 283 |
+
# Guard against all-masked rows (m_new stays -inf) -> exp(0)=1 issue
|
| 284 |
+
m_new_safe = jnp.where(m_new == -jnp.inf, 0.0, m_new)
|
| 285 |
+
|
| 286 |
+
p = jnp.exp(scores - m_new_safe) # [block_q, block_k]
|
| 287 |
+
p = jnp.where(mask, p, 0.0)
|
| 288 |
+
|
| 289 |
+
alpha = jnp.exp(jnp.where(m_i == -jnp.inf, m_new_safe, m_i) - m_new_safe)
|
| 290 |
+
l_new = l_i * alpha + jnp.sum(p, axis=-1, keepdims=True)
|
| 291 |
+
acc_new = acc * alpha + jnp.dot(p, v_blk, preferred_element_type=jnp.float32)
|
| 292 |
+
|
| 293 |
+
return m_new, l_new, acc_new
|
| 294 |
+
|
| 295 |
+
m_i, l_i, acc = jax.lax.fori_loop(0, num_k_blocks, body, (m_i, l_i, acc))
|
| 296 |
+
|
| 297 |
+
l_safe = jnp.where(l_i > 0, l_i, 1.0)
|
| 298 |
+
out = acc / l_safe
|
| 299 |
+
|
| 300 |
+
o_ref[...] = out.astype(o_ref.dtype)
|
| 301 |
+
m_ref[...] = m_i
|
| 302 |
+
l_ref[...] = l_i
|
| 303 |
+
|
| 304 |
+
|
| 305 |
+
def _pallas_fwd_single_head(q, k, v, window, block_q, block_k, scale):
|
| 306 |
+
"""
|
| 307 |
+
Runs the Pallas forward kernel for ONE query head against its KV head.
|
| 308 |
+
q: [B, S, D] k, v: [B, S, D] (already the per-head slices)
|
| 309 |
+
Returns: out [B, S, D], m [B, S, 1], l [B, S, 1] (m, l saved for bwd)
|
| 310 |
+
"""
|
| 311 |
+
B, S, D = map(int, q.shape)
|
| 312 |
+
n_q_blocks = (S + block_q - 1) // block_q
|
| 313 |
+
S_pad = n_q_blocks * block_q
|
| 314 |
+
|
| 315 |
+
q_p = jnp.pad(q, ((0, 0), (0, S_pad - S), (0, 0)))
|
| 316 |
+
# K/V padded on the right only; kernel bounds-checks k_pos < seq_len so
|
| 317 |
+
# right-padding is safe (never read past the pad due to k_hi clamp), but
|
| 318 |
+
# we still pad to a multiple of block_k so pl.load's static block shape
|
| 319 |
+
# never reads out-of-bounds memory.
|
| 320 |
+
n_k_blocks_total = (S + block_k - 1) // block_k
|
| 321 |
+
S_pad_k = n_k_blocks_total * block_k
|
| 322 |
+
k_p = jnp.pad(k, ((0, 0), (0, S_pad_k - S), (0, 0)))
|
| 323 |
+
v_p = jnp.pad(v, ((0, 0), (0, S_pad_k - S), (0, 0)))
|
| 324 |
+
|
| 325 |
+
kernel = partial(
|
| 326 |
+
_fa_fwd_kernel,
|
| 327 |
+
window=window, block_q=block_q, block_k=block_k,
|
| 328 |
+
seq_len=S, scale=scale,
|
| 329 |
+
)
|
| 330 |
+
|
| 331 |
+
out, m, l = pl.pallas_call(
|
| 332 |
+
kernel,
|
| 333 |
+
grid=(B, 1, n_q_blocks),
|
| 334 |
+
in_specs=[
|
| 335 |
+
pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),
|
| 336 |
+
pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)),
|
| 337 |
+
pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)),
|
| 338 |
+
],
|
| 339 |
+
out_specs=[
|
| 340 |
+
pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),
|
| 341 |
+
pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)),
|
| 342 |
+
pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)),
|
| 343 |
+
],
|
| 344 |
+
out_shape=[
|
| 345 |
+
jax.ShapeDtypeStruct((B, S_pad, D), q.dtype),
|
| 346 |
+
jax.ShapeDtypeStruct((B, S_pad, 1), jnp.float32),
|
| 347 |
+
jax.ShapeDtypeStruct((B, S_pad, 1), jnp.float32),
|
| 348 |
+
],
|
| 349 |
+
)(q_p, k_p, v_p)
|
| 350 |
+
|
| 351 |
+
return out[:, :S, :], m[:, :S, :], l[:, :S, :]
|
| 352 |
+
|
| 353 |
+
|
| 354 |
+
def _fa_bwd_kernel(
|
| 355 |
+
q_ref, k_ref, v_ref, o_ref, do_ref, m_ref, l_ref,
|
| 356 |
+
dq_ref, dk_ref, dv_ref,
|
| 357 |
+
*,
|
| 358 |
+
window: int,
|
| 359 |
+
block_q: int,
|
| 360 |
+
block_k: int,
|
| 361 |
+
seq_len: int,
|
| 362 |
+
scale: float,
|
| 363 |
+
):
|
| 364 |
+
"""
|
| 365 |
+
Backward kernel β one program per (batch, k_block). Recomputes scores
|
| 366 |
+
for each intersecting Q-block on the fly (from saved Q, K, V, m, l) and
|
| 367 |
+
accumulates dK/dV. dQ is accumulated via a separate pass below since it
|
| 368 |
+
is indexed by q_block, not k_block (standard FlashAttention-2 backward
|
| 369 |
+
split to avoid atomic adds across programs).
|
| 370 |
+
"""
|
| 371 |
+
k_block_idx = pl.program_id(2)
|
| 372 |
+
k_start = k_block_idx * block_k
|
| 373 |
+
|
| 374 |
+
k_blk = k_ref[...].astype(jnp.float32) # [block_k, D]
|
| 375 |
+
v_blk = v_ref[...].astype(jnp.float32) # [block_k, D]
|
| 376 |
+
|
| 377 |
+
dk_acc = jnp.zeros_like(k_blk)
|
| 378 |
+
dv_acc = jnp.zeros_like(v_blk)
|
| 379 |
+
|
| 380 |
+
# Q-blocks that can intersect this K-block: q_pos >= k_pos (causal) and
|
| 381 |
+
# q_pos - k_pos < window. q spans [k_start, seq_len-1] roughly, capped
|
| 382 |
+
# by window on the upper side: q_pos < k_start + block_k + window - 1.
|
| 383 |
+
q_lo = k_start
|
| 384 |
+
q_hi = jnp.minimum(seq_len, k_start + block_k + window - 1)
|
| 385 |
+
first_q_block = q_lo // block_q
|
| 386 |
+
num_q_blocks = (q_hi - first_q_block * block_q + block_q - 1) // block_q
|
| 387 |
+
num_q_blocks = jnp.maximum(num_q_blocks, 1)
|
| 388 |
+
|
| 389 |
+
def body(i, carry):
|
| 390 |
+
dk_acc, dv_acc = carry
|
| 391 |
+
q_start = (first_q_block + i) * block_q
|
| 392 |
+
|
| 393 |
+
q_blk = pl.load(q_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32) * scale
|
| 394 |
+
do_blk = pl.load(do_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)
|
| 395 |
+
m_blk = pl.load(m_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)
|
| 396 |
+
l_blk = pl.load(l_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)
|
| 397 |
+
o_blk = pl.load(o_ref, (pl.dslice(q_start, block_q), slice(None))).astype(jnp.float32)
|
| 398 |
+
|
| 399 |
+
scores = jnp.dot(q_blk, k_blk.T, preferred_element_type=jnp.float32)
|
| 400 |
+
|
| 401 |
+
q_pos = q_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 0)
|
| 402 |
+
k_pos = k_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 1)
|
| 403 |
+
causal_ok = k_pos <= q_pos
|
| 404 |
+
window_ok = (q_pos - k_pos) < window
|
| 405 |
+
bounds_ok = k_pos < seq_len
|
| 406 |
+
mask = causal_ok & window_ok & bounds_ok
|
| 407 |
+
|
| 408 |
+
l_safe = jnp.where(l_blk > 0, l_blk, 1.0)
|
| 409 |
+
p = jnp.where(mask, jnp.exp(scores - m_blk), 0.0) / l_safe # [block_q, block_k]
|
| 410 |
+
|
| 411 |
+
dv_acc = dv_acc + jnp.dot(p.T, do_blk, preferred_element_type=jnp.float32)
|
| 412 |
+
|
| 413 |
+
dp = jnp.dot(do_blk, v_blk.T, preferred_element_type=jnp.float32) # [block_q, block_k]
|
| 414 |
+
Di = jnp.sum(do_blk * o_blk, axis=-1, keepdims=True) # [block_q, 1]
|
| 415 |
+
dscores = p * (dp - Di)
|
| 416 |
+
dscores = jnp.where(mask, dscores, 0.0)
|
| 417 |
+
|
| 418 |
+
dk_acc = dk_acc + jnp.dot(dscores.T, q_blk, preferred_element_type=jnp.float32) * scale
|
| 419 |
+
|
| 420 |
+
return dk_acc, dv_acc
|
| 421 |
+
|
| 422 |
+
dk_acc, dv_acc = jax.lax.fori_loop(0, num_q_blocks, body, (dk_acc, dv_acc))
|
| 423 |
+
|
| 424 |
+
dk_ref[...] = dk_acc.astype(dk_ref.dtype)
|
| 425 |
+
dv_ref[...] = dv_acc.astype(dv_ref.dtype)
|
| 426 |
+
|
| 427 |
+
|
| 428 |
+
def _fa_bwd_dq_kernel(
|
| 429 |
+
q_ref, k_ref, v_ref, o_ref, do_ref, m_ref, l_ref,
|
| 430 |
+
dq_ref,
|
| 431 |
+
*,
|
| 432 |
+
window: int,
|
| 433 |
+
block_q: int,
|
| 434 |
+
block_k: int,
|
| 435 |
+
seq_len: int,
|
| 436 |
+
scale: float,
|
| 437 |
+
):
|
| 438 |
+
"""Separate pass computing dQ, one program per (batch, q_block), looping
|
| 439 |
+
over intersecting K-blocks. Kept separate from the dK/dV kernel because
|
| 440 |
+
dQ is naturally indexed by q_block and dK/dV by k_block β fusing both
|
| 441 |
+
into one kernel would need cross-program atomics, which Pallas/Triton
|
| 442 |
+
doesn't support cleanly. Recomputation cost (~2x score matmuls total
|
| 443 |
+
across both passes) is the standard FlashAttention-2 backward tradeoff."""
|
| 444 |
+
q_block_idx = pl.program_id(2)
|
| 445 |
+
q_start = q_block_idx * block_q
|
| 446 |
+
|
| 447 |
+
q_blk = q_ref[...].astype(jnp.float32) * scale
|
| 448 |
+
do_blk = do_ref[...].astype(jnp.float32)
|
| 449 |
+
m_blk = m_ref[...].astype(jnp.float32)
|
| 450 |
+
l_blk = l_ref[...].astype(jnp.float32)
|
| 451 |
+
o_blk = o_ref[...].astype(jnp.float32)
|
| 452 |
+
|
| 453 |
+
dq_acc = jnp.zeros_like(q_blk)
|
| 454 |
+
|
| 455 |
+
k_lo = jnp.maximum(0, q_start - window + 1)
|
| 456 |
+
k_hi = jnp.minimum(seq_len, q_start + block_q)
|
| 457 |
+
first_k_block = k_lo // block_k
|
| 458 |
+
num_k_blocks = (k_hi - first_k_block * block_k + block_k - 1) // block_k
|
| 459 |
+
num_k_blocks = jnp.maximum(num_k_blocks, 1)
|
| 460 |
+
|
| 461 |
+
def body(i, dq_acc):
|
| 462 |
+
k_start = (first_k_block + i) * block_k
|
| 463 |
+
k_blk = pl.load(k_ref, (pl.dslice(k_start, block_k), slice(None))).astype(jnp.float32)
|
| 464 |
+
v_blk = pl.load(v_ref, (pl.dslice(k_start, block_k), slice(None))).astype(jnp.float32)
|
| 465 |
+
|
| 466 |
+
scores = jnp.dot(q_blk, k_blk.T, preferred_element_type=jnp.float32)
|
| 467 |
+
|
| 468 |
+
q_pos = q_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 0)
|
| 469 |
+
k_pos = k_start + jax.lax.broadcasted_iota(jnp.int32, (block_q, block_k), 1)
|
| 470 |
+
causal_ok = k_pos <= q_pos
|
| 471 |
+
window_ok = (q_pos - k_pos) < window
|
| 472 |
+
bounds_ok = k_pos < seq_len
|
| 473 |
+
mask = causal_ok & window_ok & bounds_ok
|
| 474 |
+
|
| 475 |
+
l_safe = jnp.where(l_blk > 0, l_blk, 1.0)
|
| 476 |
+
p = jnp.where(mask, jnp.exp(scores - m_blk), 0.0) / l_safe
|
| 477 |
+
|
| 478 |
+
dp = jnp.dot(do_blk, v_blk.T, preferred_element_type=jnp.float32)
|
| 479 |
+
Di = jnp.sum(do_blk * o_blk, axis=-1, keepdims=True)
|
| 480 |
+
dscores = p * (dp - Di)
|
| 481 |
+
dscores = jnp.where(mask, dscores, 0.0)
|
| 482 |
+
|
| 483 |
+
dq_acc = dq_acc + jnp.dot(dscores, k_blk, preferred_element_type=jnp.float32) * scale
|
| 484 |
+
return dq_acc
|
| 485 |
+
|
| 486 |
+
dq_acc = jax.lax.fori_loop(0, num_k_blocks, body, dq_acc)
|
| 487 |
+
dq_ref[...] = dq_acc.astype(dq_ref.dtype)
|
| 488 |
+
|
| 489 |
+
|
| 490 |
+
def _pallas_bwd_single_head(q, k, v, o, do, m, l, window, block_q, block_k, scale):
|
| 491 |
+
"""Runs both backward kernels (dK/dV and dQ) for one query/KV head pair."""
|
| 492 |
+
B, S, D = map(int, q.shape)
|
| 493 |
+
n_q_blocks = (S + block_q - 1) // block_q
|
| 494 |
+
n_k_blocks = (S + block_k - 1) // block_k
|
| 495 |
+
S_pad_q = n_q_blocks * block_q
|
| 496 |
+
S_pad_k = n_k_blocks * block_k
|
| 497 |
+
|
| 498 |
+
pad_q = lambda x, fill=0.0: jnp.pad(x, ((0, 0), (0, S_pad_q - S), (0, 0)), constant_values=fill)
|
| 499 |
+
pad_k = lambda x: jnp.pad(x, ((0, 0), (0, S_pad_k - S), (0, 0)))
|
| 500 |
+
|
| 501 |
+
q_p, o_p, do_p = pad_q(q), pad_q(o), pad_q(do)
|
| 502 |
+
m_p = jnp.pad(m, ((0, 0), (0, S_pad_q - S), (0, 0)), constant_values=jnp.inf)
|
| 503 |
+
l_p = jnp.pad(l, ((0, 0), (0, S_pad_q - S), (0, 0)), constant_values=1.0)
|
| 504 |
+
k_p, v_p = pad_k(k), pad_k(v)
|
| 505 |
+
|
| 506 |
+
dkdv_kernel = partial(
|
| 507 |
+
_fa_bwd_kernel, window=window, block_q=block_q, block_k=block_k,
|
| 508 |
+
seq_len=S, scale=scale,
|
| 509 |
+
)
|
| 510 |
+
dk, dv = pl.pallas_call(
|
| 511 |
+
dkdv_kernel,
|
| 512 |
+
grid=(B, 1, n_k_blocks),
|
| 513 |
+
in_specs=[
|
| 514 |
+
pl.BlockSpec((None, S_pad_q, D), lambda b, h, j: (b, 0, 0)), # q (full, sliced inside)
|
| 515 |
+
pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)), # k block
|
| 516 |
+
pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)), # v block
|
| 517 |
+
pl.BlockSpec((None, S_pad_q, D), lambda b, h, j: (b, 0, 0)), # o (full)
|
| 518 |
+
pl.BlockSpec((None, S_pad_q, D), lambda b, h, j: (b, 0, 0)), # do (full)
|
| 519 |
+
pl.BlockSpec((None, S_pad_q, 1), lambda b, h, j: (b, 0, 0)), # m (full)
|
| 520 |
+
pl.BlockSpec((None, S_pad_q, 1), lambda b, h, j: (b, 0, 0)), # l (full)
|
| 521 |
+
],
|
| 522 |
+
out_specs=[
|
| 523 |
+
pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)),
|
| 524 |
+
pl.BlockSpec((None, block_k, D), lambda b, h, j: (b, j, 0)),
|
| 525 |
+
],
|
| 526 |
+
out_shape=[
|
| 527 |
+
jax.ShapeDtypeStruct((B, S_pad_k, D), k.dtype),
|
| 528 |
+
jax.ShapeDtypeStruct((B, S_pad_k, D), v.dtype),
|
| 529 |
+
],
|
| 530 |
+
)(q_p, k_p, v_p, o_p, do_p, m_p, l_p)
|
| 531 |
+
|
| 532 |
+
dq_kernel = partial(
|
| 533 |
+
_fa_bwd_dq_kernel, window=window, block_q=block_q, block_k=block_k,
|
| 534 |
+
seq_len=S, scale=scale,
|
| 535 |
+
)
|
| 536 |
+
dq = pl.pallas_call(
|
| 537 |
+
dq_kernel,
|
| 538 |
+
grid=(B, 1, n_q_blocks),
|
| 539 |
+
in_specs=[
|
| 540 |
+
pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)), # q block
|
| 541 |
+
pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)), # k (full)
|
| 542 |
+
pl.BlockSpec((None, S_pad_k, D), lambda b, h, i: (b, 0, 0)), # v (full)
|
| 543 |
+
pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)), # o block
|
| 544 |
+
pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)), # do block
|
| 545 |
+
pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)), # m block
|
| 546 |
+
pl.BlockSpec((None, block_q, 1), lambda b, h, i: (b, i, 0)), # l block
|
| 547 |
+
],
|
| 548 |
+
out_specs=pl.BlockSpec((None, block_q, D), lambda b, h, i: (b, i, 0)),
|
| 549 |
+
out_shape=jax.ShapeDtypeStruct((B, S_pad_q, D), q.dtype),
|
| 550 |
+
)(q_p, k_p, v_p, o_p, do_p, m_p, l_p)
|
| 551 |
+
|
| 552 |
+
return dq[:, :S, :], dk[:, :S, :], dv[:, :S, :]
|
| 553 |
+
|
| 554 |
+
|
| 555 |
+
@partial(jax.custom_vjp, nondiff_argnums=(3, 4, 5, 6))
|
| 556 |
+
def _pallas_gqa_swa_head(q, k, v, window, block_q, block_k, scale):
|
| 557 |
+
"""Single (query-head, kv-head) FlashAttention call with custom VJP.
|
| 558 |
+
q, k, v: [B, S, D] for ONE head pair (GQA broadcast handled by caller)."""
|
| 559 |
+
out, _, _ = _pallas_fwd_single_head(q, k, v, window, block_q, block_k, scale)
|
| 560 |
+
return out
|
| 561 |
+
|
| 562 |
+
|
| 563 |
+
def _pallas_gqa_swa_head_fwd(q, k, v, window, block_q, block_k, scale):
|
| 564 |
+
out, m, l = _pallas_fwd_single_head(q, k, v, window, block_q, block_k, scale)
|
| 565 |
+
return out, (q, k, v, out, m, l)
|
| 566 |
+
|
| 567 |
+
|
| 568 |
+
def _pallas_gqa_swa_head_bwd(window, block_q, block_k, scale, residuals, dout):
|
| 569 |
+
q, k, v, out, m, l = residuals
|
| 570 |
+
dq, dk, dv = _pallas_bwd_single_head(
|
| 571 |
+
q, k, v, out, dout, m, l, window, block_q, block_k, scale
|
| 572 |
+
)
|
| 573 |
+
return dq, dk, dv
|
| 574 |
+
|
| 575 |
+
|
| 576 |
+
_pallas_gqa_swa_head.defvjp(_pallas_gqa_swa_head_fwd, _pallas_gqa_swa_head_bwd)
|
| 577 |
+
|
| 578 |
+
|
| 579 |
+
def _pallas_flash_gqa_swa(
|
| 580 |
+
q: jnp.ndarray,
|
| 581 |
+
k: jnp.ndarray,
|
| 582 |
+
v: jnp.ndarray,
|
| 583 |
+
window_size: int,
|
| 584 |
+
block_q: int = _PALLAS_BLOCK_Q,
|
| 585 |
+
block_k: int = _PALLAS_BLOCK_K,
|
| 586 |
+
) -> jnp.ndarray:
|
| 587 |
+
"""
|
| 588 |
+
I/O-aware FlashAttention-style GQA SWA, entry point for the Pallas path.
|
| 589 |
+
|
| 590 |
+
q: [B, Hq, S, D]
|
| 591 |
+
k: [B, Hkv, S, D]
|
| 592 |
+
v: [B, Hkv, S, D]
|
| 593 |
+
|
| 594 |
+
GQA is handled by vmapping the single-head kernel over KV heads, and
|
| 595 |
+
within each KV head over its G query-head siblings β K/V are never
|
| 596 |
+
physically duplicated; only the (small) grid iterates over G.
|
| 597 |
+
"""
|
| 598 |
+
B, Hq, S, D = map(int, q.shape)
|
| 599 |
+
_, Hkv, Sk, Dk = map(int, k.shape)
|
| 600 |
+
if Hq % Hkv != 0:
|
| 601 |
+
raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")
|
| 602 |
+
G = Hq // Hkv
|
| 603 |
+
scale = 1.0 / math.sqrt(float(D))
|
| 604 |
+
|
| 605 |
+
# [B, Hkv, G, S, D]
|
| 606 |
+
q_g = q.reshape(B, Hkv, G, S, D)
|
| 607 |
+
|
| 608 |
+
# vmap over (Hkv, G): each call gets q[B,S,D] for one query head and the
|
| 609 |
+
# matching k/v[B,S,D] for its KV head (broadcast across G, no copy of
|
| 610 |
+
# the underlying K/V buffer beyond what vmap's batching rule does).
|
| 611 |
+
def per_kv_head(q_kv, k_h, v_h):
|
| 612 |
+
# q_kv: [G, B, S, D] k_h, v_h: [B, S, D]
|
| 613 |
+
fn = lambda qh: _pallas_gqa_swa_head(qh, k_h, v_h, window_size, block_q, block_k, scale)
|
| 614 |
+
return jax.vmap(fn)(q_kv) # [G, B, S, D]
|
| 615 |
+
|
| 616 |
+
q_g_t = q_g.transpose(1, 2, 0, 3, 4) # [Hkv, G, B, S, D]
|
| 617 |
+
k_t = k.transpose(1, 0, 2, 3) # [Hkv, B, S, D]
|
| 618 |
+
v_t = v.transpose(1, 0, 2, 3)
|
| 619 |
+
|
| 620 |
+
out = jax.vmap(per_kv_head)(q_g_t, k_t, v_t) # [Hkv, G, B, S, D]
|
| 621 |
+
out = out.transpose(2, 0, 1, 3, 4).reshape(B, Hq, S, D) # [B, Hq, S, D]
|
| 622 |
+
return out.astype(q.dtype)
|
| 623 |
+
|
| 624 |
+
|
| 625 |
# ---------------------------------------------------------------------------
|
| 626 |
# GPU path: JAX cuDNN FlashAttention via jax.nn.dot_product_attention
|
| 627 |
# ---------------------------------------------------------------------------
|
|
|
|
| 665 |
k_s = k.transpose(0, 2, 1, 3)
|
| 666 |
v_s = v.transpose(0, 2, 1, 3)
|
| 667 |
|
| 668 |
+
def _full_attn(q_, k_, v_):
|
| 669 |
+
return jax.nn.dot_product_attention(
|
| 670 |
+
q_, k_, v_,
|
| 671 |
+
scale=scale,
|
| 672 |
+
is_causal=True,
|
| 673 |
+
implementation='cudnn',
|
| 674 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 675 |
|
| 676 |
+
if use_remat:
|
| 677 |
+
# jax.checkpoint recomputes _full_attn on the backward pass; the
|
| 678 |
+
# function itself still runs exactly ONCE per forward pass.
|
| 679 |
+
result = jax.checkpoint(_full_attn)(q_s, k_s, v_s)
|
| 680 |
+
else:
|
| 681 |
+
result = _full_attn(q_s, k_s, v_s)
|
| 682 |
+
return result.transpose(0, 2, 1, 3).astype(q.dtype)
|
| 683 |
|
| 684 |
# ββ SWA path: XLA block-tiled kernel (GPU block_size=64) βββββββββββββββββ
|
| 685 |
# This is the fast path for SWA on GPU.
|
|
|
|
| 1073 |
|
| 1074 |
# ββ Dispatch ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 1075 |
if active_backend == 'gpu':
|
| 1076 |
+
B, Hq, S, D = map(int, q.shape)
|
| 1077 |
+
if _pallas_supported(D, q.dtype):
|
| 1078 |
+
try:
|
| 1079 |
+
return _pallas_flash_gqa_swa(
|
| 1080 |
+
q, k, v,
|
| 1081 |
+
window_size=int(window_size),
|
| 1082 |
+
block_q=min(_PALLAS_BLOCK_Q, S) if S < _PALLAS_BLOCK_Q else _PALLAS_BLOCK_Q,
|
| 1083 |
+
block_k=min(_PALLAS_BLOCK_K, S) if S < _PALLAS_BLOCK_K else _PALLAS_BLOCK_K,
|
| 1084 |
+
)
|
| 1085 |
+
except Exception as e:
|
| 1086 |
+
# Any Pallas/Triton compile or runtime failure (unsupported
|
| 1087 |
+
# GPU arch, block size mismatch, etc.) falls back silently to
|
| 1088 |
+
# the proven cuDNN/XLA path below β training never crashes
|
| 1089 |
+
# because of this optimization. Set
|
| 1090 |
+
# VEYLON_DEBUG_PALLAS_ATTN=1 to see what actually failed.
|
| 1091 |
+
if os.environ.get('VEYLON_DEBUG_PALLAS_ATTN', '0') == '1':
|
| 1092 |
+
print(f"[veylon_attention] Pallas path failed, falling back: "
|
| 1093 |
+
f"{type(e).__name__}: {e}")
|
| 1094 |
+
# cuDNN fused attention only accepts fp16/bf16/fp8 β fp32 inputs must
|
| 1095 |
+
# go through the plain-XLA fallback further down in
|
| 1096 |
+
# _gpu_flash_gqa_swa rather than crashing on the cuDNN dtype check.
|
| 1097 |
+
if q.dtype not in (jnp.float16, jnp.bfloat16):
|
| 1098 |
+
return _block_gqa_swa(
|
| 1099 |
+
q=q, k=k, v=v,
|
| 1100 |
+
window_size=int(window_size),
|
| 1101 |
+
block_size=int(GPU_BLOCK_SIZE),
|
| 1102 |
+
use_remat=use_remat,
|
| 1103 |
+
)
|
| 1104 |
return _gpu_flash_gqa_swa(
|
| 1105 |
q=q,
|
| 1106 |
k=k,
|
|
|
|
| 1364 |
if ot.shape == (1, 4, S_t, 16): ok(f"S={S_t} β {ot.shape}")
|
| 1365 |
else: fail(f"S={S_t} wrong shape {ot.shape}")
|
| 1366 |
|
| 1367 |
+
# ββ 8. Pallas GPU kernel (forward correctness + gradient check) βββββββββ
|
| 1368 |
+
section("8 Β· Pallas/Triton FlashAttention kernel (GPU only)")
|
| 1369 |
+
if not _PALLAS_GPU_AVAILABLE:
|
| 1370 |
+
print(" (skipped β Pallas not importable in this environment)")
|
| 1371 |
+
elif _detect_backend() != 'gpu':
|
| 1372 |
+
print(" (skipped β no GPU backend detected)")
|
| 1373 |
+
elif not _pallas_supported(16, jnp.float16):
|
| 1374 |
+
cc = _gpu_compute_capability()
|
| 1375 |
+
if cc is not None and cc < (8, 0):
|
| 1376 |
+
print(f" (skipped β GPU compute capability {cc[0]}.{cc[1]} < 8.0; "
|
| 1377 |
+
f"Triton/Pallas requires Ampere or newer. cuDNN path handles "
|
| 1378 |
+
f"FlashAttention on this GPU instead.)")
|
| 1379 |
+
else:
|
| 1380 |
+
print(" (skipped β Pallas gated off for this config; "
|
| 1381 |
+
"set VEYLON_DEBUG_PALLAS_ATTN=1 for details)")
|
| 1382 |
+
else:
|
| 1383 |
+
B_, Hq_, Hkv_, S_, D_, W_ = 1, 4, 2, 130, 16, 24
|
| 1384 |
+
ks = jax.random.split(jax.random.PRNGKey(99), 3)
|
| 1385 |
+
qp = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32) * 0.1
|
| 1386 |
+
kp = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32) * 0.1
|
| 1387 |
+
vp = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32) * 0.1
|
| 1388 |
+
|
| 1389 |
+
try:
|
| 1390 |
+
out_pallas = _pallas_flash_gqa_swa(qp, kp, vp, window_size=W_, block_q=32, block_k=32)
|
| 1391 |
+
|
| 1392 |
+
# Reference via existing XLA block-tiled kernel
|
| 1393 |
+
out_ref = _block_gqa_swa(qp, kp, vp, window_size=W_, block_size=32, use_remat=False)
|
| 1394 |
+
|
| 1395 |
+
err = float(jnp.max(jnp.abs(out_pallas - out_ref)))
|
| 1396 |
+
if err < 1e-3:
|
| 1397 |
+
ok(f"Forward matches XLA reference: max err = {err:.2e}")
|
| 1398 |
+
else:
|
| 1399 |
+
fail(f"Forward MISMATCH vs XLA reference: max err = {err:.2e}")
|
| 1400 |
+
|
| 1401 |
+
# Gradient check: compare d(sum(out))/d(q,k,v) against XLA reference
|
| 1402 |
+
def loss_pallas(q, k, v):
|
| 1403 |
+
return jnp.sum(_pallas_flash_gqa_swa(q, k, v, window_size=W_, block_q=32, block_k=32))
|
| 1404 |
+
|
| 1405 |
+
def loss_ref(q, k, v):
|
| 1406 |
+
return jnp.sum(_block_gqa_swa(q, k, v, window_size=W_, block_size=32, use_remat=False))
|
| 1407 |
+
|
| 1408 |
+
gp = jax.grad(loss_pallas, argnums=(0, 1, 2))(qp, kp, vp)
|
| 1409 |
+
gr = jax.grad(loss_ref, argnums=(0, 1, 2))(qp, kp, vp)
|
| 1410 |
+
|
| 1411 |
+
names = ['dQ', 'dK', 'dV']
|
| 1412 |
+
for name, gp_i, gr_i in zip(names, gp, gr):
|
| 1413 |
+
gerr = float(jnp.max(jnp.abs(gp_i - gr_i)))
|
| 1414 |
+
if gerr < 1e-2:
|
| 1415 |
+
ok(f"{name} matches XLA autodiff: max err = {gerr:.2e}")
|
| 1416 |
+
else:
|
| 1417 |
+
fail(f"{name} MISMATCH vs XLA autodiff: max err = {gerr:.2e}")
|
| 1418 |
+
|
| 1419 |
+
except Exception as e:
|
| 1420 |
+
fail(f"Pallas kernel raised an exception: {type(e).__name__}: {e}")
|
| 1421 |
+
|
| 1422 |
# ββ Summary ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 1423 |
print(f"\n{HDR}{'β'*62}{RST}")
|
| 1424 |
if failures == 0: print(f" {PASS} All tests passed.")
|