From 45b5695420468b0f8067e33542d7de0de449c499 Mon Sep 17 00:00:00 2001 From: Rishabh Manoj Date: Thu, 17 Sep 2026 18:51:36 +0000 Subject: [PATCH] feat(attention): 2D Ulysses+Ring custom attention with fixed-m (pre-a2a norm reduction, unpadded ring K/V) Adds fixed-m to the 2D Ulysses + Ring custom-kernel path (`_ulysses_ring_custom_attention`, attention=ulysses_ring_custom_fixed_m and the new ulysses_ring_custom_fixed_m_per_q_block), plus the supporting producer, guard and layout fixes it needed. Ring fixed-m (R > 1) - `_ring_fixed_m_norms_pre_a2a` computes every fixed-m input on the PRE-a2a shards (each device holds all heads x 1/(U*R) of the sequence): Q row norms, K max norms and the V max are reduced in ONE pmax over (ulysses, ring), then sliced to the heads this Ulysses rank owns after the a2a. Keeping these off the post-a2a layout lets them overlap the all-to-all instead of sitting between it and the `lax.cond`. With per_q_block=True a small ulysses-axis all_to_all of the Q row norms is also issued. - V safety: `v_ok = (max|V|^2 from that pmax <= 256^2) & dtype_safe`. It is derived from the pmax of V norms, so it is bit-identical on every (ulysses, ring) rank with no further collective. (The ring kernel's own pmin, which reduces global-bound eligibility, only runs when the caller omits `all_fixed_global`; the production caller always passes it.) - `all_fixed_global = all(qn * kn <= bound(N_total)) & v_ok`, also reduced over batch, is passed to the ring kernel and is the top-level accumulate-vs-LSE `lax.cond` predicate (`pregathered_mk=True`). It is mesh-uniform by construction. Blast radius: one ineligible head / Q-block on any device sends every rank to `_lse_scan` for that layer; per_q_block=True only refines dispatch inside the LSE path. How often that happens in production is not measured. - Only Q and K go through the optimization_barrier that stops XLA duplicating their producer chains into the norm reductions; V is left out so the Q/K a2a does not wait on V's producer. - Ring K/V are passed unpadded when the ring-shard KV length is 8-aligned (`kv_pad_size=1`; the kernel slices the tail, see custom_splash_unpadded_test). - Q * log2(e) is applied before the Ulysses a2a (exact; fuses into Q's producer instead of being wrapped in relayout copies after the collective). Virtual K-centering - Centering is virtual on both R == 1 and R > 1: k_mean goes to the kernel and q . k_bar is added to the fixed bound in VMEM; K is never re-written. For R > 1, k_mean = pmean(mean(K)) over (ulysses, ring) on the pre-a2a shards, and K norms are taken around that global mean. - Controlled by `use_k_centering` ("auto" | bool, new key in base_wan_27b.yml, defaulted in pyconfig). "auto" (`resolve_k_centering`) is ON for non-ring ulysses_custom_fixed_m* and OFF on every ring path, R == 1 included. The Wan pipelines do not pass this key to the attention layer in this PR, so Wan runs the layer default. Ring is kept OFF by default because, measured at the stack tip (Wan 2.2 T2V-A14B, 720p/81f/40 steps, warm AOT, DVFS unpinned), centering is never faster and it changes output: v6e-8 R == 1 127.2 s vs 127.2 s (neutral), R = 2 129.1 s vs 128.6 s (+0.4%); tpu7x-8 U=2/R=2 108.04 s vs 107.66 s (+0.35%), U=4 neutral. Turning it on changes the bf16 trajectory, with no visible quality difference. use_k_centering=True still forces it on. Other changes - `fused_rmsnorm_rope` (kernels/fused_producers.py): XLA-level fused RMSNorm + RoPE producer for self-attention Q/K, matching nnx.RMSNorm's association (x * (rsqrt(var + eps) * scale)) bit-for-bit. It is `jax.jit`-wrapped (heads/dim/eps static): inlined under the pipeline's outer jit, but when a full-size block is called eagerly (as wan_vace_transformer_test does at 75,600 tokens) XLA fuses the chain instead of materialising ~1.5 GB of extra intermediates, which OOMed that test on the v4-8 CI runner. - GQA: kv_heads plumbed through the custom Ulysses / ring paths; flash, tokamax, non-custom ulysses and ulysses_ring raise NotImplementedError on GQA. - Guards: non-ring ulysses kernels reject a ulysses_shards they cannot honour; ring kernels warn once when U == CP makes the ring degenerate (R = 1) and suggest a valid U for the mesh and head count. - Dot-product fallback: fixes the split_head_dim reshape that silently mis-read [B, H, S, D] inputs as [B, S, H*D]. - tile_size_grid_search: rejects iters < 1; `_broadcast_winner` handles bkv_compute=None. wan_block_benchmark / LTX2 / pyconfig know the new kernel names. Tests - ring_fixed_m_test.py (24): adds a U=2/R=2 integration class through `_ulysses_ring_custom_attention` vs a dense f32 reference (per_q_block x centering x GQA grid; ragged ring KV with kv_pad_size=1; a sink head that must force `_lse_scan` on all 4 devices; a single V outlier at B=2 that must make v_ok and all_fixed_global False on all 4 devices; and B=2 `_accumulate_scan` vmap + sub-128 head_dim `k_mean` padding at d=64). The branch predicates are checked through a probe of `_ring_fixed_m_norms_pre_a2a`, not the production cond itself. Needs >= 4 devices (explicit skipIf). - custom_splash_fixed_m_test.py (32), attention_config_guards_test.py (20), custom_splash_unpadded_test.py (12, TPU-only), dot_fallback_layout_test.py (5), fused_producers_test.py (6), tile_size_grid_search_test.py (16). - CI: 18 more tests skip in GitHub Actions (36 of 115 cumulative): custom_splash_unpadded_test (12), the GQA ring numerics test (1), and the 5 U=2/R=2 integration tests. Contract, guard, dot-fallback layout, fused producer and grid-search tests run in CI. run_wan_stack_tests.sh gains this PR's five new test files. Verified (final tree): TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 115 passed, 107 subtests passed. --- end_to_end/tpu/run_wan_stack_tests.sh | 6 + src/maxdiffusion/configs/base_wan_27b.yml | 22 +- src/maxdiffusion/configs/ltx2_3_video.yml | 3 +- src/maxdiffusion/configs/ltx2_video.yml | 3 +- src/maxdiffusion/kernels/fused_producers.py | 122 ++ .../splash_attention/ring_attention_kernel.py | 208 +-- src/maxdiffusion/models/attention_flax.py | 1136 ++++++++++++----- .../models/ltx2/attention_ltx2.py | 1 + src/maxdiffusion/pyconfig.py | 4 + .../tests/attention_config_guards_test.py | 260 ++++ .../tests/custom_splash_fixed_m_test.py | 42 +- .../tests/custom_splash_unpadded_test.py | 419 ++++++ .../tests/dot_fallback_layout_test.py | 165 +++ .../tests/fused_producers_test.py | 248 ++++ .../tests/ltx2/test_attention_ltx2.py | 1 + src/maxdiffusion/tests/ring_fixed_m_test.py | 394 +++++- .../tests/tile_size_grid_search_test.py | 17 + .../utils/tile_size_grid_search.py | 8 +- src/maxdiffusion/utils/wan_block_benchmark.py | 4 + 19 files changed, 2674 insertions(+), 389 deletions(-) create mode 100644 src/maxdiffusion/kernels/fused_producers.py create mode 100644 src/maxdiffusion/tests/attention_config_guards_test.py create mode 100644 src/maxdiffusion/tests/custom_splash_unpadded_test.py create mode 100644 src/maxdiffusion/tests/dot_fallback_layout_test.py create mode 100644 src/maxdiffusion/tests/fused_producers_test.py diff --git a/end_to_end/tpu/run_wan_stack_tests.sh b/end_to_end/tpu/run_wan_stack_tests.sh index 886c8f74a..850a6318e 100755 --- a/end_to_end/tpu/run_wan_stack_tests.sh +++ b/end_to_end/tpu/run_wan_stack_tests.sh @@ -28,5 +28,11 @@ TESTS=( # Fixed-m custom splash kernel and ring (feat/fixed-m-kernel) "$T/custom_splash_fixed_m_test.py" "$T/ring_fixed_m_test.py" + # Ulysses x Ring attention (feat/ring-attention) + "$T/attention_config_guards_test.py" + "$T/custom_splash_unpadded_test.py" + "$T/dot_fallback_layout_test.py" + "$T/fused_producers_test.py" + "$T/tile_size_grid_search_test.py" ) PYTHONPATH="src${PYTHONPATH:+:$PYTHONPATH}" exec "${PYTHON:-python3}" -m pytest -q -rs "${TESTS[@]}" "$@" diff --git a/src/maxdiffusion/configs/base_wan_27b.yml b/src/maxdiffusion/configs/base_wan_27b.yml index 85357b068..a1e2e0353 100644 --- a/src/maxdiffusion/configs/base_wan_27b.yml +++ b/src/maxdiffusion/configs/base_wan_27b.yml @@ -103,7 +103,7 @@ svg_low_noise_density: -1.0 # {"block_q":3328,"block_kv":2816,"block_kv_compute":256, # "block_kv_compute_in":256,"heads_per_tile":1,"vmem_limit_bytes":67108864}. svg_flash_block_sizes: {} -attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, tokamax_ring_custom, ulysses, ulysses_custom, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_bidir +attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, tokamax_ring_custom, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block, ulysses_ring_custom_bidir # @@ -112,9 +112,27 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud # CP4 (v7x-8): ulysses_shards=2 (R=2), BQ=9472 # CP8 (v7x-8): ulysses_shards=4 (R=2), BQ=9472 # CP16 (v7x-16): ulysses_shards=8 (R=2), BQ=9472 +# +# WARNING: ulysses_shards splits the context axis into (U, R = CP / U) on the +# *ring* variants and is validated (must be -1 or equal to CP) on the non-ring +# ulysses/ulysses_custom* variants: +# - Setting U == CP on a ring variant gives R=1, which is a DEGENERATE ring: +# no KV is rotated and the result is mathematically equivalent to the +# non-ring ulysses_custom* kernel (fixed-m numerics may differ because +# "auto" K-centering is off for ring variants and on for non-ring). Such a +# run must not be reported as a ring result. The attention layer logs a +# warning when this happens. +# - The non-ring ulysses/ulysses_custom* kernels always use the full context +# axis, so their Ulysses degree is fixed at CP. Passing ulysses_shards > 0 +# that does not equal CP raises ValueError instead of being silently ignored. use_base2_exp: True use_experimental_scheduler: True -# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this. +# auto: on for non-ring ulysses_custom_fixed_m* (virtual, no copy), off for ring +# paths (virtual K-centering via k_mean, plus a pmean across ring shards when R>1). +# Note: the Wan pipelines do not pass this key to the attention layer yet, so +# the layer's default ("auto") applies. +use_k_centering: "auto" +# For attention=ulysses_ring*, hidden Ulysses shard count; ring shards are context / this. ulysses_shards: -1 # Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder. # For communication-compute overlap to be effective, enable the following XLA flags: diff --git a/src/maxdiffusion/configs/ltx2_3_video.yml b/src/maxdiffusion/configs/ltx2_3_video.yml index aacd42edf..9c5e052c4 100644 --- a/src/maxdiffusion/configs/ltx2_3_video.yml +++ b/src/maxdiffusion/configs/ltx2_3_video.yml @@ -5,7 +5,8 @@ skip_jax_distributed_system: False # dot_product, flash, tokamax_flash, tokamax_ring, tokamax_ring_custom, # ulysses, ulysses_custom, ulysses_custom_fixed_m, # ulysses_custom_fixed_m_per_q_block, ulysses_ring, -# ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_bidir, +# ulysses_ring_custom, ulysses_ring_custom_fixed_m, +# ulysses_ring_custom_fixed_m_per_q_block, ulysses_ring_custom_bidir, # and cudnn_flash_te (GPU only). attention: 'flash' use_base2_exp: False diff --git a/src/maxdiffusion/configs/ltx2_video.yml b/src/maxdiffusion/configs/ltx2_video.yml index a3b77bb79..261757784 100644 --- a/src/maxdiffusion/configs/ltx2_video.yml +++ b/src/maxdiffusion/configs/ltx2_video.yml @@ -5,7 +5,8 @@ skip_jax_distributed_system: False # dot_product, flash, tokamax_flash, tokamax_ring, tokamax_ring_custom, # ulysses, ulysses_custom, ulysses_custom_fixed_m, # ulysses_custom_fixed_m_per_q_block, ulysses_ring, -# ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_bidir, +# ulysses_ring_custom, ulysses_ring_custom_fixed_m, +# ulysses_ring_custom_fixed_m_per_q_block, ulysses_ring_custom_bidir, # and cudnn_flash_te (GPU only). attention: 'flash' use_base2_exp: False diff --git a/src/maxdiffusion/kernels/fused_producers.py b/src/maxdiffusion/kernels/fused_producers.py new file mode 100644 index 000000000..795eb90a2 --- /dev/null +++ b/src/maxdiffusion/kernels/fused_producers.py @@ -0,0 +1,122 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Optimized fused producers for Wan Attention.""" + +import functools +from typing import Tuple + +import jax +import jax.numpy as jnp + + +# Static (non-array) parameters of the producer. They select shapes/epsilons and +# must be compile-time constants so the jitted wrapper below can trace them. +_STATIC_ARGNAMES = ("q_heads", "kv_heads", "dim_head", "eps", "k_eps") + + +@functools.partial(jax.jit, static_argnames=_STATIC_ARGNAMES) +def fused_rmsnorm_rope( + raw_q: jax.Array, + raw_k: jax.Array, + q_norm_scale: jax.Array, + k_norm_scale: jax.Array, + freqs_cis: jax.Array, + q_heads: int, + kv_heads: int | None = None, + dim_head: int = 128, + eps: float = 1e-6, + k_eps: float | None = None, +) -> Tuple[jax.Array, jax.Array]: + """XLA-fused FP32 RMSNorm + BF16 RoPE + Head Transposition producer. + + Performs FP32 RMSNorm normalization for numerical stability, casts to input + dtype (e.g. BF16), and applies RoPE rotation and head transposition in input + dtype precision, avoiding FP32 RoPE intermediates on long sequence lengths. + Accepts an optional `kv_heads` for GQA (`kv_heads != q_heads`); defaults to + `q_heads` (MHA, as used by `FlaxWanAttention`). + + The function is `jax.jit`-wrapped (shape/epsilon parameters static). Inside the + pipeline's outer jit this is inlined and changes nothing; when called eagerly + (e.g. a full-size transformer block exercised in a unit test without an outer + jit) it lets XLA fuse the normalize/rotate chain instead of materialising + every intermediate, which otherwise exceeds HBM on long sequences. + + Args: + raw_q: Raw query projection of shape [B, Sq, Dq] (where Dq = q_heads * dim_head). + raw_k: Raw key projection of shape [B, Sk, Dk] (where Dk = kv_heads * dim_head). + q_norm_scale: RMSNorm scale parameter for query of shape [Dq]. + k_norm_scale: RMSNorm scale parameter for key of shape [Dk]. + freqs_cis: Complex rotary embedding tensor of shape [1, 1, S, dim_head // 2]. + q_heads: Number of query attention heads. + kv_heads: Number of key/value attention heads (defaults to q_heads for MHA). + dim_head: Dimension of each attention head. + eps: Epsilon for query RMSNorm numerical stability (and key if k_eps is None). + k_eps: Optional separate epsilon for key RMSNorm numerical stability. + + Returns: + Transposed and RoPE-rotated (q_out, k_out) of shapes [B, q_heads, Sq, dim_head] + and [B, kv_heads, Sk, dim_head]. + """ + kv_heads = q_heads if kv_heads is None else kv_heads + effective_k_eps = eps if k_eps is None else k_eps + B, Sq, Dq = raw_q.shape + _, Sk, Dk = raw_k.shape + + if Dq != q_heads * dim_head: + raise ValueError(f"raw_q feature dim ({Dq}) must equal q_heads ({q_heads}) * dim_head ({dim_head})") + if Dk != kv_heads * dim_head: + raise ValueError(f"raw_k feature dim ({Dk}) must equal kv_heads ({kv_heads}) * dim_head ({dim_head})") + + # 1. FP32 RMSNorm for stability, then cast directly to target activation dtype. + # + # Association matters: Flax's `_normalize` computes `mul = rsqrt(var + eps)`, + # then `mul *= scale`, then `y = x * mul` -- i.e. x * (rsqrt * scale). Folding + # left-to-right as (x * rsqrt) * scale rounds differently and makes this path + # drift from `nnx.RMSNorm` bit-for-bit. Keep the parenthesisation below in + # step with Flax so the fused producer stays a pure fusion, not a numerical + # change. + q_fp32 = raw_q.astype(jnp.float32) + q_rms = jax.lax.rsqrt(jnp.mean(jnp.square(q_fp32), axis=-1, keepdims=True) + eps) + q_norm = (q_fp32 * (q_rms * q_norm_scale.astype(jnp.float32))).astype(raw_q.dtype) + + k_fp32 = raw_k.astype(jnp.float32) + k_rms = jax.lax.rsqrt(jnp.mean(jnp.square(k_fp32), axis=-1, keepdims=True) + effective_k_eps) + k_norm = (k_fp32 * (k_rms * k_norm_scale.astype(jnp.float32))).astype(raw_k.dtype) + + # 2. Reshape and transpose to [B, heads, S, dim_head] + q_h = q_norm.reshape(B, Sq, q_heads, dim_head).transpose(0, 2, 1, 3) + k_h = k_norm.reshape(B, Sk, kv_heads, dim_head).transpose(0, 2, 1, 3) + + # 3. Direct RoPE with freqs_cis [1, 1, S, dim_head // 2] in input dtype + cos = jnp.real(freqs_cis).astype(raw_q.dtype) + sin = jnp.imag(freqs_cis).astype(raw_q.dtype) + cos_q, sin_q = cos[:, :, :Sq, :], sin[:, :, :Sq, :] + cos_k, sin_k = cos[:, :, :Sk, :], sin[:, :, :Sk, :] + + q_pairs = q_h.reshape(B, q_heads, Sq, -1, 2) + q_0, q_1 = q_pairs[..., 0], q_pairs[..., 1] + q_out_0 = q_0 * cos_q - q_1 * sin_q + q_out_1 = q_0 * sin_q + q_1 * cos_q + q_out = jnp.stack([q_out_0, q_out_1], axis=-1).reshape(B, q_heads, Sq, dim_head) + + k_pairs = k_h.reshape(B, kv_heads, Sk, -1, 2) + k_0, k_1 = k_pairs[..., 0], k_pairs[..., 1] + k_out_0 = k_0 * cos_k - k_1 * sin_k + k_out_1 = k_0 * sin_k + k_1 * cos_k + k_out = jnp.stack([k_out_0, k_out_1], axis=-1).reshape(B, kv_heads, Sk, dim_head) + + return q_out, k_out diff --git a/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py b/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py index 0420a4dec..8f4b4277e 100644 --- a/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py +++ b/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py @@ -762,8 +762,6 @@ def _custom_bidirectional_ring_forward( axis (no sub-group perm). """ axis_size = lax.axis_size(ring_axis) - effective_kv_seq_len = orig_kv_seq_len * axis_size - recenter, ring_safe_bound = custom_splash.get_fixed_m_constants(effective_kv_seq_len) idx = lax.axis_index(ring_axis) exp_fn = jnp.exp2 if use_base2_exp else jnp.exp @@ -812,35 +810,34 @@ def _merge(m, l, o, mc, lc, oc, valid): ) # Prime buffers for t=1 (one hop each direction): device i -> KV_{i-1}, KV_{i+1}. - kr, vr = shift_r(k), shift_r(v) - kl, vl = shift_l(k), shift_l(v) - - def body(carry, t): - m, l, o, kr, vr, kl, vl = carry + if axis_size > 1: + kr, vr = shift_r(k), shift_r(v) + kl, vl = shift_l(k), shift_l(v) + + # Static Python loop over the hops rather than lax.scan: the last hop (t == + # axis_size - 1) must not issue trailing shift_r/shift_l ppermutes (which + # survive DCE and waste 4 shard transfers per layer), and issuing the next + # hop's shifts before _attn overlaps the ICI transfers with kernel compute. + for t in range(1, axis_size): + is_last_hop = t == axis_size - 1 + if not is_last_hop: + kr_n, vr_n = shift_r(kr), shift_r(vr) + kl_n, vl_n = shift_l(kl), shift_l(vl) valid_r = (idx - t) >= 0 valid_l = (idx + t) <= (axis_size - 1) # Feed real (own) K/V on invalid steps so _attn never runs on a degenerate # zero buffer (line ends receive 0 from the partial ppermute); masked below. kr_s, vr_s = jnp.where(valid_r, kr, k), jnp.where(valid_r, vr, v) kl_s, vl_s = jnp.where(valid_l, kl, k), jnp.where(valid_l, vl, v) - # Compute against the current shards (KV_{i-t}, KV_{i+t}) ... + # Compute against the current shards (KV_{i-t}, KV_{i+t}). o_r, m_r, l_r = _attn(kr_s, vr_s) m, l, o = _merge(m, l, o, m_r, l_r, o_r, valid_r) o_l, m_l, l_l = _attn(kl_s, vl_s) m, l, o = _merge(m, l, o, m_l, l_l, o_l, valid_l) - # ... and prefetch the next hop (independent of the matmuls above -> overlaps). - kr_n, vr_n = shift_r(kr), shift_r(vr) - kl_n, vl_n = shift_l(kl), shift_l(vl) - return (m, l, o, kr_n, vr_n, kl_n, vl_n), None - - (_, l_final, o_final, *_), _ = lax.scan( - body, - (m, l, o, kr, vr, kl, vl), - xs=jnp.arange(1, axis_size), - length=axis_size - 1, - unroll=True, - ) + if not is_last_hop: + kr, vr, kl, vl = kr_n, vr_n, kl_n, vl_n + l_final, o_final = l, o l_inv = jnp.where(l_final == 0.0, 0.0, 1.0 / l_final) return (o_final * l_inv[..., None]).astype(q.dtype) @@ -869,6 +866,7 @@ def _custom_ring_attention_forward( k_mean: jax.Array | None = None, uniform_fixed_m: bool | None = None, v_ok: jax.Array | bool | None = None, + all_fixed_global: jax.Array | bool | None = None, ) -> jax.Array: """Forward-only ring attention using the custom dense splash kernel. @@ -906,16 +904,23 @@ def _custom_ring_attention_forward( whether `fixed_m_norms` are squared (`True`) or unsquared (`False`). per_q_block: Whether `fixed_m_norms[0]` is per-Q-block `(num_q_heads, num_q_blocks)` (`True`) or per-head `(num_q_heads,)` (`False`). - pregathered_mk: Whether `fixed_m_norms[1]` is already gathered across the ring axis. + pregathered_mk: When `True` and `fixed_m_norms[1]` is 1D `(num_q_heads,)` + (or `(num_kv_heads,)` under GQA), treats it as already reduced across + `ring_axis` and skips the internal `lax.pmax`. k_mean: Optional per-KV-head key mean, shape `(num_kv_heads, head_dim_qk)`. When supplied (together with norms computed on the centered keys), logits - are virtually centered; when None the kernel uses raw keys. + are virtually centered; when None the kernel uses raw keys. With fixed-m, + `k_mean.shape[0]` must equal num_kv_heads (so a Q-head-indexed array + under GQA, or a sublane-padded one, raises ValueError), and head_dim is + zero-padded to q's head_dim. uniform_fixed_m: True forces the fixed-m accumulate path and bypasses the eligibility gates and `v_ok`; False forces the per-hop LSE path; None (default) dispatches on `all_fixed_global`. v_ok: Cross-ring-reduced scalar predicate asserting that value magnitudes and activation dtype are safe for fixed-m. Required when `use_fixed_m=True` and `uniform_fixed_m is not True`. + all_fixed_global: Optional precomputed mesh-uniform scalar predicate for the + accumulate-vs-LSE `lax.cond`. When supplied, skips the internal ring `pmin`. Returns: Normalized attention output, shape `(num_q_heads, q_seq_len, head_dim_v)`. @@ -1000,16 +1005,17 @@ def _custom_ring_attention_forward( exp_fn = jnp.exp2 if use_base2_exp else jnp.exp if use_fixed_m: - # Fixed-m ring: if the caller supplies `k_mean` (and matching centered - # norms), logits are virtually centered; otherwise keys are raw. Either way - # eligibility uses the two-sided floor(W/2) bound, so no row-max >= 0 - # assumption is needed. - # We gather each rank's squared K-shard norms once before the scan: mk_all_sq (R, heads), - # and form mk_global_sq = mk_all_sq.max(axis=0). The global gate uses - # floor(W(N_total)/2); when all_fixed_global holds, every hop evaluates the - # identical m_fixed, enabling direct FP32 (o_sum, l_sum) accumulation. - # In the hybrid (LSE-merge) branch each hop is gated on its own, against the - # two-sided bound for the local shard length, floor(W(N_local)/2). + # Fixed-m ring: the caller may pass already-centered K (with norms taken + # from the centered K), or a `k_mean` for in-kernel virtual centering; + # otherwise keys are raw. This kernel never computes or reduces `k_mean`. + # Either way eligibility uses the two-sided floor(W/2) bound, so no + # row-max >= 0 assumption is needed. + # The K norms are reduced once, before the ring loop, to the ring-wide max + # mk_global_sq (heads,). The global gate uses floor(W(N_total)/2); when + # all_fixed_global holds, every hop evaluates the identical m_fixed, + # enabling direct FP32 (o_sum, l_sum) accumulation. In the hybrid + # (LSE-merge) branch every hop is gated against the per-shard bound + # floor(W(N_local)/2), using the same ring-wide mk_global_sq on every hop. # All Cauchy-Schwarz gating is computed in squared-norm space # (|q|^2 * R_k^2 <= floor(W/2)^2), with no square roots in the gates. if fixed_m_norms is None: @@ -1020,10 +1026,11 @@ def _custom_ring_attention_forward( "boolean (True for squared norms, False for unsquared norms)." ) # The V-magnitude / dtype safety verdict is NOT re-derivable from Q/K norms, - # so the kernel cannot reconstruct it and must not assume it. Treating an - # omitted predicate as permission silently re-enables fixed-m for inputs it - # cannot represent -- e.g. float16 with Q=K=0 and V=1 overflows to inf. Fail - # closed unless `uniform_fixed_m=True` explicitly bypasses the gates. + # so the kernel cannot reconstruct it and must not assume it unless the + # caller explicitly forces `uniform_fixed_m=True`. Treating an omitted + # predicate as permission silently re-enables fixed-m for inputs it cannot + # represent -- e.g. float16 with Q=K=0 and V=1 overflows to inf. Fail + # closed: require the caller to state the verdict explicitly. if v_ok is None and uniform_fixed_m is not True: raise ValueError( "use_fixed_m on the ring path requires an explicit `v_ok` predicate " @@ -1050,17 +1057,18 @@ def _custom_ring_attention_forward( # exactly; a -inf init meeting an empty partial would produce inf - inf = NaN. lse_init = -1e30 - # Every rank's squared K-shard norms, gathered ONCE before the scan: (R, heads). - # A pre-gathered array keeps the per-hop gate collective-free and avoids - # serializing a third ppermute alongside K/V transfers. - if pregathered_mk or (mk_h_init_sq.ndim == 2 and mk_h_init_sq.shape[0] == axis_size): - mk_all_sq = mk_h_init_sq + # Ring-wide max K norm: (heads,). When `mk_h_init_sq` is passed as a 2D + # `(R, heads)` array, collapse the hop axis directly. When 1D `(heads,)`, + # either it is already ring-reduced (`pregathered_mk=True`) or reduce across + # `ring_axis` via `pmax` rather than `all_gather(...).max(axis=0)`. + if mk_h_init_sq.ndim == 2: + mk_global_sq = mk_h_init_sq.max(axis=0) + elif pregathered_mk: + mk_global_sq = mk_h_init_sq else: - mk_all_sq = lax.all_gather(mk_h_init_sq, ring_axis) # (axis_size, heads) - my_ring_index = lax.axis_index(ring_axis) + mk_global_sq = lax.pmax(mk_h_init_sq, ring_axis) num_q_blocks = (orig_q_seq_len + block_sizes.block_q - 1) // block_sizes.block_q - mk_global_sq = mk_all_sq.max(axis=0) # (heads,) # Validate the query-norm rank against `per_q_block`. Both gates below # multiply qn by `mk[:, None]`, so a (heads,) array supplied while @@ -1089,30 +1097,31 @@ def _custom_ring_attention_forward( # to be carried in and applied to every eligibility decision below -- # including the per-hop ones in `fixed_body`. v_gate = True if v_ok is None else v_ok - all_fixed_global = None if not per_q_block: bound_sq_1d = qn_max_sq * mk_global_sq - m_base_1d = jnp.ceil(jnp.sqrt(bound_sq_1d)) - global_recenter - m_base_expanded = jnp.broadcast_to(m_base_1d[:, None], (num_q_heads, num_q_blocks)) - fixed_ok_expanded = jnp.ones_like(m_base_expanded) - mk_arr = jnp.stack([m_base_expanded, fixed_ok_expanded], axis=0) qn_blocks_sq = jnp.broadcast_to(qn_max_sq[:, None], (num_q_heads, num_q_blocks)) - if uniform_fixed_m is None: + if uniform_fixed_m is None and all_fixed_global is None: all_fixed_local = jnp.all(bound_sq_1d <= global_centered_bound_sq) & v_gate all_fixed_global = lax.pmin(all_fixed_local, ring_axis) else: qn_blocks_sq = qn_max_sq bound_blocks_sq = qn_blocks_sq * mk_global_sq[:, None] - m_base = jnp.ceil(jnp.sqrt(bound_blocks_sq)) - global_recenter - fixed_ok_expanded = jnp.ones_like(m_base) - mk_arr = jnp.stack([m_base, fixed_ok_expanded], axis=0) # (2, heads, num_q_blocks) - if uniform_fixed_m is None: + if uniform_fixed_m is None and all_fixed_global is None: fixed_ok_local = bound_blocks_sq <= global_centered_bound_sq all_fixed_local = jnp.all(fixed_ok_local) & v_gate all_fixed_global = lax.pmin(all_fixed_local, ring_axis) def _accumulate_scan(_): + if not per_q_block: + m_base_1d = jnp.ceil(jnp.sqrt(bound_sq_1d)) - global_recenter + m_base_expanded = jnp.broadcast_to(m_base_1d[:, None], (num_q_heads, num_q_blocks)) + fixed_ok_expanded = jnp.ones_like(m_base_expanded) + mk_arr = jnp.stack([m_base_expanded, fixed_ok_expanded], axis=0) + else: + m_base = jnp.ceil(jnp.sqrt(bound_blocks_sq)) - global_recenter + fixed_ok_expanded = jnp.ones_like(m_base) + mk_arr = jnp.stack([m_base, fixed_ok_expanded], axis=0) # (2, heads, num_q_blocks) o_sum = jnp.zeros((num_q_heads, orig_q_seq_len, head_dim_v), jnp.float32) l_sum = jnp.zeros((num_q_heads, orig_q_seq_len), jnp.float32) k_current, v_current = k, v @@ -1162,7 +1171,7 @@ def _accumulate_scan(_): # the conditional, ~half of fixed-m's whole kernel win). return (o_sum * l_inv[..., None]).astype(q.dtype) - def fixed_body(carry, hop, is_last_hop): + def fixed_body(carry, is_last_hop, mk_arr): o_run, lse_run, k_current, v_current = carry # Prefetch the next shard while computing on this one. The last hop skips # it: nothing consumes the rotated shard, and the collective would still @@ -1173,18 +1182,6 @@ def fixed_body(carry, hop, is_last_hop): k_next = shift(k_current) v_next = shift(v_current) - # perm src i -> dst i+1: after `hop` shifts this rank holds the K shard - # of ring rank (my_index - hop) mod R; its norms come from the local table. - mk_h_sq = jax.lax.dynamic_index_in_dim(mk_all_sq, (my_ring_index - hop) % axis_size, keepdims=False) - bound_hop_sq = qn_blocks_sq * mk_h_sq[:, None] - # `v_gate` is load-bearing here. The Cauchy-Schwarz term is per-hop, but - # V-magnitude and dtype safety are global; recomputing eligibility from - # Q/K norms alone would re-enable fixed-m on this hop even when the - # caller's global V check already rejected it, overflowing to inf. - fixed_ok = ((bound_hop_sq <= per_shard_bound_sq) & v_gate).astype(jnp.float32) - m_base_hop = jnp.ceil(jnp.sqrt(bound_hop_sq)) - local_recenter - mk_arr = jnp.stack([m_base_hop, fixed_ok], axis=0) - o_curr, m_curr, l_curr = custom_splash._splash_attention_forward_ring( # pylint: disable=protected-access q, k_current, @@ -1216,17 +1213,42 @@ def fixed_body(carry, hop, is_last_hop): o_new = (w_run[..., None] * o_run + w_curr[..., None] * o_norm) / denom[..., None] return (o_new, lse_new + log_fn(denom), k_next, v_next), None - fixed_init = ( - jnp.zeros((num_q_heads, orig_q_seq_len, head_dim_v), jnp.float32), - jnp.full((num_q_heads, orig_q_seq_len), lse_init, jnp.float32), - k, - v, - ) - def _lse_scan(_): - carry = fixed_init + # Precompute the scalar-prefetch metadata BEFORE the ring loop so no VPU + # compute or SMEM scalar-prefetch barrier sits between `shift(k_current)` + # (`collective-permute-start`) and `_splash_attention_forward_ring`. + # + # Every hop is gated with the ring-wide `mk_global_sq` rather than the + # norm of the K shard it is processing (which would need a per-hop index + # `(my_ring_index - hop) % axis_size`): + # + # * Correctness. The max over all shards can only enlarge the + # Cauchy-Schwarz bound, so `fixed_ok` is never set where a per-shard + # bound would have cleared it: the gate is more conservative, never + # less. The cost is that one shard with large K norms forces the + # online path for that (head, Q-block) on every hop. + # * Performance. `my_ring_index` is a traced `lax.axis_index`, so a + # per-hop index put the ring index on the kernel's scalar-prefetch + # operand. The kernel could not be issued until that resolved, and the + # ~190 MiB K/V `ppermute` issued just before it -- which exists to be + # hidden behind that kernel -- was left exposed and contended with the + # output all-to-all. + # + # All hops share one metadata array, which is passed to each `fixed_body` + # call. + bound_hop_sq = qn_blocks_sq * mk_global_sq[:, None] + fixed_ok = ((bound_hop_sq <= per_shard_bound_sq) & v_gate).astype(jnp.float32) + m_base_hop = jnp.ceil(jnp.sqrt(bound_hop_sq)) - local_recenter + mk_arr_uniform = jnp.stack([m_base_hop, fixed_ok], axis=0) + + carry = ( + jnp.zeros((num_q_heads, orig_q_seq_len, head_dim_v), jnp.float32), + jnp.full((num_q_heads, orig_q_seq_len), lse_init, jnp.float32), + k, + v, + ) for hop in range(ring_size): - carry, _ = fixed_body(carry, hop, hop == ring_size - 1) + carry, _ = fixed_body(carry, hop == ring_size - 1, mk_arr_uniform) return carry[0].astype(q.dtype) if uniform_fixed_m is True: @@ -1234,7 +1256,12 @@ def _lse_scan(_): elif uniform_fixed_m is False: return _lse_scan(None) else: - return lax.cond(all_fixed_global, _accumulate_scan, _lse_scan, None) + # Collapses an explicitly shaped (e.g. per-head) predicate. It does not + # reduce a vmap batch dimension: under jax.vmap a batched predicate has + # ndim 0 per example, so the cond still lowers to a select (see the + # `all_fixed_global` note in `make_custom_ring_attention`). + cond_pred = jnp.all(all_fixed_global) if getattr(all_fixed_global, "ndim", 0) > 0 else all_fixed_global + return lax.cond(cond_pred, _accumulate_scan, _lse_scan, None) o_init = jnp.zeros((num_q_heads, orig_q_seq_len, head_dim_v), jnp.float32) l_init = jnp.zeros((num_q_heads, orig_q_seq_len), jnp.float32) @@ -1301,6 +1328,7 @@ def make_custom_ring_attention( k_mean: jax.Array | None = None, uniform_fixed_m: bool | None = None, v_ok: jax.Array | bool | None = None, + all_fixed_global: jax.Array | bool | None = None, ): """Builds a forward-only ring-attention callable around the custom kernel. @@ -1317,6 +1345,31 @@ def make_custom_ring_attention( default that reads that product as already-squared gets sqrt(2000) ~= 44.7 -- a ~45x under-estimate that wrongly admits fixed-m and overflows to inf. + `pregathered_mk`: when `True` and `fixed_m_norms[1]` is 1D `(num_q_heads,)` + (or `(num_kv_heads,)` under GQA), treats it as already ring-reduced to the + global max and skips the internal `lax.pmax` over `ring_axis`. + + `all_fixed_global` is an optional precomputed scalar predicate for the + accumulate-vs-LSE `lax.cond`. When it is given, the kernel skips its own + ring `pmin` and uses `all_fixed_global` directly as the top-level `lax.cond` + predicate, so the caller must pass a value that is identical across all + participating devices (both branches contain ring `ppermute`s) and already + includes `v_ok`. Note that `v_ok` is still required (unless `uniform_fixed_m=True`) + and is still ANDed into the per-hop `fixed_ok` gates in `_lse_scan`. When + `all_fixed_global` is omitted, the kernel derives it from `fixed_m_norms` and + `v_ok` with a ring `pmin`. + + **Under `jax.vmap`** (e.g. over the batch axis): if the predicate is batched, + `jax.vmap` lowers the `lax.cond` to a select that evaluates both the + accumulate and the LSE branch. That happens if `all_fixed_global` is omitted + and `fixed_m_norms` (or `v_ok`) are vmapped, or if a batched value is passed. + With closed-over (unbatched) norms and `v_ok`, the derived predicate is + unbatched and the cond is real. To keep a real cond with vmapped norms, + reduce the predicate over the batch outside the vmap and pass it here as an + unbatched scalar. The kernel's `jnp.all(...)` on a predicate with `ndim > 0` + does not help: inside `jax.vmap` a batched predicate has `ndim == 0` per + example, so that reduction never removes the batch dimension. + `v_ok` is a global (already cross-ring-reduced) scalar predicate asserting that the value magnitudes and activation dtype are safe for fixed-m. It is closed over rather than passed per call, since it is invariant across the batch. It is @@ -1369,6 +1422,7 @@ def _ring(q, k, v, fixed_m_norms=None, k_mean=None): k_mean=km, uniform_fixed_m=uniform_fixed_m, v_ok=v_ok, + all_fixed_global=all_fixed_global, ) return _ring diff --git a/src/maxdiffusion/models/attention_flax.py b/src/maxdiffusion/models/attention_flax.py index 6ff427b69..998d806da 100644 --- a/src/maxdiffusion/models/attention_flax.py +++ b/src/maxdiffusion/models/attention_flax.py @@ -28,6 +28,7 @@ from maxdiffusion.kernels.splash_attention import splash_attention_kernel as tokamax_splash_attention_kernel from maxdiffusion.kernels.splash_attention import ring_attention_kernel as tokamax_ring_attention_kernel from maxdiffusion.kernels.splash_attention import base as tokamax_splash_base +from maxdiffusion.kernels.fused_producers import fused_rmsnorm_rope from einops import rearrange from .. import common_types, max_logging from maxdiffusion.tpu_utils import get_tpu_type, TpuType @@ -71,6 +72,18 @@ def _coerce_tokamax_block_sizes(block_sizes): + if isinstance(block_sizes, dict): + return splash_attention_kernel.BlockSizes( + block_q=block_sizes.get("block_q", 512), + block_kv=block_sizes.get("block_kv", 512), + block_kv_compute=block_sizes.get("block_kv_compute", 512), + block_q_dkv=block_sizes.get("block_q_dkv", 512), + block_kv_dkv=block_sizes.get("block_kv_dkv", 512), + block_kv_dkv_compute=block_sizes.get("block_kv_dkv_compute", 512), + block_q_dq=block_sizes.get("block_q_dq", None), + block_kv_dq=block_sizes.get("block_kv_dq", None), + use_fused_bwd_kernel=block_sizes.get("use_fused_bwd_kernel", False), + ) # Tokamax requires fused bwd; convert if needed. if getattr(block_sizes, "use_fused_bwd_kernel", False): return block_sizes @@ -182,6 +195,126 @@ def _replace_mesh_axis_names(axis_names, old_axis: str, new_axes: tuple[str, ... return jax.sharding.PartitionSpec(*(_replace_mesh_axis(axis_name, old_axis, new_axes) for axis_name in axis_names)) +# Attention kernels are traced once per layer per transformer, so an +# unconditional log would repeat dozens of times per run and be ignored. +_WARNED_ONCE: set[str] = set() + + +def _warn_once(key: str, message: str) -> None: + """Logs `message` the first time `key` is seen in this process.""" + if key in _WARNED_ONCE: + return + _WARNED_ONCE.add(key) + max_logging.log(message) + + +def resolve_k_centering(value, *, ring: bool) -> bool: + """Resolves the `use_k_centering` setting for one attention path. + + `value` may be a bool, or "auto"/None (the config default). "auto" picks the + cheap choice per path: + * non-ring Ulysses (`ulysses_custom_fixed_m*`): ON. Centering is virtual -- + `q . k_mean` is folded into the kernel registers, with no collective and + no HBM copy -- and it tightens the fixed-m bound. + * ring paths (`ulysses_ring_custom*`): OFF. Ring centering is also + virtual (`k_mean` is passed to the kernel; R>1 adds a cross-rank pmean), + but measured end to end it is never faster (Wan 2.2 720p, 40 steps: + v6e-8 R==1 neutral, R=2 +0.4%; tpu7x-8 U=2/R=2 +0.35%, U=4 neutral) + and it changes the bf16 output trajectory, so it is opt-in via + `use_k_centering=True`. + Strings "true"/"false" (e.g. from a command-line override) are accepted; + unrecognised strings raise `ValueError`, other values are coerced with `bool()`. + """ + if value is None: + return not ring + if isinstance(value, str): + v = value.strip().lower() + if v == "auto": + return not ring + if v in ("true", "1", "yes"): + return True + if v in ("false", "0", "no"): + return False + raise ValueError(f"use_k_centering must be a bool or 'auto', got {value!r}.") + return bool(value) + + +def _validate_implicit_ulysses_degree(requested_ulysses_shards: int, context_shards: int, kernel_name: str) -> None: + """Rejects a `ulysses_shards` request the non-ring Ulysses path cannot honour. + + The non-ring kernels always shard heads across the *entire* context mesh + axis, so their Ulysses degree is implicitly `context_shards`. Silently + ignoring a different explicit request has previously caused benchmarks to + believe they were measuring U=2 while actually measuring U=4. + """ + if requested_ulysses_shards is None or requested_ulysses_shards <= 0: + return # Unset: the implicit degree is what the caller wants. + if requested_ulysses_shards == context_shards: + return # Explicit request agrees with what this path will do. + raise ValueError( + f"attention='{kernel_name}' cannot honour ulysses_shards={requested_ulysses_shards}: " + f"the non-ring Ulysses path always splits heads across the full context mesh axis, " + f"so its Ulysses degree is fixed at context_shards={context_shards}. " + f"Either set ulysses_shards={context_shards} (or leave it unset), or switch to a ring " + f"variant such as 'ulysses_ring_custom_fixed_m_per_q_block', which accepts " + f"ulysses_shards=U and forms a ring of degree R=context_shards/U." + ) + + +def _largest_ulysses_shards_for_real_ring(context_shards: int, heads: int | None = None, kv_heads: int | None = None): + """Largest Ulysses degree that still leaves a real ring (R > 1), or None if impossible. + + A usable Ulysses degree U must divide the context shard count *and* both head + counts, mirroring the constraints the ring path itself enforces. Returning the + largest such U below `context_shards` yields the smallest ring degree R > 1, + which is normally the cheapest real ring for a given mesh. + """ + for candidate in range(context_shards - 1, 0, -1): + if context_shards % candidate != 0: + continue + if heads is not None and heads % candidate != 0: + continue + if kv_heads is not None and kv_heads % candidate != 0: + continue + return candidate + return None + + +def _warn_if_ring_is_degenerate( + num_ring_shards: int, + num_ulysses_shards: int, + context_shards: int, + heads: int | None = None, + kv_heads: int | None = None, +) -> None: + """Warns when a ring variant collapses to R=1 and is really running plain Ulysses.""" + if num_ring_shards != 1: + return + + if context_shards <= 1: + advice = ( + "There is only one context shard, so no ring is possible on this mesh; " + "increase ici_context_parallelism to use a ring." + ) + else: + suggestion = _largest_ulysses_shards_for_real_ring(context_shards, heads, kv_heads) + if suggestion is None: + advice = ( + f"No ulysses_shards below context_shards={context_shards} divides both the mesh and the " + f"head counts (heads={heads}, kv_heads={kv_heads}), so this mesh cannot form a real ring." + ) + else: + advice = f"For a real ring set ulysses_shards={suggestion} (ring degree R={context_shards // suggestion})." + + _warn_once( + f"degenerate_ring:{context_shards}:{num_ulysses_shards}", + f"[attention] Ring degree R=1 (context_shards={context_shards} / ulysses_shards={num_ulysses_shards}). " + f"This ring variant is degenerate: no KV is rotated, and the result is mathematically equivalent to " + f"the corresponding non-ring Ulysses kernel (fixed-m numerics may differ: 'auto' K-centering is off " + f"for ring variants and on for non-ring). Do NOT report this as a ring-attention result. " + advice, + ) + + def _create_internal_ulysses_ring_mesh( mesh: Mesh, ring_shards: int, @@ -450,7 +583,10 @@ def _build_padding_segment_ids( kv_mask_for_batch = jnp.concatenate( [ kv_mask_for_batch, - jnp.zeros((attention_mask.shape[0], kv_padded_len - key_seq_len), jnp.int32), + jnp.zeros( + (attention_mask.shape[0], kv_padded_len - key_seq_len), + jnp.int32, + ), ], axis=1, ) @@ -567,6 +703,11 @@ def _run_chunked_ulysses_attention( Returns: The concatenated attention output tensor. """ + if query.shape[1] != key.shape[1] and ulysses_attention_chunks > 1: + raise NotImplementedError( + f"GQA (query heads {query.shape[1]} != key heads {key.shape[1]}) with " + f"ulysses_attention_chunks={ulysses_attention_chunks} > 1 is not supported." + ) head_chunk_ranges = _ulysses_head_chunk_ranges(num_heads, ulysses_shards, ulysses_attention_chunks) if len(head_chunk_ranges) > 1: chunk_outputs = [ @@ -602,8 +743,11 @@ def _tpu_flash_attention( preserve_asymmetric_block_sizes: bool = False, spatiotemporal_config: Optional[dict] = None, spatiotemporal_shape: Optional[Tuple[int, int, int]] = None, + kv_heads: Optional[int] = None, ) -> jax.Array: """TPU Flash Attention""" + if kv_heads is not None and kv_heads != heads: + raise NotImplementedError(f"{attention_kernel} does not support GQA (got heads={heads}, kv_heads={kv_heads}).") num_context_shards = mesh.shape[CONTEXT] if CONTEXT in mesh.shape else 1 query, orig_q_seq_len = _reshape_data_for_flash(query, heads, num_context_shards) @@ -754,6 +898,7 @@ def wrap_flash_attention(query, key, value, attention_mask): block_sizes=block_sizes, save_residuals=True if "ring" in attention_kernel else False, residual_checkpoint_name=residual_checkpoint_name, + interpret=(jax.default_backend() == "cpu"), ) segment_ids_in_axes = 0 if attention_mask is not None else None @@ -846,77 +991,6 @@ def ring_scan_body(carry, _): # --------------------------------------------------------------------------- -def _compute_fixed_m_metadata( - query: jax.Array, - key: jax.Array, - block_q: int, - safe_bound: float | None = None, - recenter: float | None = None, - per_q_block: bool = True, - k_mean: jax.Array | None = None, - value: jax.Array | None = None, - v_max_bound: float = 256.0, -) -> tuple[jax.Array, jax.Array]: - """Computes Cauchy-Schwarz norm bounds and per-Q-block (or per-head) fixed-m metadata.""" - batch_size, num_q_heads, q_len, _ = query.shape - num_kv_heads = key.shape[1] - if safe_bound is None or recenter is None: - rec, bnd = custom_splash.get_fixed_m_constants(key.shape[2], v_max_bound=v_max_bound) - if safe_bound is None: - safe_bound = bnd - if recenter is None: - recenter = rec - safe_bound_sq = safe_bound**2 - if k_mean is not None: - centered_k = key.astype(jnp.float32) - k_mean[:, :, None, : key.shape[-1]] - mk_h_sq = (centered_k**2).sum(axis=-1).max(axis=-1) - else: - mk_h_sq = (key.astype(jnp.float32) ** 2).sum(axis=-1).max(axis=-1) # (batch, num_kv_heads) - - if num_q_heads != num_kv_heads: - if num_q_heads % num_kv_heads != 0: - raise ValueError(f"num_q_heads ({num_q_heads}) must be divisible by num_kv_heads ({num_kv_heads}) for GQA fixed-m.") - q_heads_per_kv_head = num_q_heads // num_kv_heads - mk_h_sq = jnp.repeat(mk_h_sq, q_heads_per_kv_head, axis=1) # (batch, num_q_heads) - - dtype_safe = custom_splash.fixed_m_dtype_is_safe(query.dtype, recenter) - if dtype_safe and value is not None: - v_max_sq = (value.astype(jnp.float32) ** 2).max() - v_ok = (v_max_sq <= (v_max_bound**2)).astype(jnp.float32) - else: - v_ok = jnp.zeros((), dtype=jnp.float32) - - # The kernel's grid is ceil(q_len / block_q); callers pad Q to a multiple of - # block_q first, so floor == ceil here. Fail loudly if that contract breaks - # rather than silently dropping the ragged tail's gating metadata. - if q_len % block_q != 0: - raise ValueError( - f"_compute_fixed_m_metadata expects query padded to a multiple of block_q, got q_len={q_len}, block_q={block_q}." - ) - num_q_blocks = q_len // block_q - if per_q_block: - norm_sq = (query.astype(jnp.float32) ** 2).sum(axis=-1) # (batch, num_q_heads, q_len) - qn_max_sq = norm_sq.reshape(batch_size, num_q_heads, num_q_blocks, block_q).max( - axis=-1 - ) # (batch, num_q_heads, num_q_blocks) - bound_sq = qn_max_sq * mk_h_sq[:, :, None] - fixed_ok = (bound_sq <= safe_bound_sq).astype(jnp.float32) * v_ok - m_base = jnp.ceil(jnp.sqrt(bound_sq)) - recenter - mk_arr = jnp.stack([m_base, fixed_ok], axis=1) # (batch, 2, num_q_heads, num_q_blocks) - all_fixed = jnp.all(fixed_ok > 0.5) - else: - qn_max_sq = (query.astype(jnp.float32) ** 2).sum(axis=-1).max(axis=-1) # (batch, num_q_heads) - bound_sq_1d = qn_max_sq * mk_h_sq - fixed_ok_1d = (bound_sq_1d <= safe_bound_sq).astype(jnp.float32) * v_ok - m_base_1d = jnp.ceil(jnp.sqrt(bound_sq_1d)) - recenter - m_base_expanded = jnp.broadcast_to(m_base_1d[:, :, None], (batch_size, num_q_heads, num_q_blocks)) - fixed_ok_expanded = jnp.broadcast_to(fixed_ok_1d[:, :, None], (batch_size, num_q_heads, num_q_blocks)) - mk_arr = jnp.stack([m_base_expanded, fixed_ok_expanded], axis=1) # (batch, 2, num_q_heads, num_q_blocks) - all_fixed = jnp.all(fixed_ok_1d > 0.5) - - return mk_arr, all_fixed - - def _ulysses_attention( query: jax.Array, key: jax.Array, @@ -934,12 +1008,15 @@ def _ulysses_attention( use_base2_exp: bool = True, use_experimental_scheduler: bool = False, use_fixed_m: bool = False, - per_q_block: bool = True, ulysses_attention_chunks: int = 1, preserve_asymmetric_block_sizes: bool = False, spatiotemporal_config: Optional[dict] = None, spatiotemporal_shape: Optional[Tuple[int, int, int]] = None, - kv_heads: int | None = None, + per_q_block: bool = True, + kv_heads: Optional[int] = None, + ulysses_shards: int = -1, + kernel_name: str = "ulysses_custom", + use_k_centering: bool = True, ) -> jax.Array: """Ulysses sequence-parallel attention. @@ -947,25 +1024,42 @@ def _ulysses_attention( all-to-all collectives trade sequence shards for head shards, run local splash attention on the full sequence with a subset of heads, then all-to-all back. + + The Ulysses degree of this path is implicitly the full context mesh axis; an + explicit `ulysses_shards` that disagrees is rejected rather than ignored. """ + if kv_heads is None: + kv_heads = heads + if not use_custom_kernel and kv_heads != heads: + raise NotImplementedError(f"{kernel_name} does not support GQA (got heads={heads}, kv_heads={kv_heads}).") axis_name = CONTEXT num_shards = mesh.shape[axis_name] + _validate_implicit_ulysses_degree(ulysses_shards, num_shards, kernel_name) query, orig_q_seq_len = _reshape_data_for_flash(query, heads, num_shards) - key, orig_kv_seq_len = _reshape_data_for_flash(key, heads, num_shards) - value, _ = _reshape_data_for_flash(value, heads, num_shards) + key, orig_kv_seq_len = _reshape_data_for_flash(key, kv_heads, num_shards) + value, _ = _reshape_data_for_flash(value, kv_heads, num_shards) attention_mask = _prepare_attention_mask_for_shard_map(attention_mask, query.shape[0], key.shape[2]) if attention_mask is not None and use_custom_kernel: raise NotImplementedError( "The custom dense splash kernel (use_custom_kernel) does not support attention_mask " "(it only handles padding via orig_seq_len); got a non-None attention_mask." ) - num_heads = query.shape[1] - if num_heads % num_shards != 0: + num_q_heads = query.shape[1] + num_kv_heads = key.shape[1] + # Ulysses only redistributes existing heads across the context mesh, so + # indivisible head counts are rejected. + if num_q_heads % num_shards != 0: + raise ValueError( + "Ulysses attention requires the number of query heads to be divisible by the context shard count, " + f"got q_heads={num_q_heads} and context_shards={num_shards}." + ) + if num_kv_heads % num_shards != 0: raise ValueError( - "Ulysses attention requires the number of heads to be divisible by the context shard count, " - f"got heads={num_heads} and context_shards={num_shards}." + "Ulysses attention requires the number of KV heads to be divisible by the context shard count, " + f"got kv_heads={num_kv_heads} and context_shards={num_shards}." ) + num_heads = num_q_heads if not use_custom_kernel: block_sizes = _select_flash_block_sizes( @@ -983,6 +1077,13 @@ def _ulysses_attention( mask_needs_ulysses_gather = _mesh_axis_in_spec(kv_axis_names[2], axis_name) def wrap_ulysses_attention(query, key, value, attention_mask): + # Apply the base-2 rescale of Q *before* the all-to-all. A scalar elementwise + # multiply commutes exactly with the collective (which is pure data movement), + # so this is bit-identical. Done after the a2a it sat between the collective + # and the kernel and XLA wrapped it in relayout copies; done before, it fuses + # into the producer of Q and its 185MB round-trip disappears. + if use_custom_kernel and use_base2_exp: + query = query * LOG2E # Swap sharding: each device gives up a slice of heads and gathers # a slice of sequence, so the local kernel sees the full sequence. query = jax.lax.all_to_all(query, axis_name=axis_name, split_axis=1, concat_axis=2, tiled=True) @@ -997,11 +1098,16 @@ def wrap_ulysses_attention(query, key, value, attention_mask): "The custom dense splash kernel (use_custom_kernel) does not support attention_mask " "(it only handles padding via orig_seq_len); got a non-None attention_mask." ) - bq, bkv, bkv_compute, bkv_compute_in, heads_per_tile, vmem_limit_bytes = _extract_custom_block_sizes(flash_block_sizes) - - if use_base2_exp: - query = query * LOG2E + ( + bq, + bkv, + bkv_compute, + bkv_compute_in, + heads_per_tile, + vmem_limit_bytes, + ) = _extract_custom_block_sizes(flash_block_sizes) + # NOTE: the base-2 rescale of Q is applied before the all-to-all above. raw_key = key raw_query = query raw_value = value @@ -1013,16 +1119,26 @@ def wrap_ulysses_attention(query, key, value, attention_mask): recenter, safe_bound = custom_splash.get_fixed_m_constants(actual_kv_seq_len) query, kv_size, query_seq_len = _pad_data_for_flash(raw_query, heads, bq) - kv_pad_size = 1 if actual_kv_seq_len % 8 == 0 else bkv - key, _, key_seq_len = _pad_data_for_flash(raw_key, heads, kv_pad_size) - value, _, _ = _pad_data_for_flash(raw_value, heads, kv_pad_size) - k_mean = None - if use_fixed_m: + if use_fixed_m and use_k_centering: + # Virtual k-centering (output-invariant): project q^T \bar{k} inside the + # kernel registers without writing back / materializing (K - \bar{k}) in HBM. + # Computed strictly on real (unpadded) tokens, indexed by KV head. k_mean = jnp.mean(real_key.astype(jnp.float32), axis=2) pad_d = max(0, query.shape[-1] - k_mean.shape[-1]) if pad_d > 0: k_mean = jnp.pad(k_mean, ((0, 0), (0, 0), (0, pad_d))) + # When actual_kv_seq_len is aligned to 8 sublanes, K/V are passed with NO + # sequence padding. The fixed-m kernel slices the ragged KV tail + # (`last_compute_body_fixed` in custom_splash_attention.py) using slice + # lengths derived from the unpadded `orig_kv_seq_len`, so it never reads a + # padded K/V row; materialising the pad cost 2 x 185MB of HBM traffic per + # layer for nothing. Passing flash_block_size=1 makes only the sequence + # pad a no-op -- the head_dim->128 pad, the reshape and the returned + # (tensor, kv_size, seq_len) contract are all preserved. + kv_pad_size = 1 if actual_kv_seq_len % 8 == 0 else bkv + key, _, key_seq_len = _pad_data_for_flash(raw_key, kv_heads, kv_pad_size) + value, _, _ = _pad_data_for_flash(raw_value, kv_heads, kv_pad_size) mk_arr = None all_fixed = None @@ -1035,7 +1151,12 @@ def wrap_ulysses_attention(query, key, value, attention_mask): recenter=recenter, per_q_block=per_q_block, k_mean=k_mean, - value=value, + # Use the unpadded V: `all_fixed` gates the whole kernel through a + # lax.cond, so anything feeding it sits on the critical path. Reading + # the padded copy chained a 193MB pad + reduction behind the V + # all-to-all and left that collective fully exposed. The padding is + # zeros and the check is a max of squares, so this is output-invariant. + value=raw_value, ) bsizes = custom_splash._BlockSizes( @@ -1075,7 +1196,16 @@ def _run_uniform(q, k, v, m, km): def _run_hybrid(q, k, v, m, km): return jax.vmap(splash_kernel_hybrid, in_axes=(0, 0, 0, 0, 0))(q, k, v, m, km) - raw_out = jax.lax.cond(all_fixed, _run_uniform, _run_hybrid, query, key, value, mk_arr, k_mean) + attention_output = jax.lax.cond( + all_fixed, + _run_uniform, + _run_hybrid, + query, + key, + value, + mk_arr, + k_mean, + ) else: splash_kernel = custom_splash.make_splash_mha( block_sizes=bsizes, @@ -1088,9 +1218,18 @@ def _run_hybrid(q, k, v, m, km): use_fixed_m=False, ) vmapped_splash = jax.vmap(splash_kernel, in_axes=(0, 0, 0)) - raw_out = vmapped_splash(query, key, value) - attention_output = jnp.swapaxes(raw_out, 2, 3) - attention_output = attention_output[:, :, :query_seq_len, :kv_size].astype(query.dtype) + attention_output = vmapped_splash(query, key, value) + attention_output = attention_output[:, :, :kv_size, :context_q_seq_len].astype(query.dtype) + # Restore original layout: head-sharded/full-sequence -> sequence-sharded/full-heads. + # Sequence axis is at index 3, heads axis is at index 1. + attention_output = jax.lax.all_to_all( + attention_output, + axis_name=axis_name, + split_axis=3, + concat_axis=1, + tiled=True, + ) + return attention_output else: # Run the same local splash kernel as standard TPU flash attention, but now # on full-sequence / fewer-heads tensors produced by the all-to-all above. @@ -1124,21 +1263,22 @@ def _run_hybrid(q, k, v, m, km): block_sizes=block_sizes, save_residuals=False, residual_checkpoint_name=residual_checkpoint_name, + interpret=(jax.default_backend() == "cpu"), ) segment_ids_in_axes = 0 if attention_mask is not None else None vmapped_splash = jax.vmap(splash_kernel, in_axes=(0, 0, 0, segment_ids_in_axes)) attention_output = vmapped_splash(query, key, value, segment_ids) attention_output = attention_output[:, :, :query_seq_len, :kv_size].astype(query.dtype) - # Restore original layout: head-sharded/full-sequence -> sequence-sharded/full-heads. - attention_output = jax.lax.all_to_all( - attention_output, - axis_name=axis_name, - split_axis=2, - concat_axis=1, - tiled=True, - ) - return attention_output + # Restore original layout: head-sharded/full-sequence -> sequence-sharded/full-heads. + attention_output = jax.lax.all_to_all( + attention_output, + axis_name=axis_name, + split_axis=2, + concat_axis=1, + tiled=True, + ) + return attention_output devices_in_batch_sharding = mesh.shape["data"] * (mesh.shape["fsdp"] if "fsdp" in mesh.shape else 1) if not (query.shape[0] / devices_in_batch_sharding).is_integer(): @@ -1163,7 +1303,11 @@ def _run_hybrid(q, k, v, m, km): # Folding batch into heads destroys the one-mask-per-example association. # Keep the optimization for the common unmasked path only. fold_batch = ( - attention_mask is None and batch > 1 and devices_in_batch_sharding == 1 and (batch * num_heads) % num_shards == 0 + attention_mask is None + and batch > 1 + and devices_in_batch_sharding == 1 + and num_q_heads == num_kv_heads + and (batch * num_heads) % num_shards == 0 ) if fold_batch: query = query.reshape(1, batch * num_heads, *query.shape[2:]) @@ -1173,12 +1317,18 @@ def _run_hybrid(q, k, v, m, km): else: effective_num_heads = num_heads + out_q_axis_names = ( + jax.sharding.PartitionSpec(q_axis_names[0], q_axis_names[1], q_axis_names[3], q_axis_names[2]) + if use_custom_kernel + else q_axis_names + ) + if attention_mask is None: sharded_ulysses_attention = jax.shard_map( lambda q, k, v: wrap_ulysses_attention(q, k, v, None), mesh=mesh, in_specs=(q_axis_names, kv_axis_names, kv_axis_names), - out_specs=q_axis_names, + out_specs=out_q_axis_names, check_vma=False, ) @@ -1190,7 +1340,7 @@ def run_ulysses_attention(q, k, v): wrap_ulysses_attention, mesh=mesh, in_specs=(q_axis_names, kv_axis_names, kv_axis_names, mask_axis_names), - out_specs=q_axis_names, + out_specs=out_q_axis_names, check_vma=False, ) @@ -1207,10 +1357,19 @@ def run_ulysses_attention(q, k, v): run_ulysses_attention, ) - if fold_batch: - x = x.reshape(batch, num_heads, *x.shape[2:]) - x = x[:, :, :orig_q_seq_len, :] - x = _reshape_heads_to_head_dim(x) + if use_custom_kernel: + if fold_batch: + x = x.reshape(batch, num_heads, *x.shape[2:]) + x = x[:, :, :, :orig_q_seq_len] + b, h, d, s = x.shape + x = jnp.transpose(x, (0, 3, 1, 2)).reshape(b, -1, h * d) + axis_names = nn.logical_to_mesh_axes((BATCH, LENGTH, HEAD)) + x = jax.lax.with_sharding_constraint(x, axis_names) + else: + if fold_batch: + x = x.reshape(batch, num_heads, *x.shape[2:]) + x = x[:, :, :orig_q_seq_len, :] + x = _reshape_heads_to_head_dim(x) return x @@ -1235,6 +1394,7 @@ def _ulysses_ring_attention( ulysses_shards: int = -1, ulysses_attention_chunks: int = 1, preserve_asymmetric_block_sizes: bool = False, + kv_heads: int | None = None, ) -> jax.Array: """2D context-parallel attention using a private Ulysses x ring mesh. @@ -1243,6 +1403,10 @@ def _ulysses_ring_attention( Ulysses all-to-all over the hidden Ulysses axis, and rotates K/V over the hidden ring axis. """ + if kv_heads is None: + kv_heads = heads + if kv_heads != heads: + raise NotImplementedError(f"ulysses_ring does not support GQA (got heads={heads}, kv_heads={kv_heads}).") context_axis = CONTEXT if context_axis not in mesh.shape: @@ -1259,10 +1423,22 @@ def _ulysses_ring_attention( ) if heads % num_ulysses_shards != 0: raise ValueError( - "Ulysses ring attention requires the number of heads to be divisible by the requested Ulysses shard count, " + "Ulysses ring attention requires the number of query heads to be divisible by the requested Ulysses shard count, " f"got heads={heads} and ulysses_shards={num_ulysses_shards}." ) + if kv_heads % num_ulysses_shards != 0: + raise ValueError( + "Ulysses ring attention requires the number of KV heads to be divisible by the requested Ulysses shard count, " + f"got kv_heads={kv_heads} and ulysses_shards={num_ulysses_shards}." + ) num_ring_shards = num_context_shards // num_ulysses_shards + _warn_if_ring_is_degenerate( + num_ring_shards, + num_ulysses_shards, + num_context_shards, + heads=heads, + kv_heads=kv_heads, + ) internal_mesh = _create_internal_ulysses_ring_mesh( mesh, ring_shards=num_ring_shards, @@ -1274,8 +1450,8 @@ def _ulysses_ring_attention( num_sequence_shards = num_context_shards query, orig_q_seq_len = _reshape_data_for_flash(query, heads, num_sequence_shards) - key, _ = _reshape_data_for_flash(key, heads, num_sequence_shards) - value, _ = _reshape_data_for_flash(value, heads, num_sequence_shards) + key, _ = _reshape_data_for_flash(key, kv_heads, num_sequence_shards) + value, _ = _reshape_data_for_flash(value, kv_heads, num_sequence_shards) attention_mask = _prepare_attention_mask_for_shard_map(attention_mask, query.shape[0], key.shape[2]) num_heads = query.shape[1] @@ -1318,8 +1494,8 @@ def wrap_ulysses_ring_attention(query, key, value, attention_mask): block_q = max(*block_q_sizes) query, kv_size, query_seq_len = _pad_data_for_flash(query, heads, block_q) block_kv = max(*block_kv_sizes) - key, _, key_seq_len = _pad_data_for_flash(key, heads, block_kv) - value, _, _ = _pad_data_for_flash(value, heads, block_kv) + key, _, key_seq_len = _pad_data_for_flash(key, kv_heads, block_kv) + value, _, _ = _pad_data_for_flash(value, kv_heads, block_kv) q_padded_len = query.shape[2] kv_padded_len = key.shape[2] @@ -1380,7 +1556,11 @@ def wrap_ulysses_ring_attention(query, key, value, attention_mask): sharded_ulysses_ring_attention = jax.shard_map( lambda q, k, v: wrap_ulysses_ring_attention(q, k, v, None), mesh=internal_mesh, - in_specs=(internal_q_axis_names, internal_kv_axis_names, internal_kv_axis_names), + in_specs=( + internal_q_axis_names, + internal_kv_axis_names, + internal_kv_axis_names, + ), out_specs=internal_q_axis_names, check_vma=False, ) @@ -1392,7 +1572,12 @@ def run_ulysses_ring_attention(q, k, v): sharded_ulysses_ring_attention = jax.shard_map( wrap_ulysses_ring_attention, mesh=internal_mesh, - in_specs=(internal_q_axis_names, internal_kv_axis_names, internal_kv_axis_names, internal_mask_axis_names), + in_specs=( + internal_q_axis_names, + internal_kv_axis_names, + internal_kv_axis_names, + internal_mask_axis_names, + ), out_specs=internal_q_axis_names, check_vma=False, ) @@ -1416,11 +1601,243 @@ def run_ulysses_ring_attention(q, k, v): return x -def _max_row_norm_per_head(x: jax.Array) -> jax.Array: - """Largest row L2 norm per head of a `[B, H, S, D]` activation.""" - row_sq = jnp.square(x).sum(axis=-1, dtype=jnp.float32) - # 1.01 keeps the result an upper bound despite bf16 mantissa loss. - return jnp.sqrt(row_sq.max(axis=(0, 2))) * 1.01 +def _slice_own_ulysses_heads(x: jax.Array, ulysses_axis: str, num_ulysses_shards: int, axis: int) -> jax.Array: + """Slices an all-heads array down to the heads this rank owns after the a2a. + + `all_to_all(split_axis=1, concat_axis=2, tiled=True)` hands rank `r` the head + block `[r * H/U, (r+1) * H/U)`, so the same static block size with a + rank-dependent offset recovers exactly the heads the rank now holds. + """ + heads_per_dev = x.shape[axis] // num_ulysses_shards + start = jax.lax.axis_index(ulysses_axis) * heads_per_dev + return jax.lax.dynamic_slice_in_dim(x, start, heads_per_dev, axis=axis) + + +def _ring_fixed_m_norms_pre_a2a( + query: jax.Array, + key: jax.Array, + value: jax.Array, + *, + ulysses_axis: str, + ring_axis: str, + num_ulysses_shards: int, + num_ring_shards: int, + block_q: int, + per_q_block: bool, + use_k_centering: bool = False, +): + """Computes all R>1 fixed-m norms and global eligibility predicates *pre* a2a. + + Inputs are the shard-local activations as they arrive from the QKV + projections: `[B, H_all, S/(U*R), D]` -- every head, a 1/U slice of this ring + shard's sequence. The equivalent post-a2a arrays are `[B, H_all/U, S/R, D]`: + the same elements, redistributed. Both forms therefore admit the same + reductions, but doing them here (including `v_max_sq`, `v_ok`, and + `all_fixed_global`) is materially cheaper: + + * Every reduction reads the projection's natural output layout. After the + all-to-all the arrays carry the collective's layout, and XLA inserts + relayout copies to feed post-a2a reductions. + * All reductions and the cross-chip `pmax` (and, with virtual K-centering, + `pmean`) collectives are completely independent of `all_to_all(query, key, value)`, + allowing XLA's latency-hiding scheduler to overlap them with the + all-to-all instead of serialising reductions and collectives between the + all-to-all and the `lax.cond`. With `per_q_block=True` a small + ulysses-axis `all_to_all` of the Q row norms is also issued. + + Note: callers must apply `jax.lax.optimization_barrier((query, key))` in the + outer scope so both this function and the subsequent `all_to_all` consume the + exact same barriered Q/K (V is intentionally left out of the barrier; see the + caller). + + Returns `(qn_dev, mk_global_sq_dev, v_ok, all_fixed_global, k_mean_dev)` sliced + to the heads this Ulysses rank owns and ready for immediate `jax.lax.cond` + dispatch post-a2a. + """ + reduce_axes = (ulysses_axis, ring_axis) + key_f32 = key.astype(jnp.float32) + + q_norm_sq = (query.astype(jnp.float32) ** 2).sum(axis=-1) + qn_head_local = q_norm_sq.max(axis=-1) + vn_local = (value.astype(jnp.float32) ** 2).max() + + if use_k_centering: + # Virtual K-centering: compute global mean and centered key norms without + # modifying `key`, keeping `all_to_all(key)` independent of `pmean` and + # passing `k_mean_dev` to the ring kernel for in-kernel projection. + k_mean_all = jax.lax.pmean(jnp.mean(key_f32, axis=2), axis_name=reduce_axes) + centered_f32 = key_f32 - k_mean_all[:, :, None, :] + kn_local = jnp.sum(centered_f32**2, axis=-1).max(axis=-1) + k_mean_dev = _slice_own_ulysses_heads(k_mean_all, ulysses_axis, num_ulysses_shards, axis=1) + else: + kn_local = jnp.sum(key_f32**2, axis=-1).max(axis=-1) + k_mean_dev = None + + # Global Q/V/K max norms in a SINGLE (ulysses, ring) pmax. + # + # K needs only the ring-wide max, so no ring-axis `all_gather` of per-shard K + # norms is issued. On TPU a ring-axis gather of these small arrays forces a + # relayout in the middle of the QKV projection's fusion region, and XLA then + # fails to fuse across it. + qn_head_global, vn_global, kn_global = jax.lax.pmax((qn_head_local, vn_local, kn_local), axis_name=reduce_axes) + + if not per_q_block: + qn_dev = _slice_own_ulysses_heads(qn_head_global, ulysses_axis, num_ulysses_shards, axis=1) + else: + batch, num_q_heads, local_seq = q_norm_sq.shape + post_a2a_seq = local_seq * num_ulysses_shards + num_q_blocks = -(-post_a2a_seq // block_q) + padded_seq = num_q_blocks * block_q + q_norm_sq_dev = jax.lax.all_to_all(q_norm_sq, axis_name=ulysses_axis, split_axis=1, concat_axis=2, tiled=True) + pad_len = padded_seq - post_a2a_seq + if pad_len > 0: + q_norm_sq_dev = jnp.pad(q_norm_sq_dev, ((0, 0), (0, 0), (0, pad_len))) + qn_dev = q_norm_sq_dev.reshape(batch, num_q_heads // num_ulysses_shards, num_q_blocks, block_q).max(axis=-1) + + num_q_heads_all = qn_head_global.shape[1] + num_kv_heads_all = kn_global.shape[1] + if num_q_heads_all != num_kv_heads_all: + if num_q_heads_all % num_kv_heads_all != 0: + raise ValueError( + f"num_q_heads ({num_q_heads_all}) must be divisible by num_kv_heads ({num_kv_heads_all}) for GQA ring fixed-m." + ) + q_heads_per_kv_head = num_q_heads_all // num_kv_heads_all + kn_global_q = jnp.repeat(kn_global, q_heads_per_kv_head, axis=1) + else: + kn_global_q = kn_global + + # Slice the ring-reduced max K norm down to the heads owned by this Ulysses + # rank post-a2a: shape (batch, num_q_heads_dev), consumed directly via + # `pregathered_mk=True` in `_custom_ring_attention_forward`. + mk_global_sq_dev = _slice_own_ulysses_heads(kn_global_q, ulysses_axis, num_ulysses_shards, axis=1) + + # Evaluate global V safety and Cauchy-Schwarz fixed-m eligibility pre-a2a on + # the full-head `qn_head_global`, `kn_global_q`, and `vn_global`. Because all + # three come from the single `(ulysses, ring)` pmax above, `v_ok` and + # `all_fixed_global` are bit-identical across the entire `(ulysses, ring)` + # mesh with no further collectives. `all_fixed_global` is also reduced over + # the batch, so it is an unbatched scalar under the ring kernel's `jax.vmap`. + # + # Blast radius: because the predicate is mesh-uniform, ONE ineligible head + # (any batch element, any Q row, any device) sends EVERY rank to `_lse_scan` + # for this layer, and `per_q_block=True` does not narrow that decision -- it + # only refines the per-(head, Q-block) dispatch inside the LSE path. The + # uniform `_accumulate_scan` branch is all-or-nothing. The cost of that + # fallback in production (how often a layer trips it) has not been measured. + effective_kv_seq_len = key.shape[2] * num_ulysses_shards * num_ring_shards + global_recenter, global_bound = custom_splash.get_fixed_m_constants(effective_kv_seq_len) + global_bound_sq = global_bound**2 + dtype_safe = custom_splash.fixed_m_dtype_is_safe(query.dtype, global_recenter) + + v_max_sq = vn_global.max() + v_ok = (v_max_sq <= (custom_splash.DEFAULT_MAX_V_BOUND**2)) & dtype_safe + + bound_head_sq_all = qn_head_global * kn_global_q + all_fixed_global = jnp.all(bound_head_sq_all <= global_bound_sq) & v_ok + return qn_dev, mk_global_sq_dev, v_ok, all_fixed_global, k_mean_dev + + +def _compute_fixed_m_metadata( + query: jax.Array, + key: jax.Array, + block_q: int, + safe_bound: float | None = None, + recenter: float | None = None, + per_q_block: bool = True, + k_mean: jax.Array | None = None, + value: jax.Array | None = None, + v_max_bound: float = 256.0, +) -> tuple[jax.Array, jax.Array]: + """Computes Cauchy-Schwarz norm bounds and per-Q-block (or per-head) fixed-m metadata. + + Args: + query: Padded query activation, shape `(batch, local_q_heads, padded_q_len, head_dim)`. + key: Key activation (raw unpadded or padded), shape `(batch, local_kv_heads, kv_len, head_dim)`. + K norms are computed per KV head and repeated internally to Q heads for GQA. + block_q: Query tile block size. + safe_bound: Maximum safe norm product threshold before falling back to online softmax. + recenter: Fixed-m dynamic recenter constant C(N). + per_q_block: If True, evaluates gating independently per query tile. If False, + evaluates monolithic gating per head. + k_mean: Optional mean key vector for Virtual K-centering, KV-head indexed, + shape `(batch, >= local_kv_heads, >= head_dim)`. It is not validated: it is + sliced to `k_mean[:, :local_kv_heads, :head_dim]`, so a longer (e.g. + pre-padded or Q-head-expanded) array is silently truncated, not rejected. + value: Value activation, shape `(batch, local_kv_heads, kv_len, head_dim_v)`, used to + verify that |V| <= v_max_bound to guarantee against FP32 overflow. Omitting + `value` fails closed (`v_ok = 0.0`). + v_max_bound: Maximum safe value magnitude (default 256.0). + + Returns: + mk_arr: Gating metadata array of shape `(batch, 2, local_q_heads, num_q_blocks)` + multiplexing precomputed block base shifts and binary fixed-m gating predicates into a single + Pallas scalar prefetch memory slot: + - `mk_arr[:, 0, h, i]`: Precomputed block base shift m_B = ceil(max_i ||q_i|| * max_j ||k_j||) - C. + - `mk_arr[:, 1, h, i]`: Discrete eligibility predicate (1.0 for fixed-m, 0.0 for online). + all_fixed: Boolean scalar indicating if all elements are eligible for uniform fixed-m. + """ + batch_size, num_q_heads, q_len, _ = query.shape + num_kv_heads = key.shape[1] + if safe_bound is None or recenter is None: + rec, bnd = custom_splash.get_fixed_m_constants(key.shape[2], v_max_bound=v_max_bound) + if safe_bound is None: + safe_bound = bnd + if recenter is None: + recenter = rec + safe_bound_sq = safe_bound**2 + if k_mean is not None: + centered_k = key.astype(jnp.float32) - k_mean[:, :num_kv_heads, None, : key.shape[-1]] + mk_h_sq = (centered_k**2).sum(axis=-1).max(axis=-1) + else: + mk_h_sq = (key.astype(jnp.float32) ** 2).sum(axis=-1).max(axis=-1) # (batch, num_kv_heads) + + if num_q_heads != num_kv_heads: + if num_q_heads % num_kv_heads != 0: + raise ValueError(f"num_q_heads ({num_q_heads}) must be divisible by num_kv_heads ({num_kv_heads}) for GQA fixed-m.") + q_heads_per_kv_head = num_q_heads // num_kv_heads + mk_h_sq = jnp.repeat(mk_h_sq, q_heads_per_kv_head, axis=1) # (batch, num_q_heads) + + # Fixed-m weights reach 2**recenter before being narrowed to the activation + # dtype for the S@V matmul. If that dtype's exponent range cannot hold them + # (fp16, fp8), the FP32 bound analysis is irrelevant -- the narrowing itself + # overflows to inf -- so disqualify every head up front. Fail closed if value + # is omitted. + dtype_safe = custom_splash.fixed_m_dtype_is_safe(query.dtype, recenter) + if dtype_safe and value is not None: + v_max_sq = (value.astype(jnp.float32) ** 2).max() + v_ok = (v_max_sq <= (v_max_bound**2)).astype(jnp.float32) + else: + v_ok = jnp.zeros((), dtype=jnp.float32) + + # The kernel's grid is ceil(q_len / block_q); callers pad Q to a multiple of + # block_q first, so floor == ceil here. Fail loudly if that contract breaks + # rather than silently dropping the ragged tail's gating metadata. + if q_len % block_q != 0: + raise ValueError( + f"_compute_fixed_m_metadata expects query padded to a multiple of block_q, got q_len={q_len}, block_q={block_q}." + ) + num_q_blocks = q_len // block_q + if per_q_block: + norm_sq = (query.astype(jnp.float32) ** 2).sum(axis=-1) # (batch, num_q_heads, q_len) + qn_max_sq = norm_sq.reshape(batch_size, num_q_heads, num_q_blocks, block_q).max( + axis=-1 + ) # (batch, num_q_heads, num_q_blocks) + bound_sq = qn_max_sq * mk_h_sq[:, :, None] + fixed_ok = (bound_sq <= safe_bound_sq).astype(jnp.float32) * v_ok + m_base = jnp.ceil(jnp.sqrt(bound_sq)) - recenter + mk_arr = jnp.stack([m_base, fixed_ok], axis=1) # (batch, 2, num_q_heads, num_q_blocks) + all_fixed = jnp.all(fixed_ok > 0.5) + else: + qn_max_sq = (query.astype(jnp.float32) ** 2).sum(axis=-1).max(axis=-1) # (batch, num_q_heads) + bound_sq_1d = qn_max_sq * mk_h_sq + fixed_ok_1d = (bound_sq_1d <= safe_bound_sq).astype(jnp.float32) * v_ok + m_base_1d = jnp.ceil(jnp.sqrt(bound_sq_1d)) - recenter + m_base_expanded = jnp.broadcast_to(m_base_1d[:, :, None], (batch_size, num_q_heads, num_q_blocks)) + fixed_ok_expanded = jnp.broadcast_to(fixed_ok_1d[:, :, None], (batch_size, num_q_heads, num_q_blocks)) + mk_arr = jnp.stack([m_base_expanded, fixed_ok_expanded], axis=1) # (batch, 2, num_q_heads, num_q_blocks) + all_fixed = jnp.all(fixed_ok_1d > 0.5) + + return mk_arr, all_fixed def _ulysses_ring_custom_attention( @@ -1442,28 +1859,16 @@ def _ulysses_ring_custom_attention( bidirectional: bool = False, use_fixed_m: bool = False, ulysses_attention_chunks: int = 1, + per_q_block: bool = True, + kv_heads: int | None = None, + use_k_centering: bool = False, ) -> jax.Array: - """Hybrid Ulysses + Ring (USP) with the CUSTOM splash kernel on main's mesh. - - Uses origin/main's explicit internal `(ring, ulysses)` mesh - (`_create_internal_ulysses_ring_mesh`, commit c104db51) instead of single-axis - collective sub-groups: the public `context` axis is reshaped with the Ulysses - axis innermost, so the Ulysses all-to-all stays INTRA-chip and the ring rotates - ACROSS chips. The per-shard attention is our custom splash kernel - (`make_custom_ring_attention`), not the tokamax_ring kernel main uses. + """2D USP attention (Ulysses + Ring) using custom splash kernel with exact Fixed-m support.""" + if kv_heads is None: + kv_heads = heads - 1. all-to-all over the (intra-chip) Ulysses axis: trade sequence for heads; - 2. ring (full ppermute) over the (cross-chip) ring axis, online-softmax merge; - 3. all-to-all back to restore the sequence-sharded / full-heads layout. - - U = ulysses_shards (from config); R = context // U. U=context -> pure - Ulysses, U=1 -> pure Ring (all on the same custom kernel). - """ if attention_mask is not None: - raise NotImplementedError( - "ulysses_ring_custom does not support attention_mask (the custom splash kernels only " - "handle padding via orig_seq_len); got a non-None attention_mask." - ) + raise NotImplementedError("ulysses_ring_custom does not support attention_mask.") axis_name = "context" num_context_shards = mesh.shape[axis_name] num_ulysses_shards = ulysses_shards @@ -1476,13 +1881,32 @@ def _ulysses_ring_custom_attention( ) num_ring_shards = num_context_shards // num_ulysses_shards + # K-centering for ring variants: `use_k_centering` is resolved by the caller + # with `resolve_k_centering(..., ring=True)`, so "auto" is OFF regardless of R. + # When forced on, virtual K-centering is used (R>1: `k_mean = pmean(mean(K))` + # over (ulysses, ring) before the a2a; R==1: `k_mean = mean(K)` after the a2a), + # with norms taken from centered K and `k_mean` passed into the kernel. + _warn_if_ring_is_degenerate( + num_ring_shards, + num_ulysses_shards, + num_context_shards, + heads=heads, + kv_heads=kv_heads, + ) + query, orig_q_seq_len = _reshape_data_for_flash(query, heads, num_context_shards) - key, _ = _reshape_data_for_flash(key, heads, num_context_shards) - value, _ = _reshape_data_for_flash(value, heads, num_context_shards) + key, orig_kv_seq_len = _reshape_data_for_flash(key, kv_heads, num_context_shards) + value, _ = _reshape_data_for_flash(value, kv_heads, num_context_shards) num_heads = query.shape[1] if num_heads % num_ulysses_shards != 0: - raise ValueError(f"Ulysses+Ring requires heads divisible by U={num_ulysses_shards}, got heads={num_heads}.") - + raise ValueError(f"Ulysses+Ring requires query heads divisible by U={num_ulysses_shards}, got heads={num_heads}.") + if kv_heads % num_ulysses_shards != 0: + raise ValueError(f"Ulysses+Ring requires KV heads divisible by U={num_ulysses_shards}, got kv_heads={kv_heads}.") + if num_ring_shards > 1 and (orig_q_seq_len % num_context_shards != 0 or orig_kv_seq_len % num_context_shards != 0): + raise ValueError( + f"2D Ulysses+Ring attention requires sequence length to be divisible by context_shards={num_context_shards}, " + f"got orig_q_seq_len={orig_q_seq_len}, orig_kv_seq_len={orig_kv_seq_len}." + ) ( bq, bkv, @@ -1491,13 +1915,10 @@ def _ulysses_ring_custom_attention( heads_per_tile, vmem_limit_bytes, ) = _extract_custom_block_sizes(flash_block_sizes) - if heads_per_tile > 1: - raise NotImplementedError("ulysses_ring_custom currently supports heads_per_tile == 1 only.") - + if heads_per_tile > 1 and num_ring_shards > 1: + raise NotImplementedError("heads_per_tile > 1 is not supported for multi-shard ring attention.") internal_mesh = _create_internal_ulysses_ring_mesh(mesh, num_ring_shards, num_ulysses_shards) - ring_axis = INTERNAL_RING_AXIS - ulysses_axis = INTERNAL_ULYSSES_AXIS - + ring_axis, ulysses_axis = INTERNAL_RING_AXIS, INTERNAL_ULYSSES_AXIS q_axis_names = nn.logical_to_mesh_axes(axis_names_q) kv_axis_names = nn.logical_to_mesh_axes(axis_names_kv) internal_q_axis_names = _replace_mesh_axis_names(q_axis_names, axis_name, (ring_axis, ulysses_axis)) @@ -1515,102 +1936,98 @@ def _ulysses_ring_custom_attention( check_vma=False, ) def wrap_ulysses_ring_attention(query, key, value): - fixed_m_norms = None - v_ok = None + # Apply the base-2 rescale of Q *before* the all-to-all. A scalar elementwise + # multiply commutes exactly with the collective (which is pure data movement), + # so this is bit-identical. Done after the a2a it sat between the collective + # and the kernel and XLA wrapped it in relayout copies; done before, it fuses + # into the producer of Q and its 185MB round-trip disappears. + if use_base2_exp: + query = query * LOG2E + + # (0) R>1 fixed-m reductions and global eligibility predicates, computed + # entirely on the pre-a2a layout so zero reductions or collectives sit + # between `all_to_all` and `jax.lax.cond`. + qn_dev, mk_global_sq_dev, v_ok, all_fixed_global, k_mean_dev = None, None, None, None, None if use_fixed_m and num_ring_shards > 1: - # Fixed-m's Cauchy-Schwarz inputs and V-magnitude safety verdict, reduced - # on the PRE-a2a activations so the reductions overlap the all_to_all - # instead of stalling the first ring step (taking them after the a2a - # measured +8% end to end). - # - # The barrier is load-bearing: the norms are a second consumer of these - # activations, and without it XLA duplicates the producer chain into the - # norm fusion instead of materializing once -- worth 1.46 ms/layer, the - # difference between fixed-m breaking even and winning. - # - # Reducing them further upstream (on the flat [B, S, H*D] form, where - # head_dim is contiguous) is exact and looks cheaper, but there the array - # is still globally sharded, so the reduction becomes a per-layer - # all-reduce over the context axis: measured WORSE (+54 ms per forward). - # - # V is deliberately NOT in the barrier (the figures above were measured - # with the (query, key) barrier): putting it in would make the Q/K - # reductions and all-to-alls wait for V's producer. V has no fused - # norm/RoPE producer chain to duplicate (only its projection), so its - # max is reduced outside the barrier. + # The barrier keeps XLA from duplicating the Q/K producer chains (fused + # norm/RoPE) into the norm reductions instead of materializing once. V is + # deliberately NOT in it: V has no such producer chain to duplicate (only + # its projection), and including it would make the Q/K reductions and + # all-to-alls wait for V's producer. query, key = jax.lax.optimization_barrier((query, key)) - qn_local = (_max_row_norm_per_head(query) * (LOG2E if use_base2_exp else 1.0)) ** 2 - kn_local = _max_row_norm_per_head(key) ** 2 - # The accumulate-vs-LSE lax.cond predicate must be uniform along the RING - # axis (every ppermute participant takes the same branch). - qn_all = jax.lax.pmax(qn_local, (ring_axis, ulysses_axis)) - mk_all = jax.lax.pmax(kn_local, ulysses_axis) - heads_per_dev = qn_all.shape[0] // num_ulysses_shards - start_head = jax.lax.axis_index(ulysses_axis) * heads_per_dev - fixed_m_norms = ( - jax.lax.dynamic_slice_in_dim(qn_all, start_head, heads_per_dev), - jax.lax.dynamic_slice_in_dim(mk_all, start_head, heads_per_dev), + ( + qn_dev, + mk_global_sq_dev, + v_ok, + all_fixed_global, + k_mean_dev, + ) = _ring_fixed_m_norms_pre_a2a( + query, + key, + value, + ulysses_axis=ulysses_axis, + ring_axis=ring_axis, + num_ulysses_shards=num_ulysses_shards, + num_ring_shards=num_ring_shards, + block_q=bq, + per_q_block=per_q_block, + use_k_centering=use_k_centering, ) - effective_kv_seq_len = key.shape[2] * num_ulysses_shards * num_ring_shards - global_recenter, _ = custom_splash.get_fixed_m_constants(effective_kv_seq_len) - dtype_safe = custom_splash.fixed_m_dtype_is_safe(query.dtype, global_recenter) - v_max_sq = (value.astype(jnp.float32) ** 2).max() - v_ok_local = (v_max_sq <= (custom_splash.DEFAULT_MAX_V_BOUND**2)) & dtype_safe - v_ok = jax.lax.pmin(v_ok_local, (ring_axis, ulysses_axis)) - - # (1) Ulysses all-to-all over the (intra-chip) ulysses axis: heads -> sequence, - # so each device holds the full ring-chunk sequence with heads/U heads. + + # (1) Ulysses All-to-All: heads -> sequence a2a = functools.partial(jax.lax.all_to_all, axis_name=ulysses_axis, tiled=True) query = a2a(query, split_axis=1, concat_axis=2) key = a2a(key, split_axis=1, concat_axis=2) value = a2a(value, split_axis=1, concat_axis=2) - if use_base2_exp: - query = query * LOG2E - + # NOTE: the base-2 rescale of Q is applied before the all-to-all above. raw_key = key - query, kv_size, query_seq_len = _pad_data_for_flash(query, heads, bq) - key, _, key_seq_len = _pad_data_for_flash(key, heads, bkv) - value, _, _ = _pad_data_for_flash(value, heads, bkv) + raw_query = query + raw_value = value + context_q_seq_len = raw_query.shape[2] + actual_kv_seq_len = orig_kv_seq_len if num_ring_shards == 1 else raw_key.shape[2] + real_key = raw_key[:, :, :actual_kv_seq_len, :] + + query, kv_size, query_seq_len = _pad_data_for_flash(raw_query, heads, bq) + # When actual_kv_seq_len is aligned to 8 sublanes, K/V are passed with NO + # sequence padding. The kernel slices the ragged KV tail using slice lengths + # derived from actual_kv_seq_len, avoiding sequence pad HBM copies and + # redundant ppermute/MXU compute on padded keys. + kv_pad_size = 1 if actual_kv_seq_len % 8 == 0 else bkv + key, _, key_seq_len = _pad_data_for_flash(raw_key, kv_heads, kv_pad_size) + value, _, _ = _pad_data_for_flash(raw_value, kv_heads, kv_pad_size) + ring_kv_seq_len = actual_kv_seq_len if actual_kv_seq_len % 8 == 0 else key_seq_len k_mean = None - if use_fixed_m and num_ring_shards == 1: - k_mean = jnp.mean(raw_key.astype(jnp.float32), axis=2) + if use_fixed_m and num_ring_shards == 1 and use_k_centering: + k_mean = jnp.mean(real_key.astype(jnp.float32), axis=2) pad_d = max(0, query.shape[-1] - k_mean.shape[-1]) if pad_d > 0: k_mean = jnp.pad(k_mean, ((0, 0), (0, 0), (0, pad_d))) - mk_arr = None - all_fixed = None + mk_arr, all_fixed = None, None if use_fixed_m and num_ring_shards == 1: - recenter, safe_bound = custom_splash.get_fixed_m_constants(key_seq_len) + recenter, safe_bound = custom_splash.get_fixed_m_constants(actual_kv_seq_len) mk_arr, all_fixed = _compute_fixed_m_metadata( query, - raw_key, - block_q=bq, + real_key, + bq, safe_bound=safe_bound, recenter=recenter, - per_q_block=False, + per_q_block=per_q_block, k_mean=k_mean, - value=value, + value=raw_value, ) - bsizes = custom_splash._BlockSizes( - block_q=bq, - block_kv=bkv, - block_kv_compute=bkv_compute, - block_kv_compute_in=bkv_compute_in, - ) + bsizes = custom_splash._BlockSizes(bq, bkv, bkv_compute, bkv_compute_in) + + # (2a) R=1: Dedicated single-device splash kernel with fixed-m or online softmax if num_ring_shards == 1: - # (2a) R=1: the ring is trivial (no rotation) -> use the lighter dedicated - # splash kernel (fuse_reciprocal, no fp32 online-softmax residual windows). - # Same math as the 1-step ring, and it fits BQ=8448 where the ring kernel - # OOMs (its 3x residual windows). make_splash_mha returns [H, D, S]. if use_fixed_m: splash_kernel_uniform = custom_splash.make_splash_mha( block_sizes=bsizes, - orig_q_seq_len=query_seq_len, - orig_kv_seq_len=key_seq_len, + orig_q_seq_len=context_q_seq_len, + orig_kv_seq_len=actual_kv_seq_len, heads_per_tile=heads_per_tile, use_base2_exp=use_base2_exp, use_experimental_scheduler=use_experimental_scheduler, @@ -1620,8 +2037,8 @@ def wrap_ulysses_ring_attention(query, key, value): ) splash_kernel_hybrid = custom_splash.make_splash_mha( block_sizes=bsizes, - orig_q_seq_len=query_seq_len, - orig_kv_seq_len=key_seq_len, + orig_q_seq_len=context_q_seq_len, + orig_kv_seq_len=actual_kv_seq_len, heads_per_tile=heads_per_tile, use_base2_exp=use_base2_exp, use_experimental_scheduler=use_experimental_scheduler, @@ -1637,51 +2054,68 @@ def _run_hybrid(q, k, v, m, km): return jax.vmap(splash_kernel_hybrid, in_axes=(0, 0, 0, 0, 0))(q, k, v, m, km) raw_out = jax.lax.cond(all_fixed, _run_uniform, _run_hybrid, query, key, value, mk_arr, k_mean) - attention_output = jnp.swapaxes(raw_out, 2, 3) else: splash_kernel = custom_splash.make_splash_mha( block_sizes=bsizes, - orig_q_seq_len=query_seq_len, - orig_kv_seq_len=key_seq_len, + orig_q_seq_len=context_q_seq_len, + orig_kv_seq_len=actual_kv_seq_len, heads_per_tile=heads_per_tile, use_base2_exp=use_base2_exp, use_experimental_scheduler=use_experimental_scheduler, vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=False, ) - attention_output = jnp.swapaxes(jax.vmap(splash_kernel, in_axes=(0, 0, 0))(query, key, value), 2, 3) + raw_out = jax.vmap(splash_kernel, in_axes=(0, 0, 0))(query, key, value) + attention_output = jnp.swapaxes(raw_out, 2, 3) + + # (2b) Ring: Cross-chip ppermute schedule with custom ring kernel else: - # (2b) Ring (full ppermute over the cross-chip ring axis) with the custom kernel. - # bidirectional=True -> wrap-free schedule (streams K/V both directions one hop - # at a time), for a non-wrapping ring axis. Selected by attention=ulysses_ring_custom_bidir. - ring_kernel = tokamax_ring_attention_kernel.make_custom_ring_attention( - block_sizes=bsizes, - orig_q_seq_len=query_seq_len, - orig_kv_seq_len=key_seq_len, - use_base2_exp=use_base2_exp, - use_experimental_scheduler=use_experimental_scheduler, - vmem_limit_bytes=vmem_limit_bytes, - ring_axis=ring_axis, - ring_size=num_ring_shards, - bidirectional=bidirectional, - use_fixed_m=use_fixed_m, - fixed_m_norms=fixed_m_norms, - fixed_m_norms_squared=True if use_fixed_m else None, - v_ok=v_ok, - per_q_block=False, - ) - attention_output = jax.vmap(ring_kernel, in_axes=(0, 0, 0))(query, key, value) - attention_output = attention_output[:, :, :query_seq_len, :kv_size].astype(query.dtype) + if use_fixed_m: + ring_kernel = tokamax_ring_attention_kernel.make_custom_ring_attention( + block_sizes=bsizes, + orig_q_seq_len=query_seq_len, + orig_kv_seq_len=ring_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + ring_axis=ring_axis, + ring_size=num_ring_shards, + bidirectional=bidirectional, + use_fixed_m=True, + fixed_m_norms_squared=True, + per_q_block=per_q_block, + pregathered_mk=True, + v_ok=v_ok, + all_fixed_global=all_fixed_global, + ) + attention_output = jax.vmap(ring_kernel, in_axes=(0, 0, 0, (0, 0), 0))( + query, key, value, (qn_dev, mk_global_sq_dev), k_mean_dev + ) + else: + ring_kernel = tokamax_ring_attention_kernel.make_custom_ring_attention( + block_sizes=bsizes, + orig_q_seq_len=query_seq_len, + orig_kv_seq_len=ring_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + ring_axis=ring_axis, + ring_size=num_ring_shards, + bidirectional=bidirectional, + use_fixed_m=False, + ) + attention_output = jax.vmap(ring_kernel, in_axes=(0, 0, 0))(query, key, value) - # (3) Ulysses all-to-all back: sequence -> heads, restoring the layout. - attention_output = a2a(attention_output, split_axis=2, concat_axis=1) - return attention_output + attention_output = attention_output[:, :, :context_q_seq_len, :kv_size].astype(query.dtype) + + # (3) Ulysses All-to-All back: sequence -> heads + return a2a(attention_output, split_axis=2, concat_axis=1) x = _run_chunked_ulysses_attention( query, key, value, - num_heads, + heads, num_ulysses_shards, ulysses_attention_chunks, wrap_ulysses_ring_attention, @@ -1704,17 +2138,48 @@ def _apply_attention_dot( float32_qk_product: bool, use_memory_efficient_attention: bool, attention_mask: Array = None, + kv_heads: int | None = None, ): """Apply Attention.""" + effective_kv_heads = kv_heads if kv_heads is not None else heads if split_head_dim: - b = key.shape[0] - query_states = jnp.reshape(query, (b, -1, heads, dim_head)) - key_states = jnp.reshape(key, (b, -1, heads, dim_head)) - value_states = jnp.reshape(value, (b, -1, heads, dim_head)) + + def _to_bshd(x: Array, n_heads: int) -> Array: + """Normalise to [B, S, H, D]. + + Callers that apply rotary embeddings hand us [B, H, S, D] (see + `_unflatten_heads`), while the flat path supplies [B, S, H*D]. Only the + latter can be reshaped into [B, S, H, D]; reinterpreting [B, H, S, D] + that way keeps the shape legal but interleaves heads with tokens, so it + corrupts the output silently. Transpose the 4-D case instead. + """ + if x.ndim == 4: + return jnp.swapaxes(x, 1, 2) + return jnp.reshape(x, (x.shape[0], -1, n_heads, dim_head)) + + query_states = _to_bshd(query, heads) + key_states = _to_bshd(key, effective_kv_heads) + value_states = _to_bshd(value, effective_kv_heads) + if heads != effective_kv_heads: + num_repeats = heads // effective_kv_heads + key_states = jnp.repeat(key_states, num_repeats, axis=2) + value_states = jnp.repeat(value_states, num_repeats, axis=2) else: query_states = _reshape_heads_to_batch_dim(query, heads) - key_states = _reshape_heads_to_batch_dim(key, heads) - value_states = _reshape_heads_to_batch_dim(value, heads) + key_states = _reshape_heads_to_batch_dim(key, effective_kv_heads) + value_states = _reshape_heads_to_batch_dim(value, effective_kv_heads) + if heads != effective_kv_heads: + num_repeats = heads // effective_kv_heads + b = query.shape[0] + s_k = key_states.shape[1] + key_states = jnp.repeat(key_states.reshape(b, effective_kv_heads, s_k, -1), num_repeats, axis=1).reshape( + b * heads, s_k, -1 + ) + value_states = jnp.repeat( + value_states.reshape(b, effective_kv_heads, s_k, -1), + num_repeats, + axis=1, + ).reshape(b * heads, s_k, -1) if float32_qk_product: query_states = query_states.astype(jnp.float32) @@ -1820,6 +2285,7 @@ def dot_product_kernel(q, k, v, context): context["float32_qk_product"], context["use_memory_efficient_attention"], context["attention_mask"], + kv_heads=context.get("kv_heads", None), ) @@ -1844,6 +2310,9 @@ def ulysses_custom_kernel(q, k, v, context): ulysses_attention_chunks=context["ulysses_attention_chunks"], spatiotemporal_config=context.get("spatiotemporal_config"), spatiotemporal_shape=context.get("spatiotemporal_shape"), + kv_heads=context.get("kv_heads", None), + ulysses_shards=context.get("ulysses_shards", -1), + kernel_name="ulysses_custom", ) @@ -1866,6 +2335,7 @@ def ulysses_ring_custom_kernel(q, k, v, context): use_base2_exp=context.get("use_base2_exp", True), use_experimental_scheduler=context.get("use_experimental_scheduler", False), ulysses_attention_chunks=context["ulysses_attention_chunks"], + kv_heads=context.get("kv_heads", None), ) @@ -1896,7 +2366,37 @@ def ulysses_ring_custom_fixed_m_kernel(q, k, v, context): use_base2_exp=context.get("use_base2_exp", True), use_experimental_scheduler=context.get("use_experimental_scheduler", False), use_fixed_m=True, + per_q_block=False, ulysses_attention_chunks=context.get("ulysses_attention_chunks", 1), + kv_heads=context.get("kv_heads", None), + use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=True), + ) + + +@register_kernel("ulysses_ring_custom_fixed_m_per_q_block") +def ulysses_ring_custom_fixed_m_per_q_block_kernel(q, k, v, context): + """fixed-m variant of ulysses_ring_custom with per-Q-block gating.""" + return _ulysses_ring_custom_attention( + q, + k * context["scale"], + v, + context["heads"], + context["mesh"], + context["axis_names_q"], + context["axis_names_kv"], + context["flash_block_sizes"], + context["dtype"], + mask_padding_tokens=context["mask_padding_tokens"], + residual_checkpoint_name=context["residual_checkpoint_name"], + attention_mask=context["attention_mask"], + ulysses_shards=context["ulysses_shards"], + use_base2_exp=context.get("use_base2_exp", True), + use_experimental_scheduler=context.get("use_experimental_scheduler", False), + use_fixed_m=True, + per_q_block=True, + ulysses_attention_chunks=context.get("ulysses_attention_chunks", 1), + kv_heads=context.get("kv_heads", None), + use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=True), ) @@ -1923,6 +2423,7 @@ def ulysses_ring_custom_bidir_kernel(q, k, v, context): use_experimental_scheduler=context.get("use_experimental_scheduler", False), bidirectional=True, ulysses_attention_chunks=context["ulysses_attention_chunks"], + kv_heads=context.get("kv_heads", None), ) @@ -1946,7 +2447,11 @@ def ulysses_custom_fixed_m_kernel(q, k, v, context): use_experimental_scheduler=context.get("use_experimental_scheduler", False), use_fixed_m=True, per_q_block=False, - ulysses_attention_chunks=context.get("ulysses_attention_chunks", 1), + ulysses_attention_chunks=context["ulysses_attention_chunks"], + kv_heads=context.get("kv_heads", None), + ulysses_shards=context.get("ulysses_shards", -1), + kernel_name="ulysses_custom_fixed_m", + use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=False), ) @@ -1970,7 +2475,11 @@ def ulysses_custom_fixed_m_per_q_block_kernel(q, k, v, context): use_experimental_scheduler=context.get("use_experimental_scheduler", False), use_fixed_m=True, per_q_block=True, - ulysses_attention_chunks=context.get("ulysses_attention_chunks", 1), + ulysses_attention_chunks=context["ulysses_attention_chunks"], + kv_heads=context.get("kv_heads", None), + ulysses_shards=context.get("ulysses_shards", -1), + kernel_name="ulysses_custom_fixed_m_per_q_block", + use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=False), ) @@ -1991,6 +2500,9 @@ def ulysses_kernel(q, k, v, context): attention_mask=context["attention_mask"], ulysses_attention_chunks=context["ulysses_attention_chunks"], preserve_asymmetric_block_sizes=context.get("preserve_asymmetric_block_sizes", False), + kv_heads=context.get("kv_heads", None), + ulysses_shards=context.get("ulysses_shards", -1), + kernel_name="ulysses", ) @@ -2014,6 +2526,7 @@ def ulysses_ring_kernel(q, k, v, context): ulysses_shards=context["ulysses_shards"], ulysses_attention_chunks=context["ulysses_attention_chunks"], preserve_asymmetric_block_sizes=context.get("preserve_asymmetric_block_sizes", False), + kv_heads=context.get("kv_heads", None), ) @@ -2037,6 +2550,7 @@ def flash_kernel(q, k, v, context): use_experimental_scheduler=context["use_experimental_scheduler"], is_causal=context.get("is_causal", False), preserve_asymmetric_block_sizes=context.get("preserve_asymmetric_block_sizes", False), + kv_heads=context.get("kv_heads", None), ) @@ -2062,6 +2576,7 @@ def tokamax_flash_kernel(q, k, v, context): preserve_asymmetric_block_sizes=context.get("preserve_asymmetric_block_sizes", False), spatiotemporal_config=context.get("spatiotemporal_config"), spatiotemporal_shape=context.get("spatiotemporal_shape"), + kv_heads=context.get("kv_heads", None), ) @@ -2085,6 +2600,7 @@ def tokamax_ring_kernel(q, k, v, context): use_experimental_scheduler=context["use_experimental_scheduler"], is_causal=context.get("is_causal", False), preserve_asymmetric_block_sizes=context.get("preserve_asymmetric_block_sizes", False), + kv_heads=context.get("kv_heads", None), ) @@ -2106,6 +2622,7 @@ def tokamax_ring_custom_kernel(q, k, v, context): use_base2_exp=context.get("use_base2_exp", True), use_experimental_scheduler=context.get("use_experimental_scheduler", False), preserve_asymmetric_block_sizes=context.get("preserve_asymmetric_block_sizes", False), + kv_heads=context.get("kv_heads", None), ) @@ -2143,6 +2660,8 @@ def _apply_attention( preserve_asymmetric_block_sizes: bool = False, spatiotemporal_config: Optional[dict] = None, spatiotemporal_shape: Optional[Tuple[int, int, int]] = None, + kv_heads: Optional[int] = None, + use_k_centering: bool | str = "auto", ): """Routes to different attention kernels using a module-level registry.""" @@ -2160,6 +2679,10 @@ def _apply_attention( "ulysses_custom_fixed_m", "ulysses_custom_fixed_m_per_q_block", "ulysses_ring", + "ulysses_ring_custom", + "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", + "ulysses_ring_custom_bidir", ]: can_use_flash_attention = ( query.shape[seq_len_idx] >= flash_min_seq_length @@ -2191,6 +2714,7 @@ def _apply_attention( context = { "heads": heads, + "kv_heads": kv_heads, "mesh": mesh, "axis_names_q": axis_names_q, "axis_names_kv": axis_names_kv, @@ -2213,6 +2737,7 @@ def _apply_attention( "preserve_asymmetric_block_sizes": preserve_asymmetric_block_sizes, "spatiotemporal_config": spatiotemporal_config, "spatiotemporal_shape": spatiotemporal_shape, + "use_k_centering": use_k_centering, } if spatiotemporal_config and spatiotemporal_config.get("use_svg_attention"): @@ -2222,6 +2747,7 @@ def _apply_attention( "ulysses_custom_fixed_m_per_q_block", "ulysses_ring_custom", "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", ): raise ValueError("Head-local SVG requires a custom Ulysses attention backend.") # Dense uses its configured ring split; SVG exchanges over the full context axis. @@ -2564,12 +3090,15 @@ def __init__( use_experimental_scheduler: bool = False, ulysses_shards: int = -1, ulysses_attention_chunks: int = 1, + kv_heads: Optional[int] = None, + use_k_centering: bool | str = "auto", ): self.dpa_layer = None self.use_base2_exp = use_base2_exp self.use_experimental_scheduler = use_experimental_scheduler self.ulysses_shards = ulysses_shards self.ulysses_attention_chunks = ulysses_attention_chunks + self.use_k_centering = use_k_centering if attention_kernel == "cudnn_flash_te": from transformer_engine.jax.flax.transformer import DotProductAttention # pytype: disable=import-error @@ -2594,6 +3123,7 @@ def __init__( self.mesh = mesh self.scale = scale self.heads = heads + self.kv_heads = kv_heads self.dim_head = dim_head self.attention_kernel = attention_kernel self.use_memory_efficient_attention = use_memory_efficient_attention @@ -2646,6 +3176,8 @@ def apply_attention( preserve_asymmetric_block_sizes=preserve_asymmetric_block_sizes, spatiotemporal_config=sparse_config_override, spatiotemporal_shape=spatiotemporal_shape, + kv_heads=self.kv_heads, + use_k_centering=getattr(self, "use_k_centering", "auto"), ) @@ -2669,6 +3201,8 @@ class AttentionOp(nn.Module): ulysses_shards: int = -1 ulysses_attention_chunks: int = 1 is_causal: bool = False + kv_heads: Optional[int] = None + use_k_centering: bool | str = "auto" def setup(self): self.dpa_layer = None @@ -2731,6 +3265,8 @@ def apply_attention( preserve_asymmetric_block_sizes=preserve_asymmetric_block_sizes, spatiotemporal_config=sparse_config_override, spatiotemporal_shape=spatiotemporal_shape, + kv_heads=self.kv_heads, + use_k_centering=getattr(self, "use_k_centering", "auto"), ) @@ -2794,6 +3330,7 @@ def __init__( "svg_high_noise_density": -1.0, "svg_low_noise_density": -1.0, "svg_flash_block_sizes": None, + "use_k_centering": "auto", **(attention_config or {}), } @@ -2831,6 +3368,8 @@ def __init__( self.value_axis_names = value_axis_names self.out_axis_names = out_axis_names self.enable_jax_named_scopes = enable_jax_named_scopes + self.is_self_attention = is_self_attention + self.eps = eps cross_attention_remapped_to_flash = not is_self_attention and attention_kernel in ( "tokamax_ring", @@ -2838,6 +3377,7 @@ def __init__( "ulysses_ring", "ulysses_ring_custom", "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", "ulysses_ring_custom_bidir", "ulysses_custom", "ulysses_custom_fixed_m", @@ -2859,17 +3399,11 @@ def __init__( ) if cross_attention_remapped_to_flash: attention_kernel = "tokamax_flash" - elif attention_kernel in ("tokamax_ring", "tokamax_ring_custom", "ulysses_ring") and not is_self_attention: - attention_kernel = "tokamax_flash" # do not use ring attention for cross attention - elif ( - attention_kernel in ("ulysses_ring_custom", "ulysses_ring_custom_bidir", "ulysses_ring_custom_fixed_m") - and not is_self_attention - ): - attention_kernel = "ulysses_custom" # plain ulysses (no ring) for cross attention self.added_kv_proj_dim = added_kv_proj_dim # New for I2V self.image_seq_len = image_seq_len # New for I2V tpu_type = get_tpu_type() self.alignment = 256 if tpu_type in [TpuType.TPU_V6_LITE, TpuType.TPU_7X] else 128 + self.precision = precision self.attention_op = NNXAttentionOp( mesh=mesh, @@ -2892,6 +3426,7 @@ def __init__( use_experimental_scheduler=attention_config["use_experimental_scheduler"], ulysses_shards=attention_config["ulysses_shards"], ulysses_attention_chunks=attention_config["ulysses_attention_chunks"], + use_k_centering=attention_config["use_k_centering"], ) # None axes corresponds to the stacked weights across all blocks # because of the use of nnx.vmap and nnx.scan. @@ -3044,9 +3579,9 @@ def _apply_rope(self, xq: jax.Array, xk: jax.Array, freqs_cis: jax.Array) -> Tup xk_out_0 = xk_0 * cos - xk_1 * sin xk_out_1 = xk_0 * sin + xk_1 * cos - # 5. Stack and reshape back to original - xq_out = jnp.stack([xq_out_0, xq_out_1], axis=-1).reshape(xq.shape) - xk_out = jnp.stack([xk_out_0, xk_out_1], axis=-1).reshape(xk.shape) + # 5. Interleave the rotated pairs back into the last axis + xq_out = jnp.concatenate([xq_out_0[..., None], xq_out_1[..., None]], axis=-1).reshape(xq.shape) + xk_out = jnp.concatenate([xk_out_0[..., None], xk_out_1[..., None]], axis=-1).reshape(xk.shape) return xq_out, xk_out @@ -3068,15 +3603,15 @@ def __call__( svg_timestep: Optional[int | float | jax.Array] = None, svg_step_index: Optional[int | jax.Array] = None, ) -> jax.Array: + same_kv_source = encoder_hidden_states is None or encoder_hidden_states is hidden_states hidden_states = nn.with_logical_constraint(hidden_states, (BATCH, LENGTH, HEAD)) if encoder_hidden_states is not None: encoder_hidden_states = nn.with_logical_constraint(encoder_hidden_states, (BATCH, LENGTH, HEAD)) dtype = hidden_states.dtype - is_self_attention = getattr( - self, - "is_self_attention", - encoder_hidden_states is None or encoder_hidden_states is hidden_states, - ) + if not same_kv_source: + is_self_attention = False + else: + is_self_attention = getattr(self, "is_self_attention", True) if self.use_svg_attention and is_self_attention: if not deterministic: raise ValueError("SVG attention supports deterministic inference only.") @@ -3096,11 +3631,12 @@ def __call__( with jax.named_scope("query_proj"): query_proj = self.query(hidden_states) - if self.qk_norm: - with self.conditional_named_scope("attn_q_norm"): - query_proj = self.norm_q(query_proj) - - if not is_self_attention and cached_kv is not None and "text" in cached_kv: + if is_self_attention: + with jax.named_scope("key_proj"): + key_proj = self.key(hidden_states) + with jax.named_scope("value_proj"): + value_proj = self.value(hidden_states) + elif cached_kv is not None and "text" in cached_kv: key_proj, value_proj = cached_kv["text"] else: with jax.named_scope("key_proj"): @@ -3108,17 +3644,39 @@ def __call__( with jax.named_scope("value_proj"): value_proj = self.value(encoder_hidden_states) - if self.qk_norm: - with self.conditional_named_scope("attn_k_norm"): - key_proj = self.norm_k(key_proj) - - if rotary_emb is not None: - with self.conditional_named_scope("attn_rope"): - query_proj = _unflatten_heads(query_proj, self.heads) - key_proj = _unflatten_heads(key_proj, self.heads) + if rotary_emb is not None and self.qk_norm and is_self_attention: + with self.conditional_named_scope("fused_rmsnorm_rope"): + q_scale = self.norm_q.scale[...] + k_scale = self.norm_k.scale[...] + q_eps = getattr(self.norm_q, "epsilon", self.eps) + k_eps = getattr(self.norm_k, "epsilon", self.eps) + query_proj, key_proj = fused_rmsnorm_rope( + query_proj, + key_proj, + q_scale, + k_scale, + rotary_emb, + q_heads=self.heads, + dim_head=self.dim_head, + eps=q_eps, + k_eps=k_eps, + ) value_proj = _unflatten_heads(value_proj, self.heads) - # output of _unflatten_heads Batch, heads, seq_len, head_dim - query_proj, key_proj = self._apply_rope(query_proj, key_proj, rotary_emb) + else: + if self.qk_norm: + with self.conditional_named_scope("attn_q_norm"): + query_proj = self.norm_q(query_proj) + if is_self_attention or cached_kv is None or "text" not in cached_kv: + with self.conditional_named_scope("attn_k_norm"): + key_proj = self.norm_k(key_proj) + + if rotary_emb is not None: + with self.conditional_named_scope("attn_rope"): + query_proj = _unflatten_heads(query_proj, self.heads) + key_proj = _unflatten_heads(key_proj, self.heads) + value_proj = _unflatten_heads(value_proj, self.heads) + # output of _unflatten_heads Batch, heads, seq_len, head_dim + query_proj, key_proj = self._apply_rope(query_proj, key_proj, rotary_emb) query_proj = checkpoint_name(query_proj, "query_proj") key_proj = checkpoint_name(key_proj, "key_proj") diff --git a/src/maxdiffusion/models/ltx2/attention_ltx2.py b/src/maxdiffusion/models/ltx2/attention_ltx2.py index dc8d2c6bc..93091aba6 100644 --- a/src/maxdiffusion/models/ltx2/attention_ltx2.py +++ b/src/maxdiffusion/models/ltx2/attention_ltx2.py @@ -465,6 +465,7 @@ def __init__( "ulysses_ring", "ulysses_ring_custom", "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", "ulysses_ring_custom_bidir", "ulysses_custom", "ulysses_custom_fixed_m", diff --git a/src/maxdiffusion/pyconfig.py b/src/maxdiffusion/pyconfig.py index d3f4ed5dc..0f55436ae 100644 --- a/src/maxdiffusion/pyconfig.py +++ b/src/maxdiffusion/pyconfig.py @@ -234,6 +234,7 @@ def user_init(raw_keys): "ulysses_ring_custom", "ulysses_ring_custom_fixed_m", "ulysses_ring_custom_bidir", + "ulysses_ring_custom_fixed_m_per_q_block", } if attention in ulysses_ring_attentions and raw_keys.get("ulysses_shards", -1) <= 0: raise ValueError(f"{attention} requires ulysses_shards to be set from config or command line.") @@ -318,6 +319,9 @@ def user_init(raw_keys): if "vae_spatial" not in raw_keys: raw_keys["vae_spatial"] = -1 + if "use_k_centering" not in raw_keys: + raw_keys["use_k_centering"] = "auto" + def get_num_slices(raw_keys): if int(raw_keys["compile_topology_num_slices"]) > 0: diff --git a/src/maxdiffusion/tests/attention_config_guards_test.py b/src/maxdiffusion/tests/attention_config_guards_test.py new file mode 100644 index 000000000..b78aeb084 --- /dev/null +++ b/src/maxdiffusion/tests/attention_config_guards_test.py @@ -0,0 +1,260 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""CPU checks for the attention-config guards across TPU topologies. + +Guard 1: a non-ring ulysses kernel must REJECT a ulysses_shards it cannot honour. +Guard 2: a ring kernel must WARN when it degenerates to R=1, and the remedy it + suggests must be valid for the actual mesh and head count. +""" +import unittest +from unittest import mock + +from maxdiffusion.models import attention_flax + +# Context-parallel degrees reachable on real slices: v5e/v6e-8 (CP up to 8), +# v6e-16 / v7x-16 (CP up to 16), and larger multi-host slices. +TOPOLOGIES = [1, 2, 4, 8, 16, 32, 64] + +# Wan 2.2 T2V-A14B has 40 attention heads, which is NOT a power of two -- this +# is precisely why a naive "context_shards // 2" suggestion is unsafe. +WAN_HEADS = 40 + + +class ImplicitUlyssesDegreeGuardTest(unittest.TestCase): + + def test_unset_is_allowed_on_every_topology(self): + for cp in TOPOLOGIES: + with self.subTest(cp=cp): + attention_flax._validate_implicit_ulysses_degree(-1, cp, "ulysses_custom") + attention_flax._validate_implicit_ulysses_degree(0, cp, "ulysses_custom") + attention_flax._validate_implicit_ulysses_degree(None, cp, "ulysses_custom") + + def test_matching_request_is_allowed_on_every_topology(self): + for cp in TOPOLOGIES: + with self.subTest(cp=cp): + attention_flax._validate_implicit_ulysses_degree(cp, cp, "ulysses_custom") + + def test_mismatched_request_is_rejected_on_every_topology(self): + for cp in TOPOLOGIES: + for requested in {1, 2, cp // 2, cp * 2} - {cp, 0}: + if requested <= 0: + continue + with self.subTest(cp=cp, requested=requested): + with self.assertRaises(ValueError) as ctx: + attention_flax._validate_implicit_ulysses_degree(requested, cp, "ulysses_custom_fixed_m_per_q_block") + msg = str(ctx.exception) + self.assertIn(f"ulysses_shards={requested}", msg) + self.assertIn(f"context_shards={cp}", msg) + + def test_rejection_names_the_offending_kernel(self): + with self.assertRaises(ValueError) as ctx: + attention_flax._validate_implicit_ulysses_degree(2, 4, "ulysses_custom_fixed_m_per_q_block") + self.assertIn("ulysses_custom_fixed_m_per_q_block", str(ctx.exception)) + + +class RealRingSuggestionTest(unittest.TestCase): + + def test_suggestion_is_valid_for_every_topology(self): + """Whatever U we suggest must satisfy every constraint the ring enforces.""" + for cp in TOPOLOGIES: + with self.subTest(cp=cp, heads=WAN_HEADS): + u = attention_flax._largest_ulysses_shards_for_real_ring(cp, WAN_HEADS, WAN_HEADS) + if u is None: + # Only legitimate when no divisor below cp works. + self.assertTrue( + all(cp % c != 0 or WAN_HEADS % c != 0 for c in range(1, cp)), + f"returned None for cp={cp} despite a valid candidate existing", + ) + continue + self.assertLess(u, cp) + self.assertEqual(cp % u, 0, "suggested U must divide the context shard count") + self.assertEqual(WAN_HEADS % u, 0, "suggested U must divide the head count") + self.assertGreater(cp // u, 1, "suggested U must leave a real ring R>1") + + def test_no_suggestion_for_single_shard(self): + self.assertIsNone(attention_flax._largest_ulysses_shards_for_real_ring(1, WAN_HEADS, WAN_HEADS)) + + def test_prefers_smallest_real_ring(self): + # cp=8, heads=40 -> U=4 gives R=2, the cheapest real ring. + self.assertEqual(attention_flax._largest_ulysses_shards_for_real_ring(8, 40, 40), 4) + # cp=16, heads=40 -> U=16 and U=8 both divide 16, but only 8 divides 40. + self.assertEqual(attention_flax._largest_ulysses_shards_for_real_ring(16, 40, 40), 8) + # cp=32, heads=40 -> largest common divisor below 32 is 8, giving R=4. + self.assertEqual(attention_flax._largest_ulysses_shards_for_real_ring(32, 40, 40), 8) + + def test_respects_asymmetric_kv_heads(self): + # GQA-style: 40 query heads but 8 KV heads restricts U to divisors of 8. + u = attention_flax._largest_ulysses_shards_for_real_ring(16, 40, 8) + self.assertEqual(8 % u, 0) + self.assertEqual(40 % u, 0) + self.assertEqual(16 % u, 0) + + +class DegenerateRingWarningTest(unittest.TestCase): + + def setUp(self): + attention_flax._WARNED_ONCE.clear() + + def test_warns_on_every_topology_when_degenerate(self): + for cp in TOPOLOGIES: + with self.subTest(cp=cp): + attention_flax._WARNED_ONCE.clear() + with mock.patch.object(attention_flax.max_logging, "log") as log: + attention_flax._warn_if_ring_is_degenerate(1, cp, cp, heads=WAN_HEADS, kv_heads=WAN_HEADS) + log.assert_called_once() + msg = log.call_args[0][0] + self.assertIn("R=1", msg) + self.assertIn("Do NOT report this as a ring-attention result", msg) + + def test_suggested_remedy_in_message_is_actionable(self): + for cp in [2, 4, 8, 16, 32]: + with self.subTest(cp=cp): + attention_flax._WARNED_ONCE.clear() + with mock.patch.object(attention_flax.max_logging, "log") as log: + attention_flax._warn_if_ring_is_degenerate(1, cp, cp, heads=WAN_HEADS, kv_heads=WAN_HEADS) + msg = log.call_args[0][0] + expected_u = attention_flax._largest_ulysses_shards_for_real_ring(cp, WAN_HEADS, WAN_HEADS) + self.assertIn(f"ulysses_shards={expected_u}", msg) + self.assertIn(f"R={cp // expected_u}", msg) + + def test_single_shard_mesh_does_not_suggest_zero(self): + """Regression: context_shards//2 would have advised the impossible U=0.""" + with mock.patch.object(attention_flax.max_logging, "log") as log: + attention_flax._warn_if_ring_is_degenerate(1, 1, 1, heads=WAN_HEADS, kv_heads=WAN_HEADS) + msg = log.call_args[0][0] + self.assertNotIn("ulysses_shards=0", msg) + self.assertIn("only one context shard", msg) + + def test_silent_for_real_ring_on_every_topology(self): + for cp in [2, 4, 8, 16, 32]: + u = attention_flax._largest_ulysses_shards_for_real_ring(cp, WAN_HEADS, WAN_HEADS) + with self.subTest(cp=cp, u=u): + attention_flax._WARNED_ONCE.clear() + with mock.patch.object(attention_flax.max_logging, "log") as log: + attention_flax._warn_if_ring_is_degenerate(cp // u, u, cp, heads=WAN_HEADS, kv_heads=WAN_HEADS) + log.assert_not_called() + + def test_warns_only_once_per_configuration(self): + with mock.patch.object(attention_flax.max_logging, "log") as log: + for _ in range(80): # ~40 layers x 2 transformers + attention_flax._warn_if_ring_is_degenerate(1, 4, 4, heads=WAN_HEADS, kv_heads=WAN_HEADS) + log.assert_called_once() + + def test_distinct_configurations_each_warn(self): + with mock.patch.object(attention_flax.max_logging, "log") as log: + attention_flax._warn_if_ring_is_degenerate(1, 4, 4, heads=WAN_HEADS, kv_heads=WAN_HEADS) + attention_flax._warn_if_ring_is_degenerate(1, 8, 8, heads=WAN_HEADS, kv_heads=WAN_HEADS) + self.assertEqual(log.call_count, 2) + + +class KernelClassificationDriftTest(unittest.TestCase): + """Every registered ulysses-ring kernel must be classified for tile sizing. + + `local_tiled_seq_len` silently falls through to `return full_seq` for any + attention name it does not recognise, which yields a mesh-independent (and + therefore wrong) tile length. A newly registered ring kernel must not be able + to slip through that default. + """ + + def test_all_registered_ulysses_ring_kernels_are_classified(self): + from maxdiffusion.utils import tile_size_grid_search as tsgs + from maxdiffusion.utils import wan_block_benchmark as wbb + + registered = {name for name in attention_flax.KERNEL_REGISTRY if name.startswith("ulysses_ring")} + self.assertTrue(registered, "expected at least one registered ulysses_ring kernel") + missing = registered - tsgs.ULYSSES_RING_ATTENTION_KERNELS + self.assertEqual( + missing, + set(), + f"these ulysses_ring kernels are missing from ULYSSES_RING_ATTENTION_KERNELS and would " + f"get a mesh-independent tile length: {sorted(missing)}", + ) + missing_wbb = registered - wbb._RING_VARIANTS + self.assertEqual( + missing_wbb, + set(), + f"these ulysses_ring kernels are missing from wan_block_benchmark._RING_VARIANTS: {sorted(missing_wbb)}", + ) + + def test_classified_kernels_scale_with_topology(self): + from maxdiffusion.utils.tile_size_grid_search import local_tiled_seq_len + + full_seq = 6144 + for attention in sorted(attention_flax.KERNEL_REGISTRY): + if not attention.startswith("ulysses_ring"): + continue + for cp in [2, 4, 8, 16]: + u = attention_flax._largest_ulysses_shards_for_real_ring(cp, WAN_HEADS, WAN_HEADS) + with self.subTest(attention=attention, cp=cp, u=u): + local = local_tiled_seq_len(full_seq, attention, context_shards=cp, ulysses_shards=u) + # Ulysses gathers u chunks of the context-local sequence. + self.assertEqual(local, (full_seq // cp) * u) + self.assertLess( + local, + full_seq, + "a sharded mesh must tile less than the full sequence", + ) + + +class KCenteringResolutionTest(unittest.TestCase): + """`use_k_centering="auto"` must pick the cheap choice per attention path.""" + + def test_auto_is_on_for_non_ring_and_off_for_ring(self): + for value in ("auto", "AUTO", None): + self.assertTrue(attention_flax.resolve_k_centering(value, ring=False)) + self.assertFalse(attention_flax.resolve_k_centering(value, ring=True)) + + def test_explicit_values_are_honoured_on_both_paths(self): + for ring in (False, True): + self.assertTrue(attention_flax.resolve_k_centering(True, ring=ring)) + self.assertFalse(attention_flax.resolve_k_centering(False, ring=ring)) + self.assertTrue(attention_flax.resolve_k_centering("true", ring=ring)) + self.assertFalse(attention_flax.resolve_k_centering("False", ring=ring)) + + def test_rejects_unknown_strings(self): + with self.assertRaises(ValueError): + attention_flax.resolve_k_centering("sometimes", ring=False) + + +class GQAUnsupportedKernelsGuardTest(unittest.TestCase): + """Kernels that do not support GQA must fail loudly when kv_heads != heads.""" + + def test_flash_and_non_custom_ulysses_reject_gqa(self): + base_kwargs = { + "query": None, + "key": None, + "value": None, + "heads": 8, + "mesh": None, + "axis_names_q": None, + "axis_names_kv": None, + "flash_block_sizes": None, + "dtype": None, + "kv_heads": 2, + } + for fn, extra in [ + (attention_flax._tpu_flash_attention, {"attention_kernel": "flash"}), + (attention_flax._ulysses_attention, {"use_custom_kernel": False}), + (attention_flax._ulysses_ring_attention, {}), + ]: + with self.subTest(fn=fn.__name__): + with self.assertRaises(NotImplementedError): + fn(**base_kwargs, **extra) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/tests/custom_splash_fixed_m_test.py b/src/maxdiffusion/tests/custom_splash_fixed_m_test.py index 2a3ad8709..129e0ad7a 100644 --- a/src/maxdiffusion/tests/custom_splash_fixed_m_test.py +++ b/src/maxdiffusion/tests/custom_splash_fixed_m_test.py @@ -782,6 +782,46 @@ def test_adversarial_centered_keys_softmax_mass_loss(self): f"Fallback online softmax output must be ~15/17 (~0.8824), got {float(out[0, 0, 0]):.6f}", ) + def test_gqa_rejects_q_head_padded_k_mean(self): + """GQA (8 Q heads, 2 KV heads): a k_mean padded to 8 rows is rejected by the splash kernel. + + Only the ValueError is asserted; the correctly shaped (2, dim) k_mean is used + to build the metadata but is not run through the kernel here. + """ + from maxdiffusion.models.attention_flax import _compute_fixed_m_metadata + + batch, num_q_heads, num_kv_heads, seq_len, dim = 1, 8, 2, 1024, 128 + bq = 512 + q = jnp.ones((batch, num_q_heads, seq_len, dim), dtype=jnp.bfloat16) + # Give the two KV heads different means, so a padded array read as + # Q-head-indexed would pair Q heads with the wrong (zero) mean rows. + k = jnp.zeros((batch, num_kv_heads, seq_len, dim), dtype=jnp.bfloat16) + k = k.at[:, 0, :, 0].set(1.0) + k = k.at[:, 1, :, 0].set(50.0) + v = jnp.ones((batch, num_kv_heads, seq_len, dim), dtype=jnp.bfloat16) + + k_mean = jnp.mean(k.astype(jnp.float32), axis=2) # (batch, 2, dim) + self.assertEqual(k_mean.shape, (batch, num_kv_heads, dim)) + mk_arr, _ = _compute_fixed_m_metadata(q, k, block_q=bq, k_mean=k_mean, value=v) + block_sizes = custom_splash._BlockSizes(block_q=bq, block_kv=512, block_kv_compute=256, block_kv_compute_in=256) + + # Pre-padding k_mean from 2 rows to 8 (NUM_SUBLANES) matches num_q_heads=8 + # and must be rejected rather than misinterpreted as Q-head-expanded metadata. + padded_k_mean = jnp.pad(k_mean[0], ((0, 6), (0, 0))) + with self.assertRaisesRegex(ValueError, "indexed by KV head"): + custom_splash._splash_attention_forward( + q[0], + k[0], + v[0], + block_sizes=block_sizes, + q_seq_len=seq_len, + kv_seq_len=seq_len, + use_base2_exp=True, + use_fixed_m=True, + mk=mk_arr[0], + k_mean=padded_k_mean, + ) + @_SKIP_IN_GITHUB_ACTIONS class FixedMAttentionFlaxIntegrationTest(unittest.TestCase): @@ -1110,7 +1150,7 @@ def _verdict_seen_by_ring(self, v): def fake_make_ring(**kwargs): v_ok = kwargs["v_ok"] self.assertIsNotNone(v_ok) - return lambda q, k, val: jnp.broadcast_to(jnp.asarray(v_ok, jnp.float32), q.shape).astype(q.dtype) + return lambda q, k, val, *_: jnp.broadcast_to(jnp.asarray(v_ok, jnp.float32), q.shape).astype(q.dtype) shape = (1, self.seq_len, self.num_heads * self.head_dim) q = jnp.full(shape, 0.01, jnp.bfloat16) diff --git a/src/maxdiffusion/tests/custom_splash_unpadded_test.py b/src/maxdiffusion/tests/custom_splash_unpadded_test.py new file mode 100644 index 000000000..87b15689d --- /dev/null +++ b/src/maxdiffusion/tests/custom_splash_unpadded_test.py @@ -0,0 +1,419 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""C3 gate: may the custom splash kernel be handed PHYSICALLY UNPADDED q/k/v? + +Context +------- +Production pads Q to `block_q` with `_pad_data_for_flash`. K/V are passed +unpadded when the KV length is 8-aligned (C3a, relied on by `kv_pad_size = 1` +in `_ulysses_attention` and `_ulysses_ring_custom_attention`); otherwise they +are padded to `block_kv`. These tests gate that K/V behaviour and the +still-proposed Q unpadding (C3b). Static analysis says the padding is +unnecessary: + + * K/V: the kernel never *masks* the KV tail, it *slices* it + (`slice_k_len = kv_seq_len % bkv_compute`, using the UNPADDED length -- see + `last_compute_body_fixed` / `last_compute_body_online` in + custom_splash_attention.py and the load-bearing comment above them). No + compute path reads a padded K or V row, so their contents are irrelevant. + + * Q: rows are independent (the running max is over KV, never across the `bq` + row axis), and the kernel's OUTPUT is already ragged today -- `out_shape`'s + last dim is `actual_q_seq_len` against a `bq`-wide BlockSpec -- so Pallas is + already clipping a non-divisible last block in this very kernel. + +The one thing static analysis CANNOT settle is whether Pallas clips the ragged +last-block *input* DMA, or issues an unclipped read past the end of the array. +`compiler_params` sets `disable_bounds_checks=True`. If the DMA is unclipped, +passing an unpadded array is a silent out-of-bounds HBM read. + +**That is the only question these tests exist to answer.** Everything else about +C3 is already proven on paper. + +Why the existing coverage does not answer it +-------------------------------------------- +`custom_splash_fixed_m_test.test_non_divisible_sequence_context_padding_fixed_m` +looks like it covers this, but it pads its inputs to a block boundary before the +call (`q_in_padded = jnp.pad(...)`) and only passes a ragged *logical* length. +That is exactly today's regime. No existing test ever hands the kernel a +physically non-block-aligned array. + +Test design +----------- +Each case runs the kernel twice with IDENTICAL LOGICAL INPUTS: + reference -- physically padded to a block boundary (today's behaviour) + candidate -- physically unpadded (what C3 proposes) +and asserts the outputs are **bit-identical**. This is pure data movement: the +grid, the block sizes and every slice length are computed from the unpadded +logical length and are therefore identical between the two runs. The same +arithmetic happens in the same order on the same values, so anything other than +exact equality means memory outside the array was read. + +K/V-unpadded and Q-unpadded are separate cases so the two halves of C3 can be +gated independently. + +Two failure modes are guarded against explicitly: + 1. Passing by luck on a freshly-zeroed buffer -- see `_dirty_device_memory`, + and the repeat-run case. + 2. A harness bug making both sides equally wrong -- every case also checks the + padded reference against a dense f32 softmax reference. + +> A NOTE ON THE LIMITS OF THIS TEST. `_dirty_device_memory` is a heuristic. JAX +> gives no control over HBM placement, so we cannot *guarantee* the bytes past +> the end of the array are NaN. A PASS is therefore strong but not absolute +> evidence; a FAILURE is conclusive. Treat a pass as "no evidence of an +> unclipped DMA under adversarial conditions", not as a proof of clipping. +""" + +import gc +import math +import os +import unittest + +import jax +import jax.numpy as jnp + +from maxdiffusion.kernels import custom_splash_attention as custom_splash + +IN_GITHUB_ACTIONS = os.getenv("GITHUB_ACTIONS") == "true" +# Kernel numerics grids and multi-device tests are skipped in CI; see +# end_to_end/tpu/run_wan_stack_tests.sh. +_SKIP_IN_GITHUB_ACTIONS = unittest.skipIf( + IN_GITHUB_ACTIONS, "TPU kernel / multi-device test, skipped in GitHub Actions; run end_to_end/tpu/run_wan_stack_tests.sh" +) + +_LOG2E = math.log2(math.e) + + +def _dirty_device_memory(num_buffers: int = 8, mb_each: int = 64) -> None: + """Fills and releases device buffers of NaN to poison the allocator pool. + + The failure mode this defends against: an unclipped DMA reads whatever + happens to sit past the end of the array. On a fresh device that is often + zero, which is exactly the padding value the kernel would have seen anyway -- + so the bug would produce a correct answer and the test would pass for the + wrong reason. + + By allocating NaN buffers and then dropping them, subsequent allocations are + likely (not guaranteed) to be served from memory containing NaN. NaN is a + stronger probe than a large finite value: a large finite value could be + squashed back to something finite by a downstream mask or a saturating + operation, whereas NaN propagates through every arithmetic path in the + softmax and cannot be masked away once it enters an accumulation. + """ + elems = mb_each * 1024 * 1024 // 4 + junk = [] + for _ in range(num_buffers): + junk.append(jax.block_until_ready(jnp.full((elems,), jnp.nan, dtype=jnp.float32))) + del junk + gc.collect() + + +@_SKIP_IN_GITHUB_ACTIONS +class CustomSplashUnpaddedInputTest(unittest.TestCase): + """Bit-identity of the kernel when fed physically unpadded q/k/v.""" + + heads = 4 + head_dim = 64 + + # Deliberately non-divisible by `bq` in BOTH dimensions. + # q_len = 1008, bq = 512 -> grid_height = 2, last block covers [512, 1024) + # against an array of only 1008 rows. + # kv_len = 1001, bkv = 512 -> grid_width = 2, tail slice length 489. + q_len = 1008 + kv_len = 1001 + bq = 512 + + def setUp(self): + super().setUp() + if jax.default_backend() == "cpu": + self.skipTest("Pallas splash kernel requires TPU.") + self.scale = 1.0 / math.sqrt(self.head_dim) + self.grid_height = math.ceil(self.q_len / self.bq) + self.q_padded_len = self.grid_height * self.bq + self.kv_padded_len = math.ceil(self.kv_len / self.bq) * self.bq + + # ---------------------------------------------------------------- helpers + + def _inputs(self): + """Returns logical (unpadded) bf16 q, k, v in the kernel's own convention.""" + q = jax.random.normal( + jax.random.PRNGKey(11), + (self.heads, self.q_len, self.head_dim), + jnp.bfloat16, + ) + k = jax.random.normal( + jax.random.PRNGKey(12), + (self.heads, self.kv_len, self.head_dim), + jnp.bfloat16, + ) + v = jax.random.normal( + jax.random.PRNGKey(13), + (self.heads, self.kv_len, self.head_dim), + jnp.bfloat16, + ) + q_in = (q * _LOG2E).astype(jnp.bfloat16) + k_in = (k * self.scale).astype(jnp.bfloat16) + return q_in, k_in, v, q, k + + def _pad_seq(self, x, target): + if x.shape[1] == target: + return x + return jnp.pad(x, ((0, 0), (0, target - x.shape[1]), (0, 0))) + + def _metadata(self, q_in, k_in): + """Builds `mk` and `k_mean` from UNPADDED inputs. + + This is the C3b norm-vector trick: the per-row norms are computed on the + unpadded query and then the *norm vector* is zero-padded to the block grid, + instead of zero-padding the (far larger) activation. Zero rows contribute 0 + to a max over non-negative values, so the resulting `mk` is identical to the + one derived from a zero-padded activation -- bit-identical, not merely + close. + """ + k_mean = jnp.mean(k_in.astype(jnp.float32), axis=1) # (heads, dim) + recenter, safe_bound = custom_splash.get_fixed_m_constants(self.kv_len) + + k_centered = k_in.astype(jnp.float32) - k_mean[:, None, :] + mk_h = jnp.sqrt((k_centered**2).sum(-1)).max(axis=-1) # (heads,) + + row_norm_sq = (q_in.astype(jnp.float32) ** 2).sum(-1) # (heads, q_len) + pad = self.q_padded_len - row_norm_sq.shape[1] + if pad: + row_norm_sq = jnp.pad(row_norm_sq, ((0, 0), (0, pad))) + qn_max = jnp.sqrt(row_norm_sq.reshape(self.heads, self.grid_height, self.bq).max(axis=-1)) + + bound = qn_max * mk_h[:, None] + mk = jnp.stack( + [jnp.ceil(bound) - recenter, (bound <= safe_bound).astype(jnp.float32)], + axis=0, + ) + return mk, k_mean + + def _run( + self, + q_in, + k_in, + v_in, + mk, + k_mean, + *, + use_fixed_m, + uniform_fixed_m, + bkv_compute=None, + ): + """Invokes the kernel. Block sizes derive only from logical lengths.""" + # `uniform_fixed_m` is a sub-mode of fixed-m; the kernel rejects the + # combination outright ("uniform_fixed_m requires use_fixed_m"). Asserting + # here turns a misconfigured case into a loud harness failure instead of a + # case that aborts at kernel construction and never touches the device -- + # which would look like a kernel finding but is really a test bug. + assert not (uniform_fixed_m and not use_fixed_m), "uniform_fixed_m requires use_fixed_m" + bkv_compute = bkv_compute or self.bq + block_sizes = custom_splash._BlockSizes( + block_q=self.bq, + block_kv=self.bq, + block_kv_compute=bkv_compute, + block_kv_compute_in=bkv_compute, + ) + kernel = custom_splash.make_splash_mha( + block_sizes=block_sizes, + orig_q_seq_len=self.q_len, + orig_kv_seq_len=self.kv_len, + use_base2_exp=True, + use_fixed_m=use_fixed_m, + uniform_fixed_m=uniform_fixed_m, + ) + if use_fixed_m: + out = kernel(q_in, k_in, v_in, mk, k_mean) + else: + out = kernel(q_in, k_in, v_in) + return jnp.swapaxes(out, 1, 2).astype(jnp.float32) # (heads, seq, dim) + + def _dense_reference(self, q, k, v): + qf, kf, vf = (x.astype(jnp.float32) for x in (q, k, v)) + logits = jnp.einsum("hsd,htd->hst", qf, kf) * self.scale + return jnp.einsum("hst,htd->hsd", jax.nn.softmax(logits, axis=-1), vf) + + def _assert_reference_is_sane(self, reference, q, k, v): + """Guards against a harness bug that would make both sides equally wrong.""" + self.assertTrue(bool(jnp.all(jnp.isfinite(reference))), "padded reference is not finite") + self.assertGreater(float(jnp.max(jnp.abs(reference))), 0.0, "padded reference is all zeros") + dense = self._dense_reference(q, k, v) + diff = float(jnp.max(jnp.abs(reference[:, : self.q_len] - dense))) + self.assertLess(diff, 5e-2, f"padded reference disagrees with dense softmax: {diff=}") + + def _compare( + self, + *, + unpad_q, + unpad_kv, + use_fixed_m=True, + uniform_fixed_m=None, + bkv_compute=None, + dirty=True, + ): + """Core assertion: unpadded inputs reproduce padded inputs bit-for-bit. + + `uniform_fixed_m` defaults to tracking `use_fixed_m` rather than to a bare + `True`, because `uniform_fixed_m=True` with `use_fixed_m=False` is rejected + by the kernel at construction time. Defaulting it to `True` made the online + cases abort before reaching the device -- they looked like failures of the + kernel when they were failures of this harness. + """ + if uniform_fixed_m is None: + uniform_fixed_m = use_fixed_m + q_in, k_in, v, q_raw, k_raw = self._inputs() + mk, k_mean = self._metadata(q_in, k_in) + + reference = self._run( + self._pad_seq(q_in, self.q_padded_len), + self._pad_seq(k_in, self.kv_padded_len), + self._pad_seq(v, self.kv_padded_len), + mk, + k_mean, + use_fixed_m=use_fixed_m, + uniform_fixed_m=uniform_fixed_m, + bkv_compute=bkv_compute, + ) + jax.block_until_ready(reference) + self._assert_reference_is_sane(reference, q_raw, k_raw, v) + + if dirty: + _dirty_device_memory() + + candidate = self._run( + q_in if unpad_q else self._pad_seq(q_in, self.q_padded_len), + k_in if unpad_kv else self._pad_seq(k_in, self.kv_padded_len), + v if unpad_kv else self._pad_seq(v, self.kv_padded_len), + mk, + k_mean, + use_fixed_m=use_fixed_m, + uniform_fixed_m=uniform_fixed_m, + bkv_compute=bkv_compute, + ) + jax.block_until_ready(candidate) + + self.assertEqual(reference.shape, candidate.shape) + self.assertTrue( + bool(jnp.all(jnp.isfinite(candidate))), + "unpadded run produced non-finite values", + ) + self.assertTrue( + bool(jnp.array_equal(reference, candidate)), + "unpadded inputs changed the result -- the kernel is reading past the end " + f"of the array (max abs delta {float(jnp.max(jnp.abs(reference - candidate)))})", + ) + return reference, candidate + + # ------------------------------------------------------------ C3a: K and V + + def test_kv_unpadded_fixed_m(self): + """C3a gate. Padded K/V rows are sliced away, so removing them must be a no-op.""" + self._compare(unpad_q=False, unpad_kv=True) + + def test_kv_unpadded_hybrid(self): + """C3a on the hybrid kernel, whose tail goes through `_last_online`.""" + self._compare(unpad_q=False, unpad_kv=True, uniform_fixed_m=False) + + def test_kv_unpadded_online(self): + """C3a on the pure online path -- the simplest probe of the DMA question.""" + self._compare(unpad_q=False, unpad_kv=True, use_fixed_m=False) + + def test_kv_unpadded_multi_iteration_tail(self): + """C3a with bkv_compute < bkv, exercising the fori_loop + ragged remainder.""" + self._compare(unpad_q=False, unpad_kv=True, bkv_compute=self.bq // 2) + + def test_kv_unpadded_online_multi_iteration_tail(self): + """Same, on the online path. + + The online path tracks a *running* max instead of a pinned one, so + `_last_online` is genuinely different code from `_last_fixed` even though + both slice the tail the same way. It gets its own ragged-remainder case + rather than inheriting confidence from the fixed-m result. + """ + self._compare(unpad_q=False, unpad_kv=True, use_fixed_m=False, bkv_compute=self.bq // 2) + + # ---------------------------------------------------------------- C3b: Q + + def test_q_unpadded_fixed_m(self): + """C3b gate. Relies on the norm-vector padding in `_metadata`.""" + self._compare(unpad_q=True, unpad_kv=False) + + def test_q_unpadded_hybrid(self): + self._compare(unpad_q=True, unpad_kv=False, uniform_fixed_m=False) + + def test_q_unpadded_online(self): + self._compare(unpad_q=True, unpad_kv=False, use_fixed_m=False) + + # ------------------------------------------------------------ both halves + + def test_all_unpadded_fixed_m(self): + self._compare(unpad_q=True, unpad_kv=True) + + def test_all_unpadded_hybrid(self): + self._compare(unpad_q=True, unpad_kv=True, uniform_fixed_m=False) + + # ------------------------------------------------------- adversarial cases + + def test_all_unpadded_is_stable_across_repeats(self): + """A single cold run can pass by luck on zeroed memory; N runs cannot. + + Between repeats the allocator pool is re-poisoned with NaN. If the kernel + were reading past the end of the array, the value it reads would change + from run to run and at least one repeat would diverge. + """ + first = None + for i in range(4): + _dirty_device_memory(num_buffers=4) + _, candidate = self._compare(unpad_q=True, unpad_kv=True, dirty=False) + if first is None: + first = candidate + else: + self.assertTrue( + bool(jnp.array_equal(first, candidate)), + f"unpadded run {i} differs from run 0 -- result depends on memory " + "outside the array, which is the signature of an unclipped DMA", + ) + + def test_all_unpadded_survives_interleaved_kernel_invocation(self): + """Runs a differently-shaped kernel first so VMEM holds real, non-zero data. + + `_dirty_device_memory` poisons HBM; this poisons the VMEM scratch buffers + that a ragged block DMA would only partially overwrite. The interleaved call + uses a large-magnitude value tensor so any leakage is numerically obvious + rather than lost in rounding. + """ + q_in, k_in, v, _, _ = self._inputs() + mk, k_mean = self._metadata(q_in, k_in) + loud = (jnp.ones_like(v) * 64).astype(v.dtype) + jax.block_until_ready( + self._run( + self._pad_seq(q_in, self.q_padded_len), + self._pad_seq(k_in, self.kv_padded_len), + self._pad_seq(loud, self.kv_padded_len), + mk, + k_mean, + use_fixed_m=True, + uniform_fixed_m=True, + ) + ) + self._compare(unpad_q=True, unpad_kv=True, dirty=False) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/tests/dot_fallback_layout_test.py b/src/maxdiffusion/tests/dot_fallback_layout_test.py new file mode 100644 index 000000000..c79e71a83 --- /dev/null +++ b/src/maxdiffusion/tests/dot_fallback_layout_test.py @@ -0,0 +1,165 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Layout regression tests for the short-sequence dot-product fallback. + +Sequences below `flash_min_seq_length` bypass the flash/ulysses kernels and +run `_apply_attention_dot`. Callers that apply rotary embeddings hand the +dispatcher `[B, H, S, D]` -- the dispatcher says so itself, reading the +sequence length from axis 2 when `ndim == 4` -- but the `split_head_dim` path +reshaped those tensors as if they were `[B, S, H*D]`. + +Both layouts hold the same number of elements, so the reshape succeeded and +returned a wrong answer with no error. That silent corruption is what these +tests exist to prevent. They are pure JAX and run on CPU. +""" + +import math +import unittest + +import jax +import jax.numpy as jnp +import numpy as np + +from maxdiffusion.models.attention_flax import _apply_attention_dot + + +def _call_dot(query, key, value, heads, dim_head, kv_heads=None): + """Runs the fallback and returns [B, H, S, D]. + + `_apply_attention_dot` emits the flat `[B, S, H*D]` form, so unflatten it + here rather than at each call site. + """ + out = _apply_attention_dot( + query=query, + key=key, + value=value, + dtype=jnp.float32, + heads=heads, + dim_head=dim_head, + scale=1.0 / math.sqrt(dim_head), + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=False, + attention_mask=None, + kv_heads=kv_heads, + ) + batch, seq, _ = out.shape + return jnp.swapaxes(out.reshape(batch, seq, heads, dim_head), 1, 2) + + +def _reference_attention(query, key, value, scale): + """Dense f32 attention on explicit [B, H, S, D] inputs.""" + q, k, v = (x.astype(jnp.float32) for x in (query, key, value)) + logits = jnp.einsum("bhqd,bhkd->bhqk", q, k) * scale + return jnp.einsum("bhqk,bhkd->bhqd", jax.nn.softmax(logits, axis=-1), v) + + +class DotFallbackLayoutTest(unittest.TestCase): + """`_apply_attention_dot` must transpose, not reshape, 4-D inputs.""" + + def test_zero_logits_return_per_head_token_means(self): + """The reviewer's counterexample, reproduced exactly. + + One active head dimension, three tokens, two heads. Q = K = 0 makes every + logit zero, so softmax is uniform and each head's output is the mean of + its values over tokens: + + head 0 values [1, 2, 3] -> 2 + head 1 values [10, 20, 30] -> 20 + + Reinterpreting [B, H, S, D] as [B, S, H, D] instead walks the buffer + [1, 2, 3, 10, 20, 30] as three token-pairs (1,2), (3,10), (20,30), + producing [8, 14]. + """ + heads, seq, dim_head = 2, 3, 1 + shape = (1, heads, seq, dim_head) + query = jnp.zeros(shape, jnp.float32) + key = jnp.zeros(shape, jnp.float32) + value = jnp.array([[[[1.0], [2.0], [3.0]], [[10.0], [20.0], [30.0]]]], dtype=jnp.float32) + self.assertEqual(value.shape, shape) + + out = np.asarray(_call_dot(query, key, value, heads, dim_head)) + + np.testing.assert_allclose(out[0, 0], 2.0, rtol=1e-5, atol=1e-5) + np.testing.assert_allclose(out[0, 1], 20.0, rtol=1e-5, atol=1e-5) + # The exact wrong answer the reshape produced. + self.assertFalse(np.allclose(out[0, 0], 8.0)) + self.assertFalse(np.allclose(out[0, 1], 14.0)) + + def test_matches_reference_on_random_4d_inputs(self): + heads, seq, dim_head = 4, 8, 16 + shape = (2, heads, seq, dim_head) + query = jax.random.normal(jax.random.PRNGKey(0), shape, jnp.float32) + key = jax.random.normal(jax.random.PRNGKey(1), shape, jnp.float32) + value = jax.random.normal(jax.random.PRNGKey(2), shape, jnp.float32) + + out = np.asarray(_call_dot(query, key, value, heads, dim_head)) + expected = np.asarray(_reference_attention(query, key, value, 1.0 / math.sqrt(dim_head))) + np.testing.assert_allclose(out, expected, rtol=1e-4, atol=1e-4) + + def test_three_dim_inputs_agree_with_four_dim(self): + """The flat [B, S, H*D] contract must keep working, and agree.""" + heads, seq, dim_head = 4, 8, 16 + batch = 2 + shape = (batch, heads, seq, dim_head) + query = jax.random.normal(jax.random.PRNGKey(3), shape, jnp.float32) + key = jax.random.normal(jax.random.PRNGKey(4), shape, jnp.float32) + value = jax.random.normal(jax.random.PRNGKey(5), shape, jnp.float32) + + def flatten(x): + return jnp.swapaxes(x, 1, 2).reshape(batch, seq, heads * dim_head) + + out_4d = np.asarray(_call_dot(query, key, value, heads, dim_head)) + out_3d = np.asarray(_call_dot(flatten(query), flatten(key), flatten(value), heads, dim_head)) + np.testing.assert_allclose(out_4d, out_3d, rtol=1e-5, atol=1e-5) + + def test_gqa_repeat_still_applies_on_4d(self): + """Head repetition must happen after the transpose, on the head axis.""" + heads, kv_heads, seq, dim_head = 4, 2, 8, 16 + q = jax.random.normal(jax.random.PRNGKey(6), (1, heads, seq, dim_head), jnp.float32) + k = jax.random.normal(jax.random.PRNGKey(7), (1, kv_heads, seq, dim_head), jnp.float32) + v = jax.random.normal(jax.random.PRNGKey(8), (1, kv_heads, seq, dim_head), jnp.float32) + + out = np.asarray(_call_dot(q, k, v, heads, dim_head, kv_heads=kv_heads)) + expected = np.asarray( + _reference_attention( + q, + jnp.repeat(k, heads // kv_heads, axis=1), + jnp.repeat(v, heads // kv_heads, axis=1), + 1.0 / math.sqrt(dim_head), + ) + ) + np.testing.assert_allclose(out, expected, rtol=1e-4, atol=1e-4) + + +class DispatcherLayoutContractTest(unittest.TestCase): + """The threshold check and the dot path must agree on where S lives.""" + + def test_seq_len_axis_for_4d_is_the_third_axis(self): + """`_apply_attention` reads S from axis 2 for 4-D, i.e. [B, H, S, D]. + + That is the convention `_apply_attention_dot` now honours by transposing. + If the dispatcher ever moves to [B, S, H, D], the transpose becomes wrong + and this assertion should be revisited alongside it. + """ + query = jnp.zeros((1, 4, 128, 8), jnp.float32) # B, H, S, D + seq_len_idx = 2 if query.ndim == 4 else 1 + self.assertEqual(query.shape[seq_len_idx], 128) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/tests/fused_producers_test.py b/src/maxdiffusion/tests/fused_producers_test.py new file mode 100644 index 000000000..effab8d4f --- /dev/null +++ b/src/maxdiffusion/tests/fused_producers_test.py @@ -0,0 +1,248 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Numerical-equivalence tests for the fused attention producers. + +These run on CPU; they pin numerical contracts, not kernel performance. +""" + +import unittest + +import jax +import jax.numpy as jnp +import numpy as np +from flax import nnx + +from maxdiffusion.kernels.fused_producers import fused_rmsnorm_rope + + +def _reference_rmsnorm(x, scale, eps=1e-6): + """Mirrors flax.nnx.RMSNorm's association: x * (rsqrt(var + eps) * scale). + + Flax's `_normalize` builds `mul = rsqrt(var + eps)`, folds the scale into it + with `mul *= scale`, and only then applies `y *= mul`. Reassociating this as + `(x * rsqrt) * scale` rounds differently, so the order is load-bearing. + """ + var = jnp.mean(jnp.square(x.astype(jnp.float32)), axis=-1, keepdims=True) + return x.astype(jnp.float32) * (jax.lax.rsqrt(var + eps) * scale.astype(jnp.float32)) + + +class FusedRmsNormAssociationTest(unittest.TestCase): + """The fused producer must be a pure fusion, never a numerical change.""" + + def test_matches_flax_rmsnorm_bit_for_bit(self): + dim = 128 + key = jax.random.PRNGKey(0) + k1, k2 = jax.random.split(key) + x = jax.random.normal(k1, (2, 64, dim), jnp.float32) + scale = jax.random.normal(k2, (dim,), jnp.float32) + + layer = nnx.RMSNorm( + dim, + epsilon=1e-6, + dtype=jnp.float32, + param_dtype=jnp.float32, + rngs=nnx.Rngs(0), + ) + layer.scale.value = scale + + np.testing.assert_array_equal( + np.asarray(_reference_rmsnorm(x, scale)), + np.asarray(layer(x)), + err_msg="Reference helper must reproduce nnx.RMSNorm exactly.", + ) + + def test_left_to_right_association_is_not_equivalent(self): + """Guards the reason this test exists: the two orders genuinely differ.""" + dim = 128 + key = jax.random.PRNGKey(1) + k1, k2 = jax.random.split(key) + x = jax.random.normal(k1, (2, 64, dim), jnp.float32) + scale = jax.random.normal(k2, (dim,), jnp.float32) + + rsqrt = jax.lax.rsqrt(jnp.mean(jnp.square(x), axis=-1, keepdims=True) + 1e-6) + folded = x * (rsqrt * scale) # Flax order + left_to_right = (x * rsqrt) * scale + + self.assertFalse( + bool(jnp.all(folded == left_to_right)), + "If these ever become identical the association guard above is vacuous.", + ) + + def test_fused_producer_q_matches_flax_rmsnorm(self): + """End-to-end: the q path of the fused producer must match nnx.RMSNorm.""" + b, seq, q_heads, dim_head = 1, 8, 2, 8 + d_model = q_heads * dim_head + key = jax.random.PRNGKey(2) + k1, k2, k3 = jax.random.split(key, 3) + + raw_q = jax.random.normal(k1, (b, seq, d_model), jnp.float32) + raw_k = jax.random.normal(k2, (b, seq, d_model), jnp.float32) + q_scale = jax.random.normal(k3, (d_model,), jnp.float32) + k_scale = jnp.ones((d_model,), jnp.float32) + + # Identity rotation isolates the RMSNorm from the RoPE. + freqs_cis = jnp.ones((1, 1, seq, dim_head // 2), jnp.complex64) + + q_out, _ = fused_rmsnorm_rope( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=q_heads, + kv_heads=q_heads, + dim_head=dim_head, + ) + + expected = _reference_rmsnorm(raw_q, q_scale).reshape(b, seq, q_heads, dim_head).transpose(0, 2, 1, 3) + np.testing.assert_array_equal( + np.asarray(q_out.astype(jnp.float32)), + np.asarray(expected.astype(raw_q.dtype).astype(jnp.float32)), + err_msg="Fused RMSNorm+RoPE must be bit-identical to nnx.RMSNorm under an identity rotation.", + ) + + def test_separate_q_and_k_epsilons_are_applied(self): + """Q eps=1e-5 and K eps=1e-6 are applied separately (checked against analytic values, fp32, rtol 1e-6).""" + b, seq, q_heads, dim_head = 1, 4, 2, 8 + d_model = q_heads * dim_head + raw_q = jnp.full((b, seq, d_model), 1e-3, dtype=jnp.float32) + raw_k = jnp.full((b, seq, d_model), 1e-3, dtype=jnp.float32) + q_scale = jnp.ones((d_model,), dtype=jnp.float32) + k_scale = jnp.ones((d_model,), dtype=jnp.float32) + freqs_cis = jnp.ones((1, 1, seq, dim_head // 2), dtype=jnp.complex64) + + q_out, k_out = fused_rmsnorm_rope( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=q_heads, + kv_heads=q_heads, + dim_head=dim_head, + eps=1e-5, + k_eps=1e-6, + ) + # 1e-3 / sqrt(1e-6 + 1e-5) = 0.30151134 for Q; 1e-3 / sqrt(1e-6 + 1e-6) = 0.70710678 for K + np.testing.assert_allclose(np.asarray(q_out), 0.30151134, rtol=1e-6, atol=1e-6) + np.testing.assert_allclose(np.asarray(k_out), 0.70710678, rtol=1e-6, atol=1e-6) + + def test_nontrivial_rope_and_bfloat16_match_unfused_reference(self): + """Non-trivial complex RoPE rotation + pair interleave in bfloat16 matches unfused RMSNorm + RoPE.""" + b, seq, q_heads, dim_head = 2, 16, 4, 16 + d_model = q_heads * dim_head + key = jax.random.PRNGKey(42) + k1, k2, k3, k4, k5 = jax.random.split(key, 5) + + raw_q = jax.random.normal(k1, (b, seq, d_model), jnp.bfloat16) + raw_k = jax.random.normal(k2, (b, seq, d_model), jnp.bfloat16) + q_scale = jax.random.normal(k3, (d_model,), jnp.bfloat16) + k_scale = jax.random.normal(k4, (d_model,), jnp.bfloat16) + angles = jax.random.uniform(k5, (1, 1, seq, dim_head // 2), minval=-3.14, maxval=3.14) + freqs_cis = jnp.exp(1j * angles).astype(jnp.complex64) + + q_out, k_out = fused_rmsnorm_rope( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=q_heads, + kv_heads=q_heads, + dim_head=dim_head, + ) + self.assertEqual(q_out.dtype, jnp.bfloat16) + self.assertEqual(k_out.dtype, jnp.bfloat16) + + # The producer is jit-wrapped, so XLA keeps the bf16 RoPE chain in f32 inside + # a fusion and rounds once; op-by-op eager bf16 rounds after every multiply + # and differs by 1 ulp. Compile the reference the same way so the comparison + # pins the fusion's numerics, not eager-vs-compiled bf16 rounding. + @jax.jit + def _unfused_ref(x, scale): + normed = _reference_rmsnorm(x, scale).astype(x.dtype) + h = normed.reshape(b, seq, q_heads, dim_head).transpose(0, 2, 1, 3) + cos = jnp.real(freqs_cis).astype(x.dtype) + sin = jnp.imag(freqs_cis).astype(x.dtype) + pairs = h.reshape(b, q_heads, seq, -1, 2) + x0, x1 = pairs[..., 0], pairs[..., 1] + out0 = x0 * cos - x1 * sin + out1 = x0 * sin + x1 * cos + return jnp.stack([out0, out1], axis=-1).reshape(b, q_heads, seq, dim_head) + + np.testing.assert_array_equal( + np.asarray(q_out.astype(jnp.float32)), + np.asarray(_unfused_ref(raw_q, q_scale).astype(jnp.float32)), + ) + np.testing.assert_array_equal( + np.asarray(k_out.astype(jnp.float32)), + np.asarray(_unfused_ref(raw_k, k_scale).astype(jnp.float32)), + ) + + def test_gqa_nontrivial_rope_shapes_and_values(self): + """GQA (q_heads != kv_heads) with non-trivial RoPE produces exact reference values.""" + b, sq, sk, q_heads, kv_heads, dim_head = 2, 12, 8, 4, 2, 16 + dq = q_heads * dim_head + dk = kv_heads * dim_head + key = jax.random.PRNGKey(99) + k1, k2, k3, k4, k5 = jax.random.split(key, 5) + + raw_q = jax.random.normal(k1, (b, sq, dq), jnp.bfloat16) + raw_k = jax.random.normal(k2, (b, sk, dk), jnp.bfloat16) + q_scale = jax.random.normal(k3, (dq,), jnp.bfloat16) + k_scale = jax.random.normal(k4, (dk,), jnp.bfloat16) + angles = jax.random.uniform(k5, (1, 1, max(sq, sk), dim_head // 2), minval=-2.0, maxval=2.0) + freqs_cis = jnp.exp(1j * angles).astype(jnp.complex64) + + q_out, k_out = fused_rmsnorm_rope( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=q_heads, + kv_heads=kv_heads, + dim_head=dim_head, + ) + self.assertEqual(q_out.shape, (b, q_heads, sq, dim_head)) + self.assertEqual(k_out.shape, (b, kv_heads, sk, dim_head)) + + # Verify K against manual complex multiplication in float32. Compiled for the + # same reason as in test_nontrivial_rope_and_bfloat16_match_unfused_reference. + @jax.jit + def _k_ref(raw_k, k_scale, freqs_cis): + k_normed = _reference_rmsnorm(raw_k, k_scale).astype(jnp.bfloat16) + k_h = k_normed.reshape(b, sk, kv_heads, dim_head).transpose(0, 2, 1, 3) + cos_k = jnp.real(freqs_cis[:, :, :sk, :]).astype(jnp.bfloat16) + sin_k = jnp.imag(freqs_cis[:, :, :sk, :]).astype(jnp.bfloat16) + k_pairs = k_h.reshape(b, kv_heads, sk, -1, 2) + return jnp.stack( + [ + k_pairs[..., 0] * cos_k - k_pairs[..., 1] * sin_k, + k_pairs[..., 0] * sin_k + k_pairs[..., 1] * cos_k, + ], + axis=-1, + ).reshape(b, kv_heads, sk, dim_head) + + np.testing.assert_array_equal( + np.asarray(k_out.astype(jnp.float32)), + np.asarray(_k_ref(raw_k, k_scale, freqs_cis).astype(jnp.float32)), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/tests/ltx2/test_attention_ltx2.py b/src/maxdiffusion/tests/ltx2/test_attention_ltx2.py index f14447492..55a8c7acc 100644 --- a/src/maxdiffusion/tests/ltx2/test_attention_ltx2.py +++ b/src/maxdiffusion/tests/ltx2/test_attention_ltx2.py @@ -179,6 +179,7 @@ def test_ring_cross_attention_uses_global_kv_and_wires_kernel_flags(self): "ulysses_ring", "ulysses_ring_custom", "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", "ulysses_ring_custom_bidir", "ulysses_custom", "ulysses_custom_fixed_m", diff --git a/src/maxdiffusion/tests/ring_fixed_m_test.py b/src/maxdiffusion/tests/ring_fixed_m_test.py index 5155811c6..bc1b77e38 100644 --- a/src/maxdiffusion/tests/ring_fixed_m_test.py +++ b/src/maxdiffusion/tests/ring_fixed_m_test.py @@ -17,16 +17,18 @@ """Unit tests for the fixed-m path of the custom RING attention. The ring path gates fixed-m globally against floor(W(N_total)/2) (accumulate -merge when every (head, shard) passes) and otherwise PER (head, K-shard) +merge when every (head, Q-block) passes) and otherwise per (head, Q-block) against floor(W(N_local)/2), merging the per-hop partials in LSE space -(invariant to fixed-m's bound overshoot). These tests check, against an f32 +(invariant to fixed-m's bound overshoot). The per-hop gate uses the ring-wide +max K norm, so it is the same on every hop. These tests check, against an f32 dense-softmax reference: * the untouched online ring path (regression guard), * fixed-m with every (head, shard) eligible, * a sink head ineligible on every shard (all-online fallback), - * a head eligible on one shard but not the other -- the mixed - fixed/online partial case that requires the LSE merge. + * a head whose keys are large on one shard only (the LSE path; because the + per-hop gate uses the ring-wide K max, that head runs online on every hop), + * GQA (4 Q heads, 2 KV heads) with a ragged last KV block. """ import functools @@ -98,6 +100,10 @@ def _scaled_inputs(self, q, k): def _reference(self, q_in, k_in, v): """Dense f32 log2-domain softmax on the kernel's own bf16 inputs.""" qf, kf, vf = (x.astype(jnp.float32) for x in (q_in, k_in, v)) + if qf.shape[0] != kf.shape[0]: + q_heads_per_kv = qf.shape[0] // kf.shape[0] + kf = jnp.repeat(kf, q_heads_per_kv, axis=0) + vf = jnp.repeat(vf, q_heads_per_kv, axis=0) logits = jnp.einsum("hqd,hkd->hqk", qf, kf) # LOG2E & scale pre-folded return jax.nn.softmax(logits * math.log(2.0), axis=-1) @ vf @@ -135,7 +141,7 @@ def _body(ql, kl, vl): # The V/dtype safety verdict the production caller computes. It is # global, so it is reduced across the ring before use. v_max_sq = (vl.astype(jnp.float32) ** 2).max() - recenter, _ = custom_splash.get_fixed_m_constants(self.shard_len * _RING_SIZE) + recenter, _ = custom_splash.get_fixed_m_constants(kl.shape[1] * _RING_SIZE) dtype_safe = custom_splash.fixed_m_dtype_is_safe(ql.dtype, recenter) v_ok_local = (v_max_sq <= (custom_splash.DEFAULT_MAX_V_BOUND**2)) & dtype_safe v_ok = jax.lax.pmin(v_ok_local, axis_name=_RING_AXIS) @@ -143,8 +149,8 @@ def _body(ql, kl, vl): v_ok = v_ok_override ring = ring_attention_kernel.make_custom_ring_attention( block_sizes=self.block_sizes, - orig_q_seq_len=self.shard_len, - orig_kv_seq_len=self.shard_len, + orig_q_seq_len=ql.shape[1], + orig_kv_seq_len=kl.shape[1], use_base2_exp=True, ring_axis=_RING_AXIS, ring_size=_RING_SIZE, @@ -168,16 +174,28 @@ def _shard_max_sq_norms(self, x): return sq.max(axis=-1) def _gate_per_shard(self, q_in, k_in): - """(heads, q_rank, k_shard) per-hop eligibility, as the kernel's LSE path computes it. + """(heads, q_rank, k_shard) eligibility of each local q against each K shard on its own. - Each rank gates its stationary local q against every K shard with the - two-sided bound for the local shard length, floor(W(shard_len)/2). + Uses the two-sided bound for the local shard length, floor(W(shard_len)/2). + This describes the test input. The kernel's per-hop gate is `_gate_hop`, + which uses the ring-wide K max instead of each shard's own. """ _, per_shard_bound = custom_splash.get_fixed_m_constants(self.shard_len) qn_sq = self._shard_max_sq_norms(q_in) # (heads, q_rank) kn_sq = self._shard_max_sq_norms(k_in) # (heads, k_shard) return qn_sq[:, :, None] * kn_sq[:, None, :] <= per_shard_bound**2 + def _gate_hop(self, q_in, k_in): + """(heads, q_rank) per-hop eligibility, as the kernel's LSE path computes it. + + Each rank gates its local q against the ring-wide max K norm with + floor(W(shard_len)/2); the result is the same on every hop. + """ + _, per_shard_bound = custom_splash.get_fixed_m_constants(self.shard_len) + qn_sq = self._shard_max_sq_norms(q_in) # (heads, q_rank) + kn_sq = self._shard_max_sq_norms(k_in).max(axis=1) # (heads,) + return qn_sq * kn_sq[:, None] <= per_shard_bound**2 + def _gate_global(self, q_in, k_in): """(heads,) global eligibility; all True means the kernel takes the accumulate merge.""" _, global_bound = custom_splash.get_fixed_m_constants(self.shard_len * _RING_SIZE) @@ -229,13 +247,18 @@ def test_fixed_m_accumulate_ragged_tail(self): self.assertLess(self._run_and_compare(q, k, v, use_fixed_m=True), 2e-2) def test_mixed_fixed_online_across_shards(self): - # Amplify head 0's keys on shard 1 only: head 0 is fixed on shard 0 but - # online on shard 1 -- the mixed-partial merge the LSE space exists for. + # Amplify head 0's keys on shard 1 only. On its own, shard 0 would be + # fixed-eligible for head 0 and shard 1 would not. The kernel's per-hop + # gate uses the ring-wide K max, so head 0 runs online on every hop while + # the other heads stay fixed, and the LSE merge combines them. q, k, v = self._random_qkv(k_gain=(0, slice(self.shard_len, self.shard_len * _RING_SIZE), 40.0)) self.assertFalse(bool(self._global_gate(q, k)[0])) # forces the per-hop LSE path gate = self._gate(q, k) - self.assertTrue(bool(jnp.all(gate[0, :, 0]))) # every rank: fixed on shard 0 - self.assertFalse(bool(jnp.any(gate[0, :, 1]))) # every rank: online on shard 1 + self.assertTrue(bool(jnp.all(gate[0, :, 0]))) # input: shard 0 alone is eligible + self.assertFalse(bool(jnp.any(gate[0, :, 1]))) # input: shard 1 alone is not + hop_gate = self._gate_hop(*self._scaled_inputs(q, k)) + self.assertFalse(bool(jnp.any(hop_gate[0]))) # kernel: head 0 online on every hop + self.assertTrue(bool(jnp.all(hop_gate[1:]))) # kernel: other heads fixed self.assertLess(self._run_and_compare(q, k, v, use_fixed_m=True), 2e-2) @_SKIP_IN_GITHUB_ACTIONS @@ -267,6 +290,23 @@ def test_v_ok_false_forces_finite_output(self): self.assertTrue(bool(jnp.all(jnp.isfinite(out)))) self.assertLess(float(jnp.max(jnp.abs(out - self._reference(q_in, k_in, v)))), 2e-2) + @_SKIP_IN_GITHUB_ACTIONS + def test_gqa_ring_fixed_m_ragged_last_kv_block(self): + """GQA (4 Q heads, 2 KV heads) through the ring kernel with fixed-m. + + Drives `make_custom_ring_attention` directly (not `_ulysses_ring_custom_attention`). + The 1536-token shard is 8-aligned; the only raggedness is the last KV block + (1536 % block_kv 1024 = 512). It does not exercise `kv_pad_size = 1`. + """ + q_heads = 4 + kv_heads = 2 + ragged_len = 1536 + total_len = ragged_len * _RING_SIZE + q = jax.random.normal(jax.random.PRNGKey(42), (q_heads, total_len, self.head_dim), jnp.bfloat16) + k = jax.random.normal(jax.random.PRNGKey(43), (kv_heads, total_len, self.head_dim), jnp.bfloat16) + v = jax.random.normal(jax.random.PRNGKey(44), (kv_heads, total_len, self.head_dim), jnp.bfloat16) + self.assertLess(self._run_and_compare(q, k, v, use_fixed_m=True), 2e-2) + class RingFixedMContractTest(unittest.TestCase): """Backend-independent checks on the fixed-m ring API contract. @@ -445,7 +485,7 @@ def test_centered_logit_overflows_fp32_under_the_raw_bound(self): recenter, _ = custom_splash.get_fixed_m_constants(self.total_kv) k_mean = keys.mean(axis=0) - max_centered_logit = float(((keys - k_mean) @ query).max()) + max_centered_logit = float(jnp.dot(keys - k_mean, query, precision=jax.lax.Precision.HIGHEST).max()) fixed_m = float(jnp.linalg.norm(query)) * float(jnp.linalg.norm(keys, axis=-1).max()) # fixed-m parks the max weight at 2**recenter, so the realised exponent is @@ -465,10 +505,332 @@ def test_centering_restores_a_sound_bound(self): # Correctly rejected, so this tile takes the online-softmax path. self.assertGreater(centered_bound, safe_bound) # And had it been admitted, the bound would genuinely cap the logit. - max_centered_logit = float((centered @ query).max()) + max_centered_logit = float(jnp.dot(centered, query, precision=jax.lax.Precision.HIGHEST).max()) self.assertLessEqual(max_centered_logit, centered_bound + 1e-3) self.assertLessEqual(max_centered_logit - centered_bound + recenter, 128.0) +@unittest.skipIf( + len(jax.devices()) < 4, + f"UlyssesRingAttentionFlaxIntegrationTest requires >= 4 devices, got {len(jax.devices())}", +) +@_SKIP_IN_GITHUB_ACTIONS +class UlyssesRingAttentionFlaxIntegrationTest(unittest.TestCase): + """End-to-end 4-device (U=2, R=2) tests for `_ulysses_ring_custom_attention`. + + Exercises the R>1 production wrapper (`_ring_fixed_m_norms_pre_a2a`, + `pregathered_mk=True`, caller-supplied `all_fixed_global`, per_q_block Q-norm + all-to-all, K-centering `pmean` + virtual `k_mean_dev`, and `kv_pad_size=1` + ragged tail masking) against a dense float32 reference. + """ + + batch = 1 + heads = 4 + head_dim = 128 + scale = 1.0 / math.sqrt(128) + ulysses_shards = 2 + context_shards = 4 + + @classmethod + def setUpClass(cls): + super().setUpClass() + from flax import linen as nn + from maxdiffusion.common_types import BlockSizes + from maxdiffusion.models import attention_flax + + cls.nn = nn + cls.attention_flax = attention_flax + cls.mesh = jax.sharding.Mesh(np.asarray(jax.devices()[:4]), ("context",)) + cls.axis_names = ( + attention_flax.BATCH, + attention_flax.SELF_ATTN_HEAD, + attention_flax.SELF_ATTN_Q_LENGTH, + attention_flax.D_KV, + ) + cls.block_sizes = BlockSizes( + block_q=128, + block_kv_compute=128, + block_kv=128, + block_q_dkv=128, + block_kv_dkv=128, + block_kv_dkv_compute=128, + block_q_dq=128, + block_kv_dq=128, + use_fused_bwd_kernel=False, + ) + + def _dense_reference(self, q: jax.Array, k_scaled: jax.Array, v: jax.Array) -> jax.Array: + """Dense float32 reference on 4D `(B, H, S, D)` inputs (`k_scaled = k * scale`).""" + qf = q.astype(jnp.float32) + kf = k_scaled.astype(jnp.float32) + vf = v.astype(jnp.float32) + if qf.shape[1] != kf.shape[1]: + rep = qf.shape[1] // kf.shape[1] + kf = jnp.repeat(kf, rep, axis=1) + vf = jnp.repeat(vf, rep, axis=1) + logits = jnp.einsum("bhsd,bhtd->bhst", qf, kf, precision=jax.lax.Precision.HIGHEST) + weights = jax.nn.softmax(logits, axis=-1) + out = jnp.einsum("bhst,bhtd->bhsd", weights, vf, precision=jax.lax.Precision.HIGHEST) + # `_ulysses_ring_custom_attention` returns 3D `(B, S, H * D)`. + b, h, s, d = out.shape + return out.transpose(0, 2, 1, 3).reshape(b, s, h * d) + + def _rel_l2(self, actual: jax.Array, expected: jax.Array) -> float: + af = actual.astype(jnp.float32) + ef = expected.astype(jnp.float32) + return float(jnp.linalg.norm(af - ef) / (jnp.linalg.norm(ef) + 1e-12)) + + def _eval_pre_a2a_predicates( + self, + q: jax.Array, + k_scaled: jax.Array, + v: jax.Array, + *, + per_q_block: bool, + use_k_centering: bool, + ) -> tuple[np.ndarray, np.ndarray]: + """Runs `_ring_fixed_m_norms_pre_a2a` across all 4 (R=2, U=2) shards and returns per-device `all_fixed_global`.""" + af = self.attention_flax + internal_mesh = af._create_internal_ulysses_ring_mesh(self.mesh, 2, 2) + ring_axis, ulysses_axis = af.INTERNAL_RING_AXIS, af.INTERNAL_ULYSSES_AXIS + spec = jax.sharding.PartitionSpec(None, None, (ring_axis, ulysses_axis), None) + out_spec = jax.sharding.PartitionSpec((ring_axis, ulysses_axis)) + + @functools.partial( + jax.shard_map, + mesh=internal_mesh, + in_specs=(spec, spec, spec), + out_specs=(out_spec, out_spec), + check_vma=False, + ) + def _probe(q_s, k_s, v_s): + _, _, v_ok, all_fixed_global, _ = af._ring_fixed_m_norms_pre_a2a( + q_s * af.LOG2E, + k_s, + v_s, + ulysses_axis=ulysses_axis, + ring_axis=ring_axis, + num_ulysses_shards=2, + num_ring_shards=2, + block_q=128, + per_q_block=per_q_block, + use_k_centering=use_k_centering, + ) + return jnp.expand_dims(all_fixed_global, 0), jnp.expand_dims(v_ok, 0) + + all_fixed_per_dev, v_ok_per_dev = _probe(q, k_scaled, v) + return np.asarray(all_fixed_per_dev), np.asarray(v_ok_per_dev) + + def _run_ulysses_ring( + self, + q: jax.Array, + k_scaled: jax.Array, + v: jax.Array, + *, + kv_heads: int, + per_q_block: bool, + use_k_centering: bool, + ) -> jax.Array: + af = self.attention_flax + rules = [ + (af.BATCH, None), + (af.SELF_ATTN_HEAD, None), + (af.SELF_ATTN_Q_LENGTH, "context"), + (af.SELF_ATTN_KV_LENGTH, "context"), + (af.D_KV, None), + ] + with self.mesh, self.nn.logical_axis_rules(rules): + return af._ulysses_ring_custom_attention( + q, + k_scaled, + v, + heads=self.heads, + mesh=self.mesh, + axis_names_q=self.axis_names, + axis_names_kv=self.axis_names, + flash_block_sizes=self.block_sizes, + dtype=jnp.bfloat16, + ulysses_shards=self.ulysses_shards, + use_base2_exp=True, + use_fixed_m=True, + per_q_block=per_q_block, + kv_heads=kv_heads, + use_k_centering=use_k_centering, + ) + + def test_u2_r2_grid_per_q_block_centering_and_gqa(self): + """Sweeps per_q_block x use_k_centering x GQA at U=2, R=2 against dense f32.""" + total_seq = 512 # 128 per context shard; 256 per ring shard after Ulysses a2a (2 Q/KV blocks of 128) + k1, k2, k3 = jax.random.split(jax.random.PRNGKey(101), 3) + q = jax.random.normal(k1, (self.batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + + for kv_heads in (4, 2): + k = jax.random.normal(k2, (self.batch, kv_heads, total_seq, self.head_dim), jnp.bfloat16) + v = jax.random.normal(k3, (self.batch, kv_heads, total_seq, self.head_dim), jnp.bfloat16) + k_scaled = (k.astype(jnp.float32) * self.scale).astype(jnp.bfloat16) + ref = self._dense_reference(q, k_scaled, v) + + for per_q_block in (False, True): + for use_k_centering in (False, True): + with self.subTest(kv_heads=kv_heads, per_q_block=per_q_block, use_k_centering=use_k_centering): + all_fixed_devs, v_ok_devs = self._eval_pre_a2a_predicates( + q, + k_scaled, + v, + per_q_block=per_q_block, + use_k_centering=use_k_centering, + ) + self.assertTrue(bool(np.all(v_ok_devs))) + # Normal-scale unit inputs must take the `_accumulate_scan` branch on all 4 devices. + self.assertTrue(bool(np.all(all_fixed_devs))) + out = self._run_ulysses_ring( + q, + k_scaled, + v, + kv_heads=kv_heads, + per_q_block=per_q_block, + use_k_centering=use_k_centering, + ) + self.assertTrue(bool(jnp.all(jnp.isfinite(out)))) + self.assertLess(self._rel_l2(out, ref), 2e-2) + + def test_u2_r2_ragged_kv_length_exercises_kv_pad_size_1(self): + """Ragged ring KV shard (272 tokens = 2 * 128 + 16) exercises `kv_pad_size=1` tail masking.""" + # 4 context shards * 136 tokens/shard = 544 total tokens. + # After Ulysses (U=2) all-to-all, each ring shard has 272 tokens: + # 272 is 8-sublane aligned (`pad_kv_len == 0`) and `272 % block_kv(128) = 16 != 0`, + # so `_ulysses_ring_custom_attention` sets `kv_pad_size = 1` and masks the last KV tile. + total_seq = 544 + k1, k2, k3 = jax.random.split(jax.random.PRNGKey(202), 3) + q = jax.random.normal(k1, (self.batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + k = jax.random.normal(k2, (self.batch, 2, total_seq, self.head_dim), jnp.bfloat16) + v = jax.random.normal(k3, (self.batch, 2, total_seq, self.head_dim), jnp.bfloat16) + k_scaled = (k.astype(jnp.float32) * self.scale).astype(jnp.bfloat16) + ref = self._dense_reference(q, k_scaled, v) + + for use_k_centering in (False, True): + with self.subTest(use_k_centering=use_k_centering): + out = self._run_ulysses_ring( + q, + k_scaled, + v, + kv_heads=2, + per_q_block=True, + use_k_centering=use_k_centering, + ) + self.assertEqual(out.shape, (self.batch, total_seq, self.heads * self.head_dim)) + self.assertTrue(bool(jnp.all(jnp.isfinite(out)))) + self.assertLess(self._rel_l2(out, ref), 2e-2) + + def test_u2_r2_sink_head_forces_lse_scan_uniformly_across_mesh(self): + """A single sink head on Ulysses rank 0 forces `_lse_scan` uniformly across all 4 (U=2, R=2) devices.""" + total_seq = 512 + k1, k2, k3 = jax.random.split(jax.random.PRNGKey(303), 3) + q = jax.random.normal(k1, (self.batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + # Head 0 (owned by Ulysses rank 0 post-a2a, while Ulysses rank 1 owns heads 2..3) exceeds the bound. + q = q.at[:, 0, :64, :].multiply(jnp.bfloat16(40.0)) + k = jax.random.normal(k2, (self.batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + v = jax.random.normal(k3, (self.batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + k_scaled = (k.astype(jnp.float32) * self.scale).astype(jnp.bfloat16) + ref = self._dense_reference(q, k_scaled, v) + + for per_q_block in (False, True): + for use_k_centering in (False, True): + with self.subTest(per_q_block=per_q_block, use_k_centering=use_k_centering): + all_fixed_devs, v_ok_devs = self._eval_pre_a2a_predicates( + q, + k_scaled, + v, + per_q_block=per_q_block, + use_k_centering=use_k_centering, + ) + self.assertTrue(bool(np.all(v_ok_devs))) + # Crucial: all 4 devices (both Ulysses rank 0 and Ulysses rank 1) must agree + # that `all_fixed_global == False` so `lax.cond` takes `_lse_scan` mesh-wide. + self.assertEqual(all_fixed_devs.tolist(), [False, False, False, False]) + out = self._run_ulysses_ring( + q, + k_scaled, + v, + kv_heads=self.heads, + per_q_block=per_q_block, + use_k_centering=use_k_centering, + ) + self.assertTrue(bool(jnp.all(jnp.isfinite(out)))) + self.assertLess(self._rel_l2(out, ref), 2e-2) + + def test_u2_r2_v_outlier_forces_fallback_uniformly_batch2(self): + """One |V| > DEFAULT_MAX_V_BOUND element on one device makes v_ok False on all 4 devices (B=2).""" + total_seq = 512 + batch = 2 + k1, k2, k3 = jax.random.split(jax.random.PRNGKey(404), 3) + q = jax.random.normal(k1, (batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + k = jax.random.normal(k2, (batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + v = jax.random.normal(k3, (batch, self.heads, total_seq, self.head_dim), jnp.bfloat16) + # Last token (context shard 3), last head, second batch element only. + v = v.at[1, self.heads - 1, total_seq - 1, 0].set(jnp.bfloat16(2.0 * custom_splash.DEFAULT_MAX_V_BOUND)) + k_scaled = (k.astype(jnp.float32) * self.scale).astype(jnp.bfloat16) + ref = self._dense_reference(q, k_scaled, v) + + for per_q_block in (False, True): + with self.subTest(per_q_block=per_q_block): + all_fixed_devs, v_ok_devs = self._eval_pre_a2a_predicates( + q, + k_scaled, + v, + per_q_block=per_q_block, + use_k_centering=False, + ) + self.assertEqual(v_ok_devs.tolist(), [False, False, False, False]) + self.assertEqual(all_fixed_devs.tolist(), [False, False, False, False]) + out = self._run_ulysses_ring( + q, + k_scaled, + v, + kv_heads=self.heads, + per_q_block=per_q_block, + use_k_centering=False, + ) + self.assertEqual(out.shape, (batch, total_seq, self.heads * self.head_dim)) + self.assertTrue(bool(jnp.all(jnp.isfinite(out)))) + self.assertLess(self._rel_l2(out, ref), 2e-2) + + def test_u2_r2_accumulate_scan_batch2_and_sub_128_head_dim_k_mean_padding(self): + """Exercises vmap over batch>1 on _accumulate_scan and ring k_mean head-dim padding (d=64 -> 128).""" + total_seq = 512 + batch = 2 + sub_head_dim = 64 + sub_scale = 1.0 / math.sqrt(sub_head_dim) + k1, k2, k3 = jax.random.split(jax.random.PRNGKey(505), 3) + q = jax.random.normal(k1, (batch, self.heads, total_seq, sub_head_dim), jnp.bfloat16) + k = jax.random.normal(k2, (batch, self.heads, total_seq, sub_head_dim), jnp.bfloat16) + v = jax.random.normal(k3, (batch, self.heads, total_seq, sub_head_dim), jnp.bfloat16) + k_scaled = (k.astype(jnp.float32) * sub_scale).astype(jnp.bfloat16) + ref = self._dense_reference(q, k_scaled, v) + + for per_q_block in (False, True): + with self.subTest(per_q_block=per_q_block): + all_fixed_devs, v_ok_devs = self._eval_pre_a2a_predicates( + q, + k_scaled, + v, + per_q_block=per_q_block, + use_k_centering=True, + ) + self.assertEqual(v_ok_devs.tolist(), [True, True, True, True]) + self.assertEqual(all_fixed_devs.tolist(), [True, True, True, True]) + out = self._run_ulysses_ring( + q, + k_scaled, + v, + kv_heads=self.heads, + per_q_block=per_q_block, + use_k_centering=True, + ) + self.assertEqual(out.shape, (batch, total_seq, self.heads * sub_head_dim)) + self.assertTrue(bool(jnp.all(jnp.isfinite(out)))) + self.assertLess(self._rel_l2(out, ref), 2e-2) + + if __name__ == "__main__": unittest.main() diff --git a/src/maxdiffusion/tests/tile_size_grid_search_test.py b/src/maxdiffusion/tests/tile_size_grid_search_test.py index 3419c142f..ff2f250cc 100644 --- a/src/maxdiffusion/tests/tile_size_grid_search_test.py +++ b/src/maxdiffusion/tests/tile_size_grid_search_test.py @@ -129,6 +129,10 @@ def fake_fn(): self.assertLess(mean, 30.0) # ...and NOT in the steady-state mean self.assertEqual(len(times), 5) + def test_rejects_nonpositive_iters(self): + with self.assertRaises(ValueError): + time_callable(lambda: 1, iters=0) + class OrchestratorTest(unittest.TestCase): @@ -167,6 +171,19 @@ def test_oom_configs_pruned_not_raised(self): ) self.assertTrue(any(r.status == "oom" for r in res.results)) + def test_broadcast_winner_handles_none_bkv_compute(self): + from unittest import mock + from maxdiffusion.utils.tile_size_grid_search import _broadcast_winner + + cand = BenchResult(bq=1024, bkv=512, bkv_compute=None, status="ok", mean_ms=10.0) + with ( + mock.patch("jax.process_count", return_value=2), + mock.patch("jax.process_index", return_value=0), + mock.patch("jax.experimental.multihost_utils.broadcast_one_to_all", side_effect=lambda x, is_source: x), + ): + out = _broadcast_winner(cand, [cand]) + self.assertIs(out, cand) + if __name__ == "__main__": unittest.main() diff --git a/src/maxdiffusion/utils/tile_size_grid_search.py b/src/maxdiffusion/utils/tile_size_grid_search.py index 3b7afb90b..f1c2b0db9 100644 --- a/src/maxdiffusion/utils/tile_size_grid_search.py +++ b/src/maxdiffusion/utils/tile_size_grid_search.py @@ -46,6 +46,7 @@ "ulysses_ring", "ulysses_ring_custom", "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", "ulysses_ring_custom_bidir", }) @@ -259,6 +260,8 @@ def time_callable(fn, *, iters: int = 10, warmup: int = 2, sync=lambda x: x): import statistics import time + if iters < 1: + raise ValueError(f"iters must be >= 1, got {iters}.") t0 = time.perf_counter() sync(fn()) # call #1: compilation happens HERE, untimed compile_ms = (time.perf_counter() - t0) * 1e3 @@ -380,12 +383,13 @@ def _broadcast_winner(best: Optional[BenchResult], results: list[BenchResult]) - is_source = jax.process_index() == 0 payload = np.zeros((4,), dtype=np.int64) if is_source and best is not None: - payload[:] = (1, best.bq, best.bkv, best.bkv_compute) + bkv_cmp = -1 if best.bkv_compute is None else best.bkv_compute + payload[:] = (1, best.bq, best.bkv, bkv_cmp) payload = np.asarray(multihost_utils.broadcast_one_to_all(payload, is_source=is_source)) if payload[0] == 0: return None - winner_key = tuple(int(value) for value in payload[1:]) + winner_key = (int(payload[1]), int(payload[2]), None if int(payload[3]) < 0 else int(payload[3])) for result in results: if (result.bq, result.bkv, result.bkv_compute) == winner_key and result.status == "ok": return result diff --git a/src/maxdiffusion/utils/wan_block_benchmark.py b/src/maxdiffusion/utils/wan_block_benchmark.py index 40667f88a..348beb6af 100644 --- a/src/maxdiffusion/utils/wan_block_benchmark.py +++ b/src/maxdiffusion/utils/wan_block_benchmark.py @@ -74,9 +74,13 @@ _VAE_T, _VAE_S = 4, 8 # VAE temporal / spatial compression _PATCH_T, _PATCH_H, _PATCH_W = 1, 2, 2 _RING_VARIANTS = { + "ulysses_ring", "ulysses_ring_custom", + "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", "ulysses_ring_custom_bidir", "tokamax_ring", + "tokamax_ring_custom", "ring", }