Skip to content

feat(wan): fast serving with persistent AOT caching and tuned inference recipe - #479

Open
Perseus14 wants to merge 1 commit into
feat/ring-attentionfrom
feat/wan-fast-serving
Open

Perseus14 wants to merge 1 commit into
feat/ring-attentionfrom
feat/wan-fast-serving

Conversation

@Perseus14

@Perseus14 Perseus14 commented Sep 13, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Fast serving for Wan on TPU: a persistent per-shape AOT executable cache, a converted-weights cache, and a tuned launcher (end_to_end/tpu/run_wan_fast_inference.sh).

User-visible changes

  • AOT caching is opt-in.
    • By default generate_wan.run() is plain jax.jit.
    • A persistent cache is used only when aot_cache_dir is set and the source revision is reusable. An explicit dirty:/unversioned: revision or enable_zero_execution_warmup: True (default False in all 6 Wan ymls) uses an ephemeral temp dir, torn down in a try/finally covering setup and install.
    • aot_cache_dir must be a local path: gs:// raises ValueError regardless of revision.
  • Cache formats:
    • .aotx format v2 (v1 files are recompiled once).
    • Converted-weights cache format v3 (v1 caches are re-converted once with a warning and disk-space check; matching v2 manifests migrate header-in-place to v3).

What's in it

  • AOT signatures & metadata: keyed on scalar types, static values, a deterministic GraphDef digest, device kind, package versions, XLA/libtpu flags, PRNG/x64 settings, and a content hash of every non-test .py in maxdiffusion. Per-host .aotx files for multi-host.
  • .aotx loading: exact (module, name) pickle allowlist rejecting dotted names (defence in depth; directory write permissions are the real control).
  • Converted-weights fingerprint: repo id + HF snapshot revision (before resolving symlinks) + subfolder + index contents, fail-closed when missing, checked with local_files_only before network calls.
  • Warmup & Launcher: compile_experts compiles both experts without executing when AOT is installed; InflightWindow bounds queued steps; platform-detected v6e/v7/generic profiles with opt-in PIN_DVFS_P_STATE.

Performance

Wan 2.2 T2V-A14B, 720p / 81 frames / 40 steps, warm AOT, measured with this PR at top of stack:

TPU DVFS Generate Denoise
tpu7x-8 unpinned 115.5 s 112.8 s
tpu7x-8 pinned 102.0 s 99.4 s

Tests

aot_cache_test.py (41), converted_weights_cache_test.py (16), wan_transformer_test.py (15), wan_warmup_coverage_test.py (14).

  • Covers unpickler allowlist, run() AOT gating and teardown (including install failure), gs:// rejection, float32_qk_product keeping bf16 QK operands, fail-closed manifests, disk-space checks, revision/subfolder fingerprinting, and real TPU executable reload.
  • CI vs Local VM: 4 heavy wan_transformer tests skip in GITHUB_ACTIONS=true (40 of 201 cumulative); all 57 AOT/converted-weights and 14 warmup tests run in CI.
  • Results (TPU v6e-8): 201 passed cumulative (run_wan_stack_tests.sh, 360.5 s; 86 passed in this PR's 4 test files).

Stack: #477 → #478 → #479 → #488 → #491

@Perseus14
Perseus14 requested a review from entrpn as a code owner September 13, 2026 06:08
@github-actions

Copy link
Copy Markdown

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request introduces several optimizations and features for Wan model inference, including a fast-path signature cache in the AOT cache, support for fused LayerNorm and AdaLN kernels, and updated configuration options. The review feedback highlights a potential concurrency issue in the signature cache that could lead to race conditions, suggests moving inline imports of the fused kernel to the top of the file to reduce overhead in hot paths, and recommends avoiding os.path.join on GCS URIs to ensure cross-platform safety.

Comment thread src/maxdiffusion/aot_cache.py Outdated
Comment thread src/maxdiffusion/models/wan/transformers/transformer_wan.py Outdated
Comment thread src/maxdiffusion/models/wan/transformers/transformer_wan.py Outdated
Comment thread src/maxdiffusion/models/wan/transformers/transformer_wan.py Outdated
Comment thread src/maxdiffusion/generate_wan.py Outdated
Comment thread src/maxdiffusion/generate_wan.py Outdated
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from c9dc7b2 to 80719d9 Compare September 13, 2026 06:16
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from 80719d9 to 585247a Compare September 13, 2026 06:25
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from 585247a to 2ee9c66 Compare September 13, 2026 06:30
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from 2ee9c66 to 8756432 Compare September 13, 2026 06:49
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from 8756432 to 2957831 Compare September 13, 2026 07:15
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch 2 times, most recently from 8eb1b56 to dc78d28 Compare September 13, 2026 15:29
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from dc78d28 to e98f895 Compare September 13, 2026 15:40
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from e98f895 to ba40a77 Compare September 13, 2026 16:21
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from ba40a77 to d7361d3 Compare September 13, 2026 16:38
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from d7361d3 to 3540cb7 Compare September 13, 2026 18:29
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from 3540cb7 to 89af7b2 Compare September 13, 2026 18:53
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from 89af7b2 to a1e652a Compare September 13, 2026 19:09
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from a1e652a to 838af4b Compare September 13, 2026 19:49
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch 3 times, most recently from 09866a3 to 8631d1f Compare September 17, 2026 11:56
@Perseus14

Copy link
Copy Markdown
Collaborator Author

Thanks for testing PR #479 on v7x-8 @syhuang22! Updated the code and addressed your comments!

@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch from 8631d1f to a909245 Compare September 17, 2026 12:31
@Perseus14
Perseus14 force-pushed the feat/wan-fast-serving branch 2 times, most recently from 3a0c913 to a415017 Compare September 17, 2026 14:34
syhuang22
syhuang22 previously approved these changes Sep 17, 2026
@eltsai

eltsai commented Sep 23, 2026

Copy link
Copy Markdown
Collaborator

I was trying to reproduce the results, and I got:

Configuration Total inference Denoise (40 steps) VAE decode Warm compile
Baseline — origin/main@1bc54811 110.8 s 107.1 s 1.5 s 19.3 s
+ #477 + #478 + #479 101.2 s 98.7 s 0.7 s 7.5 s
+ #488 (full stack) 99.4 s 96.8 s 0.7 s 7.6 s

Deltas vs baseline

Configuration Total inference Denoise VAE decode
#477 + #478 + #479 −9.6 s (−8.7%) −8.4 s (−7.8%) −0.8 s (−53%)
+ #488 −11.4 s (−10.4%) −10.3 s (−9.6%) −0.8 s (−53%)

this is on v7x-8. The stats are slightly different, but I think the perf improvement is bigger than 3.8%

Comment thread src/maxdiffusion/generate_wan.py
eltsai
eltsai previously approved these changes Oct 2, 2026
…ce 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).
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants