Repository navigation
Tags: AI-Hypercomputer/maxdiffusion
Tags
fix(wan): resolve review feedback (output path fallback, libtpu finge…
…rprint, ltx2 remap, multi-host ssim, interpret guard)
- Fix format_video_output_path to fallback to wan_output_{seed}_{index}.mp4 when output_dir is default template 'sdxl-model-finetuned' or run_name is None
- Fix libtpu platform_version in _build_wan_aot_metadata via dev.client.platform_version so libtpu upgrades invalidate stale executables
- Read jax.config attributes directly in AOT metadata to avoid AttributeError on JAX >= 0.11
- Add missing ulysses_custom_fixed_m_per_q_block and ulysses_ring_custom_fixed_m_per_q_block to cross_attention_remapped_to_flash in attention_ltx2.py
- Guard print_ssim in base_wan_trainer and wan_trainer_2_2 to only run on process_index == 0 in multi-host training
- Prevent is_cpu_interpret from running in Pallas interpret mode on TPU devices in attention_flax.py
- Guard _active_token_padding with is_self_attention
feat(wan): fused RMSNorm+RoPE+transpose Pallas producer kernel, Fixed…
…-M optimizations, and serving stabilization
- Add fused_rmsnorm_rope_pallas custom Pallas TPU kernel (src/maxdiffusion/kernels/fused_rmsnorm_rope_pallas.py) combining fp32 RMSNorm scaling, complex rotary position embeddings, in-register Q/K prescaling (q * LOG2E, k * scale), and [B, S, H*D] -> [B, H, S, D] head transpose in a single VMEM pass.
- Add TPU-only platform guard in _fused_rope_producer and prevent cuDNN Flash Attention (cudnn_flash_te) from double-scaling prescaled cached keys.
- Eliminate two 387 MB standalone HBM multiply+copy round-trips per self-attention layer (774 MB/layer, 2.48 TB over 40 steps) prior to Ulysses/Ring all-to-all, worth 1.7% (tpu7x) to 3.1% (v6e) of denoise time end to end. The kernel is 0 ULP against a compiled XLA reference in isolation; inlined in the full model it is arithmetically equivalent but not bit-identical (34.1 dB PSNR over a 40-step denoise on v6e).
- Enable in-kernel VMEM transpose in custom splash attention (custom_splash_attention.py) to streamline output relayout across Ulysses sequences.
- Optimize cross-attention projection handling with split_head_dim in WanModel to eliminate unnecessary head transpose barriers.
- Update run_wan_fast_inference.sh with tuned latency-hiding scheduler options, p_state=7 on v7x, environment variable overrides, and optimized collective overlapping.
- Add JAX >= 0.11.0 compatibility shim in __init__.py and generate_wan.py for Flax NNX HiPrimitive deprecation.
- Add comprehensive test suites: fused_rmsnorm_rope_pallas_test.py (including production JIT + shard_map + f32 RoPE + prescale parity, finite/NaN assertions, and true integer ULP metrics), dot_fallback_layout_test.py, and aot_cache_test.py.
Key the AOT cache on the resolved RoPE accumulation mode.
WAN_ROPE_ACCUM changes the lowered graph's rounding but was not part of the Wan
AOT metadata fingerprint, so flipping it could silently reload an executable
compiled under the other rounding mode. Both settings produced identical
metadata.
The underlying problem was a duplicated rule: attention_flax.py computed the
platform default inline, and generate_wan.py did not model it at all. Adding a
second copy to the fingerprint would have reintroduced the same class of bug, so
the rule now lives in one place, resolve_rope_accum(mesh), called by both the
production call site and _build_wan_aot_metadata.
The fingerprint records the RESOLVED mode rather than the raw env var. Recording
the raw value would both over-invalidate (an explicit "f32" on v6e describes the
same executable as the unset default) and fail to describe what was actually
compiled. An unrecognised value now raises instead of being passed through.
Also closes a gap in parity coverage. The prescaled parity test pinned
rope_accum="f32" and checked bit-equality against another Pallas result, so the
combination that actually ships on v7 -- prescale ON, rope_accum="dtype", and the
production 1024 tile -- had no 0-ULP assertion anywhere; a regression in it could
only have surfaced as a quality drop in a 40-step denoise. The added test
resolves the mode the same way production does, uses fused_rope_block_s /
fused_rope_head_block instead of the kernel defaults, and compares against an XLA
reference whose prescale is applied inside the same jit.
Verified on tpu7x-8 and v6e-8:
* 16/16 aot_cache_test, 4/4 production-shape parity tests on both.
* The shipped config is 0 ULP against fused-in-jit, barrier-staged and post-hoc
XLA references alike (0/774,144,000 elements differ), so the strict bound is
measured, not assumed.
* Negative control: forcing WAN_ROPE_ACCUM=f32 on tpu7x fails the new test with
26.4% of elements differing, confirming it discriminates.
* Through the real AOT cache at a fixed commit, the two modes key different
executables and reverting reloads the original rather than recompiling; the
two modes produce different output, which is what made the missing key a
wrong-answer bug rather than a hygiene issue.
* E2E unchanged by this cache-keying change: denoise 96.7s (v7x) and 129.6s
(v6e), with byte-identical output video at a fixed seed for a given resolved
rope_accum mode.
Compare the kernel against a compiled reference, on both sides.
The new parity tests compared an eagerly-dispatched XLA reference against the
compiled kernel. That is not a comparison of the kernel: eager dispatch
evaluates each op separately, while XLA under jit contracts the RoPE
multiply-add into an FMA that rounds once. Measured on v6e at seq=50, the eager
and jitted references disagree on 30.6% of elements, and the kernel matches
whichever one shares its rounding -- exactly, with nothing in between. The
pairing happened to be self-consistent on v6e and tpu7x and was not on the v4-8
CI runner, where it failed with 35.3% of elements differing.
References now go through _compiled_reference and the kernel through
_compiled_kernel, with the mode from resolve_rope_accum(). One test keeps the
eager pairing deliberately: under an identity rotation both modes denote the
same arithmetic, which is the point of that test.
Compiling only the reference was not enough. norm_mode="exact" leaves the fp32
feature-axis reduction in XLA, so the kernel has an XLA prologue of its own, and
eager and jitted reductions diverge at full size: 665/96,768,000 outputs on
tpu7x and 503/96,768,000 on v6e, max |diff| 3.125e-02. Small-shape tests cannot
see it.
Also fixes resolve_rope_accum off TPU. It assumed every non-tpu7x platform
contracts into fp32 FMAs, but XLA:CPU emits the multiply and the add separately;
measured, the eager and jitted references agree to 0 ULP on CPU and only "dtype"
matches either. The f32 default is now gated on the platform actually being a
TPU, which is what makes the interpret-mode tests meaningful off TPU.
Verified: 34/34 kernel tests and 16/16 aot_cache_test on CPU (interpret), tpu7x-8
(resolves to "dtype") and v6e-8 (resolves to "f32"). Negative control: forcing
WAN_ROPE_ACCUM=dtype on v6e fails 7 tests, including every converted one.
On a parity failure the assertion now prints the reference-dispatch x
accumulation-mode grid. A bare "31% of elements differ" cannot distinguish a
real algebra bug from the kernel being built against the wrong rounding
convention for the platform, and those want opposite fixes. If some cell is
0 ULP the kernel is correct and only the resolver is wrong; if none is, the
arithmetic itself disagrees with the producer. That question has to be
answerable from a CI log on a TPU generation that is not available locally.
Scope the bit-exactness claim to the platforms where it was measured.
The v4-8 CI runner reproduces the producer under no accumulation mode at all.
Its grid, at seq=50:
ref:eager vs kernel/dtype 35.30% ref:jit vs kernel/dtype 20.88%
ref:eager vs kernel/f32 39.76% ref:jit vs kernel/f32 30.96%
The kernel is not wrong there. The normalisation is still bit-exact, the drift
is confined to the RoPE combine, and it is at most one bf16 ULP (max |diff|
3.125e-02, exactly ulp(bf16) at that magnitude); v4 rounds `a*cos + b*sin` in a
way neither mode models. Note also that the kernel is dispatch-invariant on v4 --
the eager and jit columns are identical -- so this is not reachable by another
resolver tweak.
Bit-exactness is therefore a per-platform property, and `rope_accum_is_measured`
now says where it holds: tpu7x ("dtype"), v6 ("f32") and XLA:CPU ("dtype"),
which covers both platforms Wan serves on. Tests asserting 0 ULP skip elsewhere
rather than assert something untrue. Tests with a non-zero ULP budget still run
everywhere, so v4 keeps real coverage of the algebra. An explicit
WAN_ROPE_ACCUM counts as measured, so overriding it does not quietly soften the
negative control.
Verified after the change: 34/34 kernel and 16/16 aot_cache_test on tpu7x-8 and
v6e-8, with nothing skipped on either -- the guard does not remove coverage on
the platforms that matter -- and 50/50 on CPU.
Remove the last inline copy of the platform rule. The shard_map production
parity test re-derived "dtype on tpu7x, f32 everywhere else" locally instead of
calling the resolver, so it kept asserting 0 ULP on v4 after the shared rule had
learned better. It now calls resolve_rope_accum(mesh) and is gated the same way.
feat(wan): fast serving with persistent AOT caching and tuned inferen… …ce recipe Delivers production-grade fast serving for Wan 2.2 T2V-A14B on Cloud TPU: - Persistent AOT Compilation Caching with source hash and static graphdef metadata - AOT signatures key non-static Python scalars by type, so guidance values share one executable - Shared Wan 2.2 warmup for T2V and I2V: compile both experts without executing them, bounded in-flight steps - Tuned Serving Recipe: Production launcher with platform-specific v6e and v7 profiles - Clean Key-Centering & 4-Operand Kernel (mk, q, k, v) without VMEM register pressure - Merged Pre-A2A Norm Reduction (pmax) removing ring all_gather relayout barrier - Hoisted scalar-prefetch metadata in _lse_scan eliminating collective-permute stalls - Lossless 6D unpatchify transpose optimization (collapsing p_t=1 to avoid 8D stride copies)
feat(attention): 2D Ulysses+Ring custom attention with fixed-m (pre-a… …2a norm reduction, unpadded ring K/V) Implements 2D Ulysses + Ring distributed attention with exact fixed-m accumulation: - Global Virtual K-Centering: Computes key mean on real tokens and reduces across ring axis with jax.lax.pmean - Cross-Ring Safety Fallback: Collective jax.lax.pmin reduction on v_ok across all ring ranks - Hoists Q * log2(e) before Ulysses All-to-All to fuse into Q producer - Slices ragged KV tail when R=1, eliminating dead sequence padding - Fused RMSNorm + RoPE producer with exact bit-identical nnx.RMSNorm associativity - Strict topology and configuration validation guards
feat(attention): fixed-m splash attention kernel with dynamic bounds … …and safety fallbacks Implements exact fixed-m splash attention in Pallas on TPU: - Dynamic C(N) headroom constants guaranteeing FP32 accumulator safety - fixed_m_dtype_is_safe checks rejecting FP16/FP8 exponent overflow - Value bound validation (|V| <= 256) with safe online softmax fallback - Unit tests covering all boundary conditions, dtypes, and scale factors Explicit metadata contracts on the fixed-m ring path ---------------------------------------------------- Three implicit contracts are made explicit. Each failed silently rather than loudly, and one of them produced non-finite output on TPU. 1. Norm representation is declared, not inferred. The gate previously guessed whether `fixed_m_norms` were squared with `(qn.max() * mk.max()) < 1000.0`. Magnitude cannot answer that question: legacy unsquared norms of (1000, 2) have a true bound of 2000, but read as already-squared they yield sqrt(2000) ~= 44.7 -- a ~45x under-estimate that admits fixed-m where it must fall back, and overflows. Replaced by `fixed_m_norms_squared` (default True, matching every in-tree caller); the test harness is migrated to squared norms. 2. The V-safety predicate is required, not assumed. Unlike the Cauchy-Schwarz norm bounds, the V-magnitude and dtype verdict is not re-derivable from a single hop's Q/K, so the kernel cannot reconstruct it. Omission previously meant "safe", which let fixed-m run on inputs it cannot represent: float16 with Q=K=0 and V=1 returns inf instead of 1.0. The ring path now raises unless `v_ok` is passed, mirroring the existing `fixed_m_recenter` rule, and `_ulysses_ring_custom_attention` computes it -- dtype safety plus |V| <= DEFAULT_MAX_V_BOUND, reduced with pmin over BOTH internal axes, since after the all-to-all neither axis alone observes the whole V and the fixed-m branch must be taken uniformly by every ppermute participant. 3. Norm shape is validated against per_q_block. This is the defect behind the `test_sink_head_falls_back_everywhere` TPU failure. Both gates compute `qn * mk[:, None]`, so a (num_heads,) array supplied while per_q_block=True does not raise -- it broadcasts to (num_heads, num_heads), pairing head j's query norm with head h's key norm. A sink head then inherits a small bound from an unrelated head, is wrongly marked eligible, and the kernel evaluates exp2(large_logit - small_m) -> inf. The kernel now rejects the mismatch, and the test declares per_q_block=False to match the per-head norms it supplies, as the production ring caller already did. Regression coverage: six backend-independent contract tests (both omissions raise, v_ok=False is accepted, mis-shaped norms are rejected, correctly shaped per-Q-block norms are accepted, fp16 is rejected while bf16/fp32 pass) plus two TPU tests (the two declared norm representations must agree, and an explicit unsafe verdict must force a finite fallback). Verified on v6e-8: ring_fixed_m_test 13 passed; attention_test, custom_splash_fixed_m_test, attention_block_sizes_test and ring_fixed_m_test together 58 passed.
save memory by deleting params. lint. (#277) Deletes params pytree after creating TrainState to reduce HBM usage. Lints the code. The only change is in this line: https://github.com/AI-Hypercomputer/maxdiffusion/compare/deallocate_params_tree?expand=1#diff-69cda939e98b489aca1cc8aa543ecc537d9f4bc6c58fa767cbe01cd3636aacf3R311