Arush kumar commited on
Commit
b22e951
Β·
1 Parent(s): 3554d5b

Update veylon_attention.py

Browse files
Files changed (1) hide show
  1. 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(q_s, k_s, v_s):
166
- try:
167
- out = jax.nn.dot_product_attention(
168
- q_s, k_s, v_s,
169
- scale=scale,
170
- is_causal=True,
171
- implementation='cudnn',
172
- )
173
- return out
174
- except Exception:
175
- # cuDNN unavailable: fall through to XLA path below
176
- return None
177
 
178
- result = _full_attn(q_s, k_s, v_s)
179
- if result is not None:
180
- if use_remat:
181
- result = jax.checkpoint(lambda q, k, v: _full_attn(q, k, v))(q_s, k_s, v_s)
182
- return result.transpose(0, 2, 1, 3).astype(q.dtype)
 
 
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.")