Skip to content

Wan 2.2 training - #470

Merged
copybara-service[bot] merged 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:wan-2.2-training
Sep 22, 2026
Merged

copybara-service[bot] merged 1 commit into
AI-Hypercomputer:mainfrom
Toshi-31:wan-2.2-training

Conversation

@Toshi-31

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

Copy link
Copy Markdown
Collaborator

Title:
Add Wan 2.2 Training Pipeline with Joint Timestep Routing

Description:
This PR introduces full training support for the Wan 2.2 models into MaxDiffusion.

1. Joint Timestep Routing
This implementation utilizes a unified, joint-pipeline training strategy. During each training step, the pipeline dynamically samples a timestep and routes the forward pass to either the high-noise or low-noise transformer based on the configured boundary ratio:

is_high_noise = jax.random.uniform(cond_rng) > config.boundary_ratio
This ensures a seamless and efficient training execution graph without needing separate high/low training loops.

2. Pipeline Inheritance & Code Reuse
The new wan_trainer_2_2.py relies directly on the WanPipeline2_2 class to bootstrap the models and load checkpoints. Because WanPipeline2_2 cleanly inherits from the base WanPipeline, the training loop seamlessly reuses all existing weight conversion and fast-loading logic that was previously established for inference.

The loss graphs were plotted for around 260 steps, and show a clear downward trend.

For graphs and other artefacts: https://docs.google.com/document/d/1svzC8cVZxb2XxyypeFcoJig13_1ptu6wYIe5QwnC_Lo/edit?usp=sharing

Step time: Around 39.7 seconds per device.

Through systematic roofline analysis and iterative kernel, parallelism, rematerialization, and collective optimizations (from V2 through V7), steady-state step time is reduced from ~39.7 s → 13.57 s (2.92× end-to-end speedup, -65.8% step time reduction), increasing per-device throughput from 131.6 → 384.12 TFLOP/s/device and driving Model FLOP Utilization (MFU) from ~7.5% → 40.18%.

Note on Code Changes: This PR introduces the training recipe, parallelism mesh, and XLA/Pallas config updates. All underlying kernel and model code changes (custom Pallas splash attention backward kernel with aliased dQ reduction, remat checkpoint tagging, and block-size wiring) will be raised in a companion PR in AI-Hypercomputer/maxdiffusion.

Optimization Summary & Speedup

Optimization Description Impact / Speedup
1. Selective Activation Remat & Host Overhead Removal (V2–V3) Switched from full remat (remat_policy=FULL) to selective checkpointing (MATMUL_WITHOUT_BATCH), keeping linear and attention activations in HBM (~7.7 GB across 40 layers) while recomputing only elementwise ops (RMSNorm, RoPE). Removed hot-path 14B parameter PyTree scans and host-device step counter synchronizations. Eliminated 640 redundant forward GEMMs and attention calls per step. Reduced step time from 39.7 s → 28.50 s (-28.2%, +51.35 TFLOP/s/device).
2. Pure 1D FSDP Parallelism & SparseCore Offload (V4) Reconfigured hybrid parallelism (dp=2, fsdp=8) to pure 1D FSDP (dp=1, fsdp=16, cp=4, tp=1), offloading 3D All-Gather and ND Reduce-Scatter torus collective transfers to the 32 SparseCore coprocessors while bounding collective concurrency. 100% eliminated the synchronous 27B parameter DP All-Reduce phase (saving 6.05 s of stalls). Halved model weights and Adam states in HBM (saving ~8.5 GB), reducing step time to 27.63 s.
3. Overlapped Ring Attention & Sequence Sharding (V4.2) Replaced global sequence all-gathers (attention=tokamax) with Overlapped Ring Attention (attention=tokamax_ring), keeping Key and Value states sharded locally at 18,900 tokens instead of replicating the full 75,600 sequence, and overlapped ring permutations behind compute. Resolved compilation HBM OOM (required 112.7 GB vs 94.7 GB available), eliminated 160 synchronous KV all-gathers (~3.5 s) and the duplicate splash_mha_dq pass. Dropped step time to 22.38 s (232.85 TFLOP/s).
4. High-Intensity FlashAttention Tile Scaling & Layer Scheduling (V4.1–V4.4) Scaled FlashAttention compute tiles from 1024 to 2048 (block_q=2048, block_kv=2048, block_kv_compute=1024) with native VPU base-2 exponentiation (use_base2_exp=True) to saturate TPU v7x 256×256 systolic arrays, and re-enabled Latency-Hiding Layer Scheduling (LHS) with a strict 105% memory limit. Halved inner-loop stepping overhead, doubled MXU burst lengths, and overlapped inter-layer FSDP All-Gathers/Reduce-Scatters with compute. Decreased step time to 21.82 s (238.8 TFLOP/s).
5. Global All-Reduce Elimination & Unsegmented Attention Backward (V5) Disabled Optax global L2 gradient norm clipping (opt_enable_grad_global_norm_clipping=False) to avoid blocking cross-host barrier synchronization on every step. Disabled padding token masking (mask_padding_tokens=False) to dispatch to splash_mha_dkv_no_residuals, eliminating vector unit boundary mask evaluations. Dropped collective All-Reduce overhead from 6.0 s down to 31 ms. Reduced backward attention latency from 7.40 s to 6.59 s, driving step time down to 19.37 s (269.1 TFLOP/s, 14.95% MFU).
6. Duplicate Forward Attention Remat Elimination & Chunked Ulysses (V6) Identified and eliminated duplicate forward attention recompute in the backward pass by tagging out and lse with ad_checkpoint.checkpoint_name("attn_output") in ring_attention_kernel.py and saving it in remat_policy=CUSTOM. Chunked Ulysses context parallelism into 2 chunks of 5 heads (ulysses_attention_chunks=2) with async All-to-All pipelining. Eliminated 80 duplicate splash_mha_fwd calls per step (~3.96 s savings), reduced peak HBM by 6.96 GB, and lowered step time from 19.37 s → 15.02 s (347.18 TFLOP/s, 36.31% MFU).
7. Custom Splash Attention Backward Kernel: Aliased dQ & VPU Sub-Tiling (V7) Replaced the 37-deep unaliased FP32 dQ scratch tensor (7.18 GB/layer) with a 3-slot ring buffer (dq_reduction_steps=3, input_output_aliases={6: 0}) with async DMA prefetching and in-VMEM accumulation, added inner VPU sub-tiling (bkv_compute_in=256), and fused attention output reciprocal normalization (fuse_reciprocal=True). Reduced dQ scratch HBM traffic by 12.3× (7.18 GB → 0.58 GB/layer) and cut backward attention kernel time by 33.7% (4.24 s → 2.81 s/step). Reduced step time from 15.02 s → 13.57 s (-1.45 s, 384.12 TFLOP/s, 40.18% MFU).

Key Progression Milestones

Baseline (V1): 39.70 s / step | 131.6 TFLOP/s | ~7.5% MFU
V3 (Selective Remat): 28.50 s / step | 182.95 TFLOP/s | 10.2% MFU
V5 (Unsegmented Attention + No All-Reduce): 19.37 s / step | 269.1 TFLOP/s | 14.95% MFU
V6 (Remat Fix + Chunked Ulysses): 15.02 s / step | 347.18 TFLOP/s | 36.31% MFU
V7 (Custom Splash Bwd Kernel): 13.57 s / step | 384.12 TFLOP/s | 40.18% MFU (2.92× total speedup)

@Toshi-31
Toshi-31 requested a review from entrpn as a code owner September 1, 2026 10:32

@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 training support for the Wan 2.2 model, including a new trainer (WanTrainer2_2), training script, configuration updates, and import smoke tests. It also switches GCS model downloads to use gcloud storage to prevent SSL segfaults, and optimizes disk usage during shard conversion by deleting shards after processing. The code review identified several critical issues and improvement opportunities: a logical error in the dataset validation check in WanTrainer2_2 where 'and' was used instead of 'or'; a JAX purity violation caused by mutating the input dictionary in-place inside the JIT-compiled train_step_2_2; incorrect evaluation routing for mixed-timestep batches, which should be resolved using jnp.where instead of checking only the first timestep; potential out-of-memory errors from downloading large models to /dev/shm instead of /tmp; performance overhead from recreating a ThreadPoolExecutor inside a loop during shard conversion; and the use of a mutable default argument in training_loop_2_2.

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/trainers/wan_trainer_2_2.py Outdated
Comment thread src/maxdiffusion/pyconfig.py Outdated
Comment thread src/maxdiffusion/models/wan/wan_utils.py Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
@prishajain1

Copy link
Copy Markdown
Collaborator

Can you please add the current timestep in the PR description

Comment thread src/maxdiffusion/configs/base_wan_27b.yml Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
Comment thread setup.sh
@prishajain1

Copy link
Copy Markdown
Collaborator

Please add relevant unittests

Comment thread src/maxdiffusion/pyconfig.py Outdated
Comment thread src/maxdiffusion/models/wan/wan_utils.py Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
@Toshi-31

Toshi-31 commented Sep 2, 2026

Copy link
Copy Markdown
Collaborator Author

Can you please add the current timestep in the PR description

Done!

@Toshi-31

Toshi-31 commented Sep 2, 2026

Copy link
Copy Markdown
Collaborator Author

Please add relevant unittests

Done!

@Toshi-31
Toshi-31 force-pushed the wan-2.2-training branch 3 times, most recently from ff2a971 to 4ffe7dc Compare September 2, 2026 11:23
Comment thread src/maxdiffusion/configs/training_wan_27b.yml Outdated
@Toshi-31
Toshi-31 force-pushed the wan-2.2-training branch 5 times, most recently from 8edcb2d to 1a6b32e Compare September 3, 2026 09:33
@Toshi-31
Toshi-31 requested a review from Perseus14 September 3, 2026 09:39
@Toshi-31
Toshi-31 force-pushed the wan-2.2-training branch 2 times, most recently from a1f349a to 385c2d9 Compare September 4, 2026 10:07
Comment thread src/maxdiffusion/configs/base_wan_27b.yml
@Toshi-31
Toshi-31 force-pushed the wan-2.2-training branch 2 times, most recently from 65bbdf5 to 97ebfc3 Compare September 9, 2026 05:42
@Toshi-31
Toshi-31 requested a review from entrpn September 9, 2026 06:38
Comment thread src/maxdiffusion/train_wan_2_2.py Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py
Comment thread src/maxdiffusion/train_utils.py Outdated
Comment thread src/maxdiffusion/checkpointing/wan_checkpointer_2_2.py
@Toshi-31
Toshi-31 force-pushed the wan-2.2-training branch 3 times, most recently from 0534166 to 6cc455a Compare September 18, 2026 11:11
prishajain1
prishajain1 previously approved these changes Sep 18, 2026
Perseus14
Perseus14 previously approved these changes Sep 21, 2026
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
Comment thread src/maxdiffusion/pipelines/wan/wan_pipeline.py Outdated
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py
Comment thread src/maxdiffusion/trainers/wan_trainer_2_2.py Outdated
…expert routing

- Implement WanTrainer2_2 dual-expert joint training pipeline with probabilistic batch routing based on boundary_ratio
- Unify Wan 2.2 training CLI entrypoint into train_wan.py based on model_name
- Align training timestep domain with inference shifted timesteps using inverse time shift boundary partitioning
- Pass config.flow_shift to FlaxFlowMatchScheduler for correct u_boundary partitioning and timestep scaling across resolutions
- Eliminate device-to-host sync stalls by removing np.isnan from main training loop and filtering metric keys host-side based on is_high_noise
- Make WanCheckpointer2_2 checkpoint and optimizer loading robust to both dictionary mapping and attribute access without private mock hooks
- Wire wan_config_high when restoring high-noise transformer from checkpoint in wan_pipeline.py
- Pass (state_high, state_low) functionally as operands to jax.lax.cond without outer closures
- Omit buffer donation on conditional train step to avoid XLA invalidation hazards on untouched states
- Implement batched conditional evaluation in eval_step_2_2 with exact per-sample routing and constant HLO graph complexity
- Strip redundant process_allgather on replicated evaluation metrics in eval_2_2
- Save and restore both low_noise_transformer and high_noise_transformer configurations and states in WanCheckpointer2_2
- Track active expert step counts on host to log active learning rates without pipeline stalls
- Document checkpoint_save_location as local staging cache with disk capacity considerations in base_wan_27b.yml
- Add comprehensive test suite covering training steps, eval steps, checkpointing, and resume equivalence
@copybara-service
copybara-service Bot merged commit 3b6ee39 into AI-Hypercomputer:main Sep 22, 2026
53 of 54 checks passed
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.

6 participants