Repository navigation
Conversation
There was a problem hiding this comment.
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.
c9dc7b2 to
80719d9
Compare
80719d9 to
585247a
Compare
585247a to
2ee9c66
Compare
2ee9c66 to
8756432
Compare
8756432 to
2957831
Compare
8eb1b56 to
dc78d28
Compare
dc78d28 to
e98f895
Compare
e98f895 to
ba40a77
Compare
ba40a77 to
d7361d3
Compare
d7361d3 to
3540cb7
Compare
3540cb7 to
89af7b2
Compare
89af7b2 to
a1e652a
Compare
a1e652a to
838af4b
Compare
09866a3 to
8631d1f
Compare
|
Thanks for testing PR #479 on v7x-8 @syhuang22! Updated the code and addressed your comments! |
8631d1f to
a909245
Compare
3a0c913 to
a415017
Compare
a415017 to
6b85dfc
Compare
6b85dfc to
7e38c04
Compare
7e38c04 to
1d65bfa
Compare
|
I was trying to reproduce the results, and I got:
Deltas vs baseline
this is on v7x-8. The stats are slightly different, but I think the perf improvement is bigger than 3.8% |
…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).
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
generate_wan.run()is plainjax.jit.aot_cache_diris set and the source revision is reusable. An explicitdirty:/unversioned:revision orenable_zero_execution_warmup: True(default False in all 6 Wan ymls) uses an ephemeral temp dir, torn down in atry/finallycovering setup and install.aot_cache_dirmust be a local path:gs://raisesValueErrorregardless of revision..aotxformat v2 (v1 files are recompiled once).What's in it
.pyinmaxdiffusion. Per-host.aotxfiles for multi-host..aotxloading: exact(module, name)pickle allowlist rejecting dotted names (defence in depth; directory write permissions are the real control).local_files_onlybefore network calls.compile_expertscompiles both experts without executing when AOT is installed;InflightWindowbounds queued steps; platform-detected v6e/v7/generic profiles with opt-inPIN_DVFS_P_STATE.Performance
Wan 2.2 T2V-A14B, 720p / 81 frames / 40 steps, warm AOT, measured with this PR at top of stack:
Tests
aot_cache_test.py(41),converted_weights_cache_test.py(16),wan_transformer_test.py(15),wan_warmup_coverage_test.py(14).run()AOT gating and teardown (including install failure),gs://rejection,float32_qk_productkeeping bf16 QK operands, fail-closed manifests, disk-space checks, revision/subfolder fingerprinting, and real TPU executable reload.wan_transformertests skip inGITHUB_ACTIONS=true(40 of 201 cumulative); all 57 AOT/converted-weights and 14 warmup tests run in CI.run_wan_stack_tests.sh, 360.5 s; 86 passed in this PR's 4 test files).Stack: #477 → #478 → #479 → #488 → #491