Skip to content

perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x - #487

Open
Toshi-31 wants to merge 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:wan-2.2-perf-optimizations
Open

Toshi-31 wants to merge 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:wan-2.2-perf-optimizations

Conversation

@Toshi-31

@Toshi-31 Toshi-31 commented Sep 21, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

This PR introduces hardware-optimized custom Pallas Splash Attention kernels, eliminates redundant forward self-attention rematerialization, and optimizes Ulysses sequence parallelism for Wan 2.2 T2V 14B training on TPU v7x (tpu7x-64).

Together with the companion training recipe in tpu-recipes (which carries the tuned flash_block_sizes), these changes reduce steady-state Wan 2.2 T2V step time from 39.70 s → 12.15 s (3.27× end-to-end speedup, -69.4% step time), increasing per-device throughput from 123.24 → 402.67 TFLOP/s/device and reaching 40.43% MFU against the TPU v7x BF16 peak of 996 TFLOP/s/device (and 42.10% MFU against the measured 956.4 TFLOP/s BF16 GEMM ceiling).

Companion training recipe PR: AI-Hypercomputer/tpu-recipes#321.


Key Optimizations & Technical Details

1. 3-Slot Aliased dQ Reduction in Custom Splash Bwd Kernel (dq_reduction_steps=3)

  • Problem: The stock Splash Attention fused backward kernel (splash_mha_dkv_no_residuals) iterates over 37 KV blocks on the outer loop and Q blocks on the inner loop. Because partial dQ updates could not be accumulated in VMEM across KV steps, the kernel wrote partial gradients into an unaliased 37-deep FP32 HBM scratch tensor (f32[37, 5, 75776, 128] = 7.18 GB/layer) followed by a 37-way post-kernel HLO reduction, flooding HBM with ~14.4 GB of memory traffic per layer.
  • Fix: Implemented a 3-slot FP32 ring buffer (f32[3, 5, 75776, 128] = 0.58 GB/layer, a 12.3× memory reduction) with memory aliasing input_output_aliases={6: 0}. For KV steps i >= 3, the kernel prefetches slot i % 3 via async DMA into VMEM, accumulates dQ_accum = dQ_prev + dQ_curr in-place on the VPU, and writes back to the same HBM slot, followed by an efficient 3-way reduction.
  • Impact: Cut backward attention kernel latency by 33.7% (from 4.24 s down to 2.81 s/step), reducing end-to-end step time by -1.45 s/step. Added defensive warning when unaliased scratch allocation is requested for long sequences.

2. Inner VPU Sub-Tiling (bkv_compute_in), 8× LSE Memory Reduction & Fused Reciprocal Normalization

  • Decoupled outer HBM DMA transfer tile dimensions from inner vector compute tile dimensions (block_kv_compute_in / block_kv_dkv_compute_in).
  • In the fused backward kernel, processing 256-wide inner sub-tiles keeps intermediate P, dP, and dS tiles in vector registers while maintaining large outer HBM DMA bursts (zero redundant re-reads of Q and dO).
  • Sliced row 0 of the FP32 lse buffer before saving residuals for backward pass, cutting saved residual HBM usage by 8× (TPU forward outputs 8 duplicate sublane rows, of which backward only consumes row 0).
  • Fused attention output softmax normalization reciprocal (1 / l) directly into the accumulation pass (fuse_reciprocal=True).

3. Elimination of Duplicate Forward Self-Attention Rematerialization

  • Problem: Although ring_splash_attention defines a custom VJP, the enclosing nn.remat block did not recognize ring attention activations because residual_checkpoint_name ("attn_output") was only tagged on the non-ring branch. JAX consequently re-executed forward self-attention a second time during the backward pass (4 forward kernels vs 2 backward kernels per layer).
  • Fix: Added ad_checkpoint.checkpoint_name(..., "attn_output") inside ring_attention_kernel.py and custom_splash_attention.py, and propagated residual_checkpoint_name through tokamax_ring_custom_kernel in attention_flax.py.
  • Impact: Completely eliminated 80 duplicate splash_mha_fwd invocations per step, saving ~1.41 s/step.

4. Ulysses R=1 Pure-Parallelism Fast Path & Chunked Attention

  • With ici_context_parallelism=4 and ulysses_shards=4, num_ring_shards = 1. Added a fast-path in ring_attention_kernel.py that bypasses jax.lax.scan ring permutation buffers and dispatches directly to custom_splash.make_splash_mha.
  • Split head processing into 2 chunks of 5 heads (ulysses_attention_chunks=2) to pipeline All-to-All communication behind Pallas compute while reducing peak HBM by 6.96 GB.

5. Attention Tile Tuning & Splash Kernel Cost Estimates

  • Retuned splash tiles in the recipe (block_q=2816, block_kv_compute_in=4096 forward; block_kv_dkv=4096 backward), cutting the forward attention kernel by 34% (41.3 → 27.1 ms/call) and the fused backward by 3.5%.
  • Attached pl.CostEstimate to the splash forward and fused backward pallas_calls so XLA's latency-hiding scheduler can overlap independent collectives (such as FSDP reduce-scatter) with attention, taking the step from 13.60 → 12.15 s.

Benchmark Results on TPU v7x (tpu7x-64, 8 Hosts / 32 Chips)

Workload: Wan 2.2 T2V 14B (720x1280x81, batch size per device = 0.25, ici_fsdp=16, ici_cp=4, ici_tp=1).
MFU is reported against the TPU v7x BF16 peak of 996 TFLOP/s per device. Model FLOPs: 4,892.47 TFLOP/device/step.

Optimization Stage Step Time TFLOP/s / Device MFU (% of 996 TFLOP/s Peak)
Unoptimized Baseline 39.70 s 123.24 12.37%
Selective Activation Remat & 1D FSDP Mesh 27.63 s 177.07 17.78%
Overlapped Ring Attention & 2048 Compute Tiles 21.82 s 224.22 22.51%
Unsegmented Attention & All-Reduce Elimination 19.37 s 252.58 25.36%
Forward Rematerialization Elimination & Chunked Ulysses 15.02 s 325.73 32.70%
Custom Splash Attention Backward Kernel 13.57 s 360.46 36.19%
Attention Tile Tuning + pl.CostEstimate (This PR) 12.15 s 402.67 40.43%

Files Changed

  • src/maxdiffusion/kernels/custom_splash_attention.py: Custom attention kernels with 3-slot aliased dQ reduction, bkv_compute_in sub-tiling, 8× lse memory reduction, and pl.CostEstimate on the forward and fused backward pallas_calls.
  • src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py: Ulysses $R=1$ fast path, custom VJP support, and preserved architectural documentation for fixed-m attention.
  • src/maxdiffusion/models/attention_flax.py: attn_output checkpoint tagging and custom flash block size propagation through tokamax_ring_custom_kernel.
  • src/maxdiffusion/max_utils.py: Block size parsing for custom splash kernels (retaining None fallback for block_kv_compute_in), and optimized multithreaded download_blobs.
  • src/maxdiffusion/tests/custom_splash_backward_test.py: Unit tests verifying backward gradient numerical accuracy against reference attention, and asserting that remat saves out and lse residuals without recomputing them.

Wan 2.1 14B Inference Validation on Cloud TPU v7x-8

Workload: Wan 2.1 14B T2V (720x1280, 81 frames, 50 inference steps, ulysses_shards=2, ici_data=2, ici_cp=4), Cloud TPU v7x-8.

Metric main wan-2.2-perf-optimizations
Denoise Total (50 steps) 144.8 s 144.8 s
Total Inference Time 148.3 s 148.4 s
Compilation Time 122.8 s 123.7 s
VAE Decode (TPU) 1.4 s 1.4 s
Conditioning 0.3 s 0.3 s

@Toshi-31 Toshi-31 changed the title perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU v7x [Depends on #470] perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x [Depends on #470] Sep 21, 2026
@Toshi-31
Toshi-31 marked this pull request as ready for review September 21, 2026 09:26
@Toshi-31
Toshi-31 requested a review from entrpn as a code owner September 21, 2026 09:26
@Toshi-31
Toshi-31 marked this pull request as draft September 21, 2026 09:26

@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 support for WAN 2.2 dual-expert training and checkpointing, including the implementation of custom Pallas backward kernels for splash and ring attention. Feedback on these changes highlights critical issues: a device-to-host synchronization stall when determining the active expert on the host, a multi-host evaluation bug when calling jax.device_get on globally sharded arrays, and unsafe dimension semantics in several Pallas kernels (_flash_attention_bwd_fused, _flash_attention_dq_kernel, and _flash_attention_dkv_kernel) that could cause race conditions due to loop-carried dependencies.

Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
@Toshi-31
Toshi-31 changed the base branch from main to wan-2.2-training September 21, 2026 10:27
@Toshi-31
Toshi-31 marked this pull request as ready for review September 21, 2026 10:27
@Toshi-31 Toshi-31 changed the title perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x [Depends on #470] perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x [Stacked on #470] Sep 21, 2026
@Toshi-31 Toshi-31 changed the title perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x [Stacked on #470] test Sep 21, 2026
@Toshi-31 Toshi-31 changed the title test perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x [Stacked on #470] Sep 21, 2026
@Toshi-31
Toshi-31 changed the base branch from wan-2.2-training to main September 23, 2026 05:42
@Toshi-31
Toshi-31 force-pushed the wan-2.2-perf-optimizations branch 3 times, most recently from d16343e to 7f2d7c7 Compare September 23, 2026 06:42
@Toshi-31
Toshi-31 force-pushed the wan-2.2-perf-optimizations branch 3 times, most recently from b958271 to 2213b67 Compare September 23, 2026 08:00
Comment thread src/maxdiffusion/models/attention_flax.py Outdated
Comment thread src/maxdiffusion/max_utils.py Outdated
@prishajain1

Copy link
Copy Markdown
Collaborator

Your changes will affect the inference path as well, pls run wan inference and check that no latency regressions happen between main and this branch. Also add this data in PR description

Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
@Toshi-31
Toshi-31 force-pushed the wan-2.2-perf-optimizations branch 2 times, most recently from 9f7e5b3 to fd64f2c Compare September 26, 2026 13:43
@Toshi-31

Copy link
Copy Markdown
Collaborator Author

Your changes will affect the inference path as well, pls run wan inference and check that no latency regressions happen between main and this branch. Also add this data in PR description

Done

@Perseus14 Perseus14 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

A few general items:

  1. The branch has merge conflicts with main. Please rebase now that #470 is merged.
  2. Please update the description. It still says the base branch is wan-2.2-training, and section 3 says checkpoint names were added in attention_flax.py, but they're actually added inside the kernel files.
  3. The download_blobs change and the test fix in input_pipeline_interface_test.py aren't mentioned in the description. Please add a line about them.

Comment thread src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py Outdated
Comment thread src/maxdiffusion/max_utils.py Outdated
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py
Comment thread src/maxdiffusion/tests/custom_splash_backward_test.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
@Toshi-31
Toshi-31 force-pushed the wan-2.2-perf-optimizations branch from fd64f2c to b5405a7 Compare September 30, 2026 13:28
…ses parallelism for TPU v7x

- Implement custom Pallas splash & ring attention forward and fused backward kernels
- Optimize rematerialization policy with custom checkpointing rules
- Implement Ulysses sequence parallelism with 2-phase all-to-all communications
- Eliminate device-to-host sync stall via deterministic host-side expert routing
- Support multi-host evaluation metric gathering with multihost_utils.process_allgather
- Specify sequential dimension semantics in Pallas backward kernels for safe accumulation
- Attach pl.CostEstimate to the splash forward and fused backward pallas_calls so XLA's
  latency-hiding scheduler can overlap collectives (FSDP reduce-scatter) with attention
- Address PR review comments: slice row-0 LSE to cut saved residual memory by 8x,
  warn on unaliased dQ scratch buffer allocation, document inference-only paths,
  pass residual_checkpoint_name in tokamax_ring_custom, revert default block_kv_compute_in
  to None in max_utils, restore architectural comments in ring attention, and add unit
  test for saved remat residuals.
@Toshi-31
Toshi-31 force-pushed the wan-2.2-perf-optimizations branch from b5405a7 to 069e7fb Compare October 1, 2026 05:48
@Toshi-31 Toshi-31 changed the title perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x [Stacked on #470] perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x Oct 1, 2026
@Toshi-31

Toshi-31 commented Oct 1, 2026

Copy link
Copy Markdown
Collaborator Author

A few general items:

  1. The branch has merge conflicts with main. Please rebase now that Wan 2.2 training #470 is merged.
  2. Please update the description. It still says the base branch is wan-2.2-training, and section 3 says checkpoint names were added in attention_flax.py, but they're actually added inside the kernel files.
  3. The download_blobs change and the test fix in input_pipeline_interface_test.py aren't mentioned in the description. Please add a line about them.

Done!

@Toshi-31
Toshi-31 requested a review from Perseus14 October 1, 2026 06:47
_FIXED_M_RING_SAFE_BOUND = _FIXED_M_SAFE_BOUND / 2.0


def _attention_cost_estimate(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Nice!

@Perseus14 Perseus14 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Few minor comments which can be a follow up.

LGTM!

dq_shape,
jax.ShapeDtypeStruct((num_kv_heads, active_kv_len, head_dim_qk), k.dtype),
jax.ShapeDtypeStruct((num_kv_heads, active_kv_len, head_dim_v), v.dtype),
jax.ShapeDtypeStruct((bq, head_dim_qk), jnp.float32),

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

(Optional / follow-up) The VMEM scratch buffers here (dq_scratch_ref, dk_scratch_ref, dv_scratch_ref on lines 1325–1327, and similarly in the forward/unfused kernels) are declared in out_shapes with lambda *_: (0, 0) and then discarded after pallas_call. While Pallas keeps (0, 0) blocks in VMEM during the grid loop, putting them in out_shapes still allocates HBM for them and writes them back to HBM at the end of the kernel.

No need to change this now, but in a follow-up you can pass scratch_shapes=[pltpu.VMEM((bq, head_dim_qk), jnp.float32), ...] to PrefetchScalarGridSpec (like splash_attention_kernel.py does) so they stay purely in VMEM.

grid_width = (actual_kv_seq_len + bkv - 1) // bkv
grid_height = (actual_q_seq_len + bq - 1) // bq
if save_residuals:
out_shapes.append(jax.ShapeDtypeStruct((num_q_heads, NUM_SUBLANES, active_q_len), jnp.float32))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

When save_residuals=True, cost_estimate below (line 648) doesn't include the num_q_heads * NUM_SUBLANES * active_q_len * 4 bytes written for this lse output buffer.

@@ -678,10 +808,10 @@ def v_index_map(h, i, j, *_):
vmem_limit_bytes=vmem_limit_bytes,
),
out_shape=out_shapes,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

(Optional/Follow-up) Consider also passing cost_estimate to _splash_attention_forward_ring here (and the unfused dq/dkv calls on lines 1524 and 1615) so multi-hop ring (R > 1) and unfused backward runs get the same XLA overlap scheduling.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants