Conversation
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 22, 2026 20:10
dbd6ace to
7a6e0e6
Compare
Perseus14
added this pull request to stack #486
September 22, 2026 20:11
There was a problem hiding this comment.
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.
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
9 times, most recently
from
September 23, 2026 16:48
5b75348 to
82681ec
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 23, 2026 19:19
82681ec to
b33e35e
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 23, 2026 20:01
b33e35e to
7740dcf
Compare
eltsai
reviewed
Sep 23, 2026
eltsai
reviewed
Sep 24, 2026
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 24, 2026 07:00
7740dcf to
5c0f8bb
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 24, 2026 07:50
5c0f8bb to
aee8c4a
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
2 times, most recently
from
September 24, 2026 17:13
2f8b92d to
a507d0d
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 26, 2026 14:34
5c87fcf to
a62691e
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 26, 2026 17:07
a62691e to
f3e5892
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
3 times, most recently
from
September 26, 2026 19:43
6ac9049 to
96dd50f
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 26, 2026 23:03
96dd50f to
359249b
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
2 times, most recently
from
September 27, 2026 09:50
593e337 to
4ebc6fd
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 27, 2026 14:59
4ebc6fd to
5a1fa32
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
2 times, most recently
from
September 27, 2026 16:25
a75047a to
0b1315d
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 28, 2026 16:32
0b1315d to
77afa86
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 28, 2026 16:50
77afa86 to
2e70bc8
Compare
Perseus14
force-pushed
the
feat/wan-custom-kernels
branch
from
September 29, 2026 19:22
2e70bc8 to
973594c
Compare
…, 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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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
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.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):
Set
USE_FUSED_ROPE_KERNEL=falseto turn it off.wan_rope_norm_mode,wan_fuse_qk_prescale,wan_splash_transpose_out,wan_cfg_before_unpatchify,wan_cross_attn_prescale_kvandwan_rope_accumare yml keys passed viaattention_config_entries(config)acrossWanPipeline, VACE, Animate, andwan_block_benchmark.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.Tests
fused_rmsnorm_rope_pallas_test.py(60),wan_runtime_options_test.py(11), plus updates todot_fallback_layout_test.py(11) andaot_cache_test.py(45).rope_accumskip on unvalidated TPUs (vmem_limit_is_validated,rope_accum_is_measured).GITHUB_ACTIONS=true(88 of 282 cumulative); guards, VMEM resolver, runtime options, AOT cache, and dot-fallback layout tests run in CI.run_wan_stack_tests.sh, 575.2 s).Stack: #477 → #478 → #479 → #488 → #491