perf(wan2.2): Optimize attention kernels, rematerialization, and Ulysses parallelism for WAN 2.2 Training on TPU 7x - #487
Conversation
There was a problem hiding this comment.
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.
d16343e to
7f2d7c7
Compare
b958271 to
2213b67
Compare
|
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 |
9f7e5b3 to
fd64f2c
Compare
Done |
Perseus14
left a comment
There was a problem hiding this comment.
A few general items:
- The branch has merge conflicts with
main. Please rebase now that #470 is merged. - Please update the description. It still says the base branch is
wan-2.2-training, and section 3 says checkpoint names were added inattention_flax.py, but they're actually added inside the kernel files. - The
download_blobschange and the test fix ininput_pipeline_interface_test.pyaren't mentioned in the description. Please add a line about them.
fd64f2c to
b5405a7
Compare
…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.
b5405a7 to
069e7fb
Compare
Done! |
| _FIXED_M_RING_SAFE_BOUND = _FIXED_M_SAFE_BOUND / 2.0 | ||
|
|
||
|
|
||
| def _attention_cost_estimate( |
Perseus14
left a comment
There was a problem hiding this comment.
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), |
There was a problem hiding this comment.
(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)) |
There was a problem hiding this comment.
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, | |||
There was a problem hiding this comment.
(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.
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 tunedflash_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)splash_mha_dkv_no_residuals) iterates over 37 KV blocks on the outer loop and Q blocks on the inner loop. Because partialdQupdates 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.f32[3, 5, 75776, 128]= 0.58 GB/layer, a 12.3× memory reduction) with memory aliasinginput_output_aliases={6: 0}. For KV stepsi >= 3, the kernel prefetches sloti % 3via async DMA into VMEM, accumulatesdQ_accum = dQ_prev + dQ_currin-place on the VPU, and writes back to the same HBM slot, followed by an efficient 3-way reduction.2. Inner VPU Sub-Tiling (
bkv_compute_in), 8× LSE Memory Reduction & Fused Reciprocal Normalizationblock_kv_compute_in/block_kv_dkv_compute_in).P,dP, anddStiles in vector registers while maintaining large outer HBM DMA bursts (zero redundant re-reads ofQanddO).lsebuffer 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).1 / l) directly into the accumulation pass (fuse_reciprocal=True).3. Elimination of Duplicate Forward Self-Attention Rematerialization
ring_splash_attentiondefines a custom VJP, the enclosingnn.rematblock did not recognize ring attention activations becauseresidual_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).ad_checkpoint.checkpoint_name(..., "attn_output")insidering_attention_kernel.pyandcustom_splash_attention.py, and propagatedresidual_checkpoint_namethroughtokamax_ring_custom_kernelinattention_flax.py.splash_mha_fwdinvocations per step, saving ~1.41 s/step.4. Ulysses R=1 Pure-Parallelism Fast Path & Chunked Attention
ici_context_parallelism=4andulysses_shards=4,num_ring_shards = 1. Added a fast-path inring_attention_kernel.pythat bypassesjax.lax.scanring permutation buffers and dispatches directly tocustom_splash.make_splash_mha.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
block_q=2816,block_kv_compute_in=4096forward;block_kv_dkv=4096backward), cutting the forward attention kernel by 34% (41.3 → 27.1 ms/call) and the fused backward by 3.5%.pl.CostEstimateto the splash forward and fused backwardpallas_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.
39.70 s123.2412.37%27.63 s177.0717.78%21.82 s224.2222.51%19.37 s252.5825.36%15.02 s325.7332.70%13.57 s360.4636.19%pl.CostEstimate(This PR)12.15 s402.6740.43%gs://wan2-2-training/toshipahadia-wan22-aishared-outputs/tp-wan22-v7a-e4tilescost-0930-1400/Files Changed
src/maxdiffusion/kernels/custom_splash_attention.py: Custom attention kernels with 3-slot aliaseddQreduction,bkv_compute_insub-tiling, 8×lsememory reduction, andpl.CostEstimateon the forward and fused backwardpallas_calls.src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py: Ulyssessrc/maxdiffusion/models/attention_flax.py:attn_outputcheckpoint tagging and custom flash block size propagation throughtokamax_ring_custom_kernel.src/maxdiffusion/max_utils.py: Block size parsing for custom splash kernels (retainingNonefallback forblock_kv_compute_in), and optimized multithreadeddownload_blobs.src/maxdiffusion/tests/custom_splash_backward_test.py: Unit tests verifying backward gradient numerical accuracy against reference attention, and asserting that remat savesoutandlseresiduals without recomputing them.Wan 2.1 14B Inference Validation on Cloud TPU v7x-8
mainwan-2.2-perf-optimizations144.8 s144.8 s148.3 s148.4 s122.8 s123.7 s1.4 s1.4 s0.3 s0.3 s