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", }