From 0a2db3a0b43dcecd1862589867ca4850aaaece98 Mon Sep 17 00:00:00 2001 From: Rishabh Manoj Date: Thu, 17 Sep 2026 18:51:37 +0000 Subject: [PATCH] feat(wan): fast serving with persistent AOT caching and tuned inference recipe Persistent per-shape AOT executable cache, converted-weights cache and a tuned launcher for Wan inference on TPU. User-visible changes - AOT caching is opt-in. By default `generate_wan.run()` installs no cache directory and every call is plain jax.jit. A persistent cache is used only when `aot_cache_dir` is set AND the source revision is reusable. The revision is normally the package content hash, which is always reusable; the downgrade to an ephemeral temp dir only happens when that hash cannot be computed and the git revision is dirty/unversioned, or when `aot_build_revision` is explicitly marked dirty:/unversioned:. `enable_zero_execution_warmup: True` (new key, False in all 6 Wan ymls) also uses an ephemeral dir. The ephemeral dir is removed and aot_cache uninstalled in a try/finally that also covers mkdtemp, metadata, install and the loader wait. - `aot_cache_dir` must be a local POSIX path: gs:// raises ValueError, whatever the source revision. - .aotx format bumped to v2: existing v1 executables are ignored (recompiled once). - Converted-weights cache format is v3. v1 caches (from earlier revisions of this PR) are rejected once and re-converted (and re-saved: ~28 GB per Wan 2.2 A14B expert), with a warning saying so. v2 caches whose fingerprint matches the v2 formula for the same index are accepted and their manifest header is rewritten to v3 on first load (no re-conversion). AOT cache (aot_cache.py, generate_wan.py) - Signatures key non-static Python scalars by type (jit traces them), so e.g. guidance values share one executable; statics are keyed by value; nnx GraphDef statics get a process-deterministic digest (no set-order or address dependence). - Metadata fingerprint covers model/attention/tiles/mesh/VAE/dtype/remat/ cache options, device kind, process count, jax/jaxlib/libtpu/flax/qwix/ tokamax versions, matmul precision, PRNG impl, x64, threefry, LIBTPU_INIT_ARGS and XLA_FLAGS, plus the source revision. - Source revision: a content hash of every non-test .py file in the maxdiffusion package; an explicit `aot_build_revision` has the hash folded in. `get_git_commit_hash()` now appends "-dirty" for a dirty tree; Wan and LTX2 share one reusability rule. - Multi-host: per-host `-p{idx}` .aotx files. - `_align_inputs` raises on a leaf-count mismatch instead of truncating. - The .aotx pickle envelope is loaded with an exact (module, name) allowlist (builtins containers/scalars, PyTreeDef, jax tree registries; dotted names rejected). This is defence in depth only: a .aotx holds a compiled executable, so the real control is that aot_cache_dir is writable only by the user running inference. Warmup (Wan 2.2 T2V and I2V, wan_denoise_utils.py) - When aot_cache is in warmup mode (i.e. installed on a persistent or ephemeral dir), `compile_experts` compiles both experts' forward passes without executing them, since a 2-step warmup can stay entirely on the high-noise expert. Without an installed cache, warmup is an ordinary 2-step run. Replaces the old weight-priming forward pass. - `InflightWindow` bounds queued denoise steps (MAXD_QUEUE_MAX_INFLIGHT, default 4), shared by T2V and I2V. Converted-weights cache (wan_utils.py) - Manifest carries a format version and a source fingerprint; the fingerprint is repo id + HF snapshot revision + subfolder + index-file contents, with no absolute paths: for HF hub checkpoints a different HF_HOME or mount does not invalidate it (a local-directory checkpoint has no repo id or revision, so only subfolder + index contents are hashed). The revision is read from the snapshot path before symlinks are resolved: HF snapshot entries are symlinks into blobs/, and the v2 formula (realpath first, no subfolder) always hashed an empty revision and gave Wan 2.2's transformer and transformer_2 the same fingerprint, because their index files are byte-identical. - Fail-closed: a manifest without a fingerprint (or a caller without one) is a miss; keys and shapes are validated against eval_shapes. - Warm start checks the cache against the locally cached HF index (local_files_only) before any networked hf_hub_download. Trade-off: a newer upstream revision is not picked up while a valid converted cache exists. - The invalidated dir is removed before re-saving (peak 1x, not 2-3x). Re-save is skipped with a warning if the disk lacks the tree size + 2 GB headroom; downloading missing shards raises an actionable OSError if the HF cache volume lacks space. - VACE: the cache dir and conversion use the checkpoint's own scan_layers. Other - Unpatchify: when p_t == 1, reshape/transpose through a 7D tensor instead of 8D (same result; avoids an 8D strided copy). - `use_k_centering` is plumbed from the Wan config to the attention layer; `use_k_centering`, `aot_build_revision`, `wan_debug_cond_timers` and `enable_zero_execution_warmup` added to the Wan ymls. - Dot-product attention: float32_qk_product uses preferred_element_type=f32 on the QK einsum instead of upcasting Q/K (memory-efficient path unchanged). - Video export is atomic (temp file + os.replace), process-0 only, and shared by `run()` and `inference_generate_video`; trainers print SSIM on process 0. - Launcher (end_to_end/tpu/run_wan_fast_inference.sh): platform-detected v6e / v7 profiles, generic profile for everything else (incl. v5e); DVFS pin is opt-in (PIN_DVFS_P_STATE); `set -euo pipefail` with empty-array fixes; extra key=value args forwarded. Performance, measured on tpu7x-8 with this PR at the top of the stack (without #488/#491; Wan 2.2 T2V-A14B, 720p/81f/40 steps, CP=4, DP=2, warm AOT, launcher defaults), before the review fixes to the converted-weights fingerprint and AOT planning, which do not touch the denoise loop: DVFS unpinned: generate 115.5s, denoise 112.8s DVFS pinned (PIN_DVFS_P_STATE=true): generate 102.0s, denoise 99.4s These replace the launcher's earlier reference comment (95.8s / 93.2s pinned, which was a full-stack number). v6e-8 was not re-measured at this PR. Tests - aot_cache_test.py and converted_weights_cache_test.py (57 together): restricted-unpickler allowlist, run() AOT gating with ephemeral-dir teardown (including an install failure), gs:// rejection for any revision, planning through the real revision resolver, float32_qk_product keeping bf16 QK operands, fail-closed fingerprint-less manifests, disk-space checks, warm start before the network, revision captured through HF symlinks, distinct fingerprints for subfolders with identical index files, and the one-time v2 -> v3 manifest migration. - wan/wan_transformer_test.py and wan/wan_warmup_coverage_test.py (29). - CI: the four new wan_transformer tests (fused RMSNorm+RoPE parity and the self-attention dispatch check) skip in GitHub Actions; the AOT cache, converted-weights and warmup tests all run there. run_wan_stack_tests.sh gains this PR's test files. Verified (final tree): TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 201 passed, 117 subtests passed in 360.5s (86 passed across this PR's four test files). --- end_to_end/tpu/run_wan_fast_inference.sh | 275 +++++-- end_to_end/tpu/run_wan_stack_tests.sh | 5 + src/maxdiffusion/aot_cache.py | 437 +++++++++-- src/maxdiffusion/configs/base_wan_14b.yml | 8 +- src/maxdiffusion/configs/base_wan_1_3b.yml | 6 + src/maxdiffusion/configs/base_wan_27b.yml | 7 +- src/maxdiffusion/configs/base_wan_animate.yml | 8 +- src/maxdiffusion/configs/base_wan_i2v_14b.yml | 8 +- src/maxdiffusion/configs/base_wan_i2v_27b.yml | 8 +- src/maxdiffusion/generate_ltx2.py | 7 +- src/maxdiffusion/generate_wan.py | 410 ++++++++-- src/maxdiffusion/max_utils.py | 30 +- src/maxdiffusion/models/attention_flax.py | 22 +- .../wan/transformers/transformer_wan.py | 50 +- src/maxdiffusion/models/wan/wan_utils.py | 310 +++++++- .../pipelines/wan/wan_denoise_utils.py | 82 ++ .../pipelines/wan/wan_pipeline.py | 9 +- .../pipelines/wan/wan_pipeline_2_2.py | 80 +- .../pipelines/wan/wan_pipeline_i2v_2p2.py | 30 +- .../pipelines/wan/wan_vace_pipeline_2_1.py | 4 +- src/maxdiffusion/pyconfig.py | 11 +- src/maxdiffusion/tests/aot_cache_test.py | 708 +++++++++++++++++- .../tests/converted_weights_cache_test.py | 191 ++++- .../tests/wan/wan_transformer_test.py | 345 ++++++++- .../tests/wan/wan_warmup_coverage_test.py | 251 +++++++ src/maxdiffusion/trainers/base_wan_trainer.py | 3 +- src/maxdiffusion/trainers/wan_trainer_2_2.py | 3 +- src/maxdiffusion/utils/export_utils.py | 59 +- 28 files changed, 3013 insertions(+), 354 deletions(-) create mode 100644 src/maxdiffusion/pipelines/wan/wan_denoise_utils.py create mode 100644 src/maxdiffusion/tests/wan/wan_warmup_coverage_test.py diff --git a/end_to_end/tpu/run_wan_fast_inference.sh b/end_to_end/tpu/run_wan_fast_inference.sh index 89680760b..7b8fa8b79 100755 --- a/end_to_end/tpu/run_wan_fast_inference.sh +++ b/end_to_end/tpu/run_wan_fast_inference.sh @@ -14,27 +14,57 @@ # limitations under the License. # WAN T2V fast-serving example: AOT executable cache + converted-weights -# cache + zero-exec warmup, with a tuned v7 2D-ring attention recipe. +# cache + zero-exec warmup, with tuned per-platform attention recipes +# (Ulysses on v6e, 2D Ulyssesxring on v7). +# +# The XLA flag set, attention tile and text-encoder options differ per TPU +# generation, so the platform is auto-detected from the GCE metadata server +# and the matching recipe is selected (see "TPU platform detection" below). +# Profiles: v6e and v7. Other accelerators (including v5e/v5litepod) use a +# generic profile (v6e attention recipe, 64 MiB VMEM, common libtpu flags). # # First run per (model, shape) pays one-time conversion + compile and -# populates the caches; every later process start is ~25s to ready. +# populates the caches; later process starts take ~36-43s to ready (~34-38s +# load incl. text-encoder torch.compile, ~2-5s AOT/JAX-cache compile). A +# first run with an empty JAX cache compiles for ~3 min on v6e-8. # # Usage: -# ./run_wan_fast_inference.sh [21|22] [steps] ["prompt..."] +# ./run_wan_fast_inference.sh [21|22] [steps] ["prompt..."] [key=value ...] +# (extra key=value args after the 3rd positional arg are forwarded to generate_wan.py) # Env overrides: # WAN_CACHE_ROOT cache root (default ~/.cache/maxdiffusion_wan) -# OUTPUT_DIR video/metrics output (default /tmp/wan_out) -# COMPILE_TE=true torch.compile the text encoder (adds ~30s to load, -# saves ~10s/encode; worth it for long-lived processes) -# FIXEDM=0 plain online softmax instead of fixed-m +# OUTPUT_DIR video/metrics output (default ~/maxdiffusion_wan_output) +# TMPDIR / TORCHINDUCTOR_CACHE_DIR +# scratch and TorchInductor cache dirs under WAN_CACHE_ROOT +# COMPILE_TE torch.compile the text encoder (default true; adds ~30s to +# load, saves ~10s/encode; set false for one-shot runs) +# USE_BATCHED_TE batched text encoder execution (default true) +# VAE_SPATIAL / VAE_DECODE_CHUNK +# VAE spatial tiling (default 8) and temporal chunking (default 1) +# COMMON_LIBTPU / V6E_LIBTPU / V7_LIBTPU +# replace the tuned base or per-platform libtpu flag sets +# PIN_DVFS_P_STATE pin --xla_tpu_dvfs_p_state=7 on v7 (default false; opt-in for benchmarking) # EXTRA_LIBTPU extra libtpu flags, appended to the tuned set -# -# 720p 81f / 40 steps denoise: fixed-m 105.3s, plain 109.6s. Each mode gets its -# own optimal tile below; plain degrades badly on fixed-m's. -set -u +# TPU_PROFILE force a platform profile (v6e|v7|generic), skipping +# autodetection +# ACCEL_TYPE / TPU_ACCELERATOR_TYPE +# force the raw accelerator type (e.g. v6e-8, tpu7x-8) +# ATTENTION / ULYSSES_SHARDS / BQ / BKV / BKV_COMPUTE / BKV_COMPUTE_IN / BQ_DKV / VMEM_LIMIT_BYTES +# override the per-platform attention recipe +# DP / CP / PER_DEVICE_BATCH / SEED +# override mesh parallelism, per-device batch, or RNG seed (default 12345) +set -euo pipefail MODEL=${1:-22} +case "$MODEL" in + 21 | 22) ;; + *) + echo "Usage: $0 [21|22] [steps] [\"prompt...\"] [key=value ...]" >&2 + exit 1 + ;; +esac STEPS=${2:-40} PROMPT=${3:-""} +shift $(($# > 3 ? 3 : $#)) PROJECT_ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/../.." &> /dev/null && pwd)" cd "$PROJECT_ROOT" || exit 1 @@ -42,38 +72,196 @@ export PYTHONPATH="$PROJECT_ROOT/src:${PYTHONPATH:-}" export HF_HUB_ENABLE_HF_TRANSFER=1 export JAX_DEFAULT_MATMUL_PRECISION=bfloat16 export TORCHINDUCTOR_FX_GRAPH_CACHE=1 +export XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 +# Without these the JAX persistent cache silently skips most entries, so the +# "warm" start still recompiles a large part of the graph. +export JAX_PERSISTENT_CACHE_MIN_ENTRY_SIZE_BYTES=-1 +export JAX_PERSISTENT_CACHE_MIN_COMPILE_TIME_SECS=0 CACHE_ROOT=${WAN_CACHE_ROOT:-$HOME/.cache/maxdiffusion_wan} -OUTPUT_DIR=${OUTPUT_DIR:-/tmp/wan_out} -mkdir -p "$CACHE_ROOT/jax" "$CACHE_ROOT/aot_wan$MODEL" "$CACHE_ROOT/converted" "$OUTPUT_DIR" - -# Tuned collective/scheduler flag set for v7 (from the PR #430 2D-ring -# baseline). One line: libtpu stops parsing at a literal backslash. -export LIBTPU_INIT_ARGS="--xla_tpu_spmd_rng_bit_generator_unsafe=true --xla_tpu_enable_dot_strength_reduction=true --xla_tpu_enable_async_collective_fusion_fuse_all_gather=true --xla_enable_async_collective_permute=true --xla_tpu_enable_data_parallel_all_reduce_opt=true --xla_tpu_data_parallel_opt_different_sized_ops=true --xla_tpu_enable_async_collective_fusion=true --xla_tpu_enable_async_collective_fusion_multiple_steps=true --xla_tpu_overlap_compute_collective_tc=true --xla_enable_async_all_gather=true --xla_tpu_scoped_vmem_limit_kib=65536 --xla_tpu_enable_async_all_to_all=true --xla_tpu_enable_all_experimental_scheduler_features=true --xla_tpu_enable_scheduler_memory_pressure_tracking=true --xla_tpu_host_transfer_overlap_limit=24 --xla_tpu_aggressive_opt_barrier_removal=ENABLED --xla_lhs_prioritize_async_depth_over_stall=ENABLED --xla_should_allow_loop_variant_parameter_in_chain=ENABLED --xla_should_add_loop_invariant_op_in_chain=ENABLED --xla_tpu_enable_ici_ag_pipelining=true --xla_max_concurrent_host_send_recv=100 --xla_tpu_scheduler_percent_shared_memory_limit=100 --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_enable_latency_hiding_scheduler=true --xla_tpu_enable_ag_backward_pipelining=true --xla_tpu_enable_megacore_fusion=true --xla_tpu_megacore_fusion_allow_ags=true --xla_tpu_use_single_sparse_core_for_all_gather_offload=true --xla_tpu_sparse_core_all_gather_latency_multiplier=1 --xla_tpu_sparse_core_reduce_scatter_latency_multiplier=3 --xla_tpu_enable_sparse_core_collective_aggregator=true --xla_tpu_enable_sparse_core_offload_queuing_in_lhs=true --xla_tpu_enable_sparse_core_reduce_scatter_v2=true --xla_tpu_enable_sparse_core_collective_offload_all_gather=true --xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true --xla_tpu_enable_sparse_core_collective_offload_all_reduce=true --xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true --xla_tpu_enable_sparse_core_collective_offload_3d_all_gather=true --xla_tpu_enable_concurrent_sparse_core_offloading=true --xla_tpu_assign_all_reduce_scatter_layout=true" -# Timings only compare across runs passing the same extra flags. -export LIBTPU_INIT_ARGS="${LIBTPU_INIT_ARGS} ${EXTRA_LIBTPU:-}" - -# fixed-m on by default: faster, and covered by tests/ring_fixed_m_test.py. -if [ "${FIXEDM:-1}" = "1" ]; then - ATTENTION=ulysses_ring_custom_fixed_m - BQ=6400; BKV=2048 +OUTPUT_DIR=${OUTPUT_DIR:-$HOME/maxdiffusion_wan_output} +export TMPDIR=${TMPDIR:-$CACHE_ROOT/tmp} +export TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR:-$CACHE_ROOT/torch_compile} +mkdir -p "$CACHE_ROOT/jax" "$CACHE_ROOT/aot_wan$MODEL" "$CACHE_ROOT/converted" \ + "$OUTPUT_DIR" "$TMPDIR" "$TORCHINDUCTOR_CACHE_DIR" + +# --------------------------------------------------------------------------- +# TPU platform detection +# --------------------------------------------------------------------------- +# Preference order: explicit override -> TPU_ACCELERATOR_TYPE env (set by some +# runtimes) -> GCE metadata "accelerator-type" (e.g. "v6e-8") -> the +# ACCELERATOR_TYPE line inside the "tpu-env" metadata blob. Detection is pure +# metadata/env: it must not initialise the TPU, or it would take the device +# before the real process starts. +_tpu_metadata() { + curl -s -f -m 2 -H 'Metadata-Flavor: Google' \ + "http://metadata.google.internal/computeMetadata/v1/instance/attributes/$1" 2> /dev/null || true +} + +_detect_accel_type() { + local t="${TPU_ACCELERATOR_TYPE:-}" + [ -z "$t" ] && t="$(_tpu_metadata accelerator-type)" + [ -z "$t" ] && t="$(_tpu_metadata tpu-env | sed -n "s/^ACCELERATOR_TYPE: *'\([^']*\)'.*/\1/p")" + # Guard against metadata returning an HTML error page. Real names seen in the + # wild: "v6e-8", "v5litepod-8", "tpu7x-8" (v7 reports as tpu7x, not v7x). + case "$t" in + v[0-9]* | tpu[0-9]*) printf '%s' "$t" ;; + *) printf '' ;; + esac +} + +ACCEL_TYPE=${ACCEL_TYPE:-$(_detect_accel_type)} +if [ -n "$ACCEL_TYPE" ]; then + TPU_GEN="${ACCEL_TYPE%%-*}" # v6e-8 -> v6e + TPU_CHIPS="${ACCEL_TYPE##*-}" # v6e-8 -> 8 else - ATTENTION=ulysses_ring_custom - BQ=9472; BKV=1024 + TPU_GEN="" + TPU_CHIPS="" +fi +case "$TPU_CHIPS" in + '' | *[!0-9]*) TPU_CHIPS="" ;; +esac + +if [ -z "${TPU_PROFILE:-}" ]; then + case "$TPU_GEN" in + v6e) TPU_PROFILE=v6e ;; + v7 | v7x | v7p | v7e | tpu7 | tpu7x | tpu7p | tpu7e) TPU_PROFILE=v7 ;; + *) TPU_PROFILE=generic ;; + esac +fi + +# Keep LIBTPU flags single-line (literal backslashes truncate libtpu flag parsing). +COMMON_LIBTPU=${COMMON_LIBTPU:-"--xla_tpu_spmd_rng_bit_generator_unsafe=true --xla_tpu_enable_async_collective_fusion=true --xla_tpu_enable_async_collective_fusion_fuse_all_gather=false --xla_tpu_enable_async_collective_fusion_multiple_steps=true --xla_tpu_memory_bound_loop_optimizer_options=enabled:true --xla_tpu_enable_dot_strength_reduction=true --xla_enable_async_collective_permute=true --xla_tpu_enable_data_parallel_all_reduce_opt=true --xla_tpu_data_parallel_opt_different_sized_ops=true --xla_tpu_overlap_compute_collective_tc=true --xla_enable_async_all_gather=true --xla_tpu_scoped_vmem_limit_kib=65536 --xla_tpu_enable_async_all_to_all=true --xla_tpu_enable_all_experimental_scheduler_features=true --xla_tpu_enable_scheduler_memory_pressure_tracking=true --xla_tpu_host_transfer_overlap_limit=24 --xla_tpu_aggressive_opt_barrier_removal=ENABLED --xla_lhs_prioritize_async_depth_over_stall=ENABLED --xla_should_allow_loop_variant_parameter_in_chain=ENABLED --xla_should_add_loop_invariant_op_in_chain=ENABLED --xla_tpu_enable_ici_ag_pipelining=true --xla_max_concurrent_host_send_recv=100 --xla_tpu_scheduler_percent_shared_memory_limit=100 --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_enable_latency_hiding_scheduler=true --xla_tpu_enable_ag_backward_pipelining=true --xla_tpu_use_single_sparse_core_for_all_gather_offload=true --xla_tpu_sparse_core_all_gather_latency_multiplier=1 --xla_tpu_sparse_core_reduce_scatter_latency_multiplier=3 --xla_tpu_enable_sparse_core_collective_aggregator=true --xla_tpu_enable_sparse_core_offload_queuing_in_lhs=true --xla_tpu_enable_sparse_core_reduce_scatter_v2=true --xla_tpu_enable_sparse_core_collective_offload_all_gather=true --xla_tpu_enable_sparse_core_collective_offload_2d_all_gather=true --xla_tpu_enable_sparse_core_collective_offload_all_reduce=true --xla_tpu_enable_sparse_core_collective_offload_reduce_scatter=true --xla_tpu_enable_sparse_core_collective_offload_3d_all_gather=true --xla_tpu_enable_concurrent_sparse_core_offloading=true --xla_tpu_assign_all_reduce_scatter_layout=true"} +V6E_LIBTPU=${V6E_LIBTPU:-"--xla_tpu_enable_async_collective_fusion=true --xla_tpu_enable_async_collective_fusion_fuse_all_gather=false"} +# fuse_all_gather is false in every profile (see COMMON_LIBTPU); it must stay +# false on v7, where libtpu fails backend init with it enabled ("Continuation +# fusion for AllGather ... not supported ... other than Viperlite"). +# Pinning DVFS p-state 7 (max clocks) is opt-in via PIN_DVFS_P_STATE=true so +# shared/production hosts keep normal power management by default. +PIN_DVFS_P_STATE=${PIN_DVFS_P_STATE:-false} +V7_DVFS_FLAG="" +if [ "$PIN_DVFS_P_STATE" = "true" ]; then + V7_DVFS_FLAG=" --xla_tpu_dvfs_p_state=7" fi +V7_LIBTPU=${V7_LIBTPU:-"--xla_tpu_enable_async_collective_fusion=true --xla_tpu_enable_async_collective_fusion_fuse_all_gather=false --xla_tpu_enable_megacore_fusion=true --xla_tpu_megacore_fusion_allow_ags=true${V7_DVFS_FLAG}"} + +# Reference benchmarks (Wan 2.2 T2V-A14B, 720p/81f/40-step, CP=4, DP=2, warm AOT +# cache, these launcher defaults, measured at this PR without the later PRs): +# tpu7x-8: 115.5s generate (denoise 112.8s) DVFS unpinned; +# 102.0s generate (denoise 99.4s) with PIN_DVFS_P_STATE=true +# v6e-8: not re-measured at this PR +case "$TPU_PROFILE" in + v6e) + PLATFORM_LIBTPU="$V6E_LIBTPU" + DEFAULT_ATTENTION=ulysses_custom_fixed_m_per_q_block + DEFAULT_U=4 + DEFAULT_BQ=9472 + DEFAULT_BKV=1024 + DEFAULT_BKV_COMPUTE=512 + DEFAULT_BKV_COMPUTE_IN=512 + DEFAULT_VMEM=127506841 + DEFAULT_BQ_DKV=$DEFAULT_BQ + DEFAULT_COMPILE_TE=true + DEFAULT_BATCHED_TE=true + DEFAULT_VAE_CHUNK=1 + DEFAULT_VAE_SPATIAL=8 + ;; + v7) + PLATFORM_LIBTPU="$V7_LIBTPU" + DEFAULT_ATTENTION="ulysses_ring_custom_fixed_m" + DEFAULT_U=2 + DEFAULT_BQ=6400 + DEFAULT_BKV=2048 + DEFAULT_BKV_COMPUTE=2048 + DEFAULT_BKV_COMPUTE_IN=2048 + DEFAULT_VMEM=67108864 + DEFAULT_BQ_DKV=$DEFAULT_BQ + DEFAULT_COMPILE_TE=true + DEFAULT_BATCHED_TE=true + DEFAULT_VAE_CHUNK=1 + DEFAULT_VAE_SPATIAL=8 + ;; + *) + echo "== warning: unrecognised or non-v6e/v7 accelerator '${ACCEL_TYPE:-unknown}';" \ + "using the generic 64MiB profile. Set TPU_PROFILE=v6e|v7 to override." >&2 + PLATFORM_LIBTPU="" + DEFAULT_ATTENTION="ulysses_custom_fixed_m_per_q_block" + DEFAULT_U=4 + DEFAULT_BQ=9472 + DEFAULT_BKV=1024 + DEFAULT_BKV_COMPUTE=512 + DEFAULT_BKV_COMPUTE_IN=512 + DEFAULT_VMEM=67108864 + DEFAULT_BQ_DKV=$DEFAULT_BQ + DEFAULT_COMPILE_TE=true + DEFAULT_BATCHED_TE=true + DEFAULT_VAE_CHUNK=1 + DEFAULT_VAE_SPATIAL=8 + ;; +esac + +export LIBTPU_INIT_ARGS="${COMMON_LIBTPU} ${PLATFORM_LIBTPU} ${EXTRA_LIBTPU:-}" +# A literal backslash truncates libtpu's flag parsing; fail loudly rather than +# running with silently-dropped flags. +case "$LIBTPU_INIT_ARGS" in + *\\*) + echo "ERROR: LIBTPU_INIT_ARGS contains a literal backslash; libtpu would" \ + "stop parsing there and drop the remaining flags." >&2 + exit 1 + ;; +esac + +ATTENTION=${ATTENTION:-$DEFAULT_ATTENTION} +ULYSSES_SHARDS=${ULYSSES_SHARDS:-$DEFAULT_U} +BQ=${BQ:-$DEFAULT_BQ} +BKV=${BKV:-$DEFAULT_BKV} +BKV_COMPUTE=${BKV_COMPUTE:-$DEFAULT_BKV_COMPUTE} +BKV_COMPUTE_IN=${BKV_COMPUTE_IN:-$DEFAULT_BKV_COMPUTE_IN} +BQ_DKV=${BQ_DKV:-$DEFAULT_BQ_DKV} +VMEM_LIMIT_BYTES=${VMEM_LIMIT_BYTES:-$DEFAULT_VMEM} +COMPILE_TE=${COMPILE_TE:-$DEFAULT_COMPILE_TE} +USE_BATCHED_TE=${USE_BATCHED_TE:-$DEFAULT_BATCHED_TE} +VAE_SPATIAL=${VAE_SPATIAL:-$DEFAULT_VAE_SPATIAL} +VAE_DECODE_CHUNK=${VAE_DECODE_CHUNK:-$DEFAULT_VAE_CHUNK} + +# Mesh: context parallelism carries the Ulysses shards, data parallelism takes +# whatever chips remain. Defaults to CP=4 / DP=2 on an 8-chip slice. On a slice +# smaller than CP, clamp rather than emit a mesh larger than the hardware. +CP=${CP:-4} +if [ -n "$TPU_CHIPS" ] && [ "$TPU_CHIPS" -lt "$CP" ]; then + echo "== note: $TPU_CHIPS-chip slice is smaller than CP=$CP; clamping CP to $TPU_CHIPS" >&2 + CP=$TPU_CHIPS +fi +if [ "$ULYSSES_SHARDS" -gt "$CP" ]; then + echo "== note: clamping ulysses_shards $ULYSSES_SHARDS -> $CP (cannot exceed CP)" >&2 + ULYSSES_SHARDS=$CP +fi +if [ -z "${DP:-}" ]; then + if [ -n "$TPU_CHIPS" ] && [ "$TPU_CHIPS" -ge "$CP" ]; then + DP=$((TPU_CHIPS / CP)) + else + DP=2 + fi +fi +NUM_CHIPS=$((DP * CP)) +# One global video per step: per-device batch is 1/num_chips. +PER_DEVICE_BATCH=${PER_DEVICE_BATCH:-$(awk -v c="$NUM_CHIPS" 'BEGIN { printf "%.6g", 1.0 / c }')} if [ "$MODEL" = "21" ]; then CONFIG=src/maxdiffusion/configs/base_wan_14b.yml - GUIDANCE_ARGS="" + GUIDANCE_ARGS=() else CONFIG=src/maxdiffusion/configs/base_wan_27b.yml - GUIDANCE_ARGS="guidance_scale_low=3.0 guidance_scale_high=4.0" + GUIDANCE_ARGS=(guidance_scale_low=3.0 guidance_scale_high=4.0) fi PROMPT_ARG=() [ -n "$PROMPT" ] && PROMPT_ARG=("prompt=$PROMPT") RUN_NAME="wan${MODEL}_fast_$(date +%m%d-%H%M%S)" -echo "== ${ATTENTION} | tile ${BQ}/${BKV} | ${STEPS} steps" +echo "== platform ${ACCEL_TYPE:-unknown} -> profile ${TPU_PROFILE} | mesh DP=${DP} CP=${CP}" +echo "== ${ATTENTION} | U=${ULYSSES_SHARDS} | tile ${BQ}/${BKV} (compute=${BKV_COMPUTE}, in=${BKV_COMPUTE_IN}) | ${STEPS} steps" + +FLASH_BLOCK_SIZES="{\"block_q\":$BQ,\"block_kv\":$BKV,\"block_kv_compute\":$BKV_COMPUTE,\"block_kv_compute_in\":$BKV_COMPUTE_IN,\"heads_per_tile\":1,\"vmem_limit_bytes\":$VMEM_LIMIT_BYTES,\"block_q_dkv\":$BQ_DKV,\"block_kv_dkv\":$BKV,\"block_kv_dkv_compute\":$BKV,\"block_q_dq\":$BQ_DKV,\"block_kv_dq\":$BKV}" # libtpu's XLA:CPU AOT feature-mismatch log is cosmetic and ignores every # log-level env var; filter just that message from stderr. @@ -83,25 +271,26 @@ python src/maxdiffusion/generate_wan.py "$CONFIG" \ jax_cache_dir="$CACHE_ROOT/jax" \ aot_cache_dir="$CACHE_ROOT/aot_wan$MODEL" \ converted_weights_dir="$CACHE_ROOT/converted" \ - attention=$ATTENTION \ - ulysses_shards=2 \ - ici_data_parallelism=2 ici_fsdp_parallelism=1 \ - ici_context_parallelism=4 ici_tensor_parallelism=1 \ - per_device_batch_size=0.125 \ + attention="$ATTENTION" \ + ulysses_shards="$ULYSSES_SHARDS" \ + ici_data_parallelism="$DP" ici_fsdp_parallelism=1 \ + ici_context_parallelism="$CP" ici_tensor_parallelism=1 \ + per_device_batch_size="$PER_DEVICE_BATCH" \ num_inference_steps="$STEPS" num_frames=81 width=1280 height=720 \ weights_dtype=bfloat16 activations_dtype=bfloat16 \ - vae_spatial=4 vae_decode_chunk=-1 \ + vae_spatial="$VAE_SPATIAL" vae_decode_chunk="$VAE_DECODE_CHUNK" \ vae_weights_dtype=bfloat16 vae_dtype=bfloat16 \ - text_encoder_dtype=bfloat16 compile_text_encoder="${COMPILE_TE:-false}" use_batched_text_encoder=false \ - use_base2_exp=true use_experimental_scheduler=true \ - fps=16 $GUIDANCE_ARGS \ - flash_block_sizes="{\"block_q\":$BQ,\"block_kv\":$BKV,\"block_kv_compute\":$BKV,\"block_kv_compute_in\":1024,\"heads_per_tile\":1,\"vmem_limit_bytes\":67108864,\"block_q_dkv\":$BQ,\"block_kv_dkv\":$BKV,\"block_kv_dkv_compute\":$BKV}" \ - "${PROMPT_ARG[@]}" \ + text_encoder_dtype=bfloat16 compile_text_encoder="$COMPILE_TE" use_batched_text_encoder="$USE_BATCHED_TE" \ + use_kv_cache=true use_base2_exp=true use_experimental_scheduler=true \ + fps=16 ${GUIDANCE_ARGS[@]+"${GUIDANCE_ARGS[@]}"} \ + seed="${SEED:-12345}" \ + flash_block_sizes="$FLASH_BLOCK_SIZES" \ + ${PROMPT_ARG[@]+"${PROMPT_ARG[@]}"} \ + "$@" \ 2> >(grep -vE --line-buffered 'cpu_aot_loader|machine type for execution' >&2) -mp4=$(ls -t wan_output_*.mp4 2>/dev/null | head -1) +mp4=$(find "$OUTPUT_DIR" -maxdepth 1 -name '*.mp4' -printf '%T@ %p\n' 2>/dev/null | sort -nr | head -1 | cut -d' ' -f2-) if [ -n "$mp4" ]; then - mv "$mp4" "$OUTPUT_DIR/${RUN_NAME}.mp4" echo "" - echo "=== video saved: $OUTPUT_DIR/${RUN_NAME}.mp4 ===" + echo "=== video saved: $mp4 ===" fi diff --git a/end_to_end/tpu/run_wan_stack_tests.sh b/end_to_end/tpu/run_wan_stack_tests.sh index 850a6318e..27459fa23 100755 --- a/end_to_end/tpu/run_wan_stack_tests.sh +++ b/end_to_end/tpu/run_wan_stack_tests.sh @@ -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[@]}" "$@" diff --git a/src/maxdiffusion/aot_cache.py b/src/maxdiffusion/aot_cache.py index 1c27c2c92..90c6ef7c0 100644 --- a/src/maxdiffusion/aot_cache.py +++ b/src/maxdiffusion/aot_cache.py @@ -27,8 +27,10 @@ ``install()`` is called it delegates to plain ``jax.jit`` with zero behavioral difference, so tests and trainers are unaffected. * One executable is kept PER dynamic input signature (shapes/dtypes of - array leaves + treedef + non-array leaves). Different resolutions or - frame counts never collide on disk. + array leaves + treedef + static args by value). Non-static Python + scalars are keyed by type only: jit traces them as runtime inputs, so + e.g. guidance_scale 3.0 and 4.0 share one executable. Different + resolutions or frame counts never collide on disk. * Unknown signature -> silent jit fallback; the first call's args are recorded so ``save_pending()`` (call it after warmup, synchronously) can lower + serialize that shape without touching other shapes. @@ -41,7 +43,9 @@ Usage:: - @partial(aot_cache.cached_jit, static_argnames=("guidance_scale",)) + # Static args are keyed by value (one executable each); keep per-request + # scalars such as guidance_scale dynamic so they share one executable. + @partial(aot_cache.cached_jit, static_argnames=("do_classifier_free_guidance",)) def transformer_forward_pass(...): ... @@ -53,6 +57,7 @@ def transformer_forward_pass(...): from __future__ import annotations +import collections import contextlib import glob import hashlib @@ -63,6 +68,7 @@ def transformer_forward_pass(...): import re import threading from typing import Any, Callable +import uuid import jax import jax.numpy as jnp @@ -70,7 +76,18 @@ def transformer_forward_pass(...): from maxdiffusion import max_logging -_FORMAT_VERSION = 1 +# Bump whenever the blob layout or the signature scheme changes. The version is +# folded into the filename fingerprint, so older files are never even read. +# v2: non-static Python scalars are keyed by type, not value. +_FORMAT_VERSION = 2 + +# Python scalars that jit traces as weakly typed 0-d inputs when not static. +# Exact types only: subclasses (e.g. IntEnum) conservatively stay value-keyed. +_PY_SCALAR_TYPES = (bool, int, float, complex) + +# Bound on memoized device placements of Python scalar inputs (one per distinct +# value); values such as per-request guidance can vary without limit. +_SCALAR_CACHE_MAX = 256 def _is_graphdef(x: Any) -> bool: @@ -78,21 +95,18 @@ def _is_graphdef(x: Any) -> bool: def _graphdef_desc(gd: Any) -> str: - """Extracts a process-deterministic digest of static attributes in an nnx.GraphDef.""" + """Extracts a process-deterministic digest of static attributes in an nnx.GraphDef. + + Values go through ``_format_static_val`` rather than ``json.dumps(default=repr)``: + the latter falls back to ``repr`` for sets, whose iteration order depends on + PYTHONHASHSEED, so signatures would differ across processes. + """ items = [] for k, v in getattr(gd, "attributes", ()): if str(k).startswith("_pytree__"): continue if hasattr(v, "value"): - try: - s = json.dumps( - v.value, - sort_keys=True, - default=lambda o: re.sub(r"0x[0-9a-fA-F]+", "@", repr(o)), - ) - except Exception: # noqa: BLE001 - s = re.sub(r"0x[0-9a-fA-F]+", "@", repr(v.value)) - items.append(f"{k}:{s}") + items.append(f"{k}:{_format_static_val(v.value)}") return "GraphDef(" + ",".join(items) + ")" @@ -149,31 +163,161 @@ def _metadata_fingerprint(meta: dict[str, Any]) -> str: return hashlib.sha256(serialized.encode("utf-8")).hexdigest()[:12] -def _dynamic_signature(args: tuple, kwargs: dict) -> str: +def _format_const(c: Any) -> str: + import types as _types + + if isinstance(c, _types.CodeType): + return f"" + if isinstance(c, tuple): + return "(" + ",".join(_format_const(x) for x in c) + ")" + if isinstance(c, (set, frozenset)): + return f"{type(c).__name__}({{" + ",".join(sorted(_format_const(x) for x in c)) + "})" + if isinstance(c, (int, float, bool, str, bytes, type(None))): + return repr(c) + return re.sub(r"0x[0-9a-fA-F]+", "@", repr(c)) + + +def _code_digest(code: Any) -> str: + if code is None: + return "" + hasher = hashlib.sha256(code.co_code) + consts_str = "".join(_format_const(c) for c in code.co_consts) + hasher.update(consts_str.encode("utf-8")) + hasher.update(str(code.co_names).encode("utf-8")) + return hasher.hexdigest()[:8] + + +def _format_static_val(val: Any) -> str: + if isinstance(val, dict): + items = [] + for k in sorted(val.keys(), key=_format_static_val): + v = val[k] + # `diffusers.ConfigMixin.register_to_config` stores `_use_default_values = list(set(...))`, + # whose element order depends on PYTHONHASHSEED across processes. + if k == "_use_default_values" and isinstance(v, list): + v = sorted(v, key=_format_static_val) + items.append(f"{_format_static_val(k)}:{_format_static_val(v)}") + return "{" + ",".join(items) + "}" + if isinstance(val, (set, frozenset)): + items = sorted(_format_static_val(x) for x in val) + return f"{type(val).__name__}({{" + ",".join(items) + "})" + if isinstance(val, tuple): + return "(" + ",".join(_format_static_val(x) for x in val) + ",)" + if isinstance(val, list): + return "[" + ",".join(_format_static_val(x) for x in val) + "]" + if inspect.isfunction(val) or inspect.ismethod(val): + code = getattr(val, "__code__", None) + return f"" + if isinstance(val, type): + call_code = getattr(getattr(val, "__call__", None), "__code__", None) + return f"" + return re.sub(r"0x[0-9a-fA-F]+", "@", repr(val)) + + +_GRAPHDEF_MEMO: collections.OrderedDict[int, tuple[Any, str]] = collections.OrderedDict() + + +def _extract_graphdef_statics(obj: Any, prefix: str = "") -> list[str]: + """Extracts deterministic node types, topology, and static attributes from nnx.GraphDef objects.""" + out = [] + if hasattr(obj, "nodes") and hasattr(obj, "attributes"): + obj_id = id(obj) + digest = None + if obj_id in _GRAPHDEF_MEMO: + cached_obj, cached_digest = _GRAPHDEF_MEMO[obj_id] + if cached_obj is obj: + _GRAPHDEF_MEMO.move_to_end(obj_id) + digest = cached_digest + if digest is None: + raw = [] + for i, node in enumerate(getattr(obj, "nodes", ())): + node_type = _format_static_val(getattr(node, "type", None)) + node_idx = getattr(node, "index", None) + outer_idx = getattr(node, "outer_index", None) + num_attrs = getattr(node, "num_attributes", None) + meta = _format_static_val(getattr(node, "metadata", None)) + raw.append(f"node[{i}]={node_type}:{node_idx}:{outer_idx}:{num_attrs}:{meta}") + for i, attr_item in enumerate(getattr(obj, "attributes", ())): + if isinstance(attr_item, tuple) and len(attr_item) == 2: + k, v = attr_item + if isinstance(k, str) and k.startswith("_pytree"): + continue + if hasattr(v, "value"): + desc = _format_static_val(v.value) + raw.append(f"attr[{i}].{k}={desc}") + else: + v_type = _format_static_val(getattr(v, "type", type(v))) + v_idx = getattr(v, "index", None) + v_meta = _format_static_val(getattr(v, "metadata", None)) + raw.append(f"attr[{i}].{k}=<{type(v).__name__}:{v_type}:{v_idx}:{v_meta}>") + if hasattr(v, "graphdef"): + raw.extend(_extract_graphdef_statics(v.graphdef, f"attr[{i}].{k}.graphdef")) + digest = hashlib.sha256("|".join(raw).encode("utf-8")).hexdigest()[:16] + if len(_GRAPHDEF_MEMO) >= 1024: + _GRAPHDEF_MEMO.popitem(last=False) + _GRAPHDEF_MEMO[obj_id] = (obj, digest) + out.append(f"{prefix}:graphdef={digest}") + elif isinstance(obj, (tuple, list)): + for i, item in enumerate(obj): + out.extend(_extract_graphdef_statics(item, f"{prefix}[{i}]")) + elif isinstance(obj, dict): + for k in sorted(obj.keys(), key=str): + out.extend(_extract_graphdef_statics(obj[k], f"{prefix}.{k}")) + return out + + +def _leaf_desc(leaf: Any, *, static: bool) -> str: + """Describes one flattened argument leaf for ``_dynamic_signature``.""" + if _is_graphdef(leaf): + return _graphdef_desc(leaf) + if hasattr(leaf, "shape") and hasattr(leaf, "dtype"): + weak = getattr(leaf, "weak_type", False) + return f"{tuple(leaf.shape)}:{leaf.dtype}:weak={weak}" + if not static and type(leaf) in _PY_SCALAR_TYPES: + # jit traces a non-static Python scalar as a weakly typed 0-d input: the + # executable depends on its type (bool/int/float/complex), never its value. + return f"<{type(leaf).__name__}>" + return _format_static_val(leaf) + + +def _dynamic_signature(args: tuple, kwargs: dict, static: dict | None = None) -> str: """Deterministic digest of everything that selects an executable. Structure is captured by each leaf's KEY PATH (names, order, count) -- NOT by ``repr(treedef)``: an nnx GraphDef's repr embeds object addresses and hash-order-dependent content that differ per process and - made signatures never match across restarts (measured: every array - part stable, only the treedef part unstable). Static graph metadata - in GraphDef leaves is extracted deterministically via ``_graphdef_desc``. - Array leaves contribute shape/dtype; non-array leaves (python scalars, - None flags) contribute an address-stripped repr. + made signatures never match across restarts. ``nnx.GraphDef`` leaves are + described deterministically via ``_graphdef_desc``, and their node types, + topology and static attributes are additionally captured by + ``_extract_graphdef_statics`` (address-stripped, with dicts/sets + canonicalized). Array leaves contribute shape/dtype/weak_type. Python + scalars in ``args``/``kwargs`` contribute only their type, because jit + traces them and their value is a runtime input. ``static`` holds the + static args, which are baked into the graph: they and any other non-array + leaves contribute a canonicalized representation of their value. """ - leaves_with_paths = jax.tree_util.tree_flatten_with_path((args, kwargs), is_leaf=_is_graphdef)[0] + static = static or {} parts = [] - for path, leaf in leaves_with_paths: - if _is_graphdef(leaf): - desc = _graphdef_desc(leaf) - elif hasattr(leaf, "shape") and hasattr(leaf, "dtype"): - desc = f"{tuple(leaf.shape)}:{leaf.dtype}" - else: - desc = re.sub(r"0x[0-9a-fA-F]+", "@", repr(leaf)) - parts.append(f"{jax.tree_util.keystr(path)}={desc}") + for tree, is_static in (((args, kwargs), False), (static, True)): + prefix = "static" if is_static else "" + for path, leaf in jax.tree_util.tree_flatten_with_path(tree, is_leaf=_is_graphdef)[0]: + parts.append(f"{prefix}{jax.tree_util.keystr(path)}={_leaf_desc(leaf, static=is_static)}") + parts.extend(_extract_graphdef_statics((args, kwargs, static))) return hashlib.sha256("|".join(parts).encode()).hexdigest()[:12] +def _to_aval_leaf(x: Any) -> Any: + """Replaces live jax.Array buffers with lightweight ShapeDtypeStructs.""" + if isinstance(x, jax.Array): + return jax.ShapeDtypeStruct( + x.shape, + x.dtype, + sharding=getattr(x, "sharding", None), + weak_type=getattr(x, "weak_type", False), + ) + return x + + class _AotEntry: """Executables for one wrapped fn, keyed by dynamic input signature.""" @@ -187,7 +331,11 @@ def __init__(self, name: str, fn: Callable, static_argnames: tuple): self._out_specs: dict[str, Any] = {} self._pending: dict[str, tuple] = {} self._adapters: dict[str, Any] = {} + self._sig_cache: dict[tuple, str] = {} self._on_disk: set[str] = set() + self._expected_shardings: dict[int, tuple[Any, list]] = {} + self._equiv_shardings: set[tuple[Any, Any, int]] = set() + self._scalar_cache: dict[tuple[type, Any, Any], Any] = {} self._lock = threading.Lock() def _zeros_output(self, signature: str): @@ -241,6 +389,7 @@ def _canonicalize(self, args: tuple, kwargs: dict) -> tuple[dict, dict]: paths makes the pytrees agree by construction. """ bound = self.py_signature.bind(*args, **kwargs) + bound.apply_defaults() dynamic, static = {}, {} for name, val in bound.arguments.items(): (static if name in self.static_argnames else dynamic)[name] = val @@ -260,7 +409,7 @@ def _compile_and_record(self, signature: str, leaves: list, treedef: Any, static # ---------------------------------------------------------------- call def __call__(self, *args, **kwargs): - if not _STATE.enabled: + if not _STATE.enabled and not _STATE.warmup_only: return self.jitted(*args, **kwargs) dynamic, static = self._canonicalize(args, kwargs) leaves, treedef = jax.tree_util.tree_flatten(dynamic) @@ -268,7 +417,31 @@ def __call__(self, *args, **kwargs): # Under an outer trace a deserialized executable cannot be applied # and tracers must not be recorded -- inline like a nested jit. return self.jitted(**dynamic, **static) - signature = _dynamic_signature((), {**dynamic, **static}) + + # Fast-path signature cache: avoid tree_flatten_with_path + SHA256 string hashing on repeated steps. + # Python scalars are keyed by type as in _leaf_desc, so a new value neither misses nor fills this cache. + shapes_dtypes = tuple( + (leaf.shape, leaf.dtype, getattr(leaf, "weak_type", False)) + if hasattr(leaf, "shape") and hasattr(leaf, "dtype") + else (type(leaf),) + if type(leaf) in _PY_SCALAR_TYPES + else (type(leaf), _format_static_val(leaf)) + for leaf in leaves + ) + static_items = tuple((k, _format_static_val(v)) for k, v in sorted(static.items(), key=lambda item: str(item[0]))) + graphdef_statics = tuple(_extract_graphdef_statics((args, kwargs))) + cache_key = (treedef, shapes_dtypes, static_items, graphdef_statics) + signature = getattr(self, "_sig_cache", {}).get(cache_key) + if signature is None: + signature = _dynamic_signature((), dynamic, static) + with self._lock: + if not hasattr(self, "_sig_cache"): + self._sig_cache = {} + if len(self._sig_cache) < 64: + new_cache = dict(self._sig_cache) + new_cache[cache_key] = signature + self._sig_cache = new_cache + if _STATE.warmup_only: # Compilation only needs avals; skip the (possibly seconds-long) # real execution and hand back correctly-shaped/sharded zeros so @@ -276,69 +449,95 @@ def __call__(self, *args, **kwargs): if signature not in self._compiled: self._compile_and_record(signature, leaves, treedef, static) with self._lock: - if signature not in self._on_disk: - self._pending[signature] = (leaves, treedef, static) + if _STATE.enabled and signature not in self._on_disk: + self._pending[signature] = ([_to_aval_leaf(x) for x in leaves], treedef, static) zeros = self._zeros_output(signature) if zeros is not None: return zeros compiled = self._compiled.get(signature) if compiled is not None: - flat = self._align_inputs(compiled, leaves) - if flat is not None: - return compiled(flat) - # Fewer expected shardings than leaves: XLA pruned unused inputs - # (e.g. encoder params in a decode-only executable). Compiled keeps - # the full in_tree and prunes internally, so hand it the raw leaves; - # sharding/structure problems surface as catchable Python errors. - try: - return compiled(leaves) - except Exception as e: # noqa: BLE001 - any failure means "use jit" - max_logging.log(f"[aot] {self.name}: compiled call failed ({e}); using jit") + return compiled(self._align_inputs(compiled, leaves)) with self._lock: - if signature not in self._pending and signature not in self._compiled: - self._pending[signature] = (leaves, treedef, static) + if _STATE.enabled and signature not in self._pending and signature not in self._compiled: + self._pending[signature] = ([_to_aval_leaf(x) for x in leaves], treedef, static) return self._adapter_for(signature, treedef, static)(leaves) - def _align_inputs(self, compiled: Any, leaves: list): + def _align_inputs(self, compiled: Any, leaves: list) -> list: """Reshards the flat input leaves onto the executable's shardings. jit auto-commits mismatched inputs; a deserialized Compiled does not -- a placement mismatch aborts inside PjRt (uncatchable C++). Weights already carry final shardings; in practice this only moves small - fresh-off-host activations. Returns the aligned leaf list, or None on - structural mismatch (caller falls back to jit). + fresh-off-host activations or memoized Python scalars. Pruned inputs + have expected sharding None and are passed through unchanged. Raises + loudly on sharding alignment errors rather than silently triggering a + mid-serving JIT recompile. """ - try: - flat_expected = jax.tree_util.tree_leaves(compiled.input_shardings) - if len(flat_expected) != len(leaves): - # Fewer expected shardings than leaves = XLA pruned unused inputs; - # the caller retries via Compiled's own pruning path. Not an error. - return None - aligned = [] - for leaf, expected in zip(leaves, flat_expected): - if not hasattr(leaf, "shape"): # python scalar traced as weak array - leaf = jnp.asarray(leaf) - sharding = getattr(leaf, "sharding", None) - if sharding is not None and sharding.is_equivalent_to(expected, leaf.ndim): + cid = id(compiled) + cached = self._expected_shardings.get(cid) + if cached is not None and cached[0] is compiled: + flat_expected = cached[1] + else: + flat_expected = jax.tree_util.tree_leaves(compiled.input_shardings, is_leaf=lambda x: x is None) + self._expected_shardings[cid] = (compiled, flat_expected) + + if len(leaves) != len(flat_expected): + raise ValueError( + f"{self.name}: input leaf count mismatch in _align_inputs: " + f"got {len(leaves)} leaves, expected {len(flat_expected)}." + ) + + equiv = self._equiv_shardings + # Fast path: all non-pruned leaves are already arrays with matching or verified-equivalent shardings + if all( + expected is None + or ((sh := getattr(leaf, "sharding", None)) is not None and (sh == expected or (sh, expected, leaf.ndim) in equiv)) + for leaf, expected in zip(leaves, flat_expected) + ): + return leaves + + aligned = [] + for leaf, expected in zip(leaves, flat_expected): + if expected is None: + aligned.append(leaf) + continue + if not hasattr(leaf, "shape"): # python scalar traced as weak array + skey = (type(leaf), leaf, expected) + placed = self._scalar_cache.get(skey) + if placed is None or getattr(placed, "is_deleted", lambda: False)(): + placed = jax.device_put(jnp.asarray(leaf), expected) + if skey not in self._scalar_cache and len(self._scalar_cache) >= _SCALAR_CACHE_MAX: + self._scalar_cache.pop(next(iter(self._scalar_cache))) # evict the oldest value + self._scalar_cache[skey] = placed + aligned.append(placed) + continue + sharding = getattr(leaf, "sharding", None) + if sharding is not None: + ekey = (sharding, expected, leaf.ndim) + if ekey in equiv: aligned.append(leaf) - else: - aligned.append(jax.device_put(leaf, expected)) - return aligned - except Exception as e: # noqa: BLE001 - any failure means "use jit" - max_logging.log(f"[aot] {self.name}: cannot align inputs ({e}); using jit") - return None + continue + if sharding == expected or sharding.is_equivalent_to(expected, leaf.ndim): + equiv.add(ekey) + aligned.append(leaf) + continue + aligned.append(jax.device_put(leaf, expected)) + return aligned # ---------------------------------------------------------------- disk + def _host_tag(self) -> str: + return f"-p{jax.process_index()}" if jax.process_count() > 1 else "" + def _path_for(self, signature: str) -> str: - return os.path.join(_STATE.cache_dir, f"{self.name}-{_STATE.fingerprint}-{signature}.aotx") + return os.path.join(_STATE.cache_dir, f"{self.name}-{_STATE.fingerprint}{self._host_tag()}-{signature}.aotx") def load_from_disk(self) -> None: """Deserializes every on-disk executable for this fn. Never raises.""" - pattern = os.path.join(_STATE.cache_dir, f"{self.name}-{_STATE.fingerprint}-*.aotx") + pattern = os.path.join(_STATE.cache_dir, f"{self.name}-{_STATE.fingerprint}{self._host_tag()}-*.aotx") for path in glob.glob(pattern): try: with open(path, "rb") as f: - blob = pickle.load(f) + blob = _safe_pickle_load(f) if blob["format_version"] != _FORMAT_VERSION: continue # Topology-pinned device order; the default reconstruction binds @@ -374,6 +573,7 @@ def save_pending(self) -> int: if signature in self._on_disk: # Background deserialization landed after this shape was recorded. continue + tmp_path = "" try: compiled = self._compiled.get(signature) if compiled is None: @@ -395,7 +595,7 @@ def save_pending(self) -> int: if out_spec is not None: blob["out_shapes_dtypes"] = [(list(shape), str(dtype)) for shape, dtype in out_spec[1]] path = self._path_for(signature) - tmp_path = f"{path}.tmp.{os.getpid()}" + tmp_path = f"{path}.tmp.{os.getpid()}.{uuid.uuid4().hex[:8]}" with open(tmp_path, "wb") as f: pickle.dump(blob, f) os.replace(tmp_path, path) @@ -405,10 +605,80 @@ def save_pending(self) -> int: saved += 1 max_logging.log(f"[aot] {self.name}: serialized {os.path.basename(path)} ({len(payload) / 1e6:.1f}MB)") except Exception as e: # noqa: BLE001 - saving is best-effort + if tmp_path and os.path.exists(tmp_path): + try: + os.unlink(tmp_path) + except OSError: + pass max_logging.log(f"[aot] {self.name}: serialize failed ({e}); shape stays on jit") return saved +class _RestrictedUnpickler(pickle.Unpickler): + """Unpickler for the metadata envelope of a .aotx blob. + + SECURITY: this is defence in depth, NOT the trust boundary. A .aotx file + contains a serialized compiled executable that `deserialize_and_load` hands to + the TPU runtime, so anyone who can write to `aot_cache_dir` can already run + arbitrary code on the host/device. The real control is that the cache + directory must only be writable by the user running inference (never a shared + or world-writable path). This class only stops the pickle envelope itself + from being an additional, trivially exploitable code-execution vector. + + The envelope holds builtin containers/scalars plus two PyTreeDefs (the + adapter's flat in/out trees). Only these exact (module, name) globals are + resolvable; anything else -- including dotted names, which `find_class` + would otherwise resolve attribute by attribute -- is rejected, and the + loader falls back to jit for that shape. + """ + + _ALLOWED_GLOBALS = frozenset({ + ("builtins", "dict"), + ("builtins", "list"), + ("builtins", "tuple"), + ("builtins", "set"), + ("builtins", "frozenset"), + ("builtins", "str"), + ("builtins", "bytes"), + ("builtins", "int"), + ("builtins", "float"), + ("builtins", "bool"), + # PyTreeDef and the (data-only) registry it is reconstructed against; + # serialize_executable's trees reference the tracing registry. The + # module that defines PyTreeDef moved across jaxlib versions. + ("jaxlib._jax.pytree", "PyTreeDef"), + ("jaxlib.xla_extension.pytree", "PyTreeDef"), + ("jax._src.tree_util", "default_registry"), + ("jax._src.tree_util", "none_leaf_registry"), + ("jax._src.tree_util", "dispatch_registry"), + ("jax._src.tree_util", "tracing_registry"), + }) + + def find_class(self, module: str, name: str) -> Any: + if "." not in name and (module, name) in self._ALLOWED_GLOBALS: + return super().find_class(module, name) + raise pickle.UnpicklingError(f"Disallowed global in .aotx cache blob: {module}.{name}") + + +def _safe_pickle_load(file_obj: Any) -> Any: + return _RestrictedUnpickler(file_obj).load() + + +def non_reusable_aot_revision() -> str: + """Returns a unique identity so unversioned/dirty development source can never hit old HLO.""" + return f"unversioned:{uuid.uuid4().hex}" + + +def is_reusable_aot_revision(source_revision: Any) -> bool: + """True when `source_revision` is a clean, deterministic revision safe for persistent AOT caching.""" + if source_revision is None or not str(source_revision).strip(): + return False + s = str(source_revision).strip() + if s.startswith(("dirty:", "unversioned:")) or s.endswith("-dirty"): + return False + return True + + class _State: """Process-global install state (null until install() is called).""" @@ -443,7 +713,7 @@ def install(cache_dir: str, meta: dict[str, Any], mesh: Any) -> None: """Enables the AOT cache and starts background deserialization. Args: - cache_dir: Directory for .aotx files (created if missing). + cache_dir: Local POSIX directory for .aotx files (created if missing). meta: Everything the executables depend on beyond input shapes: model path, mesh shape, sharding/attention config, jax version. Hashed into the filename so incompatible executables never load. @@ -458,19 +728,28 @@ def install(cache_dir: str, meta: dict[str, Any], mesh: Any) -> None: _STATE.fingerprint = "" _STATE.mesh = None _STATE.enabled = False + _GRAPHDEF_MEMO.clear() for entry in _REGISTRY: with entry._lock: entry._compiled.clear() entry._out_specs.clear() entry._pending.clear() entry._adapters.clear() + entry._sig_cache.clear() entry._on_disk.clear() + entry._expected_shardings.clear() + entry._equiv_shardings.clear() + entry._scalar_cache.clear() if not cache_dir: return - os.makedirs(cache_dir, exist_ok=True) + if str(cache_dir).startswith("gs://"): + raise ValueError(f"aot_cache requires a local POSIX directory path; gs:// URIs are not supported (got {cache_dir!r}).") + os.makedirs(cache_dir, mode=0o700, exist_ok=True) _STATE.cache_dir = cache_dir - _STATE.fingerprint = _metadata_fingerprint(meta) + # The format version is part of every executable's identity: files written + # under an older signature scheme never match the glob, so they are not read. + _STATE.fingerprint = _metadata_fingerprint({**meta, "aot_format_version": _FORMAT_VERSION}) _STATE.mesh = mesh _STATE.enabled = True for entry in _REGISTRY: @@ -495,6 +774,13 @@ def save_pending() -> int: return sum(entry.save_pending() for entry in _REGISTRY) +def clear_pending() -> None: + """Discards any recorded pending shapes without serializing them.""" + for entry in _REGISTRY: + with entry._lock: + entry._pending.clear() + + @contextlib.contextmanager def warmup_mode(): """Zero-execution warmup: wrapped fns lower+compile but never execute. @@ -507,11 +793,12 @@ def warmup_mode(): compile against faithful inputs. Outputs of a warmup pass are garbage by design; callers must discard them. No-op when the cache is disabled. """ + previous = _STATE.warmup_only _STATE.warmup_only = _STATE.enabled try: yield finally: - _STATE.warmup_only = False + _STATE.warmup_only = previous def in_warmup() -> bool: diff --git a/src/maxdiffusion/configs/base_wan_14b.yml b/src/maxdiffusion/configs/base_wan_14b.yml index 4c2b2b30c..6bcf0b4e1 100644 --- a/src/maxdiffusion/configs/base_wan_14b.yml +++ b/src/maxdiffusion/configs/base_wan_14b.yml @@ -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. @@ -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: '' diff --git a/src/maxdiffusion/configs/base_wan_1_3b.yml b/src/maxdiffusion/configs/base_wan_1_3b.yml index 89bcfb885..3d9a1cc95 100644 --- a/src/maxdiffusion/configs/base_wan_1_3b.yml +++ b/src/maxdiffusion/configs/base_wan_1_3b.yml @@ -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. @@ -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: '' diff --git a/src/maxdiffusion/configs/base_wan_27b.yml b/src/maxdiffusion/configs/base_wan_27b.yml index a1e2e0353..398e8d15f 100644 --- a/src/maxdiffusion/configs/base_wan_27b.yml +++ b/src/maxdiffusion/configs/base_wan_27b.yml @@ -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 @@ -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: '' diff --git a/src/maxdiffusion/configs/base_wan_animate.yml b/src/maxdiffusion/configs/base_wan_animate.yml index 5e9df7d0d..d90a7f611 100644 --- a/src/maxdiffusion/configs/base_wan_animate.yml +++ b/src/maxdiffusion/configs/base_wan_animate.yml @@ -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. @@ -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: '' diff --git a/src/maxdiffusion/configs/base_wan_i2v_14b.yml b/src/maxdiffusion/configs/base_wan_i2v_14b.yml index a129ff66c..d252b0ce7 100644 --- a/src/maxdiffusion/configs/base_wan_i2v_14b.yml +++ b/src/maxdiffusion/configs/base_wan_i2v_14b.yml @@ -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. @@ -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: '' diff --git a/src/maxdiffusion/configs/base_wan_i2v_27b.yml b/src/maxdiffusion/configs/base_wan_i2v_27b.yml index 6a28986fc..84a64e10d 100644 --- a/src/maxdiffusion/configs/base_wan_i2v_27b.yml +++ b/src/maxdiffusion/configs/base_wan_i2v_27b.yml @@ -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. @@ -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: '' diff --git a/src/maxdiffusion/generate_ltx2.py b/src/maxdiffusion/generate_ltx2.py index b60f223b8..878897272 100644 --- a/src/maxdiffusion/generate_ltx2.py +++ b/src/maxdiffusion/generate_ltx2.py @@ -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 @@ -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): @@ -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): diff --git a/src/maxdiffusion/generate_wan.py b/src/maxdiffusion/generate_wan.py index a8a184f48..dba1d0b44 100644 --- a/src/maxdiffusion/generate_wan.py +++ b/src/maxdiffusion/generate_wan.py @@ -16,6 +16,8 @@ import jax import time import os +import shutil +import tempfile from maxdiffusion.checkpointing.wan_checkpointer_2_1 import WanCheckpointer2_1 from maxdiffusion.checkpointing.wan_checkpointer_2_2 import WanCheckpointer2_2 from maxdiffusion.checkpointing.wan_checkpointer_i2v_2p1 import WanCheckpointerI2V_2_1 @@ -37,6 +39,185 @@ jax.config.update("jax_use_shardy_partitioner", True) +import functools +import hashlib + + +def _non_reusable_aot_revision(): + """Returns a unique identity so unversioned/dirty development source can never hit old HLO.""" + return aot_cache.non_reusable_aot_revision() + + +@functools.lru_cache(maxsize=1) +def _compute_wan_source_hash() -> str | None: + """Computes a deterministic SHA-256 content hash of all non-test package source files.""" + try: + pkg_dir = os.path.dirname(os.path.abspath(__file__)) + hasher = hashlib.sha256() + py_files = [] + for root, dirs, files in os.walk(pkg_dir): + dirs[:] = [d for d in dirs if d not in ("tests", "__pycache__")] + for f in files: + if f.endswith(".py"): + py_files.append(os.path.join(root, f)) + for path in sorted(set(py_files)): + rel = os.path.relpath(path, pkg_dir).replace(os.sep, "/") + hasher.update(rel.encode("utf-8")) + with open(path, "rb") as f: + hasher.update(f.read()) + return f"src:{hasher.hexdigest()[:16]}" + except Exception: # noqa: BLE001 + return None + + +def _get_pkg_version(module_name: str) -> str: + try: + import importlib.metadata + + return str(importlib.metadata.version(module_name)) + except Exception: # noqa: BLE001 + try: + mod = __import__(module_name) + return str(getattr(mod, "__version__", "unknown")) + except ImportError: + return "not_installed" + + +def _resolve_wan_aot_source_revision(config, commit_hash=None): + """Prefers explicit aot_build_revision (folded with source hash), then package source hash, then git commit hash.""" + src_hash = _compute_wan_source_hash() + explicit = getattr(config, "aot_build_revision", None) + if explicit is not None and str(explicit).strip(): + explicit_str = str(explicit).strip() + if src_hash is not None and src_hash not in explicit_str: + return f"{explicit_str}+{src_hash}" + return explicit_str + if src_hash is not None: + return src_hash + clean_commit = str(commit_hash).strip() if commit_hash is not None and str(commit_hash).strip() else None + if clean_commit is not None: + return clean_commit + return None + + +def _is_reusable_aot_revision(source_revision) -> bool: + return aot_cache.is_reusable_aot_revision(source_revision) + + +def format_video_output_path( + output_dir: str, + run_name: str, + seed: int, + index: int, + filename_prefix: str = "", +) -> str: + """Formats the target mp4 path for a generated video.""" + clean_run_name = str(run_name).strip() if run_name and str(run_name).strip() != "None" else "" + name_part = f"{clean_run_name}_{seed}_{index}" if clean_run_name else f"wan_output_{seed}_{index}" + if output_dir and not output_dir.startswith("gs://") and output_dir != "sdxl-model-finetuned": + return os.path.join(output_dir, f"{filename_prefix}{name_part}.mp4") + return f"{filename_prefix}wan_output_{seed}_{index}.mp4" + + +def _build_wan_aot_metadata(config, mesh, source_revision) -> dict[str, str]: + """Builds the install-time configuration metadata dictionary for Wan AOT caching.""" + first_dev = jax.devices()[0] if jax.devices() else None + if first_dev and hasattr(first_dev, "client") and hasattr(first_dev.client, "platform_version"): + platform_version = str(first_dev.client.platform_version) + elif first_dev and hasattr(first_dev, "platform_version") and first_dev.platform_version: + platform_version = str(first_dev.platform_version) + else: + platform_version = "unknown" + try: + val = getattr(jax.config, "jax_default_matmul_precision", None) + if val is not None: + default_matmul_precision = str(val) + else: + default_matmul_precision = os.environ.get("JAX_DEFAULT_MATMUL_PRECISION", "default") + except Exception: # noqa: BLE001 + default_matmul_precision = os.environ.get("JAX_DEFAULT_MATMUL_PRECISION", "default") + try: + val = getattr(jax.config, "jax_default_prng_impl", None) + if val is not None: + default_prng_impl = str(val) + else: + default_prng_impl = os.environ.get("JAX_DEFAULT_PRNG_IMPL", "threefry2x32") + except Exception: # noqa: BLE001 + default_prng_impl = "default" + try: + val = getattr(jax.config, "jax_enable_x64", None) + jax_enable_x64 = str(bool(val)) if val is not None else os.environ.get("JAX_ENABLE_X64", "False") + except Exception: # noqa: BLE001 + jax_enable_x64 = os.environ.get("JAX_ENABLE_X64", "False") + try: + val = getattr(jax.config, "jax_threefry_partitionable", None) + jax_threefry_partitionable = ( + str(bool(val)) if val is not None else os.environ.get("JAX_THREEFRY_PARTITIONABLE", "default") + ) + except Exception: # noqa: BLE001 + jax_threefry_partitionable = "default" + + return { + "model": str(getattr(config, "pretrained_model_name_or_path", "")), + "wan_transformer_pretrained_model_name_or_path": str( + getattr(config, "wan_transformer_pretrained_model_name_or_path", "") + ), + "attention": str(getattr(config, "attention", "")), + # Kernel block sizes change the lowered graph, not the input + # shapes — they must key the executable or a re-tuned config + # would silently hit stale binaries. + "flash_block_sizes": str(getattr(config, "flash_block_sizes", {})), + "mesh_shape": str(mesh.shape if mesh is not None else ()), + "vae_spatial": str(getattr(config, "vae_spatial", 8)), + "vae_decode_chunk": str(getattr(config, "vae_decode_chunk", 1)), + "vae_encode_chunk": str(getattr(config, "vae_encode_chunk", 0)), + "replicate_vae": str(getattr(config, "replicate_vae", False)), + "vae_logical_axis_rules": str(getattr(config, "vae_logical_axis_rules", ())), + "vae_weights_dtype": str(getattr(config, "vae_weights_dtype", "bfloat16")), + "vae_dtype": str(getattr(config, "vae_dtype", "bfloat16")), + "weights_dtype": str(getattr(config, "weights_dtype", "")), + "activations_dtype": str(getattr(config, "activations_dtype", "")), + "scan_layers": str(getattr(config, "scan_layers", True)), + "remat_policy": str(getattr(config, "remat_policy", "NONE")), + "ulysses_shards": str(getattr(config, "ulysses_shards", 1)), + "ulysses_attention_chunks": str(getattr(config, "ulysses_attention_chunks", 1)), + "use_k_centering": str(getattr(config, "use_k_centering", "auto")), + "use_kv_cache": str(getattr(config, "use_kv_cache", False)), + "use_cfg_cache": str(getattr(config, "use_cfg_cache", False)), + "use_magcache": str(getattr(config, "use_magcache", False)), + "use_sen_cache": str(getattr(config, "use_sen_cache", False)), + "use_qwix_quantization": str(getattr(config, "use_qwix_quantization", False)), + "quantization": str(getattr(config, "quantization", "")), + "qwix_module_path": str(getattr(config, "qwix_module_path", "")), + "enable_lora": str(getattr(config, "enable_lora", False)), + "lora_config": str(getattr(config, "lora_config", {})), + "flash_min_seq_length": str(getattr(config, "flash_min_seq_length", 4096)), + "mask_padding_tokens": str(getattr(config, "mask_padding_tokens", True)), + "precision": str(getattr(config, "precision", "default")), + "logical_axis_rules": str(getattr(config, "logical_axis_rules", ())), + "attention_sharding_uniform": str(getattr(config, "attention_sharding_uniform", True)), + "allow_split_physical_axes": str(getattr(config, "allow_split_physical_axes", False)), + "device_kind": str(first_dev.device_kind if first_dev else "unknown"), + "platform_version": str(platform_version), + "process_count": str(jax.process_count()), + "use_base2_exp": str(getattr(config, "use_base2_exp", True)), + "use_experimental_scheduler": str(getattr(config, "use_experimental_scheduler", False)), + "libtpu_init_args": os.environ.get("LIBTPU_INIT_ARGS", ""), + "xla_flags": os.environ.get("XLA_FLAGS", ""), + "jax": jax.__version__, + "jaxlib": _get_pkg_version("jaxlib"), + "libtpu": _get_pkg_version("libtpu"), + "flax": _get_pkg_version("flax"), + "qwix": _get_pkg_version("qwix"), + "tokamax": _get_pkg_version("tokamax"), + "default_matmul_precision": default_matmul_precision, + "default_prng_impl": default_prng_impl, + "jax_enable_x64": jax_enable_x64, + "jax_threefry_partitionable": jax_threefry_partitionable, + "source_revision": source_revision if source_revision else _non_reusable_aot_revision(), + } + + def call_pipeline(config, pipeline, prompt, negative_prompt, num_inference_steps=None): model_key = config.model_name model_type = config.model_type @@ -120,6 +301,40 @@ def call_pipeline(config, pipeline, prompt, negative_prompt, num_inference_steps raise ValueError(f"Unsupported model_name for T2V in config: {model_key}") +def _save_generated_video( + config, + video_frames, + index: int, + filename_prefix: str, + gcs_output_path: str, + saved_paths: list[str], + delete_local_after_gcs: bool = False, +) -> None: + """Exports a single video on process 0, optionally uploads to GCS, and records the path.""" + if jax.process_index() != 0: + return + import numpy as np + + output_dir = getattr(config, "output_dir", "") + video_path = format_video_output_path( + output_dir, + getattr(config, "run_name", ""), + config.seed, + index, + filename_prefix, + ) + if output_dir and not output_dir.startswith("gs://") and output_dir != "sdxl-model-finetuned": + os.makedirs(output_dir, exist_ok=True) + frames_np = np.asarray(video_frames) + export_to_video(frames_np, video_path, fps=config.fps) + saved_paths.append(video_path) + max_logging.log(f"Saved video to {video_path}") + if gcs_output_path: + max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") + if delete_local_after_gcs: + max_utils.delete_file(video_path) + + def inference_generate_video(config, pipeline, filename_prefix=""): s0 = time.perf_counter() prompt_file = getattr(config, "prompt_file", "") @@ -138,28 +353,36 @@ def inference_generate_video(config, pipeline, filename_prefix=""): if not is_multi_prompt: prompt = [prompts[0]] * batch_size negative_prompt = [config.negative_prompt] * batch_size - videos = call_pipeline(config, pipeline, prompt, negative_prompt) + outputs = call_pipeline(config, pipeline, prompt, negative_prompt) + videos = outputs[0] if isinstance(outputs, tuple) else outputs max_logging.log(f"video {filename_prefix}, generation time: {(time.perf_counter() - s0):.2f}s") for i in range(len(videos)): - video_path = f"{filename_prefix}wan_output_{config.seed}_{i}.mp4" - export_to_video(videos[i], video_path, fps=config.fps) - saved_video_paths.append(video_path) - if gcs_output_path: - max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") - max_utils.delete_file(f"./{video_path}") + _save_generated_video( + config, + videos[i], + i, + filename_prefix, + gcs_output_path, + saved_video_paths, + delete_local_after_gcs=True, + ) else: for i, padded_chunk, actual_chunk_len in max_utils.chunk_and_pad(prompts, batch_size): negative_prompt = [config.negative_prompt] * batch_size - videos = call_pipeline(config, pipeline, padded_chunk, negative_prompt) + outputs = call_pipeline(config, pipeline, padded_chunk, negative_prompt) + videos = outputs[0] if isinstance(outputs, tuple) else outputs for j in range(actual_chunk_len): prompt_idx = i + j - video_path = f"{filename_prefix}wan_output_{config.seed}_{prompt_idx}.mp4" - export_to_video(videos[j], video_path, fps=config.fps) - saved_video_paths.append(video_path) - if gcs_output_path: - max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") - max_utils.delete_file(f"./{video_path}") + _save_generated_video( + config, + videos[j], + prompt_idx, + filename_prefix, + gcs_output_path, + saved_video_paths, + delete_local_after_gcs=True, + ) max_logging.log(f"all videos {filename_prefix}, total generation time: {(time.perf_counter() - s0):.2f}s") return saved_video_paths @@ -205,6 +428,42 @@ def maybe_tune_block_sizes(config): ) +def _plan_wan_aot_cache(config, source_revision) -> tuple[str, bool]: + """Decides how `run()` installs the AOT cache. + + Returns `(persistent_dir, use_ephemeral)`: + * default (no aot_cache_dir, no opt-in): ("", False) -> aot_cache is not + installed on a directory and every call is plain jax.jit; + * aot_cache_dir with a reusable source revision: (dir, False). The + revision is normally the content hash of the package's .py files, which + is always reusable; + * aot_cache_dir downgraded because the revision is not reusable (only when + that hash cannot be computed and the git revision is dirty/unversioned, + or an explicit `aot_build_revision` is marked dirty:/unversioned:), or + `enable_zero_execution_warmup=True` without a persistent dir: + ("", True) -> a temporary directory, removed when `run()` returns. + A gs:// aot_cache_dir raises ValueError regardless of the revision. + """ + requested_aot_cache_dir = getattr(config, "aot_cache_dir", "") + if str(requested_aot_cache_dir).startswith("gs://"): + raise ValueError( + f"aot_cache_dir must be a local POSIX directory path; gs:// URIs are not supported (got {requested_aot_cache_dir!r})." + ) + enable_zero_exec_warmup = bool(getattr(config, "enable_zero_execution_warmup", False)) + aot_cache_dir = requested_aot_cache_dir + if aot_cache_dir and not _is_reusable_aot_revision(source_revision): + max_logging.log( + "[aot] Persistent Wan AOT caching is disabled for this development run; " + "using ephemeral cache for zero-execution warmup." + ) + aot_cache_dir = "" + # Only use an ephemeral cache directory when zero-execution warmup is + # explicitly opted into or when a requested aot_cache_dir was downgraded; + # otherwise leave aot_cache disabled so default callers use plain jax.jit. + use_ephemeral_cache = not aot_cache_dir and (bool(requested_aot_cache_dir) or enable_zero_exec_warmup) + return aot_cache_dir, use_ephemeral_cache + + def run(config, pipeline=None, filename_prefix="", commit_hash=None): model_key = config.model_name if pipeline is None: @@ -297,31 +556,41 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): # Per-shape AOT executable cache: deserialization starts on background # threads now and overlaps the remaining setup; unknown shapes silently - # fall back to jit and are serialized by save_pending() after warmup. - aot_cache.install( - getattr(config, "aot_cache_dir", ""), - meta={ - "model": config.pretrained_model_name_or_path, - "attention": config.attention, - # Kernel block sizes change the lowered graph, not the input - # shapes — they must key the executable or a re-tuned config - # would silently hit stale binaries. - "flash_block_sizes": str(config.flash_block_sizes), - "mesh_shape": str(pipeline.mesh.shape), - "vae_spatial": str(config.vae_spatial), - "vae_decode_chunk": str(config.vae_decode_chunk), - "weights_dtype": str(config.weights_dtype), - "activations_dtype": str(config.activations_dtype), - "scan_layers": str(config.scan_layers), - "jax": jax.__version__, - **aot_cache.extract_svg_meta(config, pipeline), - }, - mesh=pipeline.mesh, - ) - # Deserialization is seconds and warmup must see the loaded executables - # to hit them; without this the first call races the loader threads. - aot_cache.wait_for_loads() - + # fall back to jit and are serialized by save_pending() after warmup and + # again after generation (persistent cache only). + detected_revision = commit_hash if commit_hash is not None else max_utils.get_git_commit_hash() + source_revision = _resolve_wan_aot_source_revision(config, detected_revision) + aot_cache_dir, use_ephemeral_cache = _plan_wan_aot_cache(config, source_revision) + install_cache_dir = "" + try: + if use_ephemeral_cache: + install_cache_dir = tempfile.mkdtemp(prefix="wan_aot_ephemeral_") + else: + install_cache_dir = aot_cache_dir + + aot_metadata = _build_wan_aot_metadata(config, pipeline.mesh, source_revision) + aot_cache.install( + install_cache_dir, + meta={**aot_metadata, **aot_cache.extract_svg_meta(config, pipeline)}, + mesh=pipeline.mesh, + ) + # Deserialization is seconds and warmup must see the loaded executables + # to hit them; without this the first call races the loader threads. + aot_cache.wait_for_loads() + + return _warmup_and_generate(config, pipeline, filename_prefix, writer, load_time, aot_cache_dir) + finally: + # Tear the ephemeral cache down even if metadata/install/warmup/generation + # raises, so a failed run neither leaks the temp dir nor leaves aot_cache + # installed on it. + if use_ephemeral_cache: + aot_cache.install("", meta={}, mesh=None) + if install_cache_dir: + shutil.rmtree(install_cache_dir, ignore_errors=True) + + +def _warmup_and_generate(config, pipeline, filename_prefix, writer, load_time, aot_cache_dir): + """Warmup + generation + optional profiling run. `aot_cache_dir` is the PERSISTENT dir ('' if none).""" s0 = time.perf_counter() # Disable profiler for the first two runs to avoid duplicate uploads @@ -341,11 +610,15 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): f"Num steps: {config.num_inference_steps}, height: {config.height}, width: {config.width}," f" frames: {config.num_frames}, total prompts: {len(prompts)}" ) - # Warmup with 2 denoising steps instead of a full run: step 0 runs the - # high-noise transformer and step 1 crosses the boundary to the low-noise - # one (WAN 2.2), so every executable of the full run (both transformers, - # text encoder, VAE decode) gets compiled at a fraction of the cost. The - # step count only changes the Python loop trip count, not traced shapes. + # Warmup with 2 denoising steps instead of a full run. The step count only + # changes the Python loop trip count, not traced shapes. With flow_shift=12 + # (WAN 2.2 T2V) both warmup timesteps (t=999, 923) stay above the boundary + # (875), whereas flow_shift=5 (WAN 2.2 I2V 27B) produces t=[999, 833] across + # the boundary (900). In both cases, compile_experts() in run_inference_2_2 / + # run_inference_2_2_i2v explicitly compiles both high- and low-noise experts + # before the loop when warmup_mode() is active. Together this compiles every + # executable of the full run (both transformers, text encoder, VAE decode) at + # a fraction of the cost. warmup_steps = min(2, config.num_inference_steps) max_logging.log(f"Compile warmup: {warmup_steps} denoising steps") # Zero-execution warmup: wrapped transformer passes lower+compile (or @@ -353,16 +626,25 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): # the warmup pays compile time only, never real denoise compute. The # returned videos are garbage by design and are discarded below. with aot_cache.warmup_mode(): - videos = call_pipeline(config, pipeline, warmup_prompt, warmup_negative_prompt, num_inference_steps=warmup_steps) + videos = call_pipeline( + config, + pipeline, + warmup_prompt, + warmup_negative_prompt, + num_inference_steps=warmup_steps, + ) if isinstance(videos, tuple): videos, warmup_trace = videos warmup_str = ", ".join(f"{stage}={seconds:.1f}s" for stage, seconds in warmup_trace.items()) max_logging.log(f"Warmup breakdown: {warmup_str}") - # Serialize any newly-compiled shapes synchronously while still inside - # warmup-accounted time; a background save would compete with the first - # real generation (DiffusionServing PR#39 first-generation-stall lesson). - aot_cache.save_pending() + # Serialize newly-compiled shapes synchronously inside warmup-accounted time + # (a background save would stall the first request). Skipped for the + # ephemeral (non-persistent) cache dir. + if aot_cache_dir: + aot_cache.save_pending() + else: + aot_cache.clear_pending() max_logging.log("===================== Model details =======================") max_logging.log(f"model name: {config.model_name}") @@ -393,11 +675,14 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): videos = outputs trace = {} for i in range(len(videos)): - video_path = f"{filename_prefix}wan_output_{config.seed}_{i}.mp4" - export_to_video(videos[i], video_path, fps=config.fps) - saved_video_path.append(video_path) - if gcs_output_path: - max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") + _save_generated_video( + config, + videos[i], + i, + filename_prefix, + gcs_output_path, + saved_video_path, + ) else: trace = {} for i, padded_chunk, actual_chunk_len in max_utils.chunk_and_pad(prompts, batch_size): @@ -410,13 +695,20 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): videos = outputs for j in range(actual_chunk_len): prompt_idx = i + j - video_path = f"{filename_prefix}wan_output_{config.seed}_{prompt_idx}.mp4" - export_to_video(videos[j], video_path, fps=config.fps) - saved_video_path.append(video_path) - if gcs_output_path: - max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") + _save_generated_video( + config, + videos[j], + prompt_idx, + filename_prefix, + gcs_output_path, + saved_video_path, + ) generation_time = time.perf_counter() - s0 + if aot_cache_dir: + aot_cache.save_pending() + else: + aot_cache.clear_pending() max_logging.log(f"generation_time: {generation_time}") if writer and jax.process_index() == 0: writer.add_scalar("inference/generation_time", generation_time, global_step=0) diff --git a/src/maxdiffusion/max_utils.py b/src/maxdiffusion/max_utils.py index 171fb91ea..aeab46643 100644 --- a/src/maxdiffusion/max_utils.py +++ b/src/maxdiffusion/max_utils.py @@ -468,10 +468,27 @@ def delete_file(file_path: str): max_logging.log(f"The file '{file_path}' does not exist.") -def get_git_commit_hash(): +def get_git_commit_hash(check_dirty: bool = True): """Tries to get the current Git commit hash, for run provenance.""" try: - return subprocess.check_output(["git", "rev-parse", "HEAD"]).strip().decode("utf-8") + repo_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + commit = subprocess.check_output(["git", "-C", repo_dir, "rev-parse", "HEAD"]).strip().decode("utf-8") + if check_dirty: + status = ( + subprocess.check_output([ + "git", + "-C", + repo_dir, + "status", + "--porcelain", + "--untracked-files=no", + ]) + .strip() + .decode("utf-8") + ) + if status: + return f"{commit}-dirty" + return commit except subprocess.CalledProcessError: max_logging.log("Warning: 'git rev-parse HEAD' failed. Not running in a git repo?") return None @@ -586,11 +603,16 @@ def create_device_mesh(config, devices=None, logging=True): if multi_slice_env: dcn_parallelism = fill_unspecified_mesh_axes(dcn_parallelism, num_slices, "DCN") mesh = mesh_utils.create_hybrid_device_mesh( - ici_parallelism, dcn_parallelism, devices, allow_split_physical_axes=config.allow_split_physical_axes + ici_parallelism, + dcn_parallelism, + devices, + allow_split_physical_axes=config.allow_split_physical_axes, ) else: mesh = mesh_utils.create_device_mesh( - ici_parallelism, devices, allow_split_physical_axes=config.allow_split_physical_axes + ici_parallelism, + devices, + allow_split_physical_axes=config.allow_split_physical_axes, ) if logging: diff --git a/src/maxdiffusion/models/attention_flax.py b/src/maxdiffusion/models/attention_flax.py index 998d806da..9d52e3f98 100644 --- a/src/maxdiffusion/models/attention_flax.py +++ b/src/maxdiffusion/models/attention_flax.py @@ -2181,11 +2181,10 @@ def _to_bshd(x: Array, n_heads: int) -> Array: axis=1, ).reshape(b * heads, s_k, -1) - if float32_qk_product: - query_states = query_states.astype(jnp.float32) - key_states = key_states.astype(jnp.float32) - if use_memory_efficient_attention and attention_mask is None: + if float32_qk_product: + query_states = query_states.astype(jnp.float32) + key_states = key_states.astype(jnp.float32) query_states = query_states.transpose(1, 0, 2) key_states = key_states.transpose(1, 0, 2) value_states = value_states.transpose(1, 0, 2) @@ -2213,10 +2212,21 @@ def _to_bshd(x: Array, n_heads: int) -> Array: hidden_states = hidden_states.transpose(1, 0, 2) else: + preferred_element_type = jnp.float32 if float32_qk_product else None if split_head_dim: - attention_scores = jnp.einsum("b t n h, b f n h -> b n f t", key_states, query_states) + attention_scores = jnp.einsum( + "b t n h, b f n h -> b n f t", + key_states, + query_states, + preferred_element_type=preferred_element_type, + ) else: - attention_scores = jnp.einsum("b i d, b j d->b i j", query_states, key_states) + attention_scores = jnp.einsum( + "b i d, b j d->b i j", + query_states, + key_states, + preferred_element_type=preferred_element_type, + ) attention_scores = attention_scores * scale if attention_mask is not None: diff --git a/src/maxdiffusion/models/wan/transformers/transformer_wan.py b/src/maxdiffusion/models/wan/transformers/transformer_wan.py index f147e4414..fccf45c6f 100644 --- a/src/maxdiffusion/models/wan/transformers/transformer_wan.py +++ b/src/maxdiffusion/models/wan/transformers/transformer_wan.py @@ -253,8 +253,7 @@ def __init__( def __call__(self, x: jax.Array) -> jax.Array: with jax.named_scope("gelu"): - x = self.proj(x) - return nnx.gelu(x) + return nnx.gelu(self.proj(x)) class WanFeedForward(nnx.Module): @@ -322,12 +321,13 @@ def __call__( deterministic: bool = True, rngs: nnx.Rngs = None, ) -> jax.Array: - hidden_states = self.act_fn(hidden_states) # Output is (4, 75600, 13824) + hidden_states = self.act_fn(hidden_states) hidden_states = checkpoint_name(hidden_states, "ffn_activation") if self.drop_out.rate > 0: hidden_states = self.drop_out(hidden_states, deterministic=deterministic, rngs=rngs) with jax.named_scope("proj_out"): - return self.proj_out(hidden_states) # output is (4, 75600, 5120) + hidden_states = self.proj_out(hidden_states) + return hidden_states class WanTransformerBlock(nnx.Module): @@ -361,6 +361,7 @@ def __init__( "use_experimental_scheduler": False, "ulysses_shards": -1, "ulysses_attention_chunks": 1, + "use_k_centering": "auto", **(attention_config or {}), } @@ -494,7 +495,6 @@ def __call__( with self.conditional_named_scope("self_attn_attn"): attn_output = self.attn1( hidden_states=norm_hidden_states, - encoder_hidden_states=norm_hidden_states, rotary_emb=rotary_emb, deterministic=deterministic, rngs=rngs, @@ -594,6 +594,7 @@ def __init__( "use_experimental_scheduler": False, "ulysses_shards": -1, "ulysses_attention_chunks": 1, + "use_k_centering": "auto", **(attention_config or {}), } @@ -954,22 +955,37 @@ def layer_forward(hidden_states, l_kv): scale = scale.squeeze(2) # [B, sl, dim] else: shift, scale = jnp.split(self.scale_shift_table + jnp.expand_dims(temb, axis=1), 2, axis=1) + hidden_states = (self.norm_out(hidden_states.astype(jnp.float32)) * (1 + scale) + shift).astype(hidden_states.dtype) with jax.named_scope("proj_out"): hidden_states = self.proj_out(hidden_states) - hidden_states = hidden_states.reshape( - batch_size, - post_patch_num_frames, - post_patch_height, - post_patch_width, - p_t, - p_h, - p_w, - -1, - ) - hidden_states = jnp.transpose(hidden_states, (0, 7, 1, 4, 2, 5, 3, 6)) - hidden_states = hidden_states.reshape(batch_size, -1, num_frames, height, width) + if p_t == 1: + # Lossless HLO optimization: collapse p_t=1 dimension to avoid 8D non-contiguous stride copies + hidden_states = hidden_states.reshape( + batch_size, + post_patch_num_frames, + post_patch_height, + post_patch_width, + p_h, + p_w, + -1, + ) + hidden_states = jnp.transpose(hidden_states, (0, 6, 1, 2, 4, 3, 5)) + hidden_states = hidden_states.reshape(batch_size, -1, num_frames, height, width) + else: + hidden_states = hidden_states.reshape( + batch_size, + post_patch_num_frames, + post_patch_height, + post_patch_width, + p_t, + p_h, + p_w, + -1, + ) + hidden_states = jnp.transpose(hidden_states, (0, 7, 1, 4, 2, 5, 3, 6)) + hidden_states = hidden_states.reshape(batch_size, -1, num_frames, height, width) if return_residual: return hidden_states, residual_x diff --git a/src/maxdiffusion/models/wan/wan_utils.py b/src/maxdiffusion/models/wan/wan_utils.py index 4c9053219..77134f234 100644 --- a/src/maxdiffusion/models/wan/wan_utils.py +++ b/src/maxdiffusion/models/wan/wan_utils.py @@ -318,29 +318,135 @@ def _torch_tensor_to_numpy(tensor: torch.Tensor) -> np.ndarray: return tensor.numpy() +# v2: the source fingerprint no longer includes absolute paths (v1 caches are +# rejected once and re-converted), and a missing fingerprint is rejected. +# v3: the fingerprint includes the HF snapshot revision (read before symlinks +# are resolved) and the subfolder. v2 manifests whose fingerprint matches the +# v2 formula for the same index are migrated in place on first load. +_CONVERTED_WEIGHTS_FORMAT_VERSION = 3 +_LEGACY_MIGRATABLE_FORMAT_VERSION = 2 + + def _converted_key_to_filename(flax_key: tuple) -> str: return ".".join(str(k) for k in flax_key) + ".npy" -def try_load_converted_weights(cache_dir: str, eval_shapes: dict, cast_dtype_fn: Optional[Callable]) -> Optional[dict]: +def _snapshot_revision(index_file_path: str) -> str: + """HF snapshot revision for an index under `.../snapshots//`, else "". + + Read from the path as given: HF snapshot entries are symlinks into `blobs/`, + so resolving them first (realpath) would always lose the revision. Scans from + the right (and prefers a `models--*/snapshots/` pair when present) + so a parent mount directory named `snapshots` does not shadow the HF cache. + """ + parts = os.path.normpath(os.path.abspath(index_file_path)).split(os.sep) + for i in range(len(parts) - 2, 0, -1): + if parts[i] == "snapshots" and parts[i - 1].startswith("models--"): + return parts[i + 1] + for i in range(len(parts) - 2, -1, -1): + if parts[i] == "snapshots": + return parts[i + 1] + return "" + + +def _compute_source_checkpoint_fingerprint(index_file_path: str, source_id: str = "", subfolder: str = "") -> str: + """Fingerprints the source checkpoint WITHOUT depending on where it is mounted. + + Hashes `source_id` (the HF repo id, or "" for a local directory), the HF + snapshot revision when the index lives under `.../snapshots//`, + `subfolder` (Wan 2.2's transformer and transformer_2 ship byte-identical + index files), and the index file's contents (shard names, tensor->shard map, + total size). Absolute paths are deliberately excluded, so a different + HF_HOME or mount point does not invalidate a ~28 GB converted cache. + """ + import hashlib + + hasher = hashlib.sha256() + for field in (source_id, _snapshot_revision(index_file_path), subfolder): + hasher.update(field.encode("utf-8") + b"\0") + with open(index_file_path, "rb") as f: + hasher.update(f.read()) + return hasher.hexdigest()[:16] + + +def _legacy_v2_source_fingerprint(index_file_path: str, source_id: str = "") -> str: + """The format-v2 fingerprint (revision read after realpath, no subfolder); for migration only.""" + import hashlib + + parts = os.path.normpath(os.path.realpath(index_file_path)).split(os.sep) + revision = parts[parts.index("snapshots") + 1] if "snapshots" in parts[:-1] else "" + hasher = hashlib.sha256() + for field in (source_id, revision): + hasher.update(field.encode("utf-8") + b"\0") + with open(index_file_path, "rb") as f: + hasher.update(f.read()) + return hasher.hexdigest()[:16] + + +def _free_bytes(path: str) -> int: + """Free bytes on the filesystem holding `path` (or its nearest existing ancestor).""" + probe = os.path.abspath(path) + while not os.path.exists(probe) and os.path.dirname(probe) != probe: + probe = os.path.dirname(probe) + return shutil.disk_usage(probe).free + + +# Headroom kept free when writing large caches, so a nearly-full disk fails +# with an actionable message instead of at 100%. +_DISK_HEADROOM_BYTES = 2 * 1024**3 + + +def try_load_converted_weights( + cache_dir: str, + eval_shapes: dict, + cast_dtype_fn: Optional[Callable], + source_fingerprint: Optional[str], + legacy_source_fingerprint: Optional[str] = None, +) -> Optional[dict]: """Loads a converted-weights cache as mmapped arrays, or None on mismatch. The torch->flax conversion (transpose + scan-stack + cast) is a pure function of the checkpoint, so it is paid once and memoized on disk. Keys/shapes are validated against eval_shapes and dtypes against cast_dtype_fn, so a policy or model change falls back to a fresh - conversion (which re-saves). + conversion (which re-saves). Fail-closed on provenance: a manifest without + a source fingerprint, or a call without one, is treated as a miss. + + A format-v2 manifest is accepted only if its fingerprint equals + `legacy_source_fingerprint` (the v2 formula for the same index); after a + successful load its header is rewritten to the current version and + `source_fingerprint`, so the migration happens once. """ manifest_path = os.path.join(cache_dir, "manifest.json") if not os.path.isfile(manifest_path): return None try: with open(manifest_path, "r") as f: - manifest = json.load(f) - expected_keys = set(flatten_dict(eval_shapes).keys()) + raw_manifest = json.load(f) + if isinstance(raw_manifest, dict) and "__meta__" in raw_manifest: + meta_header = raw_manifest["__meta__"] + version = meta_header.get("format_version") + migrate = version == _LEGACY_MIGRATABLE_FORMAT_VERSION and bool(legacy_source_fingerprint) + if version != _CONVERTED_WEIGHTS_FORMAT_VERSION and not migrate: + raise ValueError(f"converted cache format_version {version} != {_CONVERTED_WEIGHTS_FORMAT_VERSION}") + cached_fingerprint = meta_header.get("source_fingerprint") + if not cached_fingerprint: + raise ValueError("converted cache has no source fingerprint (cannot verify its source checkpoint)") + if not source_fingerprint: + raise ValueError("no source fingerprint available for the requested checkpoint") + expected_fingerprint = legacy_source_fingerprint if migrate else source_fingerprint + if cached_fingerprint != expected_fingerprint: + raise ValueError("source checkpoint fingerprint changed") + manifest = {k: v for k, v in raw_manifest.items() if k != "__meta__"} + else: + raise ValueError("converted cache manifest missing __meta__ format header (written by an older version)") + flat_eval = flatten_dict(eval_shapes) + expected_keys = set(flat_eval.keys()) def load_one(key_str, meta): flax_key = _tuple_str_to_int(tuple(key_str.split("."))) + if flax_key not in flat_eval: + raise ValueError(f"unexpected key {key_str} in converted cache") logical_dtype = np.dtype(meta["dtype"]) if cast_dtype_fn is not None and logical_dtype != np.dtype(cast_dtype_fn(flax_key)): raise ValueError(f"dtype policy changed for {key_str}") @@ -351,8 +457,9 @@ def load_one(key_str, meta): # Non-native dtypes (bf16/fp8) are stored as same-width uints: # npy cannot resolve ml_dtypes descriptors on all paths. value = value.view(logical_dtype) - if tuple(value.shape) != tuple(meta["shape"]): - raise ValueError(f"shape changed for {key_str}") + expected_shape = tuple(flat_eval[flax_key].shape) + if tuple(value.shape) != tuple(meta["shape"]) or tuple(value.shape) != expected_shape: + raise ValueError(f"shape changed for {key_str}: got {value.shape}, expected {expected_shape}") return flax_key, value flax_state_dict = {} @@ -361,31 +468,130 @@ def load_one(key_str, meta): flax_state_dict[flax_key] = value if set(flax_state_dict.keys()) != expected_keys: return None + if migrate: + _migrate_manifest_header(manifest_path, raw_manifest, source_fingerprint) return unflatten_dict(flax_state_dict) except (OSError, ValueError, KeyError, TypeError) as e: - max_logging.log(f"Converted-weights cache unusable ({e}); reconverting") + max_logging.log( + f"WARNING: converted-weights cache at {cache_dir} is unusable ({e}). The transformer will be " + "re-converted from the source safetensors (downloading any shards missing from the HF cache) " + "and the cache re-saved; for Wan 2.2 A14B that is ~28 GB per expert of disk." + ) return None -def save_converted_weights(cache_dir: str, flat_state_dict: dict) -> None: - """Writes the converted tree as per-tensor .npy + manifest, atomically.""" - tmp_dir = f"{cache_dir}.tmp.{os.getpid()}" +def _migrate_manifest_header(manifest_path: str, raw_manifest: dict, source_fingerprint: str) -> None: + """Best-effort atomic rewrite of a validated v2 manifest header to the current format.""" + upgraded = dict(raw_manifest) + upgraded["__meta__"] = { + **raw_manifest["__meta__"], + "format_version": _CONVERTED_WEIGHTS_FORMAT_VERSION, + "source_fingerprint": source_fingerprint, + } + tmp_path = manifest_path + ".tmp" + try: + with open(tmp_path, "w") as f: + json.dump(upgraded, f) + os.replace(tmp_path, manifest_path) + max_logging.log( + f"Migrated converted-weights manifest {manifest_path} to format v{_CONVERTED_WEIGHTS_FORMAT_VERSION} " + "(fingerprint now includes the snapshot revision and subfolder)." + ) + except OSError as e: + max_logging.log(f"WARNING: could not migrate {manifest_path} ({e}); it will be re-checked on the next load.") + + +def save_converted_weights( + cache_dir: str, + flat_state_dict: dict, + source_fingerprint: str, +) -> bool: + """Writes the converted tree as per-tensor .npy + manifest, atomically. + + Returns False (after logging why) instead of writing when there is not + enough free disk for the tree plus headroom; the in-memory weights are + unaffected, only the memoization is skipped. + """ + import uuid + + if not source_fingerprint: + raise ValueError("save_converted_weights requires a source_fingerprint (unverifiable caches are never written).") + + suffix = f"{os.getpid()}.{uuid.uuid4().hex[:8]}" + # Remove any invalidated cache_dir before writing tmp_dir so peak disk usage + # stays at 1x instead of 2x-3x when re-saving a ~28GB expert. + if os.path.isdir(cache_dir): + stale_dir = f"{cache_dir}.stale.{suffix}" + try: + os.rename(cache_dir, stale_dir) + shutil.rmtree(stale_dir, ignore_errors=True) + except OSError: + shutil.rmtree(cache_dir, ignore_errors=True) + + needed = sum(int(v.nbytes) for v in flat_state_dict.values()) + _DISK_HEADROOM_BYTES + free = _free_bytes(cache_dir) + if free < needed: + max_logging.log( + f"WARNING: not saving the converted-weights cache to {cache_dir}: needs ~{needed / 1e9:.1f} GB " + f"(incl. headroom) but only {free / 1e9:.1f} GB is free. Free disk space or point " + "converted_weights_dir at a larger volume to enable the fast warm start." + ) + return False + + tmp_dir = f"{cache_dir}.tmp.{suffix}" os.makedirs(tmp_dir, exist_ok=True) - manifest = {} - uint_by_width = {1: np.uint8, 2: np.uint16, 4: np.uint32} - for flax_key, value in flat_state_dict.items(): - filename = _converted_key_to_filename(flax_key) - bitview = value.dtype.kind not in "fiub" # ml_dtypes (bf16/fp8) etc. - stored = value.view(uint_by_width[value.dtype.itemsize]) if bitview else value - np.save(os.path.join(tmp_dir, filename), stored) - key_str = ".".join(str(k) for k in flax_key) - manifest[key_str] = {"file": filename, "shape": list(value.shape), "dtype": str(value.dtype), "bitview": bitview} - with open(os.path.join(tmp_dir, "manifest.json"), "w") as f: - json.dump(manifest, f) + try: + manifest = { + "__meta__": { + "format_version": _CONVERTED_WEIGHTS_FORMAT_VERSION, + "source_fingerprint": source_fingerprint, + } + } + uint_by_width = {1: np.uint8, 2: np.uint16, 4: np.uint32} + for flax_key, value in flat_state_dict.items(): + filename = _converted_key_to_filename(flax_key) + bitview = value.dtype.kind not in "fiub" # ml_dtypes (bf16/fp8) etc. + stored = value.view(uint_by_width[value.dtype.itemsize]) if bitview else value + np.save(os.path.join(tmp_dir, filename), stored) + key_str = ".".join(str(k) for k in flax_key) + manifest[key_str] = {"file": filename, "shape": list(value.shape), "dtype": str(value.dtype), "bitview": bitview} + with open(os.path.join(tmp_dir, "manifest.json"), "w") as f: + json.dump(manifest, f) + except Exception: + shutil.rmtree(tmp_dir, ignore_errors=True) + raise try: os.rename(tmp_dir, cache_dir) except OSError: shutil.rmtree(tmp_dir, ignore_errors=True) # another process won the race + return True + + +def _check_disk_for_shard_download(repo_id: str, subfolder: str, model_files: list, index_dict: dict) -> None: + """Raises an actionable error if downloading the missing shards would fill the disk.""" + try: + from huggingface_hub import constants as hf_constants + from huggingface_hub import try_to_load_from_cache + + missing = [ + f + for f in model_files + if not isinstance(try_to_load_from_cache(repo_id, f"{subfolder}/{f}" if subfolder else f), str) + ] + total_size = int(index_dict.get("metadata", {}).get("total_size", 0)) + hub_cache = hf_constants.HF_HUB_CACHE + except Exception: # noqa: BLE001 - best effort: never block loading on the estimate itself + return + if not missing or not total_size: + return + needed = total_size * len(missing) // max(len(model_files), 1) + _DISK_HEADROOM_BYTES + free = _free_bytes(hub_cache) + if free < needed: + raise OSError( + f"Refusing to download {len(missing)}/{len(model_files)} shards of {repo_id}/{subfolder}: needs ~{needed / 1e9:.1f} " + f"GB (incl. headroom) in {hub_cache} but only {free / 1e9:.1f} GB is free. Free disk space, set HF_HOME to a " + "larger volume, or restore a valid converted_weights_dir cache." + ) def load_base_wan_transformer( @@ -416,22 +622,49 @@ def load_base_wan_transformer( Returns a nested dict of numpy arrays (host memory). """ del device # weights stay in plain host numpy until device_put by the caller - if converted_cache_dir: + filename = "diffusion_pytorch_model.safetensors.index.json" + local_files = os.path.isdir(pretrained_model_name_or_path) + source_id = "" if local_files else pretrained_model_name_or_path + + def try_cache(index_path): + if not converted_cache_dir or not index_path: + return None t_start = time.perf_counter() - cached = try_load_converted_weights(converted_cache_dir, eval_shapes, cast_dtype_fn) + cached = try_load_converted_weights( + converted_cache_dir, + eval_shapes, + cast_dtype_fn, + source_fingerprint=_compute_source_checkpoint_fingerprint(index_path, source_id, subfolder), + legacy_source_fingerprint=_legacy_v2_source_fingerprint(index_path, source_id), + ) if cached is not None: - max_logging.log( - f"Loaded converted {subfolder or 'transformer'} weights (mmap) in {time.perf_counter() - t_start:.1f}s" - ) - return cached - filename = "diffusion_pytorch_model.safetensors.index.json" - local_files = False - if os.path.isdir(pretrained_model_name_or_path): + max_logging.log(f"Loaded converted {subfolder or 'transformer'} weights in {time.perf_counter() - t_start:.1f}s") + return cached + + index_file_path = None + checked_index = None + if local_files: index_file_path = os.path.join(pretrained_model_name_or_path, subfolder, filename) if not os.path.isfile(index_file_path): raise FileNotFoundError(f"File {index_file_path} not found for local directory.") - local_files = True elif hf_download: + # Warm start without the network: if the index is already in the local HF + # cache, check the converted cache against it first. Trade-off: while a + # valid converted cache exists for the locally cached revision, a newer + # upstream revision is not picked up (delete converted_weights_dir to force). + try: + with _HF_METADATA_LOCK: + checked_index = hf_hub_download( + pretrained_model_name_or_path, + subfolder=subfolder, + filename=filename, + local_files_only=True, + ) + except Exception: # noqa: BLE001 - not in the local HF cache: fall through to the network path + checked_index = None + cached = try_cache(checked_index) + if cached is not None: + return cached # download the index file for sharded models. with _HF_METADATA_LOCK: index_file_path = hf_hub_download( @@ -439,6 +672,13 @@ def load_base_wan_transformer( subfolder=subfolder, filename=filename, ) + if index_file_path is None: + raise ValueError(f"{pretrained_model_name_or_path} is not a local directory and hf_download is False.") + if index_file_path != checked_index: + cached = try_cache(index_file_path) + if cached is not None: + return cached + source_fingerprint = _compute_source_checkpoint_fingerprint(index_file_path, source_id, subfolder) t_start = time.perf_counter() with open(index_file_path, "r") as f: index_dict = json.load(f) @@ -498,6 +738,8 @@ def convert_chunk(ckpt_shard_path, chunk_keys): # across the ~12 shard files. norm_added_q is explicitly ignored by the # diffusers implementation. chunk_size = 16 + if not local_files: + _check_disk_for_shard_download(pretrained_model_name_or_path, subfolder, model_files, index_dict) tasks = [] for model_file in model_files: ckpt_shard_path = resolve_shard_path(model_file) @@ -514,11 +756,11 @@ def convert_chunk(ckpt_shard_path, chunk_keys): future.result() # re-raise conversion errors validate_flax_state_dict(eval_shapes, flax_state_dict) - if converted_cache_dir and not os.path.isdir(converted_cache_dir): + if converted_cache_dir: t_save = time.perf_counter() if jax.process_index() == 0: - save_converted_weights(converted_cache_dir, flax_state_dict) - max_logging.log(f"Saved converted-weights cache to {converted_cache_dir} in {time.perf_counter() - t_save:.1f}s") + if save_converted_weights(converted_cache_dir, flax_state_dict, source_fingerprint=source_fingerprint): + max_logging.log(f"Saved converted-weights cache to {converted_cache_dir} in {time.perf_counter() - t_save:.1f}s") flax_state_dict = unflatten_dict(flax_state_dict) max_logging.log(f"Converted {subfolder or 'transformer'} weights to host arrays in {time.perf_counter() - t_start:.1f}s") return flax_state_dict diff --git a/src/maxdiffusion/pipelines/wan/wan_denoise_utils.py b/src/maxdiffusion/pipelines/wan/wan_denoise_utils.py new file mode 100644 index 000000000..66fd3f245 --- /dev/null +++ b/src/maxdiffusion/pipelines/wan/wan_denoise_utils.py @@ -0,0 +1,82 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Warmup compilation and dispatch-queue bounding shared by the Wan 2.2 denoise loops. + +T2V (``wan_pipeline_2_2``) and I2V (``wan_pipeline_i2v_2p2``) drive the same +two-expert Python loop and need the same warmup and queueing treatment. It +lives here so the two pipelines cannot drift apart. +""" + +import collections +import os +from typing import Any, Callable, Optional, Sequence + +import jax + +from maxdiffusion import aot_cache + + +def compile_experts(branches: Sequence[Callable[[Any], Any]], operands: Any) -> None: + """During warmup, compiles every expert's forward pass without executing it. + + Warmup runs only a couple of denoising steps, and the scheduler can put all + of them on the high-noise expert (with flow_shift=12 a 2-step schedule is + t=[999, 923], both above the boundary). Without this, the low-noise expert + would compile during the first real generation. When the experts share a + signature (the usual case), the second call is an in-memory cache hit. + + Weights are deliberately not primed with a real forward pass. On the + current stack priming added 5-7 s to every warmup, while the first + generation ran within 0.5 s of steady state without it (v6e-8 T2V denoise: + 125.7 s vs 125.6 s primed; tpu7x-8 I2V: 93.3 s either way from a warm AOT + cache, 93.8 s vs 93.4 s after a cold compile). + + No-op outside ``aot_cache.warmup_mode()``. + + Args: + branches: One callable per expert. Each takes ``operands`` and dispatches + the same cached forward pass (same signature) as the denoise loop. + operands: The branch operands, in the same layout as for ``jax.lax.cond``. + """ + if not aot_cache.in_warmup(): + return + for branch in branches: + branch(operands) + + +class InflightWindow: + """Bounds how many dispatched denoise steps are queued on the device. + + The Python loop dispatches asynchronously and otherwise queues every + remaining step. An unbounded queue once ran 40-step T2V ~20% slower (3.30 vs + 2.69 s/step). On the current stack the bound measures neutral (v6e-8 T2V, + tpu7x-8 I2V), and it stays as a cheap guard. Blocking on the output of step + (N - depth) keeps at most ``depth`` steps queued while the pipeline stays + full; a periodic FULL drain would empty the pipe and cost a bubble per drain. + + ``MAXD_QUEUE_MAX_INFLIGHT`` sets the default depth (4); 0 disables the bound. + """ + + def __init__(self, depth: Optional[int] = None): + self.depth = int(os.environ.get("MAXD_QUEUE_MAX_INFLIGHT", "4")) if depth is None else depth + self._inflight = collections.deque() + + def push(self, step_output: Any) -> None: + """Records a dispatched step's output, blocking on the oldest once more than ``depth`` are queued.""" + if self.depth <= 0: + return + self._inflight.append(step_output) + if len(self._inflight) > self.depth: + jax.block_until_ready(self._inflight.popleft()) diff --git a/src/maxdiffusion/pipelines/wan/wan_pipeline.py b/src/maxdiffusion/pipelines/wan/wan_pipeline.py index 4006e2fd1..7b7a1ce36 100644 --- a/src/maxdiffusion/pipelines/wan/wan_pipeline.py +++ b/src/maxdiffusion/pipelines/wan/wan_pipeline.py @@ -163,7 +163,7 @@ def _select_restored_transformer_state(restored_checkpoint, subfolder: str): _DEVICE_PUT_LOCK = threading.Lock() -def converted_weights_cache_dir(config, subfolder: str) -> str: +def converted_weights_cache_dir(config, subfolder: str, scan_layers: Optional[bool] = None) -> str: """Per-(model, subfolder, dtype, scan) dir for memoized converted weights.""" base = getattr(config, "converted_weights_dir", "") if not base: @@ -171,7 +171,8 @@ def converted_weights_cache_dir(config, subfolder: str) -> str: model_tag = (config.wan_transformer_pretrained_model_name_or_path or config.pretrained_model_name_or_path).replace( "/", "--" ) - return os.path.join(base, f"{model_tag}--{subfolder or 'transformer'}--{config.weights_dtype}--scan{config.scan_layers}") + scan = config.scan_layers if scan_layers is None else scan_layers + return os.path.join(base, f"{model_tag}--{subfolder or 'transformer'}--{config.weights_dtype}--scan{scan}") def put_params_into_state( @@ -393,6 +394,7 @@ def create_model(rngs: nnx.Rngs, wan_config: dict): "svg_high_noise_density": high_density, "svg_low_noise_density": low_density, "svg_flash_block_sizes": getattr(config, "svg_flash_block_sizes", None) or None, + "use_k_centering": getattr(config, "use_k_centering", "auto"), } # 2. eval_shape - will not use flops or create weights on device @@ -519,6 +521,7 @@ def __init__( # encode_prompt result cache: same-prompt calls (warmup + real run, # repeated serving requests) skip the ~10s/call CPU text encoder. self._prompt_embeds_cache = {} + self.wan_debug_cond_timers = getattr(config, "wan_debug_cond_timers", False) def check_inputs( self, @@ -1338,7 +1341,7 @@ def _prepare_model_inputs( batch_size = len(prompt) if prompt is not None else prompt_embeds.shape[0] // num_videos_per_prompt - debug_timers = bool(os.environ.get("WAN_DEBUG_COND_TIMERS")) + debug_timers = getattr(self, "wan_debug_cond_timers", False) or bool(os.environ.get("WAN_DEBUG_COND_TIMERS")) t_probe = time.perf_counter() with jax.named_scope("Encode-Prompt"): prompt_embeds, negative_prompt_embeds = self.encode_prompt( diff --git a/src/maxdiffusion/pipelines/wan/wan_pipeline_2_2.py b/src/maxdiffusion/pipelines/wan/wan_pipeline_2_2.py index f55bd79e5..80fb76e9d 100644 --- a/src/maxdiffusion/pipelines/wan/wan_pipeline_2_2.py +++ b/src/maxdiffusion/pipelines/wan/wan_pipeline_2_2.py @@ -25,7 +25,7 @@ from typing import List, Union, Optional from ...pyconfig import HyperParameters from ... import aot_cache -import os +from .wan_denoise_utils import InflightWindow, compile_experts import concurrent.futures from functools import partial import time @@ -866,51 +866,26 @@ def scan_body(carry, scan_elem): final_latents, _ = final_carry return final_latents - # Warmup only runs a couple of steps and the scheduler puts both on the - # high-noise transformer, so the low-noise weights are never executed; their - # first real call then lands mid-generation and costs 26.4s against a 2.6s - # step. Touch every weight set here instead, so the first generation already - # runs at steady-state latency. - if aot_cache.in_warmup() and do_classifier_free_guidance: - with aot_cache.real_execution(): - priming_latents = jnp.concatenate([latents] * 2) - priming_timestep = jnp.broadcast_to(timesteps[0], bsz * 2) - for _gd, _st, _rest, _gs, _kv, _mask in ( - ( - high_noise_graphdef, - high_noise_state, - high_noise_rest, - guidance_scale_high, - kv_cache_high, - encoder_attention_mask_high, - ), - ( - low_noise_graphdef, - low_noise_state, - low_noise_rest, - guidance_scale_low, - kv_cache_low, - encoder_attention_mask_low, - ), - ): - jax.block_until_ready( - transformer_forward_pass_full_cfg( - _gd, - _st, - _rest, - priming_latents, - priming_timestep, - prompt_embeds_combined, - guidance_scale=_gs, - kv_cache=_kv, - rotary_emb=rotary_emb, - encoder_attention_mask=_mask, - svg_step_index=jnp.asarray(0, dtype=jnp.int32), - ) - ) + # Warmup can miss the low-noise phase entirely, so compile both experts here + # (compile only, no execution); see compile_experts. + if aot_cache.in_warmup(): + compile_experts( + (high_noise_branch, low_noise_branch), + ( + latents, + jnp.broadcast_to(timesteps[0], bsz * 2 if do_classifier_free_guidance else bsz), + prompt_embeds_combined, + kv_cache_high, + kv_cache_low, + rotary_emb, + encoder_attention_mask_high, + encoder_attention_mask_low, + jnp.asarray(0, dtype=jnp.int32), + ), + ) profiler = None - _inflight_latents = [] + inflight = InflightWindow() for step in range(num_inference_steps): if config and max_utils.profiler_enabled(config) and step == first_profiling_step: profiler = max_utils.Profiler(config) @@ -967,22 +942,7 @@ def scan_body(carry, scan_elem): ) latents, scheduler_state = scheduler.step(scheduler_state, noise_pred, t, latents).to_tuple() - - # Bound the async dispatch queue. The python loop otherwise queues every - # remaining step on the device; past ~30 in-flight steps the runtime - # degrades ~20% (measured: 40-step fixed-m 3.30 s/step unbounded vs 2.69 - # with a mid-run drain; 4/12/24-step runs unaffected). A periodic drain - # keeps the queue shallow and costs nothing (the device stays busy). - # Sliding-window bound on in-flight dispatched steps: block on the - # output of step (N - depth) so the device queue holds at most `depth` - # steps while the pipeline stays full (a periodic FULL drain would - # empty the pipe and cost a bubble per drain). The safe depth shrinks - # as the per-step executable's footprint grows (bkv2048 tiles need - # ~2-4; bkv1024 tolerates 8+); past it the TPU runtime degrades ~20%. - _inflight_latents.append(latents) - depth = int(os.environ.get("MAXD_QUEUE_MAX_INFLIGHT", "4")) - if depth > 0 and len(_inflight_latents) > depth: - _inflight_latents.pop(0).block_until_ready() + inflight.push(latents) # Bounds the async dispatch queue; see InflightWindow. if config and max_utils.profiler_enabled(config) and step == last_profiling_step: if profiler: diff --git a/src/maxdiffusion/pipelines/wan/wan_pipeline_i2v_2p2.py b/src/maxdiffusion/pipelines/wan/wan_pipeline_i2v_2p2.py index fd249a7de..cb39277f5 100644 --- a/src/maxdiffusion/pipelines/wan/wan_pipeline_i2v_2p2.py +++ b/src/maxdiffusion/pipelines/wan/wan_pipeline_i2v_2p2.py @@ -13,7 +13,8 @@ # limitations under the License. from maxdiffusion.image_processor import PipelineImageInput -from maxdiffusion import max_logging +from maxdiffusion import aot_cache, max_logging +from .wan_denoise_utils import InflightWindow, compile_experts from .wan_pipeline import ( WanPipeline, transformer_forward_pass, @@ -1010,7 +1011,31 @@ def scan_body(carry, t): final_latents, _ = final_carry return final_latents + timesteps_np = np.array(scheduler_state.timesteps, dtype=np.int32) + step_uses_high = [bool(timesteps_np[s] >= boundary) for s in range(num_inference_steps)] + # Warmup can miss the low-noise phase entirely, so compile both experts here + # (compile only, no execution); see compile_experts. + if aot_cache.in_warmup(): + warmup_latents = jnp.concatenate([latents, latents], axis=0) if do_classifier_free_guidance else latents + warmup_latents = jnp.transpose(warmup_latents, (0, 4, 1, 2, 3)) + compile_experts( + (high_noise_branch, low_noise_branch), + ( + jnp.concatenate([warmup_latents, condition], axis=1), + jnp.broadcast_to(timesteps[0], warmup_latents.shape[0]), + prompt_embeds_combined, + image_embeds_combined, + kv_cache_high, + kv_cache_low, + rotary_emb, + encoder_attention_mask_high, + encoder_attention_mask_low, + jnp.asarray(0, dtype=jnp.int32), + ), + ) + profiler = None + inflight = InflightWindow() for step in range(num_inference_steps): if config and max_utils.profiler_enabled(config) and step == first_profiling_step: profiler = max_utils.Profiler(config) @@ -1026,7 +1051,7 @@ def scan_body(carry, t): # Timesteps are host-known: Python dispatch (like the T2V loop) avoids # tracing both 14B branches per step and keeps the AOT cache usable. - use_high_noise = bool(np.asarray(scheduler_state.timesteps)[step] >= np.asarray(boundary)) + use_high_noise = step_uses_high[step] branch = high_noise_branch if use_high_noise else low_noise_branch noise_pred = branch(( latent_model_input, @@ -1042,6 +1067,7 @@ def scan_body(carry, t): )) noise_pred = jnp.transpose(noise_pred, (0, 2, 3, 4, 1)) latents, scheduler_state = scheduler.step(scheduler_state, noise_pred, t, latents).to_tuple() + inflight.push(latents) # Bounds the async dispatch queue; see InflightWindow. if config and max_utils.profiler_enabled(config) and step == last_profiling_step: if profiler: diff --git a/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py b/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py index 8090f7590..80ab6dbb3 100644 --- a/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py +++ b/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py @@ -122,10 +122,10 @@ def create_model(rngs: nnx.Rngs, wan_config: dict): eval_shapes=params, device="cpu", num_layers=wan_config["num_layers"], - scan_layers=config.scan_layers, + scan_layers=wan_config["scan_layers"], subfolder=subfolder, cast_dtype_fn=partial(_final_param_dtype, dtype_to_cast=config.weights_dtype), - converted_cache_dir=converted_weights_cache_dir(config, subfolder), + converted_cache_dir=converted_weights_cache_dir(config, subfolder, scan_layers=wan_config["scan_layers"]), ) # No-op (returns leaves unchanged) when the loader already cast to the diff --git a/src/maxdiffusion/pyconfig.py b/src/maxdiffusion/pyconfig.py index 0f55436ae..c995e7f4e 100644 --- a/src/maxdiffusion/pyconfig.py +++ b/src/maxdiffusion/pyconfig.py @@ -76,7 +76,13 @@ def string_to_list(string_list: str) -> list: return ast.literal_eval(string_list) -_yaml_types_to_parser = {str: str, int: int, float: float, bool: string_to_bool, list: string_to_list} +_yaml_types_to_parser = { + str: str, + int: int, + float: float, + bool: string_to_bool, + list: string_to_list, +} _config = None config = None @@ -322,6 +328,9 @@ def user_init(raw_keys): if "use_k_centering" not in raw_keys: raw_keys["use_k_centering"] = "auto" + if "wan_debug_cond_timers" not in raw_keys: + raw_keys["wan_debug_cond_timers"] = False + def get_num_slices(raw_keys): if int(raw_keys["compile_topology_num_slices"]) > 0: diff --git a/src/maxdiffusion/tests/aot_cache_test.py b/src/maxdiffusion/tests/aot_cache_test.py index a0a43170d..09327c803 100644 --- a/src/maxdiffusion/tests/aot_cache_test.py +++ b/src/maxdiffusion/tests/aot_cache_test.py @@ -15,6 +15,7 @@ """Tests for the per-shape AOT executable cache (CPU backend).""" import functools +import glob import os import tempfile import unittest @@ -34,6 +35,16 @@ def _toy_fn(x, y, flag=False): return x @ y + (1.0 if flag else 0.0) +@functools.partial(aot_cache.cached_jit, static_argnames=("flag",)) +def _scaled_fn(x, scale, flag=False): + """`scale` mirrors guidance_scale: a dynamic Python float multiplying a bf16 tensor.""" + return x * scale + (1.0 if flag else 0.0) + + +def _entry(suffix): + return next(e for e in aot_cache._REGISTRY if e.name.endswith(suffix)) + + class AotCacheTest(unittest.TestCase): def setUp(self): @@ -161,6 +172,7 @@ def test_signature_deterministic_across_processes(self): "class T(nnx.Module):", " def __init__(self, rngs):", " self.lin = nnx.Linear(4, 4, rngs=rngs)", + " self.tags = {'alpha', 'beta', 'gamma'}", "", "graphdef, state = nnx.split(T(nnx.Rngs(0)))", "sig = aot_cache._dynamic_signature(", @@ -173,9 +185,14 @@ def test_signature_deterministic_across_processes(self): capture_output=True, text=True, check=True, - env={**os.environ, "JAX_PLATFORMS": "cpu"}, + env={ + **os.environ, + "JAX_PLATFORMS": "cpu", + "PYTHONPATH": os.pathsep.join(sys.path), + "PYTHONHASHSEED": str(seed), + }, ).stdout.strip() - for _ in range(2) + for seed in (1, 42) ] self.assertEqual(outs[0], outs[1]) @@ -242,6 +259,693 @@ def test_step_array_preserves_dynamic_signature_across_steps(self): } self.assertEqual(len(sigs), 1) + def test_graphdef_static_attribute_changes_dynamic_signature(self): + """Toggling static attributes on nnx.GraphDef (such as use_k_centering) changes _dynamic_signature.""" + from flax import nnx + + class DummyBlock(nnx.Module): + + def __init__(self, use_k_centering: bool): + self.attention_config = {"use_k_centering": use_k_centering} + + gd_on, _ = nnx.split(DummyBlock(use_k_centering=True)) + gd_off, _ = nnx.split(DummyBlock(use_k_centering=False)) + + sig_on = aot_cache._dynamic_signature((gd_on, jnp.ones((2, 4))), {}) + sig_off = aot_cache._dynamic_signature((gd_off, jnp.ones((2, 4))), {}) + self.assertNotEqual(sig_on, sig_off) + + def test_wan_aot_metadata_includes_use_k_centering(self): + """Toggling use_k_centering changes the Wan AOT metadata fingerprint.""" + import types + from maxdiffusion import generate_wan + + cfg_on = types.SimpleNamespace(use_k_centering=True, attention="ulysses_ring_custom_fixed_m") + cfg_off = types.SimpleNamespace(use_k_centering=False, attention="ulysses_ring_custom_fixed_m") + + meta_on = generate_wan._build_wan_aot_metadata(cfg_on, self._mesh, "rev1") + meta_off = generate_wan._build_wan_aot_metadata(cfg_off, self._mesh, "rev1") + + self.assertEqual(meta_on["use_k_centering"], "True") + self.assertEqual(meta_off["use_k_centering"], "False") + self.assertNotEqual( + aot_cache._metadata_fingerprint(meta_on), + aot_cache._metadata_fingerprint(meta_off), + ) + + def test_wan_aot_metadata_includes_compiler_flags(self): + """Raw LIBTPU_INIT_ARGS/XLA_FLAGS and runtime versions are included in the metadata.""" + import types + from unittest import mock + from maxdiffusion import generate_wan + + cfg = types.SimpleNamespace(attention="ulysses_ring_custom_fixed_m") + + def metadata(env): + with mock.patch.dict(os.environ, env, clear=False): + return generate_wan._build_wan_aot_metadata(cfg, self._mesh, "rev1") + + meta1 = metadata({"LIBTPU_INIT_ARGS": "--a=1 --b=true", "XLA_FLAGS": ""}) + meta2 = metadata({"LIBTPU_INIT_ARGS": "--a=1 --b=false", "XLA_FLAGS": ""}) + meta3 = metadata({"LIBTPU_INIT_ARGS": "--a=1 --b=true", "XLA_FLAGS": "--xla_c=2"}) + + self.assertNotEqual(aot_cache._metadata_fingerprint(meta1), aot_cache._metadata_fingerprint(meta2)) + self.assertNotEqual(aot_cache._metadata_fingerprint(meta1), aot_cache._metadata_fingerprint(meta3)) + self.assertIn("jaxlib", meta1) + self.assertIn("flax", meta1) + self.assertIn("platform_version", meta1) + self.assertIn("default_matmul_precision", meta1) + self.assertIn("default_prng_impl", meta1) + + def test_run_aot_gating_and_ephemeral_teardown(self): + """run(): default -> no cache dir; downgraded or opted-in -> ephemeral dir, removed even on error.""" + import types + from unittest import mock + from maxdiffusion import generate_wan + + def run_with(revision, fail=False, **cfg_kwargs): + cfg = types.SimpleNamespace(model_name="wan2.2", enable_lora=False, **cfg_kwargs) + pipeline = types.SimpleNamespace(mesh=None) + installs, generate_calls = [], [] + + def fake_generate(config, pipe, prefix, writer, load_time, persistent_dir): + del config, pipe, prefix, writer, load_time + generate_calls.append((persistent_dir, installs[-1])) + if fail: + raise RuntimeError("generation failed") + return ["out.mp4"] + + with ( + mock.patch.object(generate_wan, "_resolve_wan_aot_source_revision", return_value=revision), + mock.patch.object(generate_wan, "_build_wan_aot_metadata", return_value={}), + mock.patch.object(generate_wan, "_warmup_and_generate", side_effect=fake_generate), + mock.patch.object(generate_wan.max_utils, "initialize_summary_writer", return_value=None), + mock.patch.object(generate_wan.aot_cache, "extract_svg_meta", return_value={}), + mock.patch.object(generate_wan.aot_cache, "wait_for_loads"), + mock.patch.object(generate_wan.aot_cache, "install", side_effect=lambda d, **kw: installs.append(d)), + ): + try: + out = generate_wan.run(cfg, pipeline=pipeline, commit_hash="abc1234") + except RuntimeError: + out = None + return out, installs, generate_calls + + clean, dirty = "src:clean", "abc1234-dirty" + + # Default: nothing requested -> installed on "" (plain jit), no teardown install. + out, installs, calls = run_with(clean, aot_cache_dir="", enable_zero_execution_warmup=False) + self.assertEqual(out, ["out.mp4"]) + self.assertEqual(installs, [""]) + self.assertEqual(calls, [("", "")]) + + # Persistent dir + reusable revision -> used directly, never torn down. + out, installs, calls = run_with(clean, aot_cache_dir="/cache/aot", enable_zero_execution_warmup=False) + self.assertEqual(installs, ["/cache/aot"]) + self.assertEqual(calls, [("/cache/aot", "/cache/aot")]) + + # Downgraded (dirty revision) and explicit opt-in -> ephemeral dir, torn down afterwards, + # including when generation raises. + for label, revision, cfg_kwargs, fail in ( + ("downgraded", dirty, {"aot_cache_dir": "/cache/aot", "enable_zero_execution_warmup": False}, False), + ("opt_in", clean, {"aot_cache_dir": "", "enable_zero_execution_warmup": True}, False), + ("opt_in_error", clean, {"aot_cache_dir": "", "enable_zero_execution_warmup": True}, True), + ): + with self.subTest(label): + out, installs, calls = run_with(revision, fail=fail, **cfg_kwargs) + self.assertEqual(len(installs), 2) + ephemeral = installs[0] + self.assertIn("wan_aot_ephemeral_", ephemeral) + self.assertEqual(installs[1], "") # uninstalled in `finally` + self.assertFalse(os.path.exists(ephemeral)) # temp dir removed + # Generation saw the ephemeral install but no persistent dir (so nothing is saved). + self.assertEqual(calls, [("", ephemeral)]) + self.assertEqual(out, None if fail else ["out.mp4"]) + + def test_plan_rejects_gcs_and_uses_the_real_revision_resolver(self): + """gs:// raises whatever the revision; only an explicitly dirty/unversioned revision downgrades.""" + import types + from maxdiffusion import generate_wan + + for revision in ("src:clean", "abc1234-dirty", None): + with self.subTest(revision=revision): + cfg = types.SimpleNamespace(aot_cache_dir="gs://bucket/aot", enable_zero_execution_warmup=False) + with self.assertRaisesRegex(ValueError, "gs:// URIs are not supported"): + generate_wan._plan_wan_aot_cache(cfg, revision) + + def plan(aot_build_revision): + cfg = types.SimpleNamespace( + aot_cache_dir="/cache/aot", enable_zero_execution_warmup=False, aot_build_revision=aot_build_revision + ) + return generate_wan._plan_wan_aot_cache(cfg, generate_wan._resolve_wan_aot_source_revision(cfg, "abc1234-dirty")) + + # The package content hash always resolves, so a dirty git tree alone does not downgrade. + self.assertEqual(plan(None), ("/cache/aot", False)) + self.assertEqual(plan("release-1"), ("/cache/aot", False)) + self.assertEqual(plan("dirty:wip"), ("", True)) + self.assertEqual(plan("unversioned:x"), ("", True)) + + def test_run_tears_down_ephemeral_dir_when_install_fails(self): + import types + from unittest import mock + from maxdiffusion import generate_wan + + cfg = types.SimpleNamespace(model_name="wan2.2", enable_lora=False, aot_cache_dir="", enable_zero_execution_warmup=True) + installs = [] + + def fake_install(d, **kw): + del kw + installs.append(d) + if d: + raise OSError("install failed") + + with ( + mock.patch.object(generate_wan, "_resolve_wan_aot_source_revision", return_value="src:clean"), + mock.patch.object(generate_wan, "_build_wan_aot_metadata", return_value={}), + mock.patch.object(generate_wan.max_utils, "initialize_summary_writer", return_value=None), + mock.patch.object(generate_wan.aot_cache, "extract_svg_meta", return_value={}), + mock.patch.object(generate_wan.aot_cache, "install", side_effect=fake_install), + ): + with self.assertRaisesRegex(OSError, "install failed"): + generate_wan.run(cfg, pipeline=types.SimpleNamespace(mesh=None), commit_hash="abc1234") + self.assertEqual(len(installs), 2) + self.assertIn("wan_aot_ephemeral_", installs[0]) + self.assertEqual(installs[1], "") + self.assertFalse(os.path.exists(installs[0])) + + def test_wan_source_hash_includes_shared_modules_and_prefers_content_hash_over_commit(self): + """Verifies _compute_wan_source_hash hashes shared modules and prefers content hash over commit.""" + import types + from unittest import mock + from maxdiffusion import generate_wan + + base_hash = generate_wan._compute_wan_source_hash() + self.assertIsNotNone(base_hash) + self.assertTrue(base_hash.startswith("src:")) + + # Simulate modifying models/normalization_flax.py or models/embeddings_flax.py + orig_open = open + + def patched_open(path, *args, **kwargs): + f = orig_open(path, *args, **kwargs) + if str(path).endswith(("normalization_flax.py", "embeddings_flax.py")) and "rb" in args: + content = f.read() + f.close() + import io + + return io.BytesIO(content + b"\n# modified for fingerprint test\n") + return f + + generate_wan._compute_wan_source_hash.cache_clear() + with mock.patch("builtins.open", side_effect=patched_open): + mod_hash = generate_wan._compute_wan_source_hash() + generate_wan._compute_wan_source_hash.cache_clear() + + self.assertNotEqual(base_hash, mod_hash) + + # Content hash is preferred, reusable, and not invalidated by dirty status + cfg = types.SimpleNamespace(aot_build_revision=None) + rev = generate_wan._resolve_wan_aot_source_revision(cfg, commit_hash="commit_dirty-dirty") + self.assertEqual(rev, base_hash) + self.assertTrue(generate_wan._is_reusable_aot_revision(rev)) + + # When source hash is unavailable, clean commit hashes are reusable while + # "-dirty" commit hashes from get_git_commit_hash() are rejected. + self.assertTrue(generate_wan._is_reusable_aot_revision("abc1234")) + self.assertFalse(generate_wan._is_reusable_aot_revision("abc1234-dirty")) + self.assertFalse(generate_wan._is_reusable_aot_revision("dirty:abc1234")) + self.assertFalse(generate_wan._is_reusable_aot_revision("unversioned:abc1234")) + + # Explicit aot_build_revision takes highest precedence and folds in source hash + cfg_explicit = types.SimpleNamespace(aot_build_revision="build-explicit-123") + self.assertEqual( + generate_wan._resolve_wan_aot_source_revision(cfg_explicit), + f"build-explicit-123+{base_hash}", + ) + + def test_different_nnx_modules_and_unordered_sets_in_dynamic_signature(self): + from flax import nnx + + class ModMul(nnx.Module): + + def __init__(self): + self.scale = 2 + self.tags = {"alpha", "beta", "gamma"} + + def __call__(self, x): + return x * self.scale + + class ModAdd(nnx.Module): + + def __init__(self): + self.scale = 2 + self.tags = {"gamma", "alpha", "beta"} + + def __call__(self, x): + return x + self.scale + + gdef_mul, state_mul = nnx.split(ModMul()) + gdef_add, state_add = nnx.split(ModAdd()) + x = jnp.array(3.0, dtype=jnp.float32) + + sig_mul = aot_cache._dynamic_signature((gdef_mul, state_mul, x), {}) + sig_add = aot_cache._dynamic_signature((gdef_add, state_add, x), {}) + self.assertNotEqual(sig_mul, sig_add) + + # Unordered sets inside static attributes must format canonically + self.assertEqual( + aot_cache._format_static_val({"a", "b", "c"}), + "set({'a','b','c'})", + ) + + self._install() + fn = aot_cache.cached_jit(lambda g, s, val: nnx.merge(g, s)(val)) + out_mul = fn(gdef_mul, state_mul, x) + aot_cache.save_pending() + out_add = fn(gdef_add, state_add, x) + self.assertAlmostEqual(float(out_mul), 6.0) + self.assertAlmostEqual(float(out_add), 5.0) + + # Same module class with different static attributes must not collide in _sig_cache + mod_ten = ModMul() + mod_ten.scale = 10 + gdef_ten, state_ten = nnx.split(mod_ten) + out_ten = fn(gdef_ten, state_ten, x) + self.assertAlmostEqual(float(out_ten), 30.0) + + def test_compiled_call_failure_raises_loudly_without_fallback(self): + """When compiled(flat) fails, it must raise loudly instead of silently recompiling via JIT.""" + from unittest import mock + + self._install() + + @aot_cache.cached_jit + def fn(x): + return x + 1.0 + + x = jnp.array([1.0, 2.0], dtype=jnp.float32) + # Warmup and compile + with aot_cache.warmup_mode(): + fn(x) + + # Corrupt the compiled entry so compiled(flat) raises + sig = next(iter(fn._compiled.keys())) + bad_compiled = mock.MagicMock(side_effect=RuntimeError("simulated execution error")) + bad_compiled.input_shardings = fn._compiled[sig].input_shardings + fn._compiled[sig] = bad_compiled + + with self.assertRaises(RuntimeError) as ctx: + fn(x) + self.assertIn("simulated execution error", str(ctx.exception)) + + def test_warmup_mode_with_tempdir_install_returns_zeros(self): + """Installing a throwaway dir (as generate_wan does when persistence is disabled) enables zero-execution warmup.""" + import tempfile + + ephemeral_dir = tempfile.mkdtemp(prefix="aot_ephemeral_") + aot_cache.install(ephemeral_dir, meta={"test": "warmup"}, mesh=self._mesh) + self.assertTrue(aot_cache._STATE.enabled) + + with aot_cache.warmup_mode(): + res = _toy_fn(self._a, self._b, True) + self.assertEqual(res.shape, (8, 8)) + self.assertTrue(jnp.all(res == 0)) + + # Outside warmup, real execution runs and returns computed values + real = _toy_fn(self._a, self._b, True) + np.testing.assert_allclose(np.asarray(real), self._a @ self._b + 1.0) + + def test_bytecode_digest_distinguishes_constants(self): + """Bytecode digest must distinguish functions that differ only by constant values.""" + + def f1(x): + return x * 2.0 + + def f2(x): + return x * 3.0 + + val1 = aot_cache._format_static_val(f1) + val2 = aot_cache._format_static_val(f2) + self.assertNotEqual(val1, val2) + + def test_graphdef_memoization_caches_by_id(self): + """GraphDef statics extraction must memoize results by id(obj) without re-traversal.""" + from unittest import mock + from flax import nnx + + class SimpleMod(nnx.Module): + + def __init__(self): + self.param = 42 + + mod = SimpleMod() + gdef, _ = nnx.split(mod) + + aot_cache._GRAPHDEF_MEMO.clear() + out1 = aot_cache._extract_graphdef_statics(gdef) + self.assertIn(id(gdef), aot_cache._GRAPHDEF_MEMO) + self.assertEqual(len(out1), 1) + + # Second call uses cached digest in _GRAPHDEF_MEMO without calling _format_static_val + with mock.patch( + "maxdiffusion.aot_cache._format_static_val", + wraps=aot_cache._format_static_val, + ) as spy_fmt: + out2 = aot_cache._extract_graphdef_statics(gdef) + self.assertEqual(spy_fmt.call_count, 0) + self.assertEqual(out1, out2) + + def test_use_k_centering_defaults_to_auto(self): + """use_k_centering must default to 'auto' in pyconfig and WanTransformerBlock.""" + from maxdiffusion import pyconfig + from maxdiffusion.models.wan.transformers.transformer_wan import WanTransformerBlock + from flax import nnx + + pyconfig.initialize([None, "src/maxdiffusion/configs/base_wan_14b.yml", "run_name=test_k_centering"]) + self.assertEqual(pyconfig.config.use_k_centering, "auto") + + # Check WanTransformerBlock default attention_config + block = WanTransformerBlock(rngs=nnx.Rngs(0), dim=64, num_heads=4, ffn_dim=128, cross_attn_norm=True) + self.assertEqual(block.attn1.attention_op.use_k_centering, "auto") + + def test_video_output_path_preserves_prefix_and_gcs_fallback(self): + """format_video_output_path must preserve prefix and handle directory vs local naming.""" + from maxdiffusion.generate_wan import format_video_output_path + + # With output_dir + path_with_prefix = format_video_output_path("/tmp/test_wan_out", "wan_run", 42, 0, "prefix_") + self.assertEqual(path_with_prefix, "/tmp/test_wan_out/prefix_wan_run_42_0.mp4") + + path_no_prefix = format_video_output_path("/tmp/test_wan_out", "wan_run", 42, 1) + self.assertEqual(path_no_prefix, "/tmp/test_wan_out/wan_run_42_1.mp4") + + # Without output_dir (or GCS or default template output_dir) + path_empty_dir = format_video_output_path("", "wan_run", 42, 0, "prefix_") + self.assertEqual(path_empty_dir, "prefix_wan_output_42_0.mp4") + + path_gcs_dir = format_video_output_path("gs://bucket/dir", "wan_run", 42, 2, "my_") + self.assertEqual(path_gcs_dir, "my_wan_output_42_2.mp4") + + path_default_sdxl = format_video_output_path("sdxl-model-finetuned", "None", 42, 0) + self.assertEqual(path_default_sdxl, "wan_output_42_0.mp4") + + path_default_empty_run = format_video_output_path("sdxl-model-finetuned", "", 42, 0) + self.assertEqual(path_default_empty_run, "wan_output_42_0.mp4") + + def test_float32_qk_product_gate(self): + """_apply_attention_dot must use float32 preferred_element_type only when float32_qk_product=True.""" + from maxdiffusion.models.attention_flax import _apply_attention_dot + + q = jnp.ones((1, 4, 4, 16), dtype=jnp.bfloat16) + k = jnp.ones((1, 4, 4, 16), dtype=jnp.bfloat16) + v = jnp.ones((1, 4, 4, 16), dtype=jnp.bfloat16) + + def fn_false(q, k, v): + return _apply_attention_dot( + q, + k, + v, + dtype=jnp.bfloat16, + heads=4, + dim_head=16, + scale=0.25, + split_head_dim=True, + float32_qk_product=False, + use_memory_efficient_attention=False, + ) + + def fn_true(q, k, v): + return _apply_attention_dot( + q, + k, + v, + dtype=jnp.bfloat16, + heads=4, + dim_head=16, + scale=0.25, + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=False, + ) + + jaxpr_false = jax.make_jaxpr(fn_false)(q, k, v) + jaxpr_true = jax.make_jaxpr(fn_true)(q, k, v) + + preferred_false = [ + eq.params.get("preferred_element_type") for eq in jaxpr_false.eqns if eq.primitive.name == "dot_general" + ] + preferred_true = [ + eq.params.get("preferred_element_type") for eq in jaxpr_true.eqns if eq.primitive.name == "dot_general" + ] + + # The first dot_general is QK product; second is (QK)V product. + self.assertEqual(preferred_false[0], jnp.dtype("bfloat16")) + self.assertEqual(preferred_true[0], jnp.dtype("float32")) + # The QK operands stay bf16 in both modes: float32_qk_product only widens + # the accumulator. (Upcasting Q/K before the dot, the previous behaviour, + # also yields a float32 preferred_element_type, so the check above alone + # cannot tell the two apart.) + for jaxpr in (jaxpr_false, jaxpr_true): + qk_dot = next(eq for eq in jaxpr.eqns if eq.primitive.name == "dot_general") + self.assertEqual([v.aval.dtype for v in qk_dot.invars], [jnp.dtype("bfloat16")] * 2) + upcasts = [ + eq + for eq in jaxpr.eqns + if eq.primitive.name == "convert_element_type" and eq.params.get("new_dtype") == jnp.dtype("float32") + ] + first_dot_index = jaxpr.eqns.index(qk_dot) + self.assertFalse([eq for eq in upcasts if jaxpr.eqns.index(eq) < first_dot_index]) + + def test_graphdef_memo_prevents_gc_address_reuse_collision(self): + """_GRAPHDEF_MEMO must retain object references to prevent id() reuse collisions after GC.""" + from flax import nnx + from maxdiffusion.aot_cache import _extract_graphdef_statics + + class Dummy(nnx.Module): + + def __init__(self, scale: float): + self.scale = scale + + seen_digests = set() + for i in range(100): + m = Dummy(scale=float(i)) + graphdef, _ = nnx.split(m) + statics = _extract_graphdef_statics(graphdef) + digest = statics[0] + self.assertNotIn(digest, seen_digests, f"Collision detected at iter {i}: {digest}") + seen_digests.add(digest) + + def test_pending_stores_shape_dtype_structs_and_clear_pending_empties(self): + """_pending must store ShapeDtypeStructs (not live jax.Arrays) and clear_pending() must empty it.""" + self._install() + _toy_fn(self._a, self._b, True) + entry = next(e for e in aot_cache._REGISTRY if e.name.endswith("._toy_fn")) + self.assertEqual(len(entry._pending), 1) + leaves, _, _ = next(iter(entry._pending.values())) + for leaf in leaves: + self.assertIsInstance(leaf, jax.ShapeDtypeStruct) + self.assertNotIsInstance(leaf, jax.Array) + # Re-lowering from ShapeDtypeStructs in save_pending succeeds even when _compiled is empty + entry._compiled.clear() + self.assertEqual(aot_cache.save_pending(), 1) + + # Record another shape and verify clear_pending() discards it + _toy_fn(jnp.ones((4, 8)), jnp.ones((8, 4)), True) + self.assertEqual(len(entry._pending), 1) + aot_cache.clear_pending() + self.assertEqual(len(entry._pending), 0) + + def test_nested_warmup_mode_restores_previous_state(self): + """Nested warmup_mode() must restore outer warmup_only state on exit.""" + self._install() + self.assertFalse(aot_cache.in_warmup()) + with aot_cache.warmup_mode(): + self.assertTrue(aot_cache.in_warmup()) + with aot_cache.warmup_mode(): + self.assertTrue(aot_cache.in_warmup()) + self.assertTrue(aot_cache.in_warmup()) + self.assertFalse(aot_cache.in_warmup()) + + def test_format_static_val_canonicalizes_use_default_values_and_frozensets(self): + """_format_static_val must sort diffusers _use_default_values lists and _format_const must sort frozensets.""" + d1 = {"_use_default_values": ["wan_seq_pad", "wan_patch_embed_mode"], "attention": "flash"} + d2 = {"_use_default_values": ["wan_patch_embed_mode", "wan_seq_pad"], "attention": "flash"} + self.assertEqual(aot_cache._format_static_val(d1), aot_cache._format_static_val(d2)) + self.assertEqual( + aot_cache._format_const(frozenset({"b", "a", "c"})), + "frozenset({'a','b','c'})", + ) + + def test_align_inputs_memoizes_scalars_and_pruned_none_shardings(self): + """_align_inputs must preserve None-sharded pruned inputs, memoize Python scalar device placement, and reject length mismatch.""" + from unittest import mock + + self._install() + + @aot_cache.cached_jit + def fn(x, unused_param, scale=4.0): + del unused_param + return x * scale + + with aot_cache.warmup_mode(): + fn(self._a, self._b, 4.0) + + compiled = next(iter(fn._compiled.values())) + flat_expected, _ = jax.tree_util.tree_flatten(compiled.input_shardings, is_leaf=lambda s: s is None) + self.assertIsNone(flat_expected[1]) + + with self.assertRaisesRegex(ValueError, "input leaf count mismatch"): + fn._align_inputs(compiled, [self._a]) + + with mock.patch.object(aot_cache.jax, "device_put", wraps=jax.device_put) as spy_put: + out1 = fn(self._a, self._b, 4.0) + first_calls = spy_put.call_count + out2 = fn(self._a, self._b, 4.0) + second_calls = spy_put.call_count + + np.testing.assert_allclose(np.asarray(out1), self._a * 4.0) + np.testing.assert_allclose(np.asarray(out2), self._a * 4.0) + self.assertEqual(first_calls, second_calls, "Repeated calls with the same Python scalar must reuse _scalar_cache") + + def test_safe_pickle_load_and_gcs_uri_guard(self): + """_safe_pickle_load must reject arbitrary globals and install() must reject gs:// URIs.""" + import io + import pickle + + with self.assertRaisesRegex(ValueError, "gs:// URIs are not supported"): + aot_cache.install("gs://some-bucket/aot_cache", meta={}, mesh=self._mesh) + + _, treedef = jax.tree_util.tree_flatten({"a": [1, 2], "b": (3,)}) + self.assertEqual(aot_cache._safe_pickle_load(io.BytesIO(pickle.dumps(treedef))), treedef) + + malicious_payload = b"cos\nsystem\n(S'echo pwned'\ntR." + with self.assertRaises(pickle.UnpicklingError): + aot_cache._safe_pickle_load(io.BytesIO(malicious_payload)) + + def test_restricted_unpickler_is_an_exact_allowlist(self): + """Only exact (module, name) pairs resolve; loose jax.* matches and dotted names are rejected.""" + import io + import pickle + + # The full .aotx envelope shape round-trips. + _, in_tree = jax.tree_util.tree_flatten(((list(range(3)),), {})) + _, out_tree = jax.tree_util.tree_flatten((1, [2, 3])) + blob = { + "format_version": 2, + "payload": b"\x00\x01", + "in_tree": in_tree, + "out_tree": out_tree, + "dynamic_signature": "sig", + "out_shapes_dtypes": [([1, 2], "bfloat16")], + "tags": {"a", "b"}, + } + loaded = aot_cache._safe_pickle_load(io.BytesIO(pickle.dumps(blob))) + self.assertEqual(loaded["in_tree"], in_tree) + self.assertEqual(loaded["out_tree"], out_tree) + + def global_ref(module: str, name: str) -> bytes: + # Protocol-0 GLOBAL opcode followed by STOP: resolves module.name via find_class. + return f"c{module}\n{name}\n.".encode() + + rejected = [ + ("jax._src.tree_util", "register_pytree_node"), # "tree" in module used to allow any name + ("jax._src.api", "_check_callable"), # "_"-prefixed names used to be allowed + ("jaxlib._jax.pytree", "PyTreeDef.__reduce__"), # dotted: resolved attribute by attribute + ("builtins", "dict.fromkeys"), + ("builtins", "eval"), + ("os", "system"), + ] + for module, name in rejected: + with self.subTest(module=module, name=name): + with self.assertRaises(pickle.UnpicklingError): + aot_cache._safe_pickle_load(io.BytesIO(global_ref(module, name))) + + def test_dynamic_signature_keys_python_scalars_by_type_and_statics_by_value(self): + """jit traces non-static Python scalars, so only their type selects the executable.""" + + def sig(value, **static): + return aot_cache._dynamic_signature((self._a,), {"guidance_scale": value}, static) + + self.assertEqual(sig(3.0), sig(4.0)) + self.assertNotEqual(sig(3.0), sig(3)) # float and int trace to different dtypes + self.assertNotEqual(sig(1), sig(True)) # jit does not treat a bool as an int + self.assertNotEqual(sig(3.0, flag=True), sig(3.0, flag=False)) + self.assertNotEqual(sig(3.0, scale=3.0), sig(3.0, scale=4.0)) # a static float stays value-keyed + + def test_python_scalar_values_share_one_executable(self): + """Wan 2.2 guidance (4.0 high-noise, 3.0 low-noise, any per-request value) must not fork executables.""" + self._install() + entry = _entry("._scaled_fn") + x = jnp.arange(16, dtype=jnp.bfloat16).reshape(4, 4) + reference = jax.jit(_scaled_fn.fn, static_argnames=("flag",)) + with aot_cache.warmup_mode(): + _scaled_fn(x, 4.0) + for scale in (3.0, 4.0, 5.0): + out = _scaled_fn(x, scale) + self.assertEqual(out.dtype, jnp.bfloat16, "a weakly typed Python float must not promote bf16") + np.testing.assert_array_equal(np.asarray(out, np.float32), np.asarray(reference(x, scale), np.float32)) + self.assertEqual(len(entry._compiled), 1) + self.assertEqual(len(entry._adapters), 1, "a new scalar value must not fall back to jit") + self.assertEqual(aot_cache.save_pending(), 1) + self.assertEqual(len(glob.glob(os.path.join(self._tmp.name, f"{entry.name}-*.aotx"))), 1) + + def test_shared_scalar_executable_roundtrips_disk_and_honors_value(self): + """One deserialized executable serves every scalar value, with the value applied at run time.""" + x = jnp.arange(16, dtype=jnp.bfloat16).reshape(4, 4) + self._install() + with aot_cache.warmup_mode(): + _scaled_fn(x, 4.0) + self.assertEqual(aot_cache.save_pending(), 1) + + self._install() # Like a fresh process: clears all state and deserializes from disk. + entry = _entry("._scaled_fn") + self.assertEqual(len(entry._on_disk), 1) + reference = jax.jit(_scaled_fn.fn, static_argnames=("flag",)) + outs = {} + for scale in (3.0, 4.0, 5.0): + outs[scale] = np.asarray(_scaled_fn(x, scale), np.float32) + np.testing.assert_array_equal(outs[scale], np.asarray(reference(x, scale), np.float32)) + self.assertFalse(np.array_equal(outs[3.0], outs[4.0]), "the value recorded at save time must not be baked in") + self.assertEqual(_scaled_fn(x, 3.0).dtype, jnp.bfloat16) + self.assertFalse(entry._adapters, "every value must hit the deserialized executable, not jit") + self.assertFalse(entry._pending) + + def test_static_scalar_values_still_get_own_executables(self): + """Static args are baked into the graph, so each value keeps its own executable.""" + self._install() + x = jnp.ones((4, 4), jnp.float32) + np.testing.assert_allclose(np.asarray(_scaled_fn(x, 2.0, flag=True)), np.asarray(x * 2.0 + 1.0)) + np.testing.assert_allclose(np.asarray(_scaled_fn(x, 2.0, flag=False)), np.asarray(x * 2.0)) + self.assertEqual(aot_cache.save_pending(), 2) + + def test_scalar_placement_memo_is_bounded(self): + """Values that no longer fork the executable can vary without limit; their placements must not.""" + self._install() + entry = _entry("._scaled_fn") + x = jnp.ones((4, 4), jnp.float32) + with aot_cache.warmup_mode(): + _scaled_fn(x, 0.0) + for i in range(aot_cache._SCALAR_CACHE_MAX + 8): + _scaled_fn(x, float(i)) + self.assertEqual(len(entry._scalar_cache), aot_cache._SCALAR_CACHE_MAX) + # 7.0 was evicted; it must be placed again with the right value. + np.testing.assert_allclose(np.asarray(_scaled_fn(x, 7.0)), np.asarray(x * 7.0)) + + def test_format_version_bump_skips_older_executables(self): + """Executables written under an older signature scheme are never loaded.""" + from unittest import mock + + with mock.patch.object(aot_cache, "_FORMAT_VERSION", aot_cache._FORMAT_VERSION - 1): + self._install() + _toy_fn(self._a, self._b, True) + self.assertEqual(aot_cache.save_pending(), 1) + self._install() + self.assertFalse(_entry("._toy_fn")._on_disk) + + _toy_fn(self._a, self._b, True) + self.assertEqual(aot_cache.save_pending(), 1) + self._install() + self.assertEqual(len(_entry("._toy_fn")._on_disk), 1) + if __name__ == "__main__": unittest.main() diff --git a/src/maxdiffusion/tests/converted_weights_cache_test.py b/src/maxdiffusion/tests/converted_weights_cache_test.py index a34745efd..20792bf67 100644 --- a/src/maxdiffusion/tests/converted_weights_cache_test.py +++ b/src/maxdiffusion/tests/converted_weights_cache_test.py @@ -21,8 +21,14 @@ import ml_dtypes import numpy as np +import json +from unittest import mock + +from maxdiffusion.models.wan import wan_utils from maxdiffusion.models.wan.wan_utils import save_converted_weights, try_load_converted_weights +_FP = "fp_test" + def _flat_tree(): return { @@ -56,8 +62,8 @@ def tearDown(self): self._tmp.cleanup() def test_round_trip(self): - save_converted_weights(self.cache_dir, self.flat) - loaded = try_load_converted_weights(self.cache_dir, self.eval_shapes, None) + save_converted_weights(self.cache_dir, self.flat, _FP) + loaded = try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP) self.assertIsNotNone(loaded) np.testing.assert_array_equal(loaded["blocks"]["attn1"]["kernel"], self.flat[("blocks", "attn1", "kernel")]) np.testing.assert_array_equal(loaded["proj_out"][0]["bias"], self.flat[("proj_out", 0, "bias")]) @@ -66,29 +72,196 @@ def test_round_trip(self): np.testing.assert_array_equal(bf16.view(np.uint16), self.flat[("blocks", "ffn", "kernel")].view(np.uint16)) def test_missing_cache_returns_none(self): - self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None)) + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP)) def test_dtype_policy_change_invalidates(self): - save_converted_weights(self.cache_dir, self.flat) - loaded = try_load_converted_weights(self.cache_dir, self.eval_shapes, lambda key: np.dtype(np.float64)) + save_converted_weights(self.cache_dir, self.flat, _FP) + loaded = try_load_converted_weights(self.cache_dir, self.eval_shapes, lambda key: np.dtype(np.float64), _FP) self.assertIsNone(loaded) def test_key_set_change_invalidates(self): - save_converted_weights(self.cache_dir, self.flat) + save_converted_weights(self.cache_dir, self.flat, _FP) bigger = dict(self.flat) bigger[("new_param", "kernel")] = np.zeros(2, dtype=np.float32) - loaded = try_load_converted_weights(self.cache_dir, _eval_shapes(bigger), None) + loaded = try_load_converted_weights(self.cache_dir, _eval_shapes(bigger), None, _FP) self.assertIsNone(loaded) def test_dtypes_preserved(self): - save_converted_weights(self.cache_dir, self.flat) - loaded = try_load_converted_weights(self.cache_dir, self.eval_shapes, None) + save_converted_weights(self.cache_dir, self.flat, _FP) + loaded = try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP) for key, value in self.flat.items(): node = loaded for part in key: node = node[part] self.assertEqual(node.dtype, value.dtype) + def test_resave_replaces_invalidated_cache(self): + save_converted_weights(self.cache_dir, self.flat, _FP) + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, lambda key: np.dtype(np.float16), _FP)) + updated = {k: v.astype(np.float16) for k, v in self.flat.items()} + save_converted_weights(self.cache_dir, updated, _FP) + loaded = try_load_converted_weights(self.cache_dir, _eval_shapes(updated), lambda key: np.dtype(np.float16), _FP) + self.assertIsNotNone(loaded) + self.assertEqual(loaded["blocks"]["attn1"]["kernel"].dtype, np.float16) + + def test_source_fingerprint_mismatch_invalidates(self): + save_converted_weights(self.cache_dir, self.flat, source_fingerprint="fp_v1") + self.assertIsNotNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, source_fingerprint="fp_v1")) + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, source_fingerprint="fp_v2")) + + def _rewrite_meta(self, meta): + path = os.path.join(self.cache_dir, "manifest.json") + with open(path) as f: + manifest = json.load(f) + if meta is None: + manifest.pop("__meta__") + else: + manifest["__meta__"] = meta + with open(path, "w") as f: + json.dump(manifest, f) + + def test_missing_or_null_fingerprint_is_rejected(self): + """Fail-closed: a manifest that cannot prove its source is a miss.""" + save_converted_weights(self.cache_dir, self.flat, _FP) + self._rewrite_meta({"format_version": wan_utils._CONVERTED_WEIGHTS_FORMAT_VERSION, "source_fingerprint": None}) + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP)) + self._rewrite_meta({"format_version": wan_utils._CONVERTED_WEIGHTS_FORMAT_VERSION}) + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP)) + self._rewrite_meta(None) # pre-versioned manifest + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP)) + # A caller without a fingerprint cannot verify anything either. + save_converted_weights(self.cache_dir, self.flat, _FP) + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, None)) + with self.assertRaises(ValueError): + save_converted_weights(self.cache_dir, self.flat, None) + + def test_old_format_version_is_rejected(self): + save_converted_weights(self.cache_dir, self.flat, _FP) + self._rewrite_meta({"format_version": 1, "source_fingerprint": _FP}) + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP)) + + def test_fingerprint_does_not_depend_on_mount_path(self): + """Same repo/revision/index under two different HF_HOMEs -> same fingerprint; other revision -> different.""" + + def make_index(root, revision, content=b'{"weight_map": {"a": "s1.safetensors"}}'): + d = os.path.join(self._tmp.name, root, "hub", "models--org--repo", "snapshots", revision, "transformer") + os.makedirs(d) + path = os.path.join(d, "index.json") + with open(path, "wb") as f: + f.write(content) + return path + + fp = wan_utils._compute_source_checkpoint_fingerprint + a = fp(make_index("home_a", "rev1"), "org/repo") + b = fp(make_index("snapshots/mnt/home_b", "rev1"), "org/repo") + c = fp(make_index("home_c", "rev2"), "org/repo") + d = fp(make_index("home_d", "rev1", b"{}"), "org/repo") + self.assertEqual(a, b) + self.assertNotEqual(a, c) + self.assertNotEqual(a, d) + self.assertNotEqual(a, fp(make_index("home_e", "rev1"), "org/other")) + + def _hf_snapshot(self, revision, subfolders, content=b'{"weight_map": {"a": "s1.safetensors"}}'): + """Builds an HF-cache layout: snapshots///index.json symlinked to one shared blob.""" + repo = os.path.join(self._tmp.name, "hub", "models--org--repo") + blob = os.path.join(repo, "blobs", f"blob-{revision}") + os.makedirs(os.path.dirname(blob), exist_ok=True) + with open(blob, "wb") as f: + f.write(content) + paths = {} + for sub in subfolders: + d = os.path.join(repo, "snapshots", revision, sub) + os.makedirs(d) + paths[sub] = os.path.join(d, "diffusion_pytorch_model.safetensors.index.json") + os.symlink(blob, paths[sub]) + return paths + + def test_fingerprint_reads_revision_through_hf_symlinks_and_includes_subfolder(self): + """HF snapshot entries are symlinks into blobs/: the revision must still be captured, and two + subfolders with byte-identical index files (Wan 2.2 transformer / transformer_2) must differ.""" + fp = wan_utils._compute_source_checkpoint_fingerprint + rev1 = self._hf_snapshot("rev1", ("transformer", "transformer_2")) + rev2 = self._hf_snapshot("rev2", ("transformer",)) + self.assertEqual(wan_utils._snapshot_revision(rev1["transformer"]), "rev1") + self.assertNotEqual( + fp(rev1["transformer"], "org/repo", "transformer"), fp(rev1["transformer_2"], "org/repo", "transformer_2") + ) + # Same index bytes, different snapshot revision -> different fingerprint. + self.assertNotEqual( + fp(rev1["transformer"], "org/repo", "transformer"), fp(rev2["transformer"], "org/repo", "transformer") + ) + # The v2 formula lost both (this is the bug being fixed). + legacy = wan_utils._legacy_v2_source_fingerprint + self.assertEqual(legacy(rev1["transformer"], "org/repo"), legacy(rev1["transformer_2"], "org/repo")) + self.assertEqual(legacy(rev1["transformer"], "org/repo"), legacy(rev2["transformer"], "org/repo")) + + def test_v2_manifest_with_matching_legacy_fingerprint_is_migrated_once(self): + paths = self._hf_snapshot("rev1", ("transformer",)) + new_fp = wan_utils._compute_source_checkpoint_fingerprint(paths["transformer"], "org/repo", "transformer") + old_fp = wan_utils._legacy_v2_source_fingerprint(paths["transformer"], "org/repo") + save_converted_weights(self.cache_dir, self.flat, new_fp) + self._rewrite_meta({"format_version": 2, "source_fingerprint": old_fp}) + # Without the legacy fingerprint a v2 manifest is a miss. + self.assertIsNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, new_fp)) + loaded = try_load_converted_weights(self.cache_dir, self.eval_shapes, None, new_fp, legacy_source_fingerprint=old_fp) + self.assertIsNotNone(loaded) + with open(os.path.join(self.cache_dir, "manifest.json")) as f: + meta = json.load(f)["__meta__"] + self.assertEqual(meta["format_version"], wan_utils._CONVERTED_WEIGHTS_FORMAT_VERSION) + self.assertEqual(meta["source_fingerprint"], new_fp) + # Now current-format: loads with the new fingerprint alone. + self.assertIsNotNone(try_load_converted_weights(self.cache_dir, self.eval_shapes, None, new_fp)) + + def test_v2_manifest_with_other_fingerprint_is_rejected(self): + save_converted_weights(self.cache_dir, self.flat, _FP) + self._rewrite_meta({"format_version": 2, "source_fingerprint": "someone_else"}) + self.assertIsNone( + try_load_converted_weights(self.cache_dir, self.eval_shapes, None, _FP, legacy_source_fingerprint="legacy") + ) + with open(os.path.join(self.cache_dir, "manifest.json")) as f: + self.assertEqual(json.load(f)["__meta__"]["format_version"], 2) # untouched + + def test_save_is_skipped_when_disk_is_too_full(self): + with mock.patch.object(wan_utils, "_free_bytes", return_value=1024): + self.assertFalse(save_converted_weights(self.cache_dir, self.flat, _FP)) + self.assertFalse(os.path.exists(self.cache_dir)) + self.assertTrue(save_converted_weights(self.cache_dir, self.flat, _FP)) + + def test_shard_download_refuses_to_fill_the_disk(self): + import huggingface_hub + + index = {"metadata": {"total_size": 50 * 1024**3}, "weight_map": {}} + with ( + mock.patch.object(huggingface_hub, "try_to_load_from_cache", return_value=None), + mock.patch.object(wan_utils, "_free_bytes", return_value=10 * 1024**3), + ): + with self.assertRaisesRegex(OSError, "Refusing to download"): + wan_utils._check_disk_for_shard_download("org/repo", "transformer", ["s1", "s2"], index) + with mock.patch.object(huggingface_hub, "try_to_load_from_cache", return_value="/cached/path"): + wan_utils._check_disk_for_shard_download("org/repo", "transformer", ["s1", "s2"], index) # all cached: no-op + + def test_warm_start_checks_cache_before_the_network(self): + """A valid converted cache for the locally cached index is used without a networked hf_hub_download.""" + snap = os.path.join(self._tmp.name, "hub", "models--org--repo", "snapshots", "rev1", "transformer") + os.makedirs(snap) + index_path = os.path.join(snap, "diffusion_pytorch_model.safetensors.index.json") + with open(index_path, "w") as f: + json.dump({"weight_map": {}}, f) + fingerprint = wan_utils._compute_source_checkpoint_fingerprint(index_path, "org/repo", "transformer") + save_converted_weights(self.cache_dir, self.flat, fingerprint) + + def fake_download(repo_id, subfolder=None, filename=None, local_files_only=False): + del repo_id, subfolder, filename + if not local_files_only: + raise AssertionError("networked hf_hub_download called despite a valid converted cache") + return index_path + + with mock.patch.object(wan_utils, "hf_hub_download", side_effect=fake_download): + loaded = wan_utils.load_base_wan_transformer( + "org/repo", self.eval_shapes, "cpu", subfolder="transformer", converted_cache_dir=self.cache_dir + ) + np.testing.assert_array_equal(loaded["blocks"]["attn1"]["kernel"], self.flat[("blocks", "attn1", "kernel")]) + if __name__ == "__main__": unittest.main() diff --git a/src/maxdiffusion/tests/wan/wan_transformer_test.py b/src/maxdiffusion/tests/wan/wan_transformer_test.py index 69bed9a6a..946e21493 100644 --- a/src/maxdiffusion/tests/wan/wan_transformer_test.py +++ b/src/maxdiffusion/tests/wan/wan_transformer_test.py @@ -34,9 +34,11 @@ ) from maxdiffusion.models.embeddings_flax import NNXTimestepEmbedding, NNXPixArtAlphaTextProjection from maxdiffusion.models.normalization_flax import FP32LayerNorm -from maxdiffusion.models.attention_flax import FlaxWanAttention +from maxdiffusion.models.attention_flax import FlaxWanAttention, _unflatten_heads +from maxdiffusion.kernels.fused_producers import fused_rmsnorm_rope from maxdiffusion.pyconfig import HyperParameters from maxdiffusion.pipelines.wan.wan_pipeline import WanPipeline +import numpy as np import qwix import flax @@ -45,6 +47,11 @@ IN_GITHUB_ACTIONS = os.getenv("GITHUB_ACTIONS") == "true" +# Kernel numerics grids and multi-device tests are skipped in CI; see +# end_to_end/tpu/run_wan_stack_tests.sh. +_SKIP_IN_GITHUB_ACTIONS = unittest.skipIf( + IN_GITHUB_ACTIONS, "TPU kernel / multi-device test, skipped in GitHub Actions; run end_to_end/tpu/run_wan_stack_tests.sh" +) THIS_DIR = os.path.dirname(os.path.abspath(__file__)) @@ -118,12 +125,20 @@ def test_wan_time_text_embedding(self): text_embed_dim = 4096 with self.mesh, nn_partitioning.axis_rules(self.config.logical_axis_rules): layer = WanTimeTextImageEmbedding( - rngs=rngs, dim=dim, time_freq_dim=time_freq_dim, time_proj_dim=time_proj_dim, text_embed_dim=text_embed_dim + rngs=rngs, + dim=dim, + time_freq_dim=time_freq_dim, + time_proj_dim=time_proj_dim, + text_embed_dim=text_embed_dim, ) dummy_timestep = jnp.ones(batch_size) - encoder_hidden_states_shape = (batch_size, time_freq_dim * 2, text_embed_dim) + encoder_hidden_states_shape = ( + batch_size, + time_freq_dim * 2, + text_embed_dim, + ) dummy_encoder_hidden_states = jnp.ones(encoder_hidden_states_shape) temb, timestep_proj, encoder_hidden_states, _, _ = layer(dummy_timestep, dummy_encoder_hidden_states) assert temb.shape == (batch_size, dim) @@ -189,13 +204,22 @@ def test_wan_block(self): mesh=mesh, flash_block_sizes=flash_block_sizes, ) - dummy_output = wan_block(dummy_hidden_states, dummy_encoder_hidden_states, dummy_temb, dummy_rotary_emb) + dummy_output = wan_block( + dummy_hidden_states, + dummy_encoder_hidden_states, + dummy_temb, + dummy_rotary_emb, + ) assert dummy_output.shape == dummy_hidden_states.shape def test_wan_attention(self): for attention_kernel in ["flash", "tokamax_flash"]: pyconfig.initialize( - [None, os.path.join(THIS_DIR, "..", "..", "configs", "base_wan_14b.yml"), f"attention={attention_kernel}"], + [ + None, + os.path.join(THIS_DIR, "..", "..", "configs", "base_wan_14b.yml"), + f"attention={attention_kernel}", + ], unittest=True, ) config = pyconfig.config @@ -231,7 +255,9 @@ def test_wan_attention(self): dummy_hidden_states = jnp.ones(dummy_hidden_states_shape) dummy_encoder_hidden_states = jnp.ones(dummy_hidden_states_shape) dummy_output = attention( - hidden_states=dummy_hidden_states, encoder_hidden_states=dummy_encoder_hidden_states, rotary_emb=dummy_rotary_emb + hidden_states=dummy_hidden_states, + encoder_hidden_states=dummy_encoder_hidden_states, + rotary_emb=dummy_rotary_emb, ) assert dummy_output.shape == dummy_hidden_states_shape @@ -250,6 +276,303 @@ def test_wan_attention(self): except NotImplementedError: pass + @_SKIP_IN_GITHUB_ACTIONS + def test_fused_rmsnorm_rope_parity(self): + """Verifies numerical parity of fused RMSNorm + RoPE against unfused reference with complex freqs_cis.""" + key = jax.random.PRNGKey(202) + k1, k2, k3, k4, k5, k6 = jax.random.split(key, 6) + B, S, D, H, DH = 1, 1024, 5120, 40, 128 + eps = 1e-6 + + raw_q = jax.random.normal(k1, (B, S, D), dtype=jnp.bfloat16) + raw_k = jax.random.normal(k2, (B, S, D), dtype=jnp.bfloat16) + q_scale = jax.random.normal(k3, (D,), dtype=jnp.bfloat16) + k_scale = jax.random.normal(k4, (D,), dtype=jnp.bfloat16) + + freqs_real = jax.random.normal(k5, (1, 1, S, DH // 2), dtype=jnp.float32) + freqs_imag = jax.random.normal(k6, (1, 1, S, DH // 2), dtype=jnp.float32) + freqs_cis = jax.lax.complex(freqs_real, freqs_imag) + + # Unfused reference: FP32 RMSNorm -> unflatten -> RoPE + def ref_norm(x, scale): + var = jnp.mean(jnp.square(x.astype(jnp.float32)), axis=-1, keepdims=True) + return (x.astype(jnp.float32) * jax.lax.rsqrt(var + eps) * scale.astype(jnp.float32)).astype(x.dtype) + + def ref_unflatten(x, heads): + b, s, d = x.shape + return x.reshape(b, s, heads, d // heads).transpose(0, 2, 1, 3) + + def ref_apply_rope(xq, xk, freqs_cis): + cos = jnp.real(freqs_cis).astype(xq.dtype) + sin = jnp.imag(freqs_cis).astype(xq.dtype) + xq_reshaped = xq.reshape(*xq.shape[:-1], -1, 2) + xk_reshaped = xk.reshape(*xk.shape[:-1], -1, 2) + xq_0, xq_1 = xq_reshaped[..., 0], xq_reshaped[..., 1] + xk_0, xk_1 = xk_reshaped[..., 0], xk_reshaped[..., 1] + xq_out_0 = xq_0 * cos - xq_1 * sin + xq_out_1 = xq_0 * sin + xq_1 * cos + xk_out_0 = xk_0 * cos - xk_1 * sin + xk_out_1 = xk_0 * sin + xk_1 * cos + xq_out = jnp.concatenate([xq_out_0[..., None], xq_out_1[..., None]], axis=-1).reshape(xq.shape) + xk_out = jnp.concatenate([xk_out_0[..., None], xk_out_1[..., None]], axis=-1).reshape(xk.shape) + return xq_out, xk_out + + q_ref, k_ref = ref_apply_rope( + ref_unflatten(ref_norm(raw_q, q_scale), H), + ref_unflatten(ref_norm(raw_k, k_scale), H), + freqs_cis, + ) + + # Fused producer + q_fused, k_fused = fused_rmsnorm_rope( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=H, + dim_head=DH, + eps=eps, + ) + + np.testing.assert_allclose( + np.array(q_ref, dtype=np.float32), + np.array(q_fused, dtype=np.float32), + atol=0.08, + rtol=1e-2, + ) + np.testing.assert_allclose( + np.array(k_ref, dtype=np.float32), + np.array(k_fused, dtype=np.float32), + atol=0.08, + rtol=1e-2, + ) + + @_SKIP_IN_GITHUB_ACTIONS + def test_fused_rmsnorm_rope_with_wan_rotary_embed(self): + """Verifies numerical parity of fused RMSNorm + RoPE against unfused reference with real WanRotaryPosEmbed frequencies.""" + key = jax.random.PRNGKey(404) + k1, k2, k3, k4 = jax.random.split(key, 4) + B, S, D, H, DH = 1, 1024, 5120, 40, 128 + eps = 1e-6 + + # Generate real RoPE frequencies on the complex unit circle using WanRotaryPosEmbed + wan_rot_embed = WanRotaryPosEmbed(attention_head_dim=DH, patch_size=[1, 2, 2], max_seq_len=1024) + dummy_video = jnp.ones((B, 1, 64, 64, 16)) + freqs_cis = wan_rot_embed(dummy_video) # (1, 1, 1024, 64) + + raw_q = jax.random.normal(k1, (B, S, D), dtype=jnp.bfloat16) + raw_k = jax.random.normal(k2, (B, S, D), dtype=jnp.bfloat16) + q_scale = jax.random.normal(k3, (D,), dtype=jnp.bfloat16) + k_scale = jax.random.normal(k4, (D,), dtype=jnp.bfloat16) + + def ref_norm(x, scale): + var = jnp.mean(jnp.square(x.astype(jnp.float32)), axis=-1, keepdims=True) + return (x.astype(jnp.float32) * jax.lax.rsqrt(var + eps) * scale.astype(jnp.float32)).astype(x.dtype) + + def ref_unflatten(x, heads): + b, s, d = x.shape + return x.reshape(b, s, heads, d // heads).transpose(0, 2, 1, 3) + + def ref_apply_rope(xq, xk, freqs_cis): + cos = jnp.real(freqs_cis).astype(xq.dtype) + sin = jnp.imag(freqs_cis).astype(xq.dtype) + xq_reshaped = xq.reshape(*xq.shape[:-1], -1, 2) + xk_reshaped = xk.reshape(*xk.shape[:-1], -1, 2) + xq_0, xq_1 = xq_reshaped[..., 0], xq_reshaped[..., 1] + xk_0, xk_1 = xk_reshaped[..., 0], xk_reshaped[..., 1] + xq_out_0 = xq_0 * cos - xq_1 * sin + xq_out_1 = xq_0 * sin + xq_1 * cos + xk_out_0 = xk_0 * cos - xk_1 * sin + xk_out_1 = xk_0 * sin + xk_1 * cos + xq_out = jnp.concatenate([xq_out_0[..., None], xq_out_1[..., None]], axis=-1).reshape(xq.shape) + xk_out = jnp.concatenate([xk_out_0[..., None], xk_out_1[..., None]], axis=-1).reshape(xk.shape) + return xq_out, xk_out + + q_ref, k_ref = ref_apply_rope( + ref_unflatten(ref_norm(raw_q, q_scale), H), + ref_unflatten(ref_norm(raw_k, k_scale), H), + freqs_cis, + ) + + q_fused, k_fused = fused_rmsnorm_rope( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=H, + dim_head=DH, + eps=eps, + ) + + np.testing.assert_allclose( + np.array(q_ref, dtype=np.float32), + np.array(q_fused, dtype=np.float32), + atol=0.08, + rtol=1e-2, + ) + np.testing.assert_allclose( + np.array(k_ref, dtype=np.float32), + np.array(k_fused, dtype=np.float32), + atol=0.08, + rtol=1e-2, + ) + + @_SKIP_IN_GITHUB_ACTIONS + def test_fused_rmsnorm_rope_gqa_parity(self): + """Verifies that fused_rmsnorm_rope correctly handles asymmetric GQA shapes (e.g. q_heads=8, kv_heads=2).""" + key = jax.random.PRNGKey(505) + k1, k2, k3, k4, k5 = jax.random.split(key, 5) + + B = 2 + S = 64 + Q_H = 8 + KV_H = 2 + DH = 128 + D_q = Q_H * DH + D_kv = KV_H * DH + eps = 1e-6 + + raw_q = jax.random.normal(k1, (B, S, D_q), dtype=jnp.bfloat16) + raw_k = jax.random.normal(k2, (B, S, D_kv), dtype=jnp.bfloat16) + q_scale = jax.random.normal(k3, (D_q,), dtype=jnp.float32) + k_scale = jax.random.normal(k4, (D_kv,), dtype=jnp.float32) + freqs_cis = jax.random.normal(k5, (1, 1, S, DH // 2), dtype=jnp.float32) + 1j * jax.random.normal( + key, (1, 1, S, DH // 2), dtype=jnp.float32 + ) + + def ref_norm(x, scale): + x_fp32 = x.astype(jnp.float32) + rms = jax.lax.rsqrt(jnp.mean(jnp.square(x_fp32), axis=-1, keepdims=True) + eps) + return (x_fp32 * rms * scale).astype(x.dtype) + + def ref_unflatten(x, heads): + return x.reshape(B, S, heads, DH).transpose(0, 2, 1, 3) + + def ref_apply_rope(xq, xk, freqs): + cos = jnp.real(freqs).astype(xq.dtype) + sin = jnp.imag(freqs).astype(xq.dtype) + xq_reshaped = xq.reshape(*xq.shape[:-1], -1, 2) + xk_reshaped = xk.reshape(*xk.shape[:-1], -1, 2) + xq_0, xq_1 = xq_reshaped[..., 0], xq_reshaped[..., 1] + xk_0, xk_1 = xk_reshaped[..., 0], xk_reshaped[..., 1] + xq_out_0 = xq_0 * cos - xq_1 * sin + xq_out_1 = xq_0 * sin + xq_1 * cos + xk_out_0 = xk_0 * cos - xk_1 * sin + xk_out_1 = xk_0 * sin + xk_1 * cos + xq_out = jnp.concatenate([xq_out_0[..., None], xq_out_1[..., None]], axis=-1).reshape(xq.shape) + xk_out = jnp.concatenate([xk_out_0[..., None], xk_out_1[..., None]], axis=-1).reshape(xk.shape) + return xq_out, xk_out + + q_ref, k_ref = ref_apply_rope( + ref_unflatten(ref_norm(raw_q, q_scale), Q_H), + ref_unflatten(ref_norm(raw_k, k_scale), KV_H), + freqs_cis, + ) + + q_fused, k_fused = fused_rmsnorm_rope( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=Q_H, + kv_heads=KV_H, + dim_head=DH, + eps=eps, + ) + + self.assertEqual(q_fused.shape, (B, Q_H, S, DH)) + self.assertEqual(k_fused.shape, (B, KV_H, S, DH)) + + np.testing.assert_allclose( + np.array(q_ref, dtype=np.float32), + np.array(q_fused, dtype=np.float32), + atol=0.08, + rtol=1e-2, + ) + np.testing.assert_allclose( + np.array(k_ref, dtype=np.float32), + np.array(k_fused, dtype=np.float32), + atol=0.08, + rtol=1e-2, + ) + + @_SKIP_IN_GITHUB_ACTIONS + def test_wan_self_attention_is_self_attention_dispatch(self): + """Verifies that FlaxWanAttention with is_self_attention=True dispatches fused RMSNorm+RoPE and matches reference.""" + from unittest import mock + + key = jax.random.PRNGKey(303) + k1, k2, k3, k4 = jax.random.split(key, 4) + rngs = nnx.Rngs(k1) + + batch_size = 1 + seq_len = 1024 + query_dim = 5120 + heads = 40 + dim_head = 128 + + flash_block_sizes = get_flash_block_sizes(self.config) + with self.mesh, nn_partitioning.axis_rules(self.config.logical_axis_rules): + attn = FlaxWanAttention( + rngs=rngs, + query_dim=query_dim, + heads=heads, + dim_head=dim_head, + attention_kernel="dot_product", + mesh=self.mesh, + flash_block_sizes=flash_block_sizes, + is_self_attention=True, + ) + self.assertTrue(attn.is_self_attention) + + hidden_states = jax.random.normal(k2, (batch_size, seq_len, query_dim), dtype=jnp.bfloat16) + freqs_real = jax.random.normal(k3, (1, 1, seq_len, dim_head // 2), dtype=jnp.float32) + freqs_imag = jax.random.normal(k4, (1, 1, seq_len, dim_head // 2), dtype=jnp.float32) + rotary_emb = jax.lax.complex(freqs_real, freqs_imag) + + # Exercise both WanTransformerBlock's call (encoder_hidden_states=None) + # and the explicit identity call (encoder_hidden_states=hidden_states). + with mock.patch( + "maxdiffusion.models.attention_flax.fused_rmsnorm_rope", + wraps=fused_rmsnorm_rope, + ) as spy_fused: + out_none = attn( + hidden_states=hidden_states, + encoder_hidden_states=None, + rotary_emb=rotary_emb, + ) + out = attn( + hidden_states=hidden_states, + encoder_hidden_states=hidden_states, + rotary_emb=rotary_emb, + ) + self.assertEqual(spy_fused.call_count, 2) + self.assertEqual(out.shape, (batch_size, seq_len, query_dim)) + np.testing.assert_array_equal(np.asarray(out_none), np.asarray(out)) + + # Reference unfused path execution with identical weights + raw_q = attn.query(hidden_states) + raw_k = attn.key(hidden_states) + raw_v = attn.value(hidden_states) + q_norm = attn.norm_q(raw_q) + k_norm = attn.norm_k(raw_k) + q_h = _unflatten_heads(q_norm, heads) + k_h = _unflatten_heads(k_norm, heads) + v_h = _unflatten_heads(raw_v, heads) + q_rope, k_rope = attn._apply_rope(q_h, k_h, rotary_emb) + ref_attn_out = attn.attention_op.apply_attention(q_rope, k_rope, v_h, attention_mask=None) + ref_out = attn.proj_attn(ref_attn_out) + + np.testing.assert_allclose( + np.array(out, dtype=np.float32), + np.array(ref_out, dtype=np.float32), + atol=0.08, + rtol=1e-2, + ) + @pytest.mark.skipif(IN_GITHUB_ACTIONS, reason="Don't run smoke tests on Github Actions") def test_wan_model(self): pyconfig.initialize( @@ -280,14 +603,20 @@ def test_wan_model(self): num_layers = 1 with nn_partitioning.axis_rules(config.logical_axis_rules): wan_model = WanModel( - rngs=rngs, attention="flash", mesh=mesh, flash_block_sizes=flash_block_sizes, num_layers=num_layers + rngs=rngs, + attention="flash", + mesh=mesh, + flash_block_sizes=flash_block_sizes, + num_layers=num_layers, ) dummy_timestep = jnp.ones((batch_size)) dummy_encoder_hidden_states = jnp.ones((batch_size, 512, 4096)) with mesh: dummy_output = wan_model( - hidden_states=dummy_hidden_states, timestep=dummy_timestep, encoder_hidden_states=dummy_encoder_hidden_states + hidden_states=dummy_hidden_states, + timestep=dummy_timestep, + encoder_hidden_states=dummy_encoder_hidden_states, ) assert dummy_output.shape == hidden_states_shape diff --git a/src/maxdiffusion/tests/wan/wan_warmup_coverage_test.py b/src/maxdiffusion/tests/wan/wan_warmup_coverage_test.py new file mode 100644 index 000000000..d242e1bec --- /dev/null +++ b/src/maxdiffusion/tests/wan/wan_warmup_coverage_test.py @@ -0,0 +1,251 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Warmup must compile the loop's forward pass for both WAN 2.2 transformers. + +With flow_shift=12 the 2-step warmup schedule is t=[999, 923], both above the +875 (T2V) and 900 (I2V) boundaries, so the denoise loop alone never reaches the +low-noise transformer. run_inference_2_2 and run_inference_2_2_i2v must compile +it explicitly during warmup, without executing either transformer. +""" + +import contextlib +import os +import tempfile +import unittest +from unittest import mock + +import jax +import jax.numpy as jnp +import numpy as np +from jax.sharding import Mesh + +from maxdiffusion import aot_cache +from maxdiffusion.pipelines.wan import wan_denoise_utils +from maxdiffusion.pipelines.wan import wan_pipeline_2_2 as p22 +from maxdiffusion.pipelines.wan import wan_pipeline_i2v_2p2 as pi2v + + +class _Step: + + def __init__(self, latents, state): + self._out = (latents, state) + + def to_tuple(self): + return self._out + + +class _WarmupCoverageMixin: + """Shared by T2V and I2V so the two pipelines are held to the same warmup contract.""" + + module = None + + def _invoke(self, scheduler, scheduler_state, guidance_scale_low, guidance_scale_high, num_inference_steps): + raise NotImplementedError + + def _run( + self, + timesteps, + in_warmup, + guidance_scale_low=3.0, + guidance_scale_high=4.0, + max_inflight=None, + blocked=None, + ): + calls = [] + full_cfg_calls = [] + + def fake_forward(graphdef, *args, **kwargs): + calls.append(graphdef) + return jnp.zeros((1, 4, 2, 4, 4)) + + def fake_full_cfg(graphdef, *args, **kwargs): + full_cfg_calls.append(graphdef) + return jnp.zeros((1, 4, 2, 4, 4)) + + def fake_block_until_ready(x): + if blocked is not None: + blocked.append(x) + return x + + merged = mock.MagicMock() + merged.rope.return_value = jnp.zeros((1,)) + scheduler = mock.MagicMock() + scheduler.step.side_effect = lambda st, noise, t, latents: _Step(latents, st) + scheduler_state = mock.MagicMock() + scheduler_state.timesteps = timesteps + real_execution = mock.MagicMock() + + m = self.module + env = {} if max_inflight is None else {"MAXD_QUEUE_MAX_INFLIGHT": str(max_inflight)} + with contextlib.ExitStack() as stack: + stack.enter_context(mock.patch.dict(os.environ, env)) + stack.enter_context(mock.patch.object(m, "transformer_forward_pass", side_effect=fake_forward)) + stack.enter_context(mock.patch.object(m, "transformer_forward_pass_full_cfg", side_effect=fake_full_cfg)) + stack.enter_context(mock.patch.object(m.nnx, "merge", return_value=merged)) + stack.enter_context(mock.patch.object(m.aot_cache, "in_warmup", return_value=in_warmup)) + stack.enter_context(mock.patch.object(m.aot_cache, "real_execution", real_execution)) + stack.enter_context(mock.patch.object(m.jax, "block_until_ready", side_effect=fake_block_until_ready)) + self._invoke(scheduler, scheduler_state, guidance_scale_low, guidance_scale_high, len(timesteps)) + self.assertEqual(full_cfg_calls, [], "the loop must use transformer_forward_pass, not full_cfg") + real_execution.assert_not_called() # warmup compiles only; nothing executes for real + return calls + + def test_warmup_schedule_all_high_still_compiles_low(self): + blocked = [] + calls = self._run([999, 923], in_warmup=True, blocked=blocked) + self.assertEqual(calls, ["high", "low", "high", "high"]) + self.assertEqual(blocked, [], "warmup must not wait on device work") + + def test_warmup_schedule_crossing_boundary_compiles_both(self): + self.assertEqual(self._run([999, 800], in_warmup=True), ["high", "low", "high", "low"]) + + def test_warmup_no_cfg_compiles_both(self): + calls = self._run([999, 923], in_warmup=True, guidance_scale_low=1.0, guidance_scale_high=1.0) + self.assertEqual(calls, ["high", "low", "high", "high"]) + + def test_real_run_is_unchanged(self): + calls = self._run([999, 923], in_warmup=False) + self.assertEqual(calls, ["high", "high"]) + + def test_loop_bounds_inflight_steps(self): + blocked = [] + self._run([999, 923, 800, 700], in_warmup=False, max_inflight=1, blocked=blocked) + self.assertEqual(len(blocked), 3, "with depth 1, each step after the first must wait on its predecessor") + blocked.clear() + self._run([999, 923, 800, 700], in_warmup=False, max_inflight=0, blocked=blocked) + self.assertEqual(blocked, [], "depth 0 disables the bound") + + +class WarmupCoversBothTransformersTest(_WarmupCoverageMixin, unittest.TestCase): + """T2V: run_inference_2_2.""" + + module = p22 + + def _invoke(self, scheduler, scheduler_state, guidance_scale_low, guidance_scale_high, num_inference_steps): + p22.run_inference_2_2( + low_noise_graphdef="low", + low_noise_state=None, + low_noise_rest=None, + high_noise_graphdef="high", + high_noise_state=None, + high_noise_rest=None, + latents=jnp.zeros((1, 4, 2, 4, 4)), + prompt_embeds=jnp.zeros((1, 8, 16)), + negative_prompt_embeds=jnp.zeros((1, 8, 16)), + guidance_scale_low=guidance_scale_low, + guidance_scale_high=guidance_scale_high, + boundary=875, + num_inference_steps=num_inference_steps, + scheduler=scheduler, + scheduler_state=scheduler_state, + config=None, + ) + + +class I2VWarmupCoversBothTransformersTest(_WarmupCoverageMixin, unittest.TestCase): + """I2V: run_inference_2_2_i2v (latents and condition are BFHWC).""" + + module = pi2v + + def _invoke(self, scheduler, scheduler_state, guidance_scale_low, guidance_scale_high, num_inference_steps): + pi2v.run_inference_2_2_i2v( + low_noise_graphdef="low", + low_noise_state=None, + low_noise_rest=None, + high_noise_graphdef="high", + high_noise_state=None, + high_noise_rest=None, + latents=jnp.zeros((1, 2, 4, 4, 4)), + condition=jnp.zeros((1, 2, 4, 4, 5)), + prompt_embeds=jnp.zeros((1, 8, 16)), + negative_prompt_embeds=jnp.zeros((1, 8, 16)), + image_embeds=jnp.zeros((1, 4, 16)), + guidance_scale_low=guidance_scale_low, + guidance_scale_high=guidance_scale_high, + boundary=900, + num_inference_steps=num_inference_steps, + scheduler=scheduler, + scheduler_state=scheduler_state, + config=None, + ) + + +class DenoiseUtilsTest(unittest.TestCase): + + def test_inflight_window_blocks_on_oldest_beyond_depth(self): + with mock.patch.object(wan_denoise_utils.jax, "block_until_ready") as block: + window = wan_denoise_utils.InflightWindow(depth=2) + for i in range(5): + window.push(i) + self.assertEqual([c.args[0] for c in block.call_args_list], [0, 1, 2]) + + def test_inflight_window_depth_from_env_and_zero_disables(self): + with mock.patch.dict(os.environ, {"MAXD_QUEUE_MAX_INFLIGHT": "3"}): + self.assertEqual(wan_denoise_utils.InflightWindow().depth, 3) + with mock.patch.dict(os.environ): + os.environ.pop("MAXD_QUEUE_MAX_INFLIGHT", None) + self.assertEqual(wan_denoise_utils.InflightWindow().depth, 4) + with mock.patch.object(wan_denoise_utils.jax, "block_until_ready") as block: + window = wan_denoise_utils.InflightWindow(depth=0) + for i in range(5): + window.push(i) + block.assert_not_called() + self.assertFalse(window._inflight, "a disabled window must not hold step outputs") + + def test_compile_experts_is_a_noop_outside_warmup(self): + branch = mock.MagicMock() + with mock.patch.object(wan_denoise_utils.aot_cache, "in_warmup", return_value=False): + wan_denoise_utils.compile_experts((branch, branch), "operands") + branch.assert_not_called() + + def test_compile_experts_compiles_one_shared_executable_without_executing(self): + """With the real AOT cache: both experts compile once, nothing runs, and the loop reuses the executable.""" + tmp = tempfile.TemporaryDirectory() + self.addCleanup(tmp.cleanup) + self.addCleanup(aot_cache.install, "", {}, None) + aot_cache.install(tmp.name, meta={"test": "compile_experts"}, mesh=Mesh(np.array(jax.devices()[:1]), ("d",))) + + @aot_cache.cached_jit + def forward(x, guidance_scale): + return x * guidance_scale + + outs = [] + + def expert(guidance_scale): + def run(operands): + outs.append(forward(operands, guidance_scale)) + return outs[-1] + + return run + + x = jnp.arange(1, 5, dtype=jnp.bfloat16) + high, low = expert(4.0), expert(3.0) + with aot_cache.warmup_mode(): + wan_denoise_utils.compile_experts((high, low), x) + + self.assertEqual(len(forward._compiled), 1, "guidance 4.0 and 3.0 must share one executable") + self.assertEqual(len(outs), 2) + for out in outs: # compile only: warmup hands back zeros instead of executing + np.testing.assert_array_equal(np.asarray(out, np.float32), np.zeros(4)) + # The denoise loop then runs each expert on that executable with its own guidance value. + np.testing.assert_array_equal(np.asarray(high(x), np.float32), np.asarray(x * 4.0, np.float32)) + np.testing.assert_array_equal(np.asarray(low(x), np.float32), np.asarray(x * 3.0, np.float32)) + self.assertEqual(len(forward._compiled), 1, "the loop must not compile a second executable") + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/trainers/base_wan_trainer.py b/src/maxdiffusion/trainers/base_wan_trainer.py index 836c2bd2e..9356a6cb3 100644 --- a/src/maxdiffusion/trainers/base_wan_trainer.py +++ b/src/maxdiffusion/trainers/base_wan_trainer.py @@ -259,7 +259,8 @@ def start_training(self): if self.config.enable_ssim: posttrained_video_path = generate_sample(self.config, pipeline, filename_prefix="post-training-") - print_ssim(pretrained_video_path, posttrained_video_path) + if jax.process_index() == 0: + print_ssim(pretrained_video_path, posttrained_video_path) def eval(self, mesh, eval_rng_key, step, p_eval_step, state, scheduler_state, writer): eval_data_iterator = self.load_dataset(mesh, is_training=False) diff --git a/src/maxdiffusion/trainers/wan_trainer_2_2.py b/src/maxdiffusion/trainers/wan_trainer_2_2.py index 13650a0d9..8957ce245 100644 --- a/src/maxdiffusion/trainers/wan_trainer_2_2.py +++ b/src/maxdiffusion/trainers/wan_trainer_2_2.py @@ -258,7 +258,8 @@ def start_training(self): if self.config.enable_ssim: posttrained_video_path = self.generate_sample(self.config, pipeline, filename_prefix="post-training-") - print_ssim(pretrained_video_path, posttrained_video_path) + if jax.process_index() == 0: + print_ssim(pretrained_video_path, posttrained_video_path) def training_loop_2_2( self, diff --git a/src/maxdiffusion/utils/export_utils.py b/src/maxdiffusion/utils/export_utils.py index 279ad1e90..da7e8838e 100644 --- a/src/maxdiffusion/utils/export_utils.py +++ b/src/maxdiffusion/utils/export_utils.py @@ -131,7 +131,9 @@ def export_to_obj(mesh, output_obj_path: str = None): def _legacy_export_to_video( - video_frames: Union[List[np.ndarray], List[PIL.Image.Image]], output_video_path: str = None, fps: int = 10 + video_frames: Union[List[np.ndarray], List[PIL.Image.Image]], + output_video_path: str = None, + fps: int = 10, ): if is_opencv_available(): import cv2 @@ -212,21 +214,35 @@ def export_to_video( if output_video_path is None: output_video_path = tempfile.NamedTemporaryFile(suffix=".mp4").name - if isinstance(video_frames, np.ndarray): - if video_frames.dtype != np.uint8: - video_frames = (video_frames * 255).astype(np.uint8) - elif isinstance(video_frames[0], np.ndarray): - video_frames = np.stack(video_frames) - if video_frames.dtype != np.uint8: - video_frames = (video_frames * 255).astype(np.uint8) - elif isinstance(video_frames[0], PIL.Image.Image): + if isinstance(video_frames, list) and len(video_frames) > 0 and isinstance(video_frames[0], PIL.Image.Image): video_frames = np.stack([np.asarray(frame) for frame in video_frames]) + else: + video_frames = np.asarray(video_frames) + if video_frames.dtype != np.uint8: + video_frames = (video_frames * 255).clip(0, 255).astype(np.uint8) - with imageio.get_writer( - output_video_path, fps=fps, quality=quality, bitrate=bitrate, macro_block_size=macro_block_size - ) as writer: - for frame in video_frames: - writer.append_data(frame) + import os + import uuid + + tmp_video_path = f"{output_video_path}.tmp.{os.getpid()}.{uuid.uuid4().hex[:8]}.mp4" + try: + with imageio.get_writer( + tmp_video_path, + fps=fps, + quality=quality, + bitrate=bitrate, + macro_block_size=macro_block_size, + ) as writer: + for frame in video_frames: + writer.append_data(np.asarray(frame)) + os.replace(tmp_video_path, output_video_path) + except Exception: + try: + if os.path.exists(tmp_video_path): + os.remove(tmp_video_path) + except OSError: + pass + raise return output_video_path @@ -320,7 +336,12 @@ def _write_audio( def export_to_video_with_audio( - video: Any, fps: int, audio: Optional[Any], audio_sample_rate: Optional[int], output_path: str, audio_format: str = "s16" + video: Any, + fps: int, + audio: Optional[Any], + audio_sample_rate: Optional[int], + output_path: str, + audio_format: str = "s16", ) -> None: """ Encodes video (and optionally audio) to a file using PyAV. @@ -369,6 +390,12 @@ def export_to_video_with_audio( container.mux(packet) if audio is not None: - _write_audio(container, audio_stream, audio, audio_sample_rate, target_format=audio_format) + _write_audio( + container, + audio_stream, + audio, + audio_sample_rate, + target_format=audio_format, + ) container.close()