Repository navigation
feat(wan): fast inference optimizations (fused cross-attention, token patch embed, shard-major A2A, lane padding) - #491
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces several performance optimizations for Wan models on TPU, including a fused short-KV cross-attention Pallas kernel, token-sharded patch embedding, inverse Ulysses all-to-all optimizations, and token sequence padding. A critical issue was identified in the new Pallas cross-attention kernel where the BlockSpec index mapping fails to scale grid indices by their block sizes, which would result in incorrect memory slicing and wrong outputs.
2ba8966 to
e86f174
Compare
e86f174 to
1d4c9f6
Compare
0d91078 to
4703103
Compare
295c3b5 to
7a60c3d
Compare
08d405c to
7db8313
Compare
7db8313 to
6fdc94e
Compare
ebaac14 to
58f889e
Compare
58f889e to
68a580f
Compare
68a580f to
395f932
Compare
53fa702 to
6a446d5
Compare
6a446d5 to
d11ccb6
Compare
888462a to
cdc584f
Compare
|
thanks for these stacks of PRs @Perseus14 ! I wonder if we can do an experiment on videos with lengths of 2.5 sec, 5 sec (already covered), 7.5 sec, and 10 sec to see if we would see any regression on other seq length? |
cdc584f to
2eba645
Compare
08ccfb2 to
68c3732
Compare
|
Thanks @eltsai — ran the sweep at 41 / 81 / 121 / 161 frames (2.5 / 5 / 7.5 / 10 s @ 16 fps), 720p, 40 steps, seed 12345, same prompt, DP=2×CP=4, DVFS unpinned, on tpu7x-8 and v6e-8. A =
No regression vs Note that all runs use the launcher's 81 f-tuned flash block sizes (v6e: BQ 9472 / BKV 1024; tpu7x: BQ 4736 / BKV 2048). We did not do a block-size sweep per length: the sequence length (and therefore the per-shard length and its alignment to BQ) differs for each frame count, so the optimal block sizes need to be found separately for each frame choice. The non-81 f numbers above are therefore conservative for all three configurations. * v6e @ 121 f is the one length where the launcher's v6e profile (Ulysses U=4, BQ 9472) is the wrong pick: the per-shard length (27,900 tokens) is not block-aligned, and the 2D-ring recipe that is also in this stack ( All runs were bit-reproducible across cold/warm AOT (identical md5s), and the outputs are visually clean on both platforms. |
0414334 to
34789ff
Compare
… patch embed, shard-major A2A, lane padding)
All four switches default off ("xla" / "conv" / "flat" / "off") in the Wan
YAMLs and in wan_runtime_options. They are plain config keys, plumbed to
FlaxWanAttention / WanModel through `attention_config_entries(config)` (the
#488 mechanism) and resolved once at module build; nothing is read from the
environment at trace time. The unconditional pyconfig default for
`wan_cross_attn_cpu_interpret` was removed (it is an attention_config entry,
default False, used only by CPU tests).
- Short-KV text cross-attention Pallas kernel (wan_cross_attn_kernel="pallas"):
kernels/cross_attention_pallas.py keeps a head block's whole 512-token text
KV resident in VMEM and computes an exact single-pass softmax per query tile
(full-row max, exp2, one numerator/denominator, normalised after PV; no
running max / online rescaling). It replaces only what would otherwise be
the unmasked XLA dot-product fallback. FlaxWanAttention falls back to XLA
(logging the reason once) when the inputs or sharding don't fit, and now also
when the VMEM budget is not validated or too small: 64 MiB is requested only
when every mesh TPU is v6e or tpu7x (otherwise Mosaic's default), and
`estimate_vmem_bytes` (Q/out/K/V double buffers + one head's FP32 scores) is
checked against it (~42 MiB at the Wan shard shape; MAX_KV_LEN KV does not
fit and falls back rather than failing to compile). The estimate has only
been checked against real compiles at Wan's KV length (512). The kernel
itself raises a ValueError for an over-budget working set only when a budget
resolves (an explicit vmem_limit_bytes, or the v6e/tpu7x default); under
Mosaic's default it does no check of its own, and callers that can fall back
use `vmem_unusable_reason`. The two fallback reasons (unusable inputs, VMEM)
log under separate warn-once keys.
- Token-sharded patch embedding (wan_patch_embed_mode="tokens"): the Conv3d
patch embedding becomes a patchify + matmul on the token-sharded layout,
avoiding the activation all-gather.
- Shard-major Ulysses A2A (wan_ulysses_out_a2a="chunked" / "shard_major"): a
single layout-preserving all_to_all over a reshaped (num_shards, chunk)
leading axis (not chunked compute/comm overlap), so XLA no longer inserts
relayout copies around the collective. With the custom splash kernel (and
S_local % block_q == 0, D == 128, heads_per_tile == 1) the forward Q
exchange also stays shard-major (4-D q) and the kernel writes the
shard-major output directly (out_num_shards), on both the 1D Ulysses path
(v6e) and the 2D Ulysses+Ring path (tpu7x).
- Lane sequence padding (wan_seq_pad="lane"): pads video tokens at the tail to
a per-shard multiple of 128 (75,600 -> 75,776 at 720p/CP=4) so Ulysses
reshapes are copy-free; TokenPadding masks the pad keys in the custom
ulysses / ulysses_ring_custom kernels, and fixed-m metadata now also
excludes pad query rows from the Q-norm bound. Every configuration that
disables it logs why, once per cause. Behaviour change: if padding is
active, the paths that cannot mask it (flash and non-custom ulysses via
_tpu_flash_attention, ulysses_ring, dot_product, cudnn_flash_te,
head_local_svg) raise NotImplementedError via
`_reject_active_token_padding` instead of silently attending to pads.
- run_wan_fast_inference.sh: the v6e and v7 profiles set
WAN_CROSS_ATTN_KERNEL=pallas, WAN_PATCH_EMBED_MODE=tokens,
WAN_ULYSSES_OUT_A2A=chunked, WAN_SEQ_PAD=lane (all overridable); the
generic profile sets tokens only. The fused RMSNorm+RoPE kernel from #488
stays on in the v6e and v7 profiles and off in the generic one.
LIBTPU: v6e now disables
--xla_tpu_enable_ici_ag_pipelining / --xla_tpu_enable_ag_backward_pipelining,
v7 enables them and adds --xla_tpu_enable_offloading_copy_to_sparsecore=false
and BQ=4736. These two flags moved out of COMMON_LIBTPU, so the generic
(non-v6e/v7) profile no longer sets them at all.
Performance, measured at the stack tip (Wan 2.2 T2V-A14B, 720p/81f/40 steps,
warm AOT, DVFS unpinned unless noted; CP=4, DP=2, launcher defaults including
the fused RoPE kernel):
v6e-8: denoise 125.9s, generate 128.05s
tpu7x-8: denoise 105.1s, generate 107.6s
tpu7x-8, DVFS pinned (PIN_DVFS_P_STATE=true), fused RoPE kernel off:
denoise 93.6s, generate 96.1s (pinned with it on was not measured)
Per-feature deltas were not measured separately and are omitted. In
isolation the cross-attention kernel is 0.51 ms vs 3.94 ms (v6e) and
0.49 ms vs 1.81 ms (tpu7x) per call for the XLA path inside the model.
The launcher's reference-benchmark comment now carries these numbers.
Tests:
- cross_attention_pallas_test.py: golden parity (row/tile max, mxu/vpu sum,
ragged tiles, k_prescaled, batch), guards, VMEM budget resolution /
working-set estimate / fallback reasons, and FlaxWanAttention dispatch
(CPU interpret integration, fallbacks, I2V text-vs-image split). The three
tests that expect the kernel to be used skip on TPUs outside the validated
kinds (there FlaxWanAttention falls back by design); that fallback is
covered by test_falls_back_to_xla_on_unvalidated_tpu_kind.
- wan_seq_pad_test.py: padding plan and per-cause disable logging, pad-key
masking for ulysses custom (unaligned and aligned real_len, flat and
shard-major out A2A, every fixed-m mode) and ulysses_ring_custom (k-centring
on/off; and the tpu7x production path scaled down: U=2 x R=2, aligned
real_len, chunked out A2A, asserting the shard-major layout was taken;
needs >= 4 devices), fixed-m metadata ignoring pad query rows, and the
NotImplementedError guards.
- wan_patch_embed_test.py, ulysses_out_a2a_test.py: tokens vs conv parity and
config-driven mode; flat vs chunked vs shard_major parity.
- wan_runtime_options_test.py: the new switches' YAML defaults and enums.
- CI: 29 tests skip in GitHub Actions: kernel parity and the Pallas-active
dispatch / unaligned-KV fallback tests in cross_attention_pallas_test
(19 of 39), the three kernel bit-exact tests in ulysses_out_a2a_test, the
two WanModel patch-embed parity tests and the five masking / end-to-end
tests in wan_seq_pad_test. Guards, VMEM budget logic, the remaining XLA
fallback dispatch tests (including the unvalidated-TPU fallback),
padding-plan logic and runtime options run in CI.
- end_to_end/tpu/run_wan_stack_tests.sh: adds this PR's four test files.
Verified (final tree):
- TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 350 passed, 1 skipped (off-TPU fallback only), 259 subtests passed.
- TPU v6e-8 with GITHUB_ACTIONS=true (all 17 stack test files): 231 passed, 120 skipped, 202 subtests passed.
- Emulated unvalidated TPU (device_kind="TPU v4", rope_accum_is_measured=False): 0 failures across all 17 stack test files.
34789ff to
87c36b2
Compare
Summary
Four Wan 2.2 inference optimizations. Each is a config switch that defaults off in the yml; the launcher's v6e and v7 profiles turn them on.
wan_cross_attn_kernel="pallas"v6e/tpu7x), orestimate_vmem_bytes(~42 MiB at Wan shard shape, validated at KV=512) exceeds the 64 MiB budget.wan_patch_embed_mode="tokens"wan_ulysses_out_a2a="shard_major"(aliaschunked)(num_shards, chunk)axis, avoiding XLA relayout copies. With the custom splash kernel, writes shard-major output directly on both v6e (1D Ulysses) and tpu7x (2D Ulysses+Ring).wan_seq_pad="lane"NotImplementedError.Other changes:
attention_config_entries(config);pyconfigno longer sets Wan options for non-Wan models.ici_ag_pipelining/ag_backward_pipelining.Performance
Wan 2.2 T2V-A14B, 720p / 81 frames / 40 steps, launcher defaults (fused RoPE on), warm AOT, DVFS unpinned:
PIN_DVFS_P_STATE=true, fused RoPE off): 93.6 s denoise / 96.1 s generate.Tests
cross_attention_pallas_test.py(39),wan_seq_pad_test.py(17, including U=2×R=2 alignedreal_len+chunkedshard-major layout),wan_patch_embed_test.py(7),ulysses_out_a2a_test.py(5),wan_runtime_options_test.py(12).vmem_budget_is_validated), whiletest_falls_back_to_xla_on_unvalidated_tpu_kindcovers the fallback.end_to_end/tpu/run_wan_stack_tests.sh):824.7 s).GITHUB_ACTIONS=true(all 17 stack test files): 231 passed, 120 skipped in ~2m 16s.device_kind="TPU v4",rope_accum_is_measured=False): 0 failures across all 17 stack test files.Stack: #477 → #478 → #479 → #488 → #491