Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 70 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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

Expand All @@ -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)
Expand Down Expand Up @@ -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
Comment thread
Toshi-31 marked this conversation as resolved.
```

### 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:
Expand Down
Loading