Skip to content

feat(attention): 2D Ulysses+Ring custom attention with fixed-m (pre-a2a norm reduction, unpadded ring K/V) - #478

Open
Perseus14 wants to merge 1 commit into
feat/fixed-m-kernelfrom
feat/ring-attention
Open

Perseus14 wants to merge 1 commit into
feat/fixed-m-kernelfrom
feat/ring-attention

Conversation

@Perseus14

@Perseus14 Perseus14 commented Sep 13, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Adds fixed-m to the 2D Ulysses + Ring custom-kernel path (ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block), plus the producer, guard and layout fixes it needed.

What's in it

  • Pre-all-to-all norm reduction: Q row norms, K max norms and the V max are reduced in one pmax over (ulysses, ring) on the pre-all-to-all shards, so they overlap the collective.
    • v_ok comes from that pmax, so it is identical on every rank.
    • all_fixed_global is mesh-uniform and drives the accumulate-vs-LSE lax.cond.
    • Blast radius: one ineligible head or Q-block on any device sends every rank to _lse_scan for that layer.
  • Unpadded ring K/V: passed unpadded when the ring-shard KV length is 8-aligned (kv_pad_size=1). Q·log2(e) is folded in before the all-to-all.
  • Virtual K-centering: works for R == 1 and R > 1. k_mean is computed as pmean over (ulysses, ring) and passed to the kernel; K is never rewritten.
    • Off by default on ring. Measured at the stack tip (Wan 2.2 720p/81f/40 steps, warm AOT, DVFS unpinned), it is never faster and it changes the bf16 output:

      TPU Ring setup Off On Delta
      v6e-8 R == 1 127.2 s 127.2 s neutral
      v6e-8 R = 2 128.6 s 129.1 s +0.4%
      tpu7x-8 U=2/R=2 107.66 s 108.04 s +0.35%
    • use_k_centering=True forces it on.

  • Other changes:
    • XLA fused RMSNorm+RoPE producer (fused_producers.py), bit-exact with nnx.RMSNorm. It is jax.jit-wrapped (heads/dim/eps static) so eager callers also get fusion — this is what fixed the wan_vace_transformer_test OOM on the v4-8 CI runner (the test calls a full 75,600-token block eagerly; unfused, the producer materialised ~1.5 GB of extra intermediates). Inlined, i.e. a no-op, under the pipeline's outer jit.
    • GQA plumbing, plus NotImplementedError on paths that don't support GQA.
    • ulysses_shards guards, dot-fallback split_head_dim reshape fix, and tile-search hardening.

Tests

  • New: U=2/R=2 integration against a dense f32 reference (per_q_block × centering × GQA, ragged KV, sink head forcing _lse_scan, V outlier forcing v_ok=False, and B=2 _accumulate_scan vmap + d=64 k_mean head-dim padding), plus attention_config_guards_test (20), custom_splash_unpadded_test (12), dot_fallback_layout_test (5), fused_producers_test (6), tile_size_grid_search_test (16).
  • CI vs Local VM: 18 additional TPU kernel/multi-device tests skip under GITHUB_ACTIONS=true (36 of 115 cumulative); end_to_end/tpu/run_wan_stack_tests.sh runs all of them.
  • Results (TPU v6e-8): 115 passed (run_wan_stack_tests.sh, 284.2 s).

Stack: #477 → #478 → #479 → #488 → #491

@github-actions

Copy link
Copy Markdown

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request introduces optimized fused producers (fused_ln_adaln and fused_rmsnorm_rope) and implements exact Fixed-m support with Global Virtual K-Centering for Ulysses and Ring attention. It also adds support for Grouped Query Attention (GQA) across these attention kernels and includes comprehensive unit tests. The review feedback suggests moving an inline import of fused_rmsnorm_rope in attention_flax.py to the top of the file to avoid performance overhead in a hot path, and simplifying a double-negation conditional expression to improve code readability.

Comment thread src/maxdiffusion/models/attention_flax.py Outdated
Comment thread src/maxdiffusion/models/attention_flax.py Outdated
@Perseus14
Perseus14 force-pushed the feat/ring-attention branch 2 times, most recently from 4ace67b to 0c191db Compare September 13, 2026 15:40
@Perseus14
Perseus14 force-pushed the feat/ring-attention branch 4 times, most recently from ea0148a to 837cebe Compare September 13, 2026 18:53
@Perseus14
Perseus14 force-pushed the feat/ring-attention branch 3 times, most recently from b1e2b13 to e6db513 Compare September 13, 2026 19:51
@Perseus14
Perseus14 force-pushed the feat/ring-attention branch 2 times, most recently from b749c64 to 3268a70 Compare September 14, 2026 05:37
@Perseus14
Perseus14 added this pull request to stack #486 September 17, 2026 19:10
@Perseus14
Perseus14 force-pushed the feat/ring-attention branch 2 times, most recently from 0548265 to dbb51e1 Compare September 23, 2026 20:01
@eltsai

eltsai commented Sep 23, 2026

Copy link
Copy Markdown
Collaborator

for the e2e verification, did you check the output video quality as well?

Comment thread src/maxdiffusion/configs/base_wan_27b.yml Outdated
@Perseus14

Copy link
Copy Markdown
Collaborator Author

Yes, I verified manually as well as both pixel-level fidelity (PSNR / SSIM against exact FP32 online softmax) and full VBench quality metrics on Wan 2.2 T2V-A14B (1280x720, 81 frames, 40 steps). I have updated the results in #491

eltsai
eltsai previously approved these changes Oct 2, 2026
…2a norm reduction, unpadded ring K/V)

Adds fixed-m to the 2D Ulysses + Ring custom-kernel path
(`_ulysses_ring_custom_attention`, attention=ulysses_ring_custom_fixed_m and
the new ulysses_ring_custom_fixed_m_per_q_block), plus the supporting
producer, guard and layout fixes it needed.

Ring fixed-m (R > 1)
- `_ring_fixed_m_norms_pre_a2a` computes every fixed-m input on the PRE-a2a
  shards (each device holds all heads x 1/(U*R) of the sequence): Q row norms,
  K max norms and the V max are reduced in ONE pmax over (ulysses, ring), then
  sliced to the heads this Ulysses rank owns after the a2a. Keeping these off
  the post-a2a layout lets them overlap the all-to-all instead of sitting
  between it and the `lax.cond`. With per_q_block=True a small ulysses-axis
  all_to_all of the Q row norms is also issued.
- V safety: `v_ok = (max|V|^2 from that pmax <= 256^2) & dtype_safe`. It is
  derived from the pmax of V norms, so it is bit-identical on every
  (ulysses, ring) rank with no further collective. (The ring kernel's own pmin,
  which reduces global-bound eligibility, only runs when the caller omits
  `all_fixed_global`; the production caller always passes it.)
- `all_fixed_global = all(qn * kn <= bound(N_total)) & v_ok`, also reduced over
  batch, is passed to the ring kernel and is the top-level
  accumulate-vs-LSE `lax.cond` predicate (`pregathered_mk=True`). It is
  mesh-uniform by construction. Blast radius: one ineligible head / Q-block on
  any device sends every rank to `_lse_scan` for that layer; per_q_block=True
  only refines dispatch inside the LSE path. How often that happens in
  production is not measured.
- Only Q and K go through the optimization_barrier that stops XLA duplicating
  their producer chains into the norm reductions; V is left out so the Q/K
  a2a does not wait on V's producer.
- Ring K/V are passed unpadded when the ring-shard KV length is 8-aligned
  (`kv_pad_size=1`; the kernel slices the tail, see custom_splash_unpadded_test).
- Q * log2(e) is applied before the Ulysses a2a (exact; fuses into Q's
  producer instead of being wrapped in relayout copies after the collective).

Virtual K-centering
- Centering is virtual on both R == 1 and R > 1: k_mean goes to the kernel and
  q . k_bar is added to the fixed bound in VMEM; K is never re-written. For
  R > 1, k_mean = pmean(mean(K)) over (ulysses, ring) on the pre-a2a shards,
  and K norms are taken around that global mean.
- Controlled by `use_k_centering` ("auto" | bool, new key in
  base_wan_27b.yml, defaulted in pyconfig). "auto" (`resolve_k_centering`) is
  ON for non-ring ulysses_custom_fixed_m* and OFF on every ring path, R == 1
  included. The Wan pipelines do not pass this key to the attention layer in
  this PR, so Wan runs the layer default.
  Ring is kept OFF by default because, measured at the stack tip (Wan 2.2
  T2V-A14B, 720p/81f/40 steps, warm AOT, DVFS unpinned), centering is never faster and it changes output: v6e-8 R == 1
  127.2 s vs 127.2 s (neutral), R = 2 129.1 s vs 128.6 s (+0.4%); tpu7x-8
  U=2/R=2 108.04 s vs 107.66 s (+0.35%), U=4 neutral. Turning it on changes
  the bf16 trajectory, with no visible quality difference. use_k_centering=True
  still forces it on.

Other changes
- `fused_rmsnorm_rope` (kernels/fused_producers.py): XLA-level fused
  RMSNorm + RoPE producer for self-attention Q/K, matching nnx.RMSNorm's
  association (x * (rsqrt(var + eps) * scale)) bit-for-bit. It is
  `jax.jit`-wrapped (heads/dim/eps static): inlined under the pipeline's
  outer jit, but when a full-size block is called eagerly (as
  wan_vace_transformer_test does at 75,600 tokens) XLA fuses the chain
  instead of materialising ~1.5 GB of extra intermediates, which OOMed that
  test on the v4-8 CI runner.
- GQA: kv_heads plumbed through the custom Ulysses / ring paths; flash,
  tokamax, non-custom ulysses and ulysses_ring raise NotImplementedError on GQA.
- Guards: non-ring ulysses kernels reject a ulysses_shards they cannot honour;
  ring kernels warn once when U == CP makes the ring degenerate (R = 1) and
  suggest a valid U for the mesh and head count.
- Dot-product fallback: fixes the split_head_dim reshape that silently
  mis-read [B, H, S, D] inputs as [B, S, H*D].
- tile_size_grid_search: rejects iters < 1; `_broadcast_winner` handles
  bkv_compute=None. wan_block_benchmark / LTX2 / pyconfig know the new
  kernel names.

Tests
- ring_fixed_m_test.py (24): adds a U=2/R=2 integration class through
  `_ulysses_ring_custom_attention` vs a dense f32 reference (per_q_block x
  centering x GQA grid; ragged ring KV with kv_pad_size=1; a sink head that
  must force `_lse_scan` on all 4 devices; a single V outlier at B=2 that must
  make v_ok and all_fixed_global False on all 4 devices; and B=2 `_accumulate_scan`
  vmap + sub-128 head_dim `k_mean` padding at d=64). The branch
  predicates are checked through a probe of `_ring_fixed_m_norms_pre_a2a`, not
  the production cond itself. Needs >= 4 devices (explicit skipIf).
- custom_splash_fixed_m_test.py (32), attention_config_guards_test.py (20),
  custom_splash_unpadded_test.py (12, TPU-only), dot_fallback_layout_test.py
  (5), fused_producers_test.py (6), tile_size_grid_search_test.py (16).
- CI: 18 more tests skip in GitHub Actions (36 of 115 cumulative):
  custom_splash_unpadded_test (12), the GQA ring numerics test (1), and the
  5 U=2/R=2 integration tests. Contract, guard, dot-fallback layout,
  fused producer and grid-search tests run in CI. run_wan_stack_tests.sh
  gains this PR's five new test files.
Verified (final tree): TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 115 passed, 107 subtests passed.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants