Skip to content
Draft
Show file tree
Hide file tree
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
134 changes: 121 additions & 13 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -115,7 +115,7 @@ Minimum requirements: Ubuntu Version 22.04 and Python 3.12.

## Getting Started

For your first time running MaxDiffusion, we provide specific [instructions](docs/getting_started/first_run.md).
For your first time running MaxDiffusion, we provide specific [instructions](docs/getting_started/first_run.md). To run multi-host jobs on GKE, see [Running MaxDiffusion with Cluster Toolkit](docs/getting_started/run_maxdiffusion_via_cluster_toolkit.md).

## NVIDIA DGX Spark

Expand All @@ -127,7 +127,7 @@ After installation completes, run the training script.

## Wan 2.1 Training

In the first part, we'll run on a single host VM to get familiar with the workflow, then run on xpk for large scale training.
In the first part, we'll run on a single host VM to get familiar with the workflow, then run on a GKE cluster with Cluster Toolkit for large scale training.

Although not required, attaching an external disk is recommended as weights take up a lot of disk space. [Follow these instructions if you would like to attach an external disk](https://cloud.google.com/tpu/docs/attach-durable-block-storage).

Expand Down Expand Up @@ -211,7 +211,7 @@ After installation completes, run the training script.
```

```bash
export LIBTPU_INIT_ARGS='--xla_tpu_enable_async_collective_fusion_fuse_all_gather=true \
export LIBTPU_INIT_ARGS="--xla_tpu_enable_async_collective_fusion_fuse_all_gather=true \
--xla_tpu_megacore_fusion_allow_ags=false \
--xla_enable_async_collective_permute=true \
--xla_tpu_enable_ag_backward_pipelining=true \
Expand All @@ -235,7 +235,7 @@ After installation completes, run the training script.
--xla_latency_hiding_scheduler_rerun=2 \
--xla_tpu_use_minor_sharding_for_major_trivial_input=true \
--xla_tpu_relayout_group_size_threshold_for_reduce_scatter=1 \
--xla_tpu_assign_all_reduce_scatter_layout=true'
--xla_tpu_assign_all_reduce_scatter_layout=true"
```

```bash
Expand Down Expand Up @@ -273,7 +273,7 @@ After installation completes, run the training script.

It is important to note a couple of things:
- per_device_batch_size can be fractional, but must be a whole number when multiplied by number of devices. In this example, 0.25 * 4 (devices) = effective global batch size = 1.
- The step time in v5p-8 with global batch size = 1 is large due to the aggressive remat policy (`HIDDEN_STATE_WITH_OFFLOAD`). On a larger number of chips we can run larger batch sizes, greatly increasing MFU, as we will see in the next section on deploying with xpk.
- The step time in v5p-8 with global batch size = 1 is large due to the aggressive remat policy (`HIDDEN_STATE_WITH_OFFLOAD`). On a larger number of chips we can run larger batch sizes, greatly increasing MFU, as we will see in the next section on deploying with Cluster Toolkit.
- To enable eval during training set `eval_every` to a value > 0.
- In Wan2.1, the ici_fsdp_parallelism axis is used for sequence parallelism, the ici_tensor_parallelism axis is used for head parallelism.
- You can enable both, keeping in mind that Wan2.1 has 40 heads and 40 must be evenly divisible by ici_tensor_parallelism.
Expand Down Expand Up @@ -315,11 +315,14 @@ After installation completes, run the training script.
completed step: 6, seconds: 36.006, TFLOP/s/device: 135.900, loss: 0.169
```

### Deploying with XPK
### Deploying with Cluster Toolkit

This assumes the user has already created an xpk cluster, installed all dependencies and also created the dataset from the step above. For getting started with MaxDiffusion and xpk see [this guide](docs/getting_started/run_maxdiffusion_via_xpk.md).
This assumes the user has already created a GKE cluster with TPUs, installed `gcluster` (the [Cluster Toolkit](https://github.com/GoogleCloudPlatform/cluster-toolkit) CLI), pushed a MaxDiffusion image built from this checkout to Artifact Registry and also created the dataset from the step above. For getting started with MaxDiffusion and Cluster Toolkit see [this guide](docs/getting_started/run_maxdiffusion_via_cluster_toolkit.md).

Using v5p-256, the command to run on xpk is as follows:
> [!NOTE]
> [XPK is deprecated](https://github.com/AI-Hypercomputer/xpk) in favor of Cluster Toolkit. The `gcluster` commands below were translated from the previous XPK commands using the [official migration guide](https://github.com/GoogleCloudPlatform/cluster-toolkit/blob/main/docs/migration/xpk_to_clustertoolkit.md) and have not yet been validated end-to-end on a cluster. The legacy XPK commands are kept in the collapsed sections for users with existing XPK clusters.

Using v5p-256, the command to run with `gcluster` is as follows:

```bash
RUN_NAME=jfacevedo-wan-v5p-8-${RANDOM}
Expand All @@ -330,7 +333,7 @@ After installation completes, run the training script.
```

```bash
LIBTPU_INIT_ARGS='--xla_tpu_enable_async_collective_fusion_fuse_all_gather=true \
LIBTPU_INIT_ARGS="--xla_tpu_enable_async_collective_fusion_fuse_all_gather=true \
--xla_tpu_megacore_fusion_allow_ags=false \
--xla_enable_async_collective_permute=true \
--xla_tpu_enable_ag_backward_pipelining=true \
Expand All @@ -354,9 +357,76 @@ After installation completes, run the training script.
--xla_latency_hiding_scheduler_rerun=2 \
--xla_tpu_use_minor_sharding_for_major_trivial_input=true \
--xla_tpu_relayout_group_size_threshold_for_reduce_scatter=1 \
--xla_tpu_assign_all_reduce_scatter_layout=true'
--xla_tpu_assign_all_reduce_scatter_layout=true"
```

```bash
# Machine type and topology of the slice. Examples with 128 JAX devices, matching
# ici_data_parallelism=32 x ici_fsdp_parallelism=4 below (TPU7x exposes two devices per chip):
# v5p-256 -> COMPUTE_TYPE=ct5p-hightpu-4t TOPOLOGY=4x4x8 (128 chips)
# v6e-128 -> COMPUTE_TYPE=ct6e-standard-4t TOPOLOGY=8x16 (128 chips)
# tpu7x-4x4x4 -> COMPUTE_TYPE=tpu7x-standard-4t TOPOLOGY=4x4x4 (64 chips)
COMPUTE_TYPE=ct5p-hightpu-4t
TOPOLOGY=4x4x8
# Image built from this checkout and pushed to Artifact Registry in step 3 of
# docs/getting_started/run_maxdiffusion_via_cluster_toolkit.md. PROJECT_ID, LOCATION, CLUSTER_NAME,
# REGION and AR_REPO are the variables set in step 2 of that guide.
IMAGE=${REGION}-docker.pkg.dev/${PROJECT_ID}/${AR_REPO}/maxdiffusion_runner:latest

gcluster job submit \
--name=${RUN_NAME} \
--cluster=$CLUSTER_NAME \
--project=$PROJECT_ID \
--location=$LOCATION \
--compute-type=${COMPUTE_TYPE} \
--topology=${TOPOLOGY} \
--num-slices=1 \
--image=${IMAGE} \
--env="LIBTPU_INIT_ARGS=${LIBTPU_INIT_ARGS}" \
--verbose \
--command=" \
pip install --no-deps . && HF_HUB_CACHE=/mnt/disks/external_disk/maxdiffusion_hf_cache/ python src/maxdiffusion/train_wan.py \
src/maxdiffusion/configs/base_wan_14b.yml \
attention='flash' \
weights_dtype=bfloat16 \
activations_dtype=bfloat16 \
guidance_scale=5.0 \
flow_shift=5.0 \
fps=16 \
skip_jax_distributed_system=False \
run_name=${RUN_NAME} \
output_dir=${OUTPUT_DIR} \
train_data_dir=${DATASET_DIR} \
load_tfrecord_cached=True \
height=1280 \
width=720 \
num_frames=81 \
num_inference_steps=50 \
jax_cache_dir=${OUTPUT_DIR}/jax_cache/ \
enable_profiler=True \
dataset_save_location=${SAVE_DATASET_DIR} \
remat_policy='HIDDEN_STATE_WITH_OFFLOAD' \
flash_min_seq_length=0 \
seed=$RANDOM \
skip_first_n_steps_for_profiler=3 \
profiler_steps=3 \
per_device_batch_size=0.25 \
ici_data_parallelism=32 \
ici_fsdp_parallelism=4 \
ici_tensor_parallelism=1 \
max_train_steps=5000 \
eval_every=100 \
eval_data_dir=${EVAL_DATA_DIR} \
enable_generate_video_for_eval=True" \
--priority=medium \
--restarts=0
```

Environment variables are not forwarded from your shell automatically, which is why `LIBTPU_INIT_ARGS` is passed with `--env` above (alternatively set them inline in `--command` as done for `HF_HUB_CACHE`). `--verbose` enables the same TPU debug logging as XPK's `--enable-debug-logs`. `pip install --no-deps .` reinstalls `maxdiffusion` from the checkout shipped in the image so that your local changes are the ones imported; see the [guide](docs/getting_started/run_maxdiffusion_via_cluster_toolkit.md#4-submit-your-first-workload) for details.

<details>
<summary>Legacy XPK command (deprecated)</summary>

```bash
python3 ~/xpk/xpk.py workload create \
--cluster=$CLUSTER_NAME \
Expand Down Expand Up @@ -405,6 +475,8 @@ After installation completes, run the training script.
--max-restarts=0
```

</details>

## 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.
Expand Down Expand Up @@ -434,9 +506,43 @@ After installation completes, run the training script.
save_final_checkpoint=True
```

### Multi-Host Training with XPK
### Multi-Host Training with Cluster Toolkit

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-256`, `v6e-128`, or `tpu7x-4x4x4`):
This section covers large-scale multi-host training across TPU pods (e.g. v5p, v6e, v7x). The following example is configured for a slice with 128 JAX devices (such as `v5p-256`, `v6e-128`, or `tpu7x-4x4x4`; TPU7x exposes two devices per chip) and reuses the `COMPUTE_TYPE`, `TOPOLOGY` and `IMAGE` variables from the [Wan 2.1 Cluster Toolkit example](#deploying-with-cluster-toolkit). Pick a new `RUN_NAME` for each submission:

```bash
gcluster job submit \
--name=${RUN_NAME} \
--cluster=$CLUSTER_NAME \
--project=$PROJECT_ID \
--location=$LOCATION \
--compute-type=${COMPUTE_TYPE} \
--topology=${TOPOLOGY} \
--num-slices=1 \
--image=${IMAGE} \
--command=" \
pip install --no-deps . && 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" \
--priority=medium \
--restarts=0
```

<details>
<summary>Legacy XPK command (deprecated)</summary>

```bash
python3 ~/xpk/xpk.py workload create \
Expand Down Expand Up @@ -468,6 +574,8 @@ After installation completes, run the training script.
--max-restarts=0
```

</details>

### 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.
Expand Down Expand Up @@ -971,7 +1079,7 @@ The optimal attention tile sizes (`block_q` / `block_kv`) depend on the sequence
```

## Getting Started: Multihost development
Multihost training for Stable Diffusion 2 base can be run using the following command:
For multi-host training on GKE, see [Running MaxDiffusion with Cluster Toolkit](docs/getting_started/run_maxdiffusion_via_cluster_toolkit.md). On a multi-host TPU VM slice, multihost training for Stable Diffusion 2 base can be run using the following command:
```bash
TPU_NAME=<your-tpu-name>
ZONE=<your-zone>
Expand Down
3 changes: 2 additions & 1 deletion docs/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,8 @@ This folder contains documentation for getting started with and using MaxDiffusi
## Getting Started

* **[First Run](getting_started/first_run.md)** - Provides instructions for setting up and running MaxDiffusion for the first time.
* **[Running MaxDiffusion via XPK](getting_started/run_maxdiffusion_via_xpk.md)** - Explains how to run MaxDiffusion on GKE using XPK.
* **[Running MaxDiffusion via Cluster Toolkit](getting_started/run_maxdiffusion_via_cluster_toolkit.md)** - Recommended way to run MaxDiffusion at scale on GKE using Cluster Toolkit's `gcluster` CLI.
* **[Running MaxDiffusion via XPK](getting_started/run_maxdiffusion_via_xpk.md)** - Explains how to run MaxDiffusion on GKE using XPK (deprecated, kept for existing XPK clusters).
* **[NVIDIA DGX Spark](dgx_spark.md)** - Explains how to run MaxDiffusion on an NVIDIA DGX Spark.

## Contributing & Community
Expand Down
4 changes: 3 additions & 1 deletion docs/getting_started/first_run.md
Original file line number Diff line number Diff line change
Expand Up @@ -33,4 +33,6 @@ source ~/$venv_name/bin/activate

## Getting Started: Multihost development

[GKE, recommended] [Running MaxDiffusion with xpk](run_maxdiffusion_via_xpk.md) - Quick Experimentation and Production support
[GKE, recommended] [Running MaxDiffusion with Cluster Toolkit](run_maxdiffusion_via_cluster_toolkit.md) - Quick Experimentation and Production support

[GKE, deprecated] [Running MaxDiffusion with XPK](run_maxdiffusion_via_xpk.md) - Kept for existing XPK clusters only; XPK is deprecated in favor of Cluster Toolkit
Loading
Loading