Repository navigation
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces optimized fused producers (fused_ln_adaln and fused_rmsnorm_rope) and implements exact Fixed-m support with Global Virtual K-Centering for Ulysses and Ring attention. It also adds support for Grouped Query Attention (GQA) across these attention kernels and includes comprehensive unit tests. The review feedback suggests moving an inline import of fused_rmsnorm_rope in attention_flax.py to the top of the file to avoid performance overhead in a hot path, and simplifying a double-negation conditional expression to improve code readability.
5745a10 to
5c674d1
Compare
5c674d1 to
52593cc
Compare
52593cc to
643ca72
Compare
643ca72 to
c58cdad
Compare
4ace67b to
0c191db
Compare
ea0148a to
837cebe
Compare
b1e2b13 to
e6db513
Compare
b749c64 to
3268a70
Compare
3268a70 to
66ada3d
Compare
66ada3d to
080bcca
Compare
8589576 to
e8ed02e
Compare
e8ed02e to
e8e4dbc
Compare
0548265 to
dbb51e1
Compare
|
for the e2e verification, did you check the output video quality as well? |
321569b to
e09db23
Compare
e09db23 to
406593e
Compare
406593e to
10372f8
Compare
7739cae to
472bcf3
Compare
472bcf3 to
72e024a
Compare
|
Yes, I verified manually as well as both pixel-level fidelity (PSNR / SSIM against exact FP32 online softmax) and full VBench quality metrics on Wan 2.2 T2V-A14B (1280x720, 81 frames, 40 steps). I have updated the results in #491 |
…2a norm reduction, unpadded ring K/V)
Adds fixed-m to the 2D Ulysses + Ring custom-kernel path
(`_ulysses_ring_custom_attention`, attention=ulysses_ring_custom_fixed_m and
the new ulysses_ring_custom_fixed_m_per_q_block), plus the supporting
producer, guard and layout fixes it needed.
Ring fixed-m (R > 1)
- `_ring_fixed_m_norms_pre_a2a` computes every fixed-m input on the PRE-a2a
shards (each device holds all heads x 1/(U*R) of the sequence): Q row norms,
K max norms and the V max are reduced in ONE pmax over (ulysses, ring), then
sliced to the heads this Ulysses rank owns after the a2a. Keeping these off
the post-a2a layout lets them overlap the all-to-all instead of sitting
between it and the `lax.cond`. With per_q_block=True a small ulysses-axis
all_to_all of the Q row norms is also issued.
- V safety: `v_ok = (max|V|^2 from that pmax <= 256^2) & dtype_safe`. It is
derived from the pmax of V norms, so it is bit-identical on every
(ulysses, ring) rank with no further collective. (The ring kernel's own pmin,
which reduces global-bound eligibility, only runs when the caller omits
`all_fixed_global`; the production caller always passes it.)
- `all_fixed_global = all(qn * kn <= bound(N_total)) & v_ok`, also reduced over
batch, is passed to the ring kernel and is the top-level
accumulate-vs-LSE `lax.cond` predicate (`pregathered_mk=True`). It is
mesh-uniform by construction. Blast radius: one ineligible head / Q-block on
any device sends every rank to `_lse_scan` for that layer; per_q_block=True
only refines dispatch inside the LSE path. How often that happens in
production is not measured.
- Only Q and K go through the optimization_barrier that stops XLA duplicating
their producer chains into the norm reductions; V is left out so the Q/K
a2a does not wait on V's producer.
- Ring K/V are passed unpadded when the ring-shard KV length is 8-aligned
(`kv_pad_size=1`; the kernel slices the tail, see custom_splash_unpadded_test).
- Q * log2(e) is applied before the Ulysses a2a (exact; fuses into Q's
producer instead of being wrapped in relayout copies after the collective).
Virtual K-centering
- Centering is virtual on both R == 1 and R > 1: k_mean goes to the kernel and
q . k_bar is added to the fixed bound in VMEM; K is never re-written. For
R > 1, k_mean = pmean(mean(K)) over (ulysses, ring) on the pre-a2a shards,
and K norms are taken around that global mean.
- Controlled by `use_k_centering` ("auto" | bool, new key in
base_wan_27b.yml, defaulted in pyconfig). "auto" (`resolve_k_centering`) is
ON for non-ring ulysses_custom_fixed_m* and OFF on every ring path, R == 1
included. The Wan pipelines do not pass this key to the attention layer in
this PR, so Wan runs the layer default.
Ring is kept OFF by default because, measured at the stack tip (Wan 2.2
T2V-A14B, 720p/81f/40 steps, warm AOT, DVFS unpinned), centering is never faster and it changes output: v6e-8 R == 1
127.2 s vs 127.2 s (neutral), R = 2 129.1 s vs 128.6 s (+0.4%); tpu7x-8
U=2/R=2 108.04 s vs 107.66 s (+0.35%), U=4 neutral. Turning it on changes
the bf16 trajectory, with no visible quality difference. use_k_centering=True
still forces it on.
Other changes
- `fused_rmsnorm_rope` (kernels/fused_producers.py): XLA-level fused
RMSNorm + RoPE producer for self-attention Q/K, matching nnx.RMSNorm's
association (x * (rsqrt(var + eps) * scale)) bit-for-bit. It is
`jax.jit`-wrapped (heads/dim/eps static): inlined under the pipeline's
outer jit, but when a full-size block is called eagerly (as
wan_vace_transformer_test does at 75,600 tokens) XLA fuses the chain
instead of materialising ~1.5 GB of extra intermediates, which OOMed that
test on the v4-8 CI runner.
- GQA: kv_heads plumbed through the custom Ulysses / ring paths; flash,
tokamax, non-custom ulysses and ulysses_ring raise NotImplementedError on GQA.
- Guards: non-ring ulysses kernels reject a ulysses_shards they cannot honour;
ring kernels warn once when U == CP makes the ring degenerate (R = 1) and
suggest a valid U for the mesh and head count.
- Dot-product fallback: fixes the split_head_dim reshape that silently
mis-read [B, H, S, D] inputs as [B, S, H*D].
- tile_size_grid_search: rejects iters < 1; `_broadcast_winner` handles
bkv_compute=None. wan_block_benchmark / LTX2 / pyconfig know the new
kernel names.
Tests
- ring_fixed_m_test.py (24): adds a U=2/R=2 integration class through
`_ulysses_ring_custom_attention` vs a dense f32 reference (per_q_block x
centering x GQA grid; ragged ring KV with kv_pad_size=1; a sink head that
must force `_lse_scan` on all 4 devices; a single V outlier at B=2 that must
make v_ok and all_fixed_global False on all 4 devices; and B=2 `_accumulate_scan`
vmap + sub-128 head_dim `k_mean` padding at d=64). The branch
predicates are checked through a probe of `_ring_fixed_m_norms_pre_a2a`, not
the production cond itself. Needs >= 4 devices (explicit skipIf).
- custom_splash_fixed_m_test.py (32), attention_config_guards_test.py (20),
custom_splash_unpadded_test.py (12, TPU-only), dot_fallback_layout_test.py
(5), fused_producers_test.py (6), tile_size_grid_search_test.py (16).
- CI: 18 more tests skip in GitHub Actions (36 of 115 cumulative):
custom_splash_unpadded_test (12), the GQA ring numerics test (1), and the
5 U=2/R=2 integration tests. Contract, guard, dot-fallback layout,
fused producer and grid-search tests run in CI. run_wan_stack_tests.sh
gains this PR's five new test files.
Verified (final tree): TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 115 passed, 107 subtests passed.
Summary
Adds fixed-m to the 2D Ulysses + Ring custom-kernel path (
ulysses_ring_custom_fixed_m,ulysses_ring_custom_fixed_m_per_q_block), plus the producer, guard and layout fixes it needed.What's in it
pmaxover(ulysses, ring)on the pre-all-to-all shards, so they overlap the collective.v_okcomes from thatpmax, so it is identical on every rank.all_fixed_globalis mesh-uniform and drives the accumulate-vs-LSElax.cond._lse_scanfor that layer.kv_pad_size=1).Q·log2(e)is folded in before the all-to-all.k_meanis computed aspmeanover(ulysses, ring)and passed to the kernel; K is never rewritten.Off by default on ring. Measured at the stack tip (Wan 2.2 720p/81f/40 steps, warm AOT, DVFS unpinned), it is never faster and it changes the bf16 output:
use_k_centering=Trueforces it on.fused_producers.py), bit-exact withnnx.RMSNorm. It isjax.jit-wrapped (heads/dim/eps static) so eager callers also get fusion — this is what fixed thewan_vace_transformer_testOOM on the v4-8 CI runner (the test calls a full 75,600-token block eagerly; unfused, the producer materialised ~1.5 GB of extra intermediates). Inlined, i.e. a no-op, under the pipeline's outer jit.NotImplementedErroron paths that don't support GQA.ulysses_shardsguards, dot-fallbacksplit_head_dimreshape fix, and tile-search hardening.Tests
per_q_block × centering × GQA, ragged KV, sink head forcing_lse_scan, V outlier forcingv_ok=False, andB=2_accumulate_scanvmap+d=64k_meanhead-dim padding), plusattention_config_guards_test(20),custom_splash_unpadded_test(12),dot_fallback_layout_test(5),fused_producers_test(6),tile_size_grid_search_test(16).GITHUB_ACTIONS=true(36 of 115 cumulative);end_to_end/tpu/run_wan_stack_tests.shruns all of them.run_wan_stack_tests.sh, 284.2 s).Stack: #477 → #478 → #479 → #488 → #491