Skip to content

feat(wan): fast inference optimizations (fused cross-attention, token patch embed, shard-major A2A, lane padding) - #491

Open
Perseus14 wants to merge 1 commit into
feat/wan-custom-kernelsfrom
feat/wan-fast-inference-optimizations
Open

Perseus14 wants to merge 1 commit into
feat/wan-custom-kernelsfrom
feat/wan-fast-inference-optimizations

Conversation

@Perseus14

@Perseus14 Perseus14 commented Sep 25, 2026 •

Copy link
Copy Markdown
Collaborator

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.

Switch What it does
wan_cross_attn_kernel="pallas" Text cross-attention kernel keeping one head block's 512-token KV resident in VMEM with an exact single-pass softmax (no online rescaling). Falls back to XLA (logging once under separate warn keys for input vs VMEM) when shapes/sharding don't fit, the TPU isn't validated (v6e/tpu7x), or estimate_vmem_bytes (~42 MiB at Wan shard shape, validated at KV=512) exceeds the 64 MiB budget.
wan_patch_embed_mode="tokens" Patch embedding as patchify + matmul on the token-sharded layout, avoiding the activation all-gather.
wan_ulysses_out_a2a="shard_major" (alias chunked) One layout-preserving all-to-all over a (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" Pads video tokens to a per-shard multiple of 128 (75,600 → 75,776). Pad keys are masked in custom Ulysses/ring kernels and pad query rows are excluded from the fixed-m Q-norm bound. Paths that cannot mask padding raise NotImplementedError.

Other changes:

  • All switches plumb through attention_config_entries(config); pyconfig no longer sets Wan options for non-Wan models.
  • The launcher's generic profile no longer sets 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:

TPU Denoise Generate
v6e-8 125.9 s 128.05 s
tpu7x-8 105.1 s 107.6 s
  • tpu7x-8 with DVFS pinned (PIN_DVFS_P_STATE=true, fused RoPE off): 93.6 s denoise / 96.1 s generate.
  • Output videos were visually inspected on both platforms: clean, sharp, no artifacts.

Tests

cross_attention_pallas_test.py (39), wan_seq_pad_test.py (17, including U=2×R=2 aligned real_len + chunked shard-major layout), wan_patch_embed_test.py (7), ulysses_out_a2a_test.py (5), wan_runtime_options_test.py (12).

  • Pallas-active dispatch tests skip on unvalidated TPUs (vmem_budget_is_validated), while test_falls_back_to_xla_on_unvalidated_tpu_kind covers the fallback.
  • CI vs Local VM (end_to_end/tpu/run_wan_stack_tests.sh):
    • Full local runner on TPU v6e-8: 350 passed, 1 skipped (off-TPU fallback only) in 13m 44s (824.7 s).
    • With GITHUB_ACTIONS=true (all 17 stack test files): 231 passed, 120 skipped in ~2m 16s.
    • Emulated unvalidated TPU (device_kind="TPU v4", rope_accum_is_measured=False): 0 failures across all 17 stack test files.

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

@Perseus14
Perseus14 requested a review from entrpn as a code owner September 25, 2026 16:26
@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 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.

Comment thread src/maxdiffusion/kernels/cross_attention_pallas.py
@Perseus14
Perseus14 added this pull request to stack #486 September 25, 2026 16:30
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from 2ba8966 to e86f174 Compare September 25, 2026 16:34
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from e86f174 to 1d4c9f6 Compare September 25, 2026 17:20
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch 2 times, most recently from 0d91078 to 4703103 Compare September 25, 2026 18:28
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch 3 times, most recently from 295c3b5 to 7a60c3d Compare September 25, 2026 22:13
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch 2 times, most recently from 08d405c to 7db8313 Compare September 26, 2026 10:31
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from 7db8313 to 6fdc94e Compare September 26, 2026 13:13
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch 2 times, most recently from ebaac14 to 58f889e Compare September 26, 2026 14:34
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from 58f889e to 68a580f Compare September 26, 2026 17:07
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from 68a580f to 395f932 Compare September 26, 2026 18:21
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from 53fa702 to 6a446d5 Compare September 28, 2026 16:32
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from 6a446d5 to d11ccb6 Compare September 28, 2026 16:51
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch 2 times, most recently from 888462a to cdc584f Compare September 29, 2026 19:29
@eltsai

eltsai commented Sep 30, 2026

Copy link
Copy Markdown
Collaborator

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?

@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from cdc584f to 2eba645 Compare October 1, 2026 07:31
@Perseus14 Perseus14 changed the title feat(wan): Wan 2.2 fast inference optimizations (Pallas cross-attention, token patch embed, chunked A2A, lane padding) feat(wan): fast inference optimizations (fused cross-attention, token patch embed, shard-major A2A, lane padding) Oct 1, 2026
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch 2 times, most recently from 08ccfb2 to 68c3732 Compare October 1, 2026 10:35
@Perseus14

Copy link
Copy Markdown
Collaborator Author

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 = main (10af306), B = this stack with the #491 switches off (= #477→#488), C = this stack with the launcher defaults. Denoise seconds for 40 steps (load/compile/VAE excluded):

frames tpu7x A tpu7x B tpu7x C C vs A v6e A v6e B v6e C C vs A
41 56.5 43.0 41.5 1.36× 58.8 53.0 49.6 1.19×
81 123.0 109.1 105.0 1.17× 144.6 132.8 125.9 1.15×
121 234.2 216.2 212.7 1.10× 290.8* 308.1 313.0 / 254.6* 1.14×*
161 385.0 362.1 356.5 1.08× 454.7 430.4 419.2 1.08×

No regression vs main at any length; the gain narrows as attention dominates the step. The #491 switches themselves (B→C) are +1.6–3.9 % on tpu7x and +2.7–6.9 % on v6e, except v6e @ 121 f (−1.6 %).

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 (ATTENTION=ulysses_ring_custom_fixed_m ULYSSES_SHARDS=2 BQ=4736 BKV=2048) is 19 % faster there (254.6 s; main with its own ring recipe: 290.8 s, with its Ulysses kernel 657.6 s). At 41 / 81 / 161 f Ulysses remains the better v6e choice by 1–37 %.

All runs were bit-reproducible across cold/warm AOT (identical md5s), and the outputs are visually clean on both platforms.

@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch 2 times, most recently from 0414334 to 34789ff Compare October 2, 2026 17:15
eltsai
eltsai previously approved these changes Oct 2, 2026
… 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.
@Perseus14
Perseus14 force-pushed the feat/wan-fast-inference-optimizations branch from 34789ff to 87c36b2 Compare October 5, 2026 03:57
@Perseus14
Perseus14 requested a review from eltsai October 5, 2026 16:32
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants