SIGN IN SIGN UP

feat(wan2.2): Implement Wan 2.2 joint training pipeline with dynamic 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
T
Toshi Pahadia committed
758e6896ebc6286d61763fb692da81ec514ea7e4
Parent: 1bc5481