Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions end_to_end/tpu/run_wan_stack_tests.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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[@]}" "$@"
22 changes: 20 additions & 2 deletions src/maxdiffusion/configs/base_wan_27b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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


#
Expand All @@ -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:
Expand Down
3 changes: 2 additions & 1 deletion src/maxdiffusion/configs/ltx2_3_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 2 additions & 1 deletion src/maxdiffusion/configs/ltx2_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
122 changes: 122 additions & 0 deletions src/maxdiffusion/kernels/fused_producers.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading