From a9c36e68833c5a3c4ce0577abb8bf00e82ca1620 Mon Sep 17 00:00:00 2001 From: Toshi Pahadia Date: Mon, 28 Sep 2026 10:45:22 +0530 Subject: [PATCH] docs: Add Wan 2.2 text2vid training documentation to README --- README.md | 71 ++++++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 70 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index 4d208830d..911d859b7 100755 --- a/README.md +++ b/README.md @@ -17,6 +17,7 @@ [![Unit Tests](https://github.com/AI-Hypercomputer/maxdiffusion/actions/workflows/UnitTests.yml/badge.svg)](https://github.com/AI-Hypercomputer/maxdiffusion/actions/workflows/UnitTests.yml) # What's new? +- **`2026/09/23`**: Wan2.2 text2vid dual-expert training is now supported. - **`2026/08/28`**: Flux2.Klein text to image and image editing (w/ KV Cache) is now supported. - **`2026/07/14`**: Automatic attention tile-size (`block_q`/`block_kv`) search for Wan is now supported. - **`2026/06/26`**: 2D ring (USP) attention with a custom splash kernel is now supported for Wan (`ulysses_ring_custom`), splitting context parallelism into an intra-chip Ulysses axis and a cross-chip ring axis. @@ -59,7 +60,7 @@ MaxDiffusion supports * LTX-Video text2vid, img2vid (inference). * LTX-2 Video text2vid (inference). * Wan2.1 text2vid (training and inference). -* Wan2.2 text2vid (inference). +* Wan2.2 text2vid (training and inference). **Note on GPU Support:** GPU support is not actively maintained, but contributions are welcome @@ -74,6 +75,7 @@ MaxDiffusion supports - [NVIDIA DGX Spark](#nvidia-dgx-spark) - [Training](#training) - [Wan2.1](#wan-21-training) + - [Wan2.2](#wan-22-training) - [Flux](#flux-training) - [SDXL](#stable-diffusion-xl-training) - [SD 2 base](#stable-diffusion-2-base-training) @@ -398,6 +400,73 @@ After installation completes, run the training script. --max-restarts=0 ``` + ## Wan 2.2 Training + + Wan 2.2 introduces a **dual-expert DiT architecture** (High-Noise Expert and Low-Noise Expert, ~27B total parameters). MaxDiffusion supports joint dual-expert training where samples are dynamically routed to the appropriate expert based on `boundary_ratio` (default `0.875`) using the Flow Match time shift schedule. + + ### Single Host Training + + Wan 2.2 uses the same dataset format (TFRecords or local/GCS dataset directories) prepared in the [Wan 2.1 dataset preparation step](#dataset-preparation). + + Run single-host training using `train_wan.py` with `base_wan_27b.yml`: + + ```bash + python src/maxdiffusion/train_wan.py \ + src/maxdiffusion/configs/base_wan_27b.yml \ + run_name=${RUN_NAME} \ + output_dir=${OUTPUT_DIR} \ + train_data_dir=${DATASET_DIR} \ + dataset_save_location=${SAVE_DATASET_DIR} \ + boundary_ratio=0.875 \ + weights_dtype=bfloat16 \ + activations_dtype=bfloat16 \ + per_device_batch_size=0.25 \ + ici_fsdp_parallelism=4 \ + ici_data_parallelism=1 \ + remat_policy='HIDDEN_STATE_WITH_OFFLOAD' \ + max_train_steps=1000 \ + checkpoint_every=500 \ + save_final_checkpoint=True + ``` + + ### Multi-Host Training with XPK + + For large-scale multi-host training across TPU pods or clusters (e.g. v5p, v6e, v7x). The following example is configured for a 128-chip slice (such as `v5p-128`, `v6e-128`, or `tpu7x-4x4x4`): + + ```bash + python3 ~/xpk/xpk.py workload create \ + --cluster=$CLUSTER_NAME \ + --project=$PROJECT \ + --zone=$ZONE \ + --device-type=$DEVICE_TYPE \ + --num-slices=1 \ + --workload=${RUN_NAME} \ + --command=" \ + python3 src/maxdiffusion/train_wan.py \ + src/maxdiffusion/configs/base_wan_27b.yml \ + run_name=${RUN_NAME} \ + output_dir=${OUTPUT_DIR} \ + train_data_dir=${DATASET_DIR} \ + dataset_save_location=${SAVE_DATASET_DIR} \ + boundary_ratio=0.875 \ + weights_dtype=bfloat16 \ + activations_dtype=bfloat16 \ + per_device_batch_size=1 \ + ici_fsdp_parallelism=4 \ + ici_data_parallelism=32 \ + remat_policy='HIDDEN_STATE_WITH_OFFLOAD' \ + max_train_steps=5000 \ + checkpoint_every=1000 \ + save_final_checkpoint=True" \ + --base-docker-image=${IMAGE_DIR} \ + --priority=medium \ + --max-restarts=0 + ``` + + ### Checkpointing & Export + + Wan 2.2 training checkpoints are managed with Orbax via `WanCheckpointer2_2`. Checkpoints store optimizer states and weights for both experts (`low_noise_transformer_state` and `high_noise_transformer_state`) alongside model configurations. MaxDiffusion inference pipelines (`WanPipeline2_2`) can load directly from these Orbax checkpoints for downstream sampling. + ## Flux Training Expected results on 1024 x 1024 images with flash attention and bfloat16: