Skip to content

feat(wan): fused RMSNorm+RoPE Pallas producer; Wan switches in config, not env - #488

Open
Perseus14 wants to merge 1 commit into
feat/wan-fast-servingfrom
feat/wan-custom-kernels
Open

Perseus14 wants to merge 1 commit into
feat/wan-fast-servingfrom
feat/wan-custom-kernels

Conversation

@Perseus14

@Perseus14 Perseus14 commented Sep 22, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Adds a fused RMSNorm + RoPE + head-transpose Pallas producer for Wan self-attention Q/K, and moves Wan's graph-changing switches out of environment variables into config.

What's in it

  • Pallas producer (kernels/fused_rmsnorm_rope_pallas.py): computes FP32 RMSNorm + RoPE + the [B,S,H*D] → [B,H,S,D] transpose, with optional log2(e) / 1/√d prescale folding.
    • norm_mode="exact" is 0 ULP against the separately jitted XLA producer.
    • In isolation on v6e it is 1.86×–2.06× faster than XLA (0.964 vs 1.789 ms at seq 9450; 1.953 vs 4.018 ms at seq 18900).
  • Defaults:
    • Off by default in the yml (use_fused_rope_kernel: False) and in the launcher's generic profile: once inlined into the 40-layer graph, XLA rounds the unfused producer differently, so the 40-step v6e trajectory diverges (~15 dB PSNR against the XLA-path video, visually clean; bit-identical on tpu7x).

    • Turned on in the launcher's v6e and v7 profiles. Measured at the stack tip with feat(wan): fast inference optimizations (fused cross-attention, token patch embed, shard-major A2A, lane padding) #491 on top (720p / 81 frames / 40 steps, warm AOT, DVFS unpinned):

      TPU Denoise, on Denoise, off
      v6e-8 125.9 s 131.2 s
      tpu7x-8 105.1 s 106.7 s
    • Set USE_FUSED_ROPE_KERNEL=false to turn it off.

  • Config, not env:
    • wan_rope_norm_mode, wan_fuse_qk_prescale, wan_splash_transpose_out, wan_cfg_before_unpatchify, wan_cross_attn_prescale_kv and wan_rope_accum are yml keys passed via attention_config_entries(config) across WanPipeline, VACE, Animate, and wan_block_benchmark.
    • No process-global store and no trace-time reads. Config-less module builds use built-in defaults.
  • wan_cfg_before_unpatchify (on by default): applies CFG before unpatchify. A compiled parity test shows bit-identical output in f32 everywhere and in bf16 on CPU/v6e, and within 1 bf16 ULP on tpu7x.
  • VMEM: 64 MiB is requested only when every mesh TPU is v6e or tpu7x; other TPUs get Mosaic's default.

Tests

fused_rmsnorm_rope_pallas_test.py (60), wan_runtime_options_test.py (11), plus updates to dot_fallback_layout_test.py (11) and aot_cache_test.py (45).

  • TPU tests relying on 64 MiB VMEM or measured rope_accum skip on unvalidated TPUs (vmem_limit_is_validated, rope_accum_is_measured).
  • CI vs Local VM: 48 additional kernel-grid/CFG tests skip under GITHUB_ACTIONS=true (88 of 282 cumulative); guards, VMEM resolver, runtime options, AOT cache, and dot-fallback layout tests run in CI.
  • Results (TPU v6e-8): 282 passed cumulative (run_wan_stack_tests.sh, 575.2 s).

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

@Perseus14
Perseus14 requested a review from entrpn as a code owner September 22, 2026 20:09
@github-actions

Copy link
Copy Markdown

@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from dbd6ace to 7a6e0e6 Compare September 22, 2026 20:10
@Perseus14
Perseus14 added this pull request to stack #486 September 22, 2026 20:11

@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 a fused RMSNorm + RoPE + head-transposition Pallas kernel for TPU, integrates it into the Flax attention pipeline, and implements several optimizations such as in-kernel VMEM transpose and Classifier-Free Guidance (CFG) before unpatchify. Feedback on these changes highlights critical runtime issues, including an AttributeError from using jax.shard_map instead of jax.experimental.shard_map.shard_map and a potential NameError due to an undefined _warn_once helper. Additionally, the reviewer recommends adding defensive checks to validate the sequence length of freqs_cis and removing redundant JAX compatibility shims that are already handled in the package initialization.

Comment thread src/maxdiffusion/models/attention_flax.py Outdated
Comment thread src/maxdiffusion/tests/fused_rmsnorm_rope_pallas_test.py Outdated
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/kernels/fused_rmsnorm_rope_pallas.py
Comment thread src/maxdiffusion/generate_wan.py Outdated
Comment thread src/maxdiffusion/tests/fused_rmsnorm_rope_pallas_test.py Outdated
Comment thread src/maxdiffusion/tests/aot_cache_test.py
Comment thread src/maxdiffusion/kernels/fused_rmsnorm_rope_pallas.py Outdated
@Perseus14 Perseus14 changed the title feat(wan): fused RMSNorm+RoPE+transpose Pallas producer kernel, Fixed-M optimizations, and serving stabilization feat(wan): fused RMSNorm+RoPE+transpose Pallas producer kernel, CFG pre-unpatchify, and runtime options Sep 26, 2026
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from 5c87fcf to a62691e Compare September 26, 2026 14:34
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from a62691e to f3e5892 Compare September 26, 2026 17:07
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch 3 times, most recently from 6ac9049 to 96dd50f Compare September 26, 2026 19:43
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from 96dd50f to 359249b Compare September 26, 2026 23:03
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch 2 times, most recently from 593e337 to 4ebc6fd Compare September 27, 2026 09:50
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from 4ebc6fd to 5a1fa32 Compare September 27, 2026 14:59
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch 2 times, most recently from a75047a to 0b1315d Compare September 27, 2026 16:25
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from 0b1315d to 77afa86 Compare September 28, 2026 16:32
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from 77afa86 to 2e70bc8 Compare September 28, 2026 16:50
@Perseus14
Perseus14 force-pushed the feat/wan-custom-kernels branch from 2e70bc8 to 973594c Compare September 29, 2026 19:22
…, not env

Fused producer (kernels/fused_rmsnorm_rope_pallas.py)
- A Pallas kernel computing FP32 RMSNorm + RoPE + the [B,S,H*D] -> [B,H,S,D]
  head transpose for self-attention Q/K. It optionally folds log2(e) into Q
  and 1/sqrt(d) into K.
  - norm_mode="exact" leaves the feature reduction to XLA. It is 0 ULP
    against the separately jitted XLA producer (fused_rmsnorm_rope).
  - norm_mode="fused" does the reduction in-kernel. It stays within 2x
    pair-relative bf16 eps.
  - rope_accum ("dtype" | "f32") matches how the local compiler rounds the
    RoPE multiply-add. "auto" uses the per-platform measured mode: tpu7x
    "dtype", v6e "f32".
- Off by default in the YAML (`use_fused_rope_kernel: False`) and in the
  launcher's generic profile. Once inlined into the 40-layer graph, XLA rounds
  the unfused producer differently, so the output is equivalent but not
  identical (v6e 720p/81f, same seed: 53.8 dB PSNR after 1 step, 34.1 dB after
  40). The kernel is refused on TPUs whose rounding was not measured, and when
  the sequence split is uneven.
- run_wan_fast_inference.sh turns it ON in the v6e and v7 profiles despite
  the YAML default, because it is measurably faster there. Measured at the
  stack tip, i.e. with #491 on top (Wan 2.2 T2V-A14B, 720p/81f/40 steps, warm
  AOT, DVFS unpinned), kernel on vs off: denoise 125.9s vs 131.2s on v6e-8,
  105.1s vs 106.7s on tpu7x-8. Not measured with this PR alone. Full-run output vs
  the XLA producer: on v6e the bf16 trajectory diverges (~15 dB PSNR against
  the XLA-path video, visually clean); on tpu7x it was bit-identical.
  USE_FUSED_ROPE_KERNEL=false restores the XLA producer.
- Isolated speed, one v6e chip, jax 0.11.2, 40 heads, d=128, bf16, exact
  mode, block_s=1024, median of 50 runs: per-shard seq 9450 takes 0.964 ms
  vs 1.789 ms for XLA (1.86x); seq 18900 takes 1.953 ms vs 4.018 ms (2.06x).
  This replaces the unbenchmarked "1.9x" and "2.5-3.0x" figures.
- VMEM budget: 64 MiB is requested only when every TPU in the mesh is a v6e or
  tpu7x (the validated generations). Any other TPU gets Mosaic's default
  scoped limit. The mesh is now passed through to the resolver.
- Two fused RMSNorm+RoPE producers now exist: the XLA one from the ring PR and
  this Pallas one. The XLA one is the reference and the fallback.

Wan graph-changing switches
- wan_rope_norm_mode, wan_fuse_qk_prescale, wan_splash_transpose_out,
  wan_cfg_before_unpatchify, wan_cross_attn_prescale_kv and wan_rope_accum are
  YAML keys, declared in every base_wan*.yml with the code default. The
  pipelines resolve them once, at model-build time, with
  `resolve_from_config`: the config value if the key is set, else the legacy
  WAN_* env var, else the default. Since every Wan YAML sets every key, the
  env vars are effectively ignored on normal runs (a differing one is logged
  once).
  - WanPipeline, VACE, Animate and wan_block_benchmark all pass them in
    `attention_config` (`attention_config_entries(config)`).
  - FlaxWanAttention and WanModel built without an entry use the built-in
    default and ignore the environment.
  - So every value lives on the GraphDef, and therefore in the AOT key.
- There is no process-global store. Nothing reads these switches at trace
  time, and the generic Ulysses / ring wrappers take
  `transpose_out: bool = False`. A Wan setting can no longer leak into
  another model in the same process. pyconfig only coerces the values
  (e.g. "false" from the CLI).
- wan_cfg_before_unpatchify (default True) applies CFG on packed tokens
  before unpatchify. Wan 2.1 no-cache CFG now goes through the same
  transformer_forward_pass path instead of transformer_forward_pass_full_cfg.
  A jitted test through a real WanModel (p_t = 1 and 2) shows both are
  bit-identical to the CFG-after-unpatchify path and to full_cfg in f32
  everywhere, and in bf16 on CPU and v6e. In bf16 on tpu7x, XLA fuses the CFG
  combine differently, so ~30% of elements differ by one bf16 ULP of the
  output scale; the test allows that there.
- splash `transpose_out`: an in-kernel output layout [H,S,D] instead of
  [H,D,S], for the custom splash kernel and the non-ring Ulysses wrapper. Off
  by default.

Other changes
- split_head_dim is now plumbed from the config into WanModel and its
  attention layers. The six base_wan*.yml files flip `split_head_dim` from
  True to False. Before this change no Wan code read that key, so the
  effective behaviour (False) is unchanged.
- _apply_attention_dot / cudnn_flash_te: correct 4-D [B,H,S,D] handling and
  correct handling of prescaled Q/K. Prescaled Q/K is refused for
  unsupported kernels and for head-local SVG attention.
- `_fused_rope_producer` always returns ((q, k), qk_prescaled), falling back
  to XLA when the kernel cannot run.

Tests
- fused_rmsnorm_rope_pallas_test.py:
  - parity grid, guards, backward, production shape (TPU);
  - VMEM resolver;
  - FlaxWanAttention sequence-sharded over `context` via logical axis rules,
    against the XLA producer: bit-identical on CPU, within 2e-2 on TPU where
    the inlined XLA producer rounds differently (dot_product kernel, so the
    prescale fold is not exercised there);
  - transpose_out for the plain, fixed-m and k-centred kernels, MHPT (TPU),
    the Ulysses wrapper and the Ulysses+Ring wrapper at R == 1.
  - TPU tests that rely on the 64 MiB VMEM budget or a measured rope_accum
    (production-shape parity, in-register prescale, fused-mode drift, the
    mesh test's TPU branch) skip on TPUs outside the validated kinds
    (`vmem_limit_is_validated`, `rope_accum_is_measured`).
- dot_fallback_layout_test.py: 4-D layout, prescale, cross-attention K
  prescale via config, and the compiled CFG parity test.
- wan_runtime_options_test.py: YAML/default agreement, precedence, coercion,
  no global store, and module-level resolution (FlaxWanAttention, WanModel).
- aot_cache_test.py: Wan switches and the resolved rope_accum key the
  executable.
- CI: 48 tests skip in GitHub Actions: 47 of the 60 in
  fused_rmsnorm_rope_pallas_test (only the guard and VMEM-resolver classes
  run) and the compiled CFG parity test. wan_runtime_options, aot_cache and
  the dot-fallback layout tests run in CI.
- end_to_end/tpu/run_wan_stack_tests.sh: adds this PR's two new test files.
Verified (final tree):
- TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 282 passed, 168 subtests passed in 575.2s.
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.

2 participants