Repository navigation
Wan 2.2 training - #470
Wan 2.2 training#470
Conversation
There was a problem hiding this comment.
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.
|
Can you please add the current timestep in the PR description |
|
Please add relevant unittests |
Done! |
Done! |
ff2a971 to
4ffe7dc
Compare
8edcb2d to
1a6b32e
Compare
a1f349a to
385c2d9
Compare
65bbdf5 to
97ebfc3
Compare
97ebfc3 to
c95b607
Compare
0534166 to
6cc455a
Compare
6cc455a to
2b40a0c
Compare
…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
758e689
2b40a0c to
758e689
Compare
3b6ee39
into
AI-Hypercomputer:main
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
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.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.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.splash_mha_dqpass. Dropped step time to 22.38 s (232.85 TFLOP/s).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.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 tosplash_mha_dkv_no_residuals, eliminating vector unit boundary mask evaluations.outandlsewithad_checkpoint.checkpoint_name("attn_output")inring_attention_kernel.pyand saving it inremat_policy=CUSTOM. Chunked Ulysses context parallelism into 2 chunks of 5 heads (ulysses_attention_chunks=2) with async All-to-All pipelining.splash_mha_fwdcalls 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).dQscratch 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).dQscratch 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)