Skip to content
Open
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
275 changes: 232 additions & 43 deletions end_to_end/tpu/run_wan_fast_inference.sh

Large diffs are not rendered by default.

5 changes: 5 additions & 0 deletions end_to_end/tpu/run_wan_stack_tests.sh
Original file line number Diff line number Diff line change
Expand Up @@ -34,5 +34,10 @@ TESTS=(
"$T/dot_fallback_layout_test.py"
"$T/fused_producers_test.py"
"$T/tile_size_grid_search_test.py"
# Wan fast serving / AOT cache (feat/wan-fast-serving)
"$T/aot_cache_test.py"
"$T/converted_weights_cache_test.py"
"$T/wan/wan_transformer_test.py"
"$T/wan/wan_warmup_coverage_test.py"
)
PYTHONPATH="src${PYTHONPATH:+:$PYTHONPATH}" exec "${PYTHON:-python3}" -m pytest -q -rs "${TESTS[@]}" "$@"
437 changes: 362 additions & 75 deletions src/maxdiffusion/aot_cache.py

Large diffs are not rendered by default.

8 changes: 7 additions & 1 deletion src/maxdiffusion/configs/base_wan_14b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -101,10 +101,11 @@ svg_low_noise_density: -1.0
# compute blocks and the sparse kernel for small ones. Must stay a dict so the
# command line can override it with JSON.
svg_flash_block_sizes: {}
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_ring
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, tokamax_ring_custom, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block, ulysses_ring_custom_bidir

use_base2_exp: True
use_experimental_scheduler: True
use_k_centering: "auto" # auto: on for non-ring ulysses_custom_fixed_m* (virtual, no copy), off for ring paths (virtual via k_mean, plus pmean when R>1)
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down Expand Up @@ -309,6 +310,11 @@ dataset_config_name: ''
jax_cache_dir: ''
# Directory for per-shape AOT serialized executables ('' = disabled).
aot_cache_dir: ''
# Zero-execution warmup (warmup compiles the transformer passes without executing them)
# needs the AOT cache; with aot_cache_dir empty it installs an ephemeral one. Off = plain jax.jit.
enable_zero_execution_warmup: False
aot_build_revision: ''
wan_debug_cond_timers: False
# Directory for memoized torch->flax converted weights ('' = disabled).
converted_weights_dir: ''
hf_data_dir: ''
Expand Down
6 changes: 6 additions & 0 deletions src/maxdiffusion/configs/base_wan_1_3b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,7 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud

use_base2_exp: True
use_experimental_scheduler: True
use_k_centering: "auto"
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down Expand Up @@ -262,6 +263,11 @@ dataset_config_name: ''
jax_cache_dir: ''
# Directory for per-shape AOT serialized executables ('' = disabled).
aot_cache_dir: ''
# Zero-execution warmup (warmup compiles the transformer passes without executing them)
# needs the AOT cache; with aot_cache_dir empty it installs an ephemeral one. Off = plain jax.jit.
enable_zero_execution_warmup: False
aot_build_revision: ''
wan_debug_cond_timers: False
# Directory for memoized torch->flax converted weights ('' = disabled).
converted_weights_dir: ''
hf_data_dir: ''
Expand Down
7 changes: 5 additions & 2 deletions src/maxdiffusion/configs/base_wan_27b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -129,8 +129,6 @@ use_base2_exp: True
use_experimental_scheduler: True
# auto: on for non-ring ulysses_custom_fixed_m* (virtual, no copy), off for ring
# paths (virtual K-centering via k_mean, plus a pmean across ring shards when R>1).
# Note: the Wan pipelines do not pass this key to the attention layer yet, so
# the layer's default ("auto") applies.
use_k_centering: "auto"
# For attention=ulysses_ring*, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
Expand Down Expand Up @@ -304,6 +302,11 @@ dataset_config_name: ''
jax_cache_dir: ''
# Directory for per-shape AOT serialized executables ('' = disabled).
aot_cache_dir: ''
# Zero-execution warmup (warmup compiles the transformer passes without executing them)
# needs the AOT cache; with aot_cache_dir empty it installs an ephemeral one. Off = plain jax.jit.
enable_zero_execution_warmup: False
aot_build_revision: ''
wan_debug_cond_timers: False
# Directory for memoized torch->flax converted weights ('' = disabled).
converted_weights_dir: ''
hf_data_dir: ''
Expand Down
8 changes: 7 additions & 1 deletion src/maxdiffusion/configs/base_wan_animate.yml
Original file line number Diff line number Diff line change
Expand Up @@ -81,9 +81,10 @@ jit_initializers: True
# Set true to load weights from pytorch
from_pt: True
split_head_dim: True
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_ring
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block
use_base2_exp: True
use_experimental_scheduler: True
use_k_centering: "auto"
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down Expand Up @@ -253,6 +254,11 @@ dataset_config_name: ''
jax_cache_dir: '.jax_cache'
# Directory for per-shape AOT serialized executables ('' = disabled).
aot_cache_dir: ''
# Zero-execution warmup (warmup compiles the transformer passes without executing them)
# needs the AOT cache; with aot_cache_dir empty it installs an ephemeral one. Off = plain jax.jit.
enable_zero_execution_warmup: False
aot_build_revision: ''
wan_debug_cond_timers: False
# Directory for memoized torch->flax converted weights ('' = disabled).
converted_weights_dir: ''
hf_data_dir: ''
Expand Down
8 changes: 7 additions & 1 deletion src/maxdiffusion/configs/base_wan_i2v_14b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -83,9 +83,10 @@ jit_initializers: True
# Set true to load weights from pytorch
from_pt: True
split_head_dim: True
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_ring
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block
use_base2_exp: True
use_experimental_scheduler: True
use_k_centering: "auto"
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down Expand Up @@ -256,6 +257,11 @@ dataset_config_name: ''
jax_cache_dir: ''
# Directory for per-shape AOT serialized executables ('' = disabled).
aot_cache_dir: ''
# Zero-execution warmup (warmup compiles the transformer passes without executing them)
# needs the AOT cache; with aot_cache_dir empty it installs an ephemeral one. Off = plain jax.jit.
enable_zero_execution_warmup: False
aot_build_revision: ''
wan_debug_cond_timers: False
# Directory for memoized torch->flax converted weights ('' = disabled).
converted_weights_dir: ''
hf_data_dir: ''
Expand Down
8 changes: 7 additions & 1 deletion src/maxdiffusion/configs/base_wan_i2v_27b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -83,9 +83,10 @@ jit_initializers: True
# Set true to load weights from pytorch
from_pt: True
split_head_dim: True
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_ring
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block
use_base2_exp: True
use_experimental_scheduler: True
use_k_centering: "auto"
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down Expand Up @@ -257,6 +258,11 @@ dataset_config_name: ''
jax_cache_dir: ''
# Directory for per-shape AOT serialized executables ('' = disabled).
aot_cache_dir: ''
# Zero-execution warmup (warmup compiles the transformer passes without executing them)
# needs the AOT cache; with aot_cache_dir empty it installs an ephemeral one. Off = plain jax.jit.
enable_zero_execution_warmup: False
aot_build_revision: ''
wan_debug_cond_timers: False
# Directory for memoized torch->flax converted weights ('' = disabled).
converted_weights_dir: ''
hf_data_dir: ''
Expand Down
7 changes: 2 additions & 5 deletions src/maxdiffusion/generate_ltx2.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,6 @@
import time
import os
import subprocess
import uuid
from maxdiffusion.checkpointing.ltx2_checkpointer import LTX2Checkpointer
from maxdiffusion import aot_cache, pyconfig, max_logging, max_utils
from absl import app
Expand Down Expand Up @@ -148,7 +147,7 @@ def _canonical_aot_value(value):

def _non_reusable_aot_revision():
"""Returns a unique identity so unversioned source can never hit old HLO."""
return f"unversioned:{uuid.uuid4().hex}"
return aot_cache.non_reusable_aot_revision()


def _resolve_ltx2_aot_source_revision(config, commit_hash=None):
Expand All @@ -160,9 +159,7 @@ def _resolve_ltx2_aot_source_revision(config, commit_hash=None):


def _is_reusable_aot_revision(source_revision) -> bool:
if source_revision is None or not str(source_revision).strip():
return False
return not str(source_revision).startswith(("dirty:", "unversioned:"))
return aot_cache.is_reusable_aot_revision(source_revision)


def ltx2_aot_metadata(config, pipeline, source_revision=None):
Expand Down
Loading
Loading