From e4a38599d2c260b7dadb4d266067b70c73a24fe8 Mon Sep 17 00:00:00 2001 From: Rishabh Manoj Date: Tue, 22 Sep 2026 20:10:17 +0000 Subject: [PATCH] feat(wan): fused RMSNorm+RoPE Pallas producer; Wan switches in config, 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. --- benchmarks/vbench/run_tpu_generation.sh | 11 + end_to_end/tpu/run_wan_fast_inference.sh | 42 +- end_to_end/tpu/run_wan_stack_tests.sh | 3 + src/maxdiffusion/__init__.py | 1 + src/maxdiffusion/configs/base_wan_14b.yml | 31 +- src/maxdiffusion/configs/base_wan_1_3b.yml | 29 +- src/maxdiffusion/configs/base_wan_27b.yml | 38 +- src/maxdiffusion/configs/base_wan_animate.yml | 29 +- src/maxdiffusion/configs/base_wan_i2v_14b.yml | 29 +- src/maxdiffusion/configs/base_wan_i2v_27b.yml | 29 +- src/maxdiffusion/generate_wan.py | 23 +- .../kernels/custom_splash_attention.py | 115 +- .../kernels/fused_rmsnorm_rope_pallas.py | 701 +++++++ src/maxdiffusion/models/attention_flax.py | 536 +++++- .../wan/transformers/transformer_wan.py | 40 +- .../pipelines/wan/wan_pipeline.py | 85 +- .../pipelines/wan/wan_pipeline_2_1.py | 19 +- .../pipelines/wan/wan_pipeline_animate.py | 3 +- .../pipelines/wan/wan_vace_pipeline_2_1.py | 2 + src/maxdiffusion/pyconfig.py | 9 + src/maxdiffusion/tests/aot_cache_test.py | 127 +- .../tests/dot_fallback_layout_test.py | 298 +++ .../tests/fused_rmsnorm_rope_pallas_test.py | 1679 +++++++++++++++++ .../tests/wan_runtime_options_test.py | 166 ++ src/maxdiffusion/utils/wan_block_benchmark.py | 9 +- src/maxdiffusion/wan_runtime_options.py | 145 ++ 26 files changed, 4061 insertions(+), 138 deletions(-) create mode 100644 src/maxdiffusion/kernels/fused_rmsnorm_rope_pallas.py create mode 100644 src/maxdiffusion/tests/fused_rmsnorm_rope_pallas_test.py create mode 100644 src/maxdiffusion/tests/wan_runtime_options_test.py create mode 100644 src/maxdiffusion/wan_runtime_options.py diff --git a/benchmarks/vbench/run_tpu_generation.sh b/benchmarks/vbench/run_tpu_generation.sh index 893853981..18d084353 100644 --- a/benchmarks/vbench/run_tpu_generation.sh +++ b/benchmarks/vbench/run_tpu_generation.sh @@ -57,6 +57,8 @@ WAN_OVERRIDES=( "use_batched_text_encoder|USE_BATCHED_TEXT_ENCODER|true" "flash_block_sizes|FLASH_BLOCK_SIZES|" "prompt_file|PROMPT_FILE|./benchmarks/vbench/prompts_110.txt" + # Matches the shipped config default; set USE_FUSED_ROPE_KERNEL=true to test it. + "use_fused_rope_kernel|USE_FUSED_ROPE_KERNEL|false" ) usage() { @@ -70,6 +72,11 @@ Common options: RUN_NAME Generation run name (default: wan-inference; videos are saved to /videos) PROMPT_FILE Prompt file path (default: ./benchmarks/vbench/prompts_110.txt) CONFIG_FILE WAN config file (default: src/maxdiffusion/configs/base_wan_27b.yml) + USE_FUSED_ROPE_KERNEL + Enable the fused RMSNorm+RoPE Pallas producer (default: false). + Only measured on TPU v6e and v7; refused on other platforms. + FUSED_ROPE_BLOCK_S / FUSED_ROPE_HEAD_BLOCK + Optional sequence and head tile overrides for the fused RoPE Pallas producer. EXTERNAL_DISK Mounted disk root for large local files (default: /mnt/disks/external_disk) HF_CACHE_ROOT Hugging Face cache root (default: \$EXTERNAL_DISK/hf_cache) HF_HOME Hugging Face home directory (default: \$HF_CACHE_ROOT) @@ -247,6 +254,8 @@ emit_remote_args() { emit_remote_arg "${var}" done emit_remote_arg_if_explicit VENV_DIR + emit_remote_arg_if_explicit FUSED_ROPE_BLOCK_S + emit_remote_arg_if_explicit FUSED_ROPE_HEAD_BLOCK for item in "${WAN_OVERRIDES[@]}"; do IFS='|' read -r key var value <<< "${item}" emit_remote_arg "${var}" @@ -375,6 +384,8 @@ run_generation() { IFS='|' read -r key var value <<< "${item}" args+=("${key}=${!var}") done + [[ -n "${FUSED_ROPE_BLOCK_S:-}" ]] && args+=("fused_rope_block_s=${FUSED_ROPE_BLOCK_S}") + [[ -n "${FUSED_ROPE_HEAD_BLOCK:-}" ]] && args+=("fused_rope_head_block=${FUSED_ROPE_HEAD_BLOCK}") args+=("seed=12345" "base_output_directory=gs://${GCS_BUCKET}") "${args[@]}" } diff --git a/end_to_end/tpu/run_wan_fast_inference.sh b/end_to_end/tpu/run_wan_fast_inference.sh index 7b8fa8b79..d82227f7a 100755 --- a/end_to_end/tpu/run_wan_fast_inference.sh +++ b/end_to_end/tpu/run_wan_fast_inference.sh @@ -49,8 +49,20 @@ # autodetection # ACCEL_TYPE / TPU_ACCELERATOR_TYPE # force the raw accelerator type (e.g. v6e-8, tpu7x-8) -# ATTENTION / ULYSSES_SHARDS / BQ / BKV / BKV_COMPUTE / BKV_COMPUTE_IN / BQ_DKV / VMEM_LIMIT_BYTES +# ATTENTION / USE_K_CENTERING / ULYSSES_SHARDS / BQ / BKV / BKV_COMPUTE / BKV_COMPUTE_IN / BQ_DKV / VMEM_LIMIT_BYTES # override the per-platform attention recipe +# USE_FUSED_ROPE_KERNEL / FUSED_ROPE_BLOCK_S / FUSED_ROPE_HEAD_BLOCK +# override the fused RMSNorm+RoPE Pallas producer settings. +# The v6e and v7 profiles turn the kernel ON, unlike the +# YAML default (off) and the generic profile (off), because +# it is measurably faster there: at the full 5-PR stack +# (40 steps, 720p/81f, warm AOT, DVFS unpinned) denoise is +# 125.9s vs 131.2s on v6e-8 and 105.1s vs 106.7s on tpu7x-8. +# Its output is equivalent but not bit-identical to the XLA +# producer: on v6e the bf16 trajectory diverges (PSNR ~15 dB +# against the XLA-path video, visually clean); on tpu7x the +# output was bit-identical in that measurement. +# USE_FUSED_ROPE_KERNEL=false restores the XLA producer. # DP / CP / PER_DEVICE_BATCH / SEED # override mesh parallelism, per-device batch, or RNG seed (default 12345) set -euo pipefail @@ -73,8 +85,6 @@ export HF_HUB_ENABLE_HF_TRANSFER=1 export JAX_DEFAULT_MATMUL_PRECISION=bfloat16 export TORCHINDUCTOR_FX_GRAPH_CACHE=1 export XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 -# Without these the JAX persistent cache silently skips most entries, so the -# "warm" start still recompiles a large part of the graph. export JAX_PERSISTENT_CACHE_MIN_ENTRY_SIZE_BYTES=-1 export JAX_PERSISTENT_CACHE_MIN_COMPILE_TIME_SECS=0 @@ -85,14 +95,7 @@ export TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR:-$CACHE_ROOT/torch_comp mkdir -p "$CACHE_ROOT/jax" "$CACHE_ROOT/aot_wan$MODEL" "$CACHE_ROOT/converted" \ "$OUTPUT_DIR" "$TMPDIR" "$TORCHINDUCTOR_CACHE_DIR" -# --------------------------------------------------------------------------- -# TPU platform detection -# --------------------------------------------------------------------------- -# Preference order: explicit override -> TPU_ACCELERATOR_TYPE env (set by some -# runtimes) -> GCE metadata "accelerator-type" (e.g. "v6e-8") -> the -# ACCELERATOR_TYPE line inside the "tpu-env" metadata blob. Detection is pure -# metadata/env: it must not initialise the TPU, or it would take the device -# before the real process starts. +# Detect TPU generation via env or GCE metadata without initializing the device. _tpu_metadata() { curl -s -f -m 2 -H 'Metadata-Flavor: Google' \ "http://metadata.google.internal/computeMetadata/v1/instance/attributes/$1" 2> /dev/null || true @@ -102,8 +105,6 @@ _detect_accel_type() { local t="${TPU_ACCELERATOR_TYPE:-}" [ -z "$t" ] && t="$(_tpu_metadata accelerator-type)" [ -z "$t" ] && t="$(_tpu_metadata tpu-env | sed -n "s/^ACCELERATOR_TYPE: *'\([^']*\)'.*/\1/p")" - # Guard against metadata returning an HTML error page. Real names seen in the - # wild: "v6e-8", "v5litepod-8", "tpu7x-8" (v7 reports as tpu7x, not v7x). case "$t" in v[0-9]* | tpu[0-9]*) printf '%s' "$t" ;; *) printf '' ;; @@ -112,8 +113,8 @@ _detect_accel_type() { ACCEL_TYPE=${ACCEL_TYPE:-$(_detect_accel_type)} if [ -n "$ACCEL_TYPE" ]; then - TPU_GEN="${ACCEL_TYPE%%-*}" # v6e-8 -> v6e - TPU_CHIPS="${ACCEL_TYPE##*-}" # v6e-8 -> 8 + TPU_GEN="${ACCEL_TYPE%%-*}" + TPU_CHIPS="${ACCEL_TYPE##*-}" else TPU_GEN="" TPU_CHIPS="" @@ -165,6 +166,7 @@ case "$TPU_PROFILE" in DEFAULT_BATCHED_TE=true DEFAULT_VAE_CHUNK=1 DEFAULT_VAE_SPATIAL=8 + DEFAULT_FUSED_ROPE=true ;; v7) PLATFORM_LIBTPU="$V7_LIBTPU" @@ -180,6 +182,7 @@ case "$TPU_PROFILE" in DEFAULT_BATCHED_TE=true DEFAULT_VAE_CHUNK=1 DEFAULT_VAE_SPATIAL=8 + DEFAULT_FUSED_ROPE=true ;; *) echo "== warning: unrecognised or non-v6e/v7 accelerator '${ACCEL_TYPE:-unknown}';" \ @@ -197,6 +200,7 @@ case "$TPU_PROFILE" in DEFAULT_BATCHED_TE=true DEFAULT_VAE_CHUNK=1 DEFAULT_VAE_SPATIAL=8 + DEFAULT_FUSED_ROPE=false ;; esac @@ -212,6 +216,7 @@ case "$LIBTPU_INIT_ARGS" in esac ATTENTION=${ATTENTION:-$DEFAULT_ATTENTION} +USE_K_CENTERING=${USE_K_CENTERING:-auto} # auto: on for non-ring (v6e), off for ring (v7) ULYSSES_SHARDS=${ULYSSES_SHARDS:-$DEFAULT_U} BQ=${BQ:-$DEFAULT_BQ} BKV=${BKV:-$DEFAULT_BKV} @@ -223,6 +228,10 @@ COMPILE_TE=${COMPILE_TE:-$DEFAULT_COMPILE_TE} USE_BATCHED_TE=${USE_BATCHED_TE:-$DEFAULT_BATCHED_TE} VAE_SPATIAL=${VAE_SPATIAL:-$DEFAULT_VAE_SPATIAL} VAE_DECODE_CHUNK=${VAE_DECODE_CHUNK:-$DEFAULT_VAE_CHUNK} +USE_FUSED_ROPE_KERNEL=${USE_FUSED_ROPE_KERNEL:-$DEFAULT_FUSED_ROPE} +FUSED_ROPE_ARGS=() +[ -n "${FUSED_ROPE_BLOCK_S:-}" ] && FUSED_ROPE_ARGS+=("fused_rope_block_s=$FUSED_ROPE_BLOCK_S") +[ -n "${FUSED_ROPE_HEAD_BLOCK:-}" ] && FUSED_ROPE_ARGS+=("fused_rope_head_block=$FUSED_ROPE_HEAD_BLOCK") # Mesh: context parallelism carries the Ulysses shards, data parallelism takes # whatever chips remain. Defaults to CP=4 / DP=2 on an 8-chip slice. On a slice @@ -272,6 +281,7 @@ python src/maxdiffusion/generate_wan.py "$CONFIG" \ aot_cache_dir="$CACHE_ROOT/aot_wan$MODEL" \ converted_weights_dir="$CACHE_ROOT/converted" \ attention="$ATTENTION" \ + use_k_centering="$USE_K_CENTERING" \ ulysses_shards="$ULYSSES_SHARDS" \ ici_data_parallelism="$DP" ici_fsdp_parallelism=1 \ ici_context_parallelism="$CP" ici_tensor_parallelism=1 \ @@ -282,6 +292,8 @@ python src/maxdiffusion/generate_wan.py "$CONFIG" \ vae_weights_dtype=bfloat16 vae_dtype=bfloat16 \ text_encoder_dtype=bfloat16 compile_text_encoder="$COMPILE_TE" use_batched_text_encoder="$USE_BATCHED_TE" \ use_kv_cache=true use_base2_exp=true use_experimental_scheduler=true \ + use_fused_rope_kernel="$USE_FUSED_ROPE_KERNEL" \ + ${FUSED_ROPE_ARGS[@]+"${FUSED_ROPE_ARGS[@]}"} \ fps=16 ${GUIDANCE_ARGS[@]+"${GUIDANCE_ARGS[@]}"} \ seed="${SEED:-12345}" \ flash_block_sizes="$FLASH_BLOCK_SIZES" \ diff --git a/end_to_end/tpu/run_wan_stack_tests.sh b/end_to_end/tpu/run_wan_stack_tests.sh index 27459fa23..05ef84332 100755 --- a/end_to_end/tpu/run_wan_stack_tests.sh +++ b/end_to_end/tpu/run_wan_stack_tests.sh @@ -39,5 +39,8 @@ TESTS=( "$T/converted_weights_cache_test.py" "$T/wan/wan_transformer_test.py" "$T/wan/wan_warmup_coverage_test.py" + # Fused RMSNorm+RoPE producer, Wan switches in config (feat/wan-custom-kernels) + "$T/fused_rmsnorm_rope_pallas_test.py" + "$T/wan_runtime_options_test.py" ) PYTHONPATH="src${PYTHONPATH:+:$PYTHONPATH}" exec "${PYTHON:-python3}" -m pytest -q -rs "${TESTS[@]}" "$@" diff --git a/src/maxdiffusion/__init__.py b/src/maxdiffusion/__init__.py index d008f60c0..bb7d5001a 100644 --- a/src/maxdiffusion/__init__.py +++ b/src/maxdiffusion/__init__.py @@ -55,6 +55,7 @@ "schedulers": [], "tpu_utils": [], "train_utils": [], + "wan_runtime_options": [], "utils": [ "OptionalDependencyNotAvailable", "is_flax_available", diff --git a/src/maxdiffusion/configs/base_wan_14b.yml b/src/maxdiffusion/configs/base_wan_14b.yml index 6bcf0b4e1..3465cb184 100644 --- a/src/maxdiffusion/configs/base_wan_14b.yml +++ b/src/maxdiffusion/configs/base_wan_14b.yml @@ -82,7 +82,7 @@ jit_initializers: True # Set true to load weights from pytorch from_pt: True -split_head_dim: True +split_head_dim: False # Sparse VideoGen (SVG) Attention configuration use_svg_attention: False svg_spatial_density: 0.25 @@ -105,7 +105,34 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud use_base2_exp: True use_experimental_scheduler: True -use_k_centering: "auto" # auto: on for non-ring ulysses_custom_fixed_m* (virtual, no copy), off for ring paths (virtual via k_mean, plus pmean when R>1) +# use_k_centering: "auto" turns K-centering on for the non-ring +# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no +# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a +# pmean across ring shards when R>1). Set True/False to force either path. +use_k_centering: "auto" +# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK. +# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml. +use_fused_rope_kernel: False +# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 = +# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model. +fused_rope_block_s: 1024 +fused_rope_head_block: -1 +# Wan inference switches. These change the compiled graph, so they are part of +# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since +# this file sets every key, the env vars are ignored, with a log line if they differ.) +# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction, +# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps). +wan_rope_norm_mode: "exact" +# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer. +wan_fuse_qk_prescale: True +# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default. +wan_splash_transpose_out: False +# Apply classifier-free guidance before unpatchify (halves the unpatchify work). +wan_cfg_before_unpatchify: True +# Pre-scale cached cross-attention K by 1/sqrt(d) once per video. +wan_cross_attn_prescale_kv: False +# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32". +wan_rope_accum: "auto" # For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this. ulysses_shards: -1 # Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder. diff --git a/src/maxdiffusion/configs/base_wan_1_3b.yml b/src/maxdiffusion/configs/base_wan_1_3b.yml index 3d9a1cc95..e358e33a8 100644 --- a/src/maxdiffusion/configs/base_wan_1_3b.yml +++ b/src/maxdiffusion/configs/base_wan_1_3b.yml @@ -79,7 +79,7 @@ jit_initializers: True # Set true to load weights from pytorch from_pt: True -split_head_dim: True +split_head_dim: False # Sparse VideoGen (SVG) Attention configuration use_svg_attention: False svg_spatial_density: 0.25 @@ -102,7 +102,34 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud use_base2_exp: True use_experimental_scheduler: True +# use_k_centering: "auto" turns K-centering on for the non-ring +# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no +# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a +# pmean across ring shards when R>1). Set True/False to force either path. use_k_centering: "auto" +# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK. +# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml. +use_fused_rope_kernel: False +# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 = +# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model. +fused_rope_block_s: 1024 +fused_rope_head_block: -1 +# Wan inference switches. These change the compiled graph, so they are part of +# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since +# this file sets every key, the env vars are ignored, with a log line if they differ.) +# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction, +# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps). +wan_rope_norm_mode: "exact" +# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer. +wan_fuse_qk_prescale: True +# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default. +wan_splash_transpose_out: False +# Apply classifier-free guidance before unpatchify (halves the unpatchify work). +wan_cfg_before_unpatchify: True +# Pre-scale cached cross-attention K by 1/sqrt(d) once per video. +wan_cross_attn_prescale_kv: False +# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32". +wan_rope_accum: "auto" # For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this. ulysses_shards: -1 # Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder. diff --git a/src/maxdiffusion/configs/base_wan_27b.yml b/src/maxdiffusion/configs/base_wan_27b.yml index 398e8d15f..43227799f 100644 --- a/src/maxdiffusion/configs/base_wan_27b.yml +++ b/src/maxdiffusion/configs/base_wan_27b.yml @@ -82,7 +82,7 @@ jit_initializers: True # Set true to load weights from pytorch from_pt: True -split_head_dim: True +split_head_dim: False # Sparse VideoGen (SVG) Attention configuration use_svg_attention: False svg_spatial_density: 0.35 @@ -127,9 +127,41 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud # that does not equal CP raises ValueError instead of being silently ignored. use_base2_exp: True use_experimental_scheduler: True -# auto: on for non-ring ulysses_custom_fixed_m* (virtual, no copy), off for ring -# paths (virtual K-centering via k_mean, plus a pmean across ring shards when R>1). +# use_k_centering: "auto" turns K-centering on for the non-ring +# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no +# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a +# pmean across ring shards when R>1). Set True/False to force either path. use_k_centering: "auto" +# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK. +# In isolation on one v6e chip it is 1.9x-2.1x faster than the XLA producer +# (40 heads, per-shard seq 9450 / 18900). End-to-end (40 steps, 720p/81f, full +# stack, DVFS unpinned) it cuts denoise 131.2s -> 125.9s on v6e-8 and +# 106.7s -> 105.1s on tpu7x-8; run_wan_fast_inference.sh enables it there. +# Off by default: inlined into the model 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 denoise step, 34.1 dB after 40). Refused on +# platforms whose RoPE rounding is unmeasured; tuned and measured for inference. +use_fused_rope_kernel: False +# Sequence block size for the fused RoPE kernel grid (clamped to 256 in "fused" norm mode). +fused_rope_block_s: 1024 +# Heads processed per grid step; -1 means "all heads in one step". +fused_rope_head_block: -1 +# Wan inference switches. These change the compiled graph, so they are part of +# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since +# this file sets every key, the env vars are ignored, with a log line if they differ.) +# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction, +# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps). +wan_rope_norm_mode: "exact" +# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer. +wan_fuse_qk_prescale: True +# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default. +wan_splash_transpose_out: False +# Apply classifier-free guidance before unpatchify (halves the unpatchify work). +wan_cfg_before_unpatchify: True +# Pre-scale cached cross-attention K by 1/sqrt(d) once per video. +wan_cross_attn_prescale_kv: False +# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32". +wan_rope_accum: "auto" # For attention=ulysses_ring*, hidden Ulysses shard count; ring shards are context / this. ulysses_shards: -1 # Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder. diff --git a/src/maxdiffusion/configs/base_wan_animate.yml b/src/maxdiffusion/configs/base_wan_animate.yml index d90a7f611..80831d6a5 100644 --- a/src/maxdiffusion/configs/base_wan_animate.yml +++ b/src/maxdiffusion/configs/base_wan_animate.yml @@ -80,11 +80,38 @@ jit_initializers: True # Set true to load weights from pytorch from_pt: True -split_head_dim: True +split_head_dim: False attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block use_base2_exp: True use_experimental_scheduler: True +# use_k_centering: "auto" turns K-centering on for the non-ring +# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no +# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a +# pmean across ring shards when R>1). Set True/False to force either path. use_k_centering: "auto" +# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK. +# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml. +use_fused_rope_kernel: False +# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 = +# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model. +fused_rope_block_s: 1024 +fused_rope_head_block: -1 +# Wan inference switches. These change the compiled graph, so they are part of +# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since +# this file sets every key, the env vars are ignored, with a log line if they differ.) +# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction, +# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps). +wan_rope_norm_mode: "exact" +# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer. +wan_fuse_qk_prescale: True +# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default. +wan_splash_transpose_out: False +# Apply classifier-free guidance before unpatchify (halves the unpatchify work). +wan_cfg_before_unpatchify: True +# Pre-scale cached cross-attention K by 1/sqrt(d) once per video. +wan_cross_attn_prescale_kv: False +# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32". +wan_rope_accum: "auto" # For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this. ulysses_shards: -1 # Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder. diff --git a/src/maxdiffusion/configs/base_wan_i2v_14b.yml b/src/maxdiffusion/configs/base_wan_i2v_14b.yml index d252b0ce7..834c40592 100644 --- a/src/maxdiffusion/configs/base_wan_i2v_14b.yml +++ b/src/maxdiffusion/configs/base_wan_i2v_14b.yml @@ -82,11 +82,38 @@ jit_initializers: True # Set true to load weights from pytorch from_pt: True -split_head_dim: True +split_head_dim: False attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block use_base2_exp: True use_experimental_scheduler: True +# use_k_centering: "auto" turns K-centering on for the non-ring +# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no +# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a +# pmean across ring shards when R>1). Set True/False to force either path. use_k_centering: "auto" +# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK. +# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml. +use_fused_rope_kernel: False +# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 = +# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model. +fused_rope_block_s: 1024 +fused_rope_head_block: -1 +# Wan inference switches. These change the compiled graph, so they are part of +# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since +# this file sets every key, the env vars are ignored, with a log line if they differ.) +# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction, +# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps). +wan_rope_norm_mode: "exact" +# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer. +wan_fuse_qk_prescale: True +# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default. +wan_splash_transpose_out: False +# Apply classifier-free guidance before unpatchify (halves the unpatchify work). +wan_cfg_before_unpatchify: True +# Pre-scale cached cross-attention K by 1/sqrt(d) once per video. +wan_cross_attn_prescale_kv: False +# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32". +wan_rope_accum: "auto" # For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this. ulysses_shards: -1 # Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder. diff --git a/src/maxdiffusion/configs/base_wan_i2v_27b.yml b/src/maxdiffusion/configs/base_wan_i2v_27b.yml index 84a64e10d..2bd82a40f 100644 --- a/src/maxdiffusion/configs/base_wan_i2v_27b.yml +++ b/src/maxdiffusion/configs/base_wan_i2v_27b.yml @@ -82,11 +82,38 @@ jit_initializers: True # Set true to load weights from pytorch from_pt: True -split_head_dim: True +split_head_dim: False attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block use_base2_exp: True use_experimental_scheduler: True +# use_k_centering: "auto" turns K-centering on for the non-ring +# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no +# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a +# pmean across ring shards when R>1). Set True/False to force either path. use_k_centering: "auto" +# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK. +# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml. +use_fused_rope_kernel: False +# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 = +# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model. +fused_rope_block_s: 1024 +fused_rope_head_block: -1 +# Wan inference switches. These change the compiled graph, so they are part of +# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since +# this file sets every key, the env vars are ignored, with a log line if they differ.) +# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction, +# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps). +wan_rope_norm_mode: "exact" +# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer. +wan_fuse_qk_prescale: True +# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default. +wan_splash_transpose_out: False +# Apply classifier-free guidance before unpatchify (halves the unpatchify work). +wan_cfg_before_unpatchify: True +# Pre-scale cached cross-attention K by 1/sqrt(d) once per video. +wan_cross_attn_prescale_kv: False +# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32". +wan_rope_accum: "auto" # For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this. ulysses_shards: -1 # Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder. diff --git a/src/maxdiffusion/generate_wan.py b/src/maxdiffusion/generate_wan.py index dba1d0b44..3345a1487 100644 --- a/src/maxdiffusion/generate_wan.py +++ b/src/maxdiffusion/generate_wan.py @@ -14,6 +14,7 @@ from typing import Sequence import jax + import time import os import shutil @@ -22,7 +23,8 @@ from maxdiffusion.checkpointing.wan_checkpointer_2_2 import WanCheckpointer2_2 from maxdiffusion.checkpointing.wan_checkpointer_i2v_2p1 import WanCheckpointerI2V_2_1 from maxdiffusion.checkpointing.wan_checkpointer_i2v_2p2 import WanCheckpointerI2V_2_2 -from maxdiffusion import aot_cache, pyconfig, max_logging, max_utils +from maxdiffusion import aot_cache, pyconfig, max_logging, max_utils, wan_runtime_options +from maxdiffusion.kernels.fused_rmsnorm_rope_pallas import resolve_rope_accum from absl import app from maxdiffusion.train_utils import transformer_engine_context from maxdiffusion.utils import export_to_video @@ -194,6 +196,7 @@ def _build_wan_aot_metadata(config, mesh, source_revision) -> dict[str, str]: "flash_min_seq_length": str(getattr(config, "flash_min_seq_length", 4096)), "mask_padding_tokens": str(getattr(config, "mask_padding_tokens", True)), "precision": str(getattr(config, "precision", "default")), + "split_head_dim": str(getattr(config, "split_head_dim", False)), "logical_axis_rules": str(getattr(config, "logical_axis_rules", ())), "attention_sharding_uniform": str(getattr(config, "attention_sharding_uniform", True)), "allow_split_physical_axes": str(getattr(config, "allow_split_physical_axes", False)), @@ -202,8 +205,26 @@ def _build_wan_aot_metadata(config, mesh, source_revision) -> dict[str, str]: "process_count": str(jax.process_count()), "use_base2_exp": str(getattr(config, "use_base2_exp", True)), "use_experimental_scheduler": str(getattr(config, "use_experimental_scheduler", False)), + "use_fused_rope_kernel": str(getattr(config, "use_fused_rope_kernel", False)), + "fused_rope_block_s": str(getattr(config, "fused_rope_block_s", 1024)), + "fused_rope_head_block": str(getattr(config, "fused_rope_head_block", -1)), "libtpu_init_args": os.environ.get("LIBTPU_INIT_ARGS", ""), "xla_flags": os.environ.get("XLA_FLAGS", ""), + # Graph-changing Wan switches (YAML `wan_*` keys, see wan_runtime_options). + **{ + name: value for name, value in wan_runtime_options.snapshot_from_config(config).items() if name != "wan_rope_accum" + }, + # The RoPE accumulation mode changes the lowered graph's rounding, and + # its default is platform-dependent, so the RESOLVED mode is recorded + # rather than the raw env var. Recording the raw value would both + # over-invalidate (an explicit "f32" on v6e is the same executable as + # the unset default) and fail to describe what was actually compiled. + "wan_rope_accum": str( + resolve_rope_accum( + mesh, + wan_runtime_options.resolve_from_config(config, "wan_rope_accum"), + ) + ), "jax": jax.__version__, "jaxlib": _get_pkg_version("jaxlib"), "libtpu": _get_pkg_version("libtpu"), diff --git a/src/maxdiffusion/kernels/custom_splash_attention.py b/src/maxdiffusion/kernels/custom_splash_attention.py index 5473686dc..a233b3e2e 100644 --- a/src/maxdiffusion/kernels/custom_splash_attention.py +++ b/src/maxdiffusion/kernels/custom_splash_attention.py @@ -26,6 +26,7 @@ from jax.experimental import pallas as pl from jax.experimental.pallas import tpu as pltpu + DEFAULT_MASK_VALUE = -0.7 * float(np.finfo(np.dtype("float32")).max) NUM_LANES = 128 NUM_SUBLANES = 8 @@ -158,6 +159,7 @@ def _flash_attention_kernel_impl( use_fixed_m: bool = False, uniform_fixed_m: bool = False, use_k_centering: bool = False, + transpose_out: bool = False, ): """Pallas Mosaic TPU flash attention kernel with fixed-m support. @@ -446,12 +448,19 @@ def end(): l = l_scratch_ref[...] if fuse_reciprocal: l_inv = jnp.tile(1.0 / l, (head_dim_v_repeats, 1)) - o_ref[...] = (o_scratch_ref[...] * l_inv).astype(o_ref.dtype) + o_val = o_scratch_ref[...] * l_inv + if transpose_out: + o_ref[...] = o_val.T.astype(o_ref.dtype) + else: + o_ref[...] = o_val.astype(o_ref.dtype) else: # Ring path: emit the un-normalized numerator plus the running softmax # stats (max logit `m` and linear denominator `l`) so the outer ring loop # can merge shard contributions and normalize only once at the very end. - o_ref[...] = o_scratch_ref[...].astype(o_ref.dtype) + if transpose_out: + o_ref[...] = o_scratch_ref[...].T.astype(o_ref.dtype) + else: + o_ref[...] = o_scratch_ref[...].astype(o_ref.dtype) if l_ring_ref is not None: l_ring_ref[...] = l.astype(l_ring_ref.dtype) if m_ring_ref is not None: @@ -481,6 +490,7 @@ def _flash_attention_kernel( fuse_reciprocal: bool = True, use_fixed_m: bool = False, uniform_fixed_m: bool = False, + transpose_out: bool = False, ): return _flash_attention_kernel_impl( mk_ref, @@ -506,6 +516,7 @@ def _flash_attention_kernel( use_fixed_m=use_fixed_m, uniform_fixed_m=uniform_fixed_m, use_k_centering=False, + transpose_out=transpose_out, ) @@ -534,6 +545,7 @@ def _flash_attention_kernel_kcentered( use_fixed_m: bool = False, uniform_fixed_m: bool = False, use_k_centering: bool = True, + transpose_out: bool = False, ): return _flash_attention_kernel_impl( mk_ref, @@ -559,6 +571,7 @@ def _flash_attention_kernel_kcentered( use_fixed_m=use_fixed_m, uniform_fixed_m=uniform_fixed_m, use_k_centering=use_k_centering, + transpose_out=transpose_out, ) @@ -580,6 +593,7 @@ def _flash_attention_kernel_mhpt( kv_seq_len: int, heads_per_tile: int, use_base2_exp: bool = True, + transpose_out: bool = False, ): float32 = jnp.float32 head_dim_v_repeats, rem = divmod(head_dim_v, NUM_SUBLANES) @@ -711,7 +725,11 @@ def end(): for h_local in range(heads_per_tile): l = l_scratch_ref[h_local] l_inv = jnp.tile(1.0 / l, (head_dim_v_repeats, 1)) - o_ref[h_local] = (o_scratch_ref[h_local] * l_inv).astype(o_ref.dtype) + o_val = o_scratch_ref[h_local] * l_inv + if transpose_out: + o_ref[h_local] = o_val.T.astype(o_ref.dtype) + else: + o_ref[h_local] = o_val.astype(o_ref.dtype) def _splash_attention_forward( @@ -729,9 +747,12 @@ def _splash_attention_forward( uniform_fixed_m: bool = False, k_mean: jax.Array | None = None, interpret: bool | None = None, + transpose_out: bool | None = None, ): if interpret is None: interpret = jax.default_backend() == "cpu" + if transpose_out is None: + transpose_out = False num_q_heads, padded_q_seq_len, head_dim_qk = q.shape head_dim_v = v.shape[-1] bq, bkv = block_sizes.block_q, block_sizes.block_kv @@ -777,18 +798,32 @@ def k_index_map(h, i, j, *_): def v_index_map(h, i, j, *_): return (h // q_heads_per_kv_head, j, 0) - out_shapes = [ - jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), - jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), - jax.ShapeDtypeStruct((head_dim_v, bq), jnp.float32), - jax.ShapeDtypeStruct((num_q_heads, head_dim_v, actual_q_seq_len), q.dtype), - ] - out_specs = [ - pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), - pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), - pl.BlockSpec((head_dim_v, bq), lambda *_: (0, 0)), - pl.BlockSpec((None, head_dim_v, bq), out_index_map), - ] + if transpose_out: + out_shapes = [ + jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((head_dim_v, bq), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, actual_q_seq_len, head_dim_v), q.dtype), + ] + out_specs = [ + pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), + pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), + pl.BlockSpec((head_dim_v, bq), lambda *_: (0, 0)), + pl.BlockSpec((None, bq, head_dim_v), lambda h, i, j, *_: (h, i, 0)), + ] + else: + out_shapes = [ + jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((head_dim_v, bq), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, head_dim_v, actual_q_seq_len), q.dtype), + ] + out_specs = [ + pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), + pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), + pl.BlockSpec((head_dim_v, bq), lambda *_: (0, 0)), + pl.BlockSpec((None, head_dim_v, bq), out_index_map), + ] if k_mean is None: in_specs = [ @@ -808,6 +843,7 @@ def v_index_map(h, i, j, *_): use_base2_exp=use_base2_exp, use_fixed_m=use_fixed_m, uniform_fixed_m=uniform_fixed_m, + transpose_out=transpose_out, ) kernel_args = (mk, q, k, v) else: @@ -834,6 +870,7 @@ def v_index_map(h, i, j, *_): use_fixed_m=use_fixed_m, uniform_fixed_m=uniform_fixed_m, use_k_centering=True, + transpose_out=transpose_out, ) kernel_args = (mk, q, k, v, k_mean_tiled) @@ -1040,7 +1077,11 @@ def _splash_attention_forward_mhpt( use_base2_exp: bool = True, use_experimental_scheduler: bool = False, vmem_limit_bytes: int | None = None, + transpose_out: bool | None = None, ): + if transpose_out is None: + transpose_out = False + num_q_heads, padded_q_seq_len, head_dim_qk = q.shape head_dim_v = v.shape[-1] bq, bkv = block_sizes.block_q, block_sizes.block_kv @@ -1073,18 +1114,32 @@ def out_index_map(h, i, j, *_): pl.BlockSpec((hpt, bkv, head_dim_qk), k_index_map), pl.BlockSpec((hpt, bkv, head_dim_v), v_index_map), ] - out_shapes = [ - jax.ShapeDtypeStruct((hpt, NUM_SUBLANES, bq), jnp.float32), - jax.ShapeDtypeStruct((hpt, NUM_SUBLANES, bq), jnp.float32), - jax.ShapeDtypeStruct((hpt, head_dim_v, bq), jnp.float32), - jax.ShapeDtypeStruct((num_q_heads, head_dim_v, actual_q_seq_len), q.dtype), - ] - out_specs = [ - pl.BlockSpec((hpt, NUM_SUBLANES, bq), lambda *_: (0, 0, 0)), - pl.BlockSpec((hpt, NUM_SUBLANES, bq), lambda *_: (0, 0, 0)), - pl.BlockSpec((hpt, head_dim_v, bq), lambda *_: (0, 0, 0)), - pl.BlockSpec((hpt, head_dim_v, bq), out_index_map), - ] + if transpose_out: + out_shapes = [ + jax.ShapeDtypeStruct((hpt, NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((hpt, NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((hpt, head_dim_v, bq), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, actual_q_seq_len, head_dim_v), q.dtype), + ] + out_specs = [ + pl.BlockSpec((hpt, NUM_SUBLANES, bq), lambda *_: (0, 0, 0)), + pl.BlockSpec((hpt, NUM_SUBLANES, bq), lambda *_: (0, 0, 0)), + pl.BlockSpec((hpt, head_dim_v, bq), lambda *_: (0, 0, 0)), + pl.BlockSpec((hpt, bq, head_dim_v), lambda h, i, j, *_: (h, i, 0)), + ] + else: + out_shapes = [ + jax.ShapeDtypeStruct((hpt, NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((hpt, NUM_SUBLANES, bq), jnp.float32), + jax.ShapeDtypeStruct((hpt, head_dim_v, bq), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, head_dim_v, actual_q_seq_len), q.dtype), + ] + out_specs = [ + pl.BlockSpec((hpt, NUM_SUBLANES, bq), lambda *_: (0, 0, 0)), + pl.BlockSpec((hpt, NUM_SUBLANES, bq), lambda *_: (0, 0, 0)), + pl.BlockSpec((hpt, head_dim_v, bq), lambda *_: (0, 0, 0)), + pl.BlockSpec((hpt, head_dim_v, bq), out_index_map), + ] grid_width = (actual_kv_seq_len + bkv - 1) // bkv grid_height = (actual_q_seq_len + bq - 1) // bq grid = (num_q_heads // hpt, grid_height, grid_width) @@ -1101,6 +1156,7 @@ def out_index_map(h, i, j, *_): kv_seq_len=actual_kv_seq_len, heads_per_tile=hpt, use_base2_exp=use_base2_exp, + transpose_out=transpose_out, ), grid_spec=pltpu.PrefetchScalarGridSpec( num_scalar_prefetch=0, @@ -1131,7 +1187,10 @@ def make_splash_mha( use_fixed_m: bool = False, uniform_fixed_m: bool = False, interpret: bool | None = None, + transpose_out: bool | None = None, ): + if transpose_out is None: + transpose_out = False if use_fixed_m and not use_base2_exp: raise NotImplementedError( "fixed-m softmax bounds are derived strictly for base-2 exponents. Please set use_base2_exp=True." @@ -1154,6 +1213,7 @@ def _splash_attention(q, k, v, mk=None, k_mean=None): use_base2_exp=use_base2_exp, use_experimental_scheduler=use_experimental_scheduler, vmem_limit_bytes=vmem_limit_bytes, + transpose_out=transpose_out, ) return _splash_attention_forward( q, @@ -1170,6 +1230,7 @@ def _splash_attention(q, k, v, mk=None, k_mean=None): uniform_fixed_m=uniform_fixed_m, k_mean=k_mean, interpret=interpret, + transpose_out=transpose_out, ) return _splash_attention diff --git a/src/maxdiffusion/kernels/fused_rmsnorm_rope_pallas.py b/src/maxdiffusion/kernels/fused_rmsnorm_rope_pallas.py new file mode 100644 index 000000000..73b85cb84 --- /dev/null +++ b/src/maxdiffusion/kernels/fused_rmsnorm_rope_pallas.py @@ -0,0 +1,701 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Fused RMSNorm + RoPE + head-transposition Pallas kernel for TPU. + +Why +--- +The unfused producer in `fused_producers.fused_rmsnorm_rope` makes three HBM +passes over the same activation (193.5 MB per batch element, or 387 MB at CFG +batch=2, per tensor per device at 18,900 local tokens x 5,120 features in bf16): +FP32 RMSNorm over the 5120-wide feature axis; the `[B, S, H*D] -> [B, H, S, D]` +head transposition; then the RoPE elementwise chain reading the result back. + +The transposition is not a free bitcast -- `[S, H*D]` is tiled `(8, 128)` over +`(S, H*D)` while `[H, S, D]` is tiled over `(S, D)` -- so XLA emits a physical +relayout copy: 0.35 ms x 4 ops per 2 steps on 16 cores, ~1.75 s per 40-step +denoise. + +XLA cannot fuse the three, because RMSNorm reduces along the feature axis (which +spans *all* heads) while RoPE and attention need a head-major layout. Commuting +RoPE ahead of the transpose in pure JAX (measured in an out-of-tree experiment) +is bit-exact but 3.8-4.3 s *slower*: broadcasting `cos`/`sin` across an interior +`heads` axis destroys sublane broadcast coalescing. + +Normalisation modes +------------------- +`norm_mode="exact"` (default) + The FP32 `mean(x**2)` reduction stays in XLA; the kernel gets only its + `[B, S, 1]` result, then does scale + RoPE + transpose in one pass. Every op + is the same HLO op in the same order as the reference, so the kernel is + bit-identical to the *separately jitted* `fused_rmsnorm_rope` producer + (asserted at 0 ULP in the tests). That is a kernel-level guarantee only: + inside the full 40-layer graph XLA fuses the unfused producer with its + neighbours and rounds it differently, so end-to-end output is equivalent + but not identical (v6e 720p/81f, same seed: 53.8 dB PSNR after 1 denoise + step, 34.1 dB after 40; see `_fused_rope_producer`). Costs one extra + streaming read. + +`norm_mode="fused"` + The reduction moves into the kernel, so the activation is read once. + Mosaic's reduction tree need not match XLA's, which can flip the final bf16 + rounding: within 2x pair-relative bf16 machine epsilon + (`2 * eps(bf16) * ||(q_{2i}, q_{2i+1})||_2`, affecting ~0.001% of elements + with max `|diff| <= 3.125e-2` at the Wan shard shape). Use only where a + hash-identical video is not required. + +Numerical contract +------------------ +Apart from the reduction tree and the RoPE combine's rounding (selected by +`rope_accum`, see `resolve_rope_accum`), every op matches the reference: + + * RMSNorm keeps Flax's association `x * (rsqrt(var + eps) * scale)`; folding + left-to-right as `(x * rsqrt) * scale` rounds differently. + * RoPE is evaluated either in the activation dtype (`rope_accum="dtype"`) or + in FP32 with a single rounding (`rope_accum="f32"`) as + `out[2i] = q[2i] *cos[i] + q[2i+1]*(-sin[i])` and + `out[2i+1] = q[2i+1]*cos[i] + q[2i] *(+sin[i])`, which reproduces the + reference's rounding on measured platforms (`v6e`, `tpu7x`, `XLA:CPU`): + `a + (-b) == a - b` and `a + b == b + a` hold in IEEE-754. + * `cos`/`sin` are rounded to the activation dtype *before* lane duplication, + so each lane holds precisely the value the reference multiplies by. + +The pairwise `(q[2i], q[2i+1]) -> (q[2i+1], q[2i])` swap uses two circular lane +rotations and an even-lane select, not a strided `x[..., 0::2]` gather, which +Mosaic cannot lower efficiently. +""" + +import functools +from typing import Tuple + +import jax +import jax.numpy as jnp +from jax.experimental import pallas as pl +from jax.experimental.pallas import tpu as pltpu + +NUM_LANES = 128 + +# Kernel-level default sequence tile sizes (production overrides this with +# `fused_rope_block_s = 1024` from the Wan YAMLs; in `fused` mode `block_s` is +# capped at 256). Both paths stream a full-width `[block_s, feature]` row tile, +# so VMEM scales as `block_s * feature * dtype_size` (times two for double +# buffering). The FP32 temporaries differ: "exact" builds them after the +# per-head slice so they are `[block_s, dim_head]`, while "fused" must hold a +# `[block_s, feature]` FP32 intermediate for the reduction. +DEFAULT_BLOCK_S_EXACT = 512 +DEFAULT_BLOCK_S_FUSED = 256 + +# v6e exposes 128 MiB of VMEM per core and tpu7x exposes 64 MiB. 64 MiB fits +# both (the full tpu7x budget). Other TPU generations get Mosaic's own default; +# see `_resolve_vmem_limit_bytes`. +DEFAULT_VMEM_LIMIT_BYTES = 64 * 1024 * 1024 + +NORM_MODES = ("exact", "fused") +ROPE_ACCUM_MODES = ("dtype", "f32") + + +# Rounding conventions that have actually been measured, as +# (device_kind substring, accumulation mode). A platform is on this list only +# if the kernel was shown to be bit-identical to the XLA producer under that +# mode at the production shape. Anything absent gets a best-effort default and +# is reported as unmeasured by `rope_accum_is_measured`, because the honest +# claim there is "close", not "identical": on TPU v4, for instance, no mode is +# bit-identical -- the closest is "dtype" against a compiled reference, which +# still differs on 20.9% of elements by up to one bf16 ULP. +_MEASURED_TPU_ROUNDING = ( + ("7x", "dtype"), + ("v6", "f32"), +) +_UNMEASURED_TPU_DEFAULT = "f32" + + +def _rope_accum_devices(mesh=None): + """Returns the device list backing the resolution, or None if no backend.""" + try: + return list(mesh.devices.flat) if mesh is not None else jax.devices() + except RuntimeError: + return None + + +# Device kinds on which this kernel has been compiled and run with +# DEFAULT_VMEM_LIMIT_BYTES (v6e: 128 MiB VMEM per core; tpu7x: 64 MiB). +_VMEM_LIMIT_VALIDATED_KINDS = ("7x", "v6") + + +def _resolve_vmem_limit_bytes(vmem_limit_bytes: int | None = None, mesh=None) -> int | None: + """Scoped VMEM budget for the kernel. + + An explicit value wins. Otherwise DEFAULT_VMEM_LIMIT_BYTES is requested only + when every TPU in `mesh` (or `jax.devices()`) is a generation this kernel was + validated on. Any other TPU returns None, so Mosaic keeps its own + conservative scoped default instead of being asked for a large fraction (or + all) of a smaller physical VMEM. Off TPU (interpret mode) the value is + unused, and the default is returned. + """ + if vmem_limit_bytes is not None: + return vmem_limit_bytes + devices = _rope_accum_devices(mesh) + kinds = [getattr(d, "device_kind", "").lower() for d in (devices or ()) if getattr(d, "platform", "") == "tpu"] + if not kinds: + return DEFAULT_VMEM_LIMIT_BYTES + if all(any(tag in k for tag in _VMEM_LIMIT_VALIDATED_KINDS) for k in kinds): + return DEFAULT_VMEM_LIMIT_BYTES + return None + + +def vmem_limit_is_validated(mesh=None) -> bool: + """True when every device in `mesh` (or `jax.devices()`) is a TPU this kernel was validated on. + + Tests that only hold on the validated generations (production shapes that + need DEFAULT_VMEM_LIMIT_BYTES) gate on this, so they skip on other TPUs. + """ + devices = _rope_accum_devices(mesh) or () + kinds = [getattr(d, "device_kind", "").lower() for d in devices if getattr(d, "platform", "") == "tpu"] + return ( + bool(kinds) and len(kinds) == len(devices) and all(any(tag in k for tag in _VMEM_LIMIT_VALIDATED_KINDS) for k in kinds) + ) + + +def resolve_rope_accum(mesh=None, rope_accum: str | None = None) -> str: + """Returns the RoPE accumulation mode this process will actually compile with. + + The default is platform-dependent, because the mode has to match however the + local compiler chooses to round the reference's RoPE multiply-add. Matching + that contraction is what yields 0-ULP output against the XLA producer. + + * tpu7x: XLA emits a native bf16 multiply-add -> "dtype". + * TPU v6: XLA contracts the multiply-add into an FP32 FMA -> "f32". + * Off TPU: XLA:CPU emits the multiply and the add as separate rounded + ops -> "dtype". Measured, not assumed: on CPU the eager and jitted + references agree to 0 ULP and only "dtype" matches either. This branch is + reached only by Pallas interpret-mode tests, since production falls back + to the unfused XLA producer off TPU. + * Other TPU generations: "f32" as a best effort. See + `rope_accum_is_measured` before relying on bit-exactness there. + + A non-"auto" `rope_accum` argument or `wan_rope_accum` config value (legacy + env: `WAN_ROPE_ACCUM`) overrides all of it. + + This resolver is shared by the production call site and by the AOT cache + fingerprint. Keep it that way: the mode changes the lowered graph, so a + duplicated copy of this rule that drifts from the real one would let an + executable compiled under one rounding mode be served for another. + + Args: + mesh: Mesh whose devices determine the platform default. Falls back to + `jax.devices()` when None. Pass the same mesh the kernel runs under. + rope_accum: Optional explicit override from config (`"auto"`, `"dtype"`, or + `"f32"`). None means `"auto"`. + + Returns: + One of `ROPE_ACCUM_MODES`. + """ + override = "auto" if rope_accum is None else str(rope_accum) + if override != "auto": + if override not in ROPE_ACCUM_MODES: + raise ValueError(f"wan_rope_accum must be 'auto' or one of {ROPE_ACCUM_MODES}, got {override!r}") + return override + devices = _rope_accum_devices(mesh) + if devices is None: + # No backend (e.g. metadata built before device init); assume the + # conservative FP32 contraction. + return _UNMEASURED_TPU_DEFAULT + if not any(getattr(d, "platform", "") == "tpu" for d in devices): + return "dtype" + kinds = [getattr(d, "device_kind", "").lower() for d in devices] + for fragment, mode in _MEASURED_TPU_ROUNDING: + if any(fragment in kind for kind in kinds): + return mode + return _UNMEASURED_TPU_DEFAULT + + +def rope_accum_is_measured(mesh=None) -> bool: + """Whether the current hardware platform's RoPE rounding has been measured. + + False means the kernel still computes the right value but may not be + *bit-identical* to the XLA producer, because no accumulation mode has been + shown to reproduce this hardware's rounding (e.g. on TPU v4). Callers that + require verified hardware (such as `_fused_rope_producer`) consult this guard, + and an explicit `wan_rope_accum` does not bypass it on unmeasured platforms. + """ + devices = _rope_accum_devices(mesh) + if devices is None: + return False + if not any(getattr(d, "platform", "") == "tpu" for d in devices): + # Only XLA:CPU was measured; a GPU backend has not been. + return all(getattr(d, "platform", "") == "cpu" for d in devices) + kinds = [getattr(d, "device_kind", "").lower() for d in devices] + return any(fragment in kind for fragment, _ in _MEASURED_TPU_ROUNDING for kind in kinds) + + +def with_xla_backward(fused_fn, xla_fn): + """Makes a Pallas producer differentiable by transposing an equivalent XLA graph. + + `pallas_call` has no transpose rule, so a fused producer in a training graph + fails at `jax.grad` time, from inside the autodiff machinery and with no + indication that the fusion is what broke. This keeps the kernel on the forward + pass and routes the backward pass through `xla_fn`, which must compute the same + function by unfused primitives. The substitution is exact wherever the two + forwards agree -- the property `rope_accum_is_measured` tracks. It costs one + extra unfused forward per step, since the backward re-runs it to linearise. + + Args: + fused_fn: Callable over differentiable array arguments only (bind static + configuration with `functools.partial` first). + xla_fn: Same signature and same mathematics, differentiable. + + Returns: + A callable equivalent to `fused_fn` whose VJP is that of `xla_fn`. + """ + + @jax.custom_vjp + def wrapped(*args): + return fused_fn(*args) + + def _fwd(*args): + return fused_fn(*args), args + + def _bwd(residual, cotangents): + return jax.vjp(xla_fn, *residual)[1](cotangents) + + wrapped.defvjp(_fwd, _bwd) + return wrapped + + +def _apply_rope( + normed: jax.Array, + cos: jax.Array, + sin_signed: jax.Array, + accum: str, + even_mask: jax.Array | None = None, +) -> jax.Array: + """Combines `normed * cos + swap(normed) * sin` under a chosen rounding mode. + + The reference is `q0*cos - q1*sin`. Whether XLA contracts that into an FP32 + fused multiply-add (rounding once) or emits separate rounded ops is + platform-dependent: `v6e` contracts under `jit`, while `tpu7x` and `XLA:CPU` + do not (and eager vs `jit` can also differ; see `resolve_rope_accum`). Both + roundings are legitimate; only one matches any given reference build, so the + caller must be able to pick. + + `accum="dtype"`: round each product to the activation dtype, then add. + Matches an uncontracted reference. + `accum="f32"`: evaluate both products and the sum in FP32 and round once. + Matches an FMA-contracted reference. + """ + if accum == "f32": + wide = normed.astype(jnp.float32) + out = wide * cos.astype(jnp.float32) + _pair_swap(wide, even_mask=even_mask) * sin_signed.astype(jnp.float32) + return out.astype(normed.dtype) + return normed * cos + _pair_swap(normed, even_mask=even_mask) * sin_signed + + +def _pair_swap(x: jax.Array, even_mask: jax.Array | None = None) -> jax.Array: + """Swaps adjacent lane pairs: `[a0, a1, a2, a3, ...] -> [a1, a0, a3, a2, ...]`. + + Implemented with two circular rotations and an even-lane select. The strided + `x[..., 0::2]` / `x[..., 1::2]` gather that the reference relies on XLA to + handle is a lane-stride-2 relayout which Mosaic lowers very poorly, whereas + `tpu.DynamicRotate` is a single cheap lane rotation. + + Rotation is done at 32-bit width on all platforms (required by + `tpu.DynamicRotate` on TPU 7x); the round-trip cast is exact. + + The rotation is circular, but because `dim_head` is even the wrap-around + lanes land exactly where the swap needs them: + * lane 0 takes `roll_left[0] = x[1]`, the partner of lane 0. + * lane D-1 takes `roll_right[D-1] = x[D-2]`, the partner of lane D-1. + """ + orig_dtype = x.dtype + if orig_dtype != jnp.float32 and orig_dtype != jnp.int32: + x = x.astype(jnp.float32) + dim = x.shape[-1] + axis = x.ndim - 1 + roll_right = pltpu.roll(x, 1, axis) # roll_right[i] = x[i - 1] + roll_left = pltpu.roll(x, dim - 1, axis) # roll_left[i] = x[i + 1] + if even_mask is None: + lane = jax.lax.broadcasted_iota(jnp.int32, x.shape, axis) + even_mask = jax.lax.rem(lane, 2) == 0 + res = jnp.where(even_mask, roll_left, roll_right) + return res.astype(orig_dtype) + + +def _rope_tables(freqs_cis: jax.Array, seq_len: int, dtype: jnp.dtype) -> Tuple[jax.Array, jax.Array]: + """Expands `freqs_cis` into lane-aligned `cos` / signed-`sin` tables. + + Args: + freqs_cis: Complex rotary embedding, shape `[1, 1, S, dim_head // 2]`. + seq_len: Number of sequence positions to keep. + dtype: Activation dtype. The tables are rounded to it *before* the + duplication so every lane holds exactly the value the reference + implementation multiplies by. + + Returns: + `(cos_full, sin_signed)`, each `[1, seq_len, dim_head]`, with + `cos_full[..., 2i] == cos_full[..., 2i+1] == cos[i]`, + `sin_signed[..., 2i] == -sin[i]` and `sin_signed[..., 2i+1] == +sin[i]`. + """ + cos = jnp.real(freqs_cis)[0, :, :seq_len, :].astype(dtype) + sin = jnp.imag(freqs_cis)[0, :, :seq_len, :].astype(dtype) + cos_full = jnp.repeat(cos, 2, axis=-1) + sin_full = jnp.repeat(sin, 2, axis=-1) + # [-1, +1, -1, +1, ...]; scaling a float by +-1 is exact. + sign = jnp.tile(jnp.array([-1.0, 1.0], dtype=dtype), (sin_full.shape[-1] // 2,)) + return cos_full, sin_full * sign + + +def _exact_kernel( + x_ref, + scale_ref, + rsqrt_ref, + cos_ref, + sin_ref, + o_ref, + *, + dim_head: int, + rope_accum: str, + head_block: int, + prescale: float = 1.0, +): + """Scale + RoPE + transpose for one `(batch, sequence tile, head block)` step. + + The activation tile is fetched at *full feature width* and the per-head slice + is taken in VMEM. Slicing the head out of HBM instead (via the input + `BlockSpec`) would read `dim_head * 2 = 256` contiguous bytes per row out of + a 10,240-byte row, i.e. a strided burst pattern that runs at a fraction of + HBM speed. At full width the read is sequential, and because the block index + does not depend on the head, Pallas fetches each tile exactly once. + + `head_block` heads are emitted per grid step. With one head per step the grid + is `B * ceil(S/block_s) * heads` and each output DMA is only `block_s * dim_head` + elements; the per-step overhead then dominates, costing ~45% over emitting + all heads at once. The loop is unrolled in Python so every store addresses a + statically known sub-block. + + The FP32 temporaries are created *after* the slice, so they are + `[block_s, dim_head]` rather than `[block_s, heads * dim_head]`. + """ + blk = pl.program_id(2) + rsqrt_val = rsqrt_ref[0] + cos = cos_ref[0].astype(jnp.float32) if rope_accum == "f32" else cos_ref[0] + sin = sin_ref[0].astype(jnp.float32) if rope_accum == "f32" else sin_ref[0] + lane = jax.lax.broadcasted_iota(jnp.int32, cos.shape, cos.ndim - 1) + even_mask = jax.lax.rem(lane, 2) == 0 + prescale_val = jnp.asarray(prescale, o_ref.dtype) if prescale != 1.0 else None + + for i in range(head_block): + # Pure lane-tile selection: `dim_head` is a multiple of NUM_LANES, so this + # picks whole lane tiles and needs no relayout. `pl.multiple_of` supplies + # the alignment fact Mosaic cannot infer from a dynamic product. + offset = pl.multiple_of((blk * head_block + i) * dim_head, dim_head) + x = x_ref[0, :, pl.ds(offset, dim_head)] + + # Flax's association: fold the scale into the reciprocal before applying it. + mul = rsqrt_val * scale_ref[i, 0].astype(jnp.float32) + normed = (x.astype(jnp.float32) * mul).astype(x.dtype) + out = _apply_rope(normed, cos, sin, rope_accum, even_mask=even_mask) + if prescale_val is not None: + # Intentionally multiply in out.dtype (bfloat16) after _apply_rope so that + # in-register prescaling is 0-ULP bit-identical on measured platforms + # (v6e, tpu7x) to the unfused attention path (`query * LOG2E` and + # `k * scale` on the bfloat16 RoPE output). + out = out * prescale_val + o_ref[0, i] = out + + +def _fused_kernel( + x_ref, + scale_ref, + cos_ref, + sin_ref, + o_ref, + *, + dim_head: int, + eps: float, + rope_accum: str, + head_block: int, + prescale: float = 1.0, +): + """As `_exact_kernel`, but also computes the feature-axis reduction in VMEM.""" + blk = pl.program_id(2) + x_full_f32 = x_ref[0].astype(jnp.float32) + var = jnp.mean(x_full_f32 * x_full_f32, axis=-1, keepdims=True) + inv_rms = jax.lax.rsqrt(var + eps) + cos = cos_ref[0].astype(jnp.float32) if rope_accum == "f32" else cos_ref[0] + sin = sin_ref[0].astype(jnp.float32) if rope_accum == "f32" else sin_ref[0] + lane = jax.lax.broadcasted_iota(jnp.int32, cos.shape, cos.ndim - 1) + even_mask = jax.lax.rem(lane, 2) == 0 + prescale_val = jnp.asarray(prescale, o_ref.dtype) if prescale != 1.0 else None + + for i in range(head_block): + offset = pl.multiple_of((blk * head_block + i) * dim_head, dim_head) + x = x_ref[0, :, pl.ds(offset, dim_head)] + mul = inv_rms * scale_ref[i, 0].astype(jnp.float32) + normed = (x.astype(jnp.float32) * mul).astype(x.dtype) + out = _apply_rope(normed, cos, sin, rope_accum, even_mask=even_mask) + if prescale_val is not None: + out = out * prescale_val + o_ref[0, i] = out + + +def _run_exact( + x, + scale, + cos_full, + sin_signed, + *, + heads, + dim_head, + eps, + block_s, + vmem_limit_bytes, + interpret, + rope_accum, + head_block, + prescale: float = 1.0, +): + batch, seq_len, feature = x.shape + block_s = min(block_s, seq_len) + + # The reduction is left to XLA, expressed exactly as the reference expresses + # it. That makes it bit-identical *provided* XLA emits the same reduction for + # both graphs; see `rope_accum_is_measured` on why that is checked empirically + # rather than assumed. + x_f32 = x.astype(jnp.float32) + rsqrt_val = jax.lax.rsqrt(jnp.mean(jnp.square(x_f32), axis=-1, keepdims=True) + eps) + + return pl.pallas_call( + functools.partial(_exact_kernel, dim_head=dim_head, rope_accum=rope_accum, head_block=head_block, prescale=prescale), + grid=(batch, pl.cdiv(seq_len, block_s), heads // head_block), + in_specs=[ + # Full-width and head-invariant: one sequential fetch per tile. + pl.BlockSpec((1, block_s, feature), lambda b, s, h: (b, s, 0)), + # Mosaic requires the second-minor block dimension to be a multiple + # of 8 or to equal the array's. `scale` is carried as + # `[heads, 1, dim_head]` so the unit axis satisfies the latter; a + # `[heads, dim_head]` layout with a `(1, dim_head)` block does not + # lower at all. + pl.BlockSpec((head_block, 1, dim_head), lambda b, s, h: (h, 0, 0)), + pl.BlockSpec((1, block_s, 1), lambda b, s, h: (b, s, 0)), + pl.BlockSpec((1, block_s, dim_head), lambda b, s, h: (0, s, 0)), + pl.BlockSpec((1, block_s, dim_head), lambda b, s, h: (0, s, 0)), + ], + out_specs=pl.BlockSpec((1, head_block, block_s, dim_head), lambda b, s, h: (b, h, s, 0)), + out_shape=jax.ShapeDtypeStruct((batch, heads, seq_len, dim_head), x.dtype), + compiler_params=pltpu.CompilerParams( + dimension_semantics=("parallel", "parallel", "arbitrary"), vmem_limit_bytes=vmem_limit_bytes + ), + interpret=interpret, + )(x, scale.reshape(heads, 1, dim_head), rsqrt_val, cos_full, sin_signed) + + +def _run_fused( + x, + scale, + cos_full, + sin_signed, + *, + heads, + dim_head, + eps, + block_s, + vmem_limit_bytes, + interpret, + rope_accum, + head_block, + prescale: float = 1.0, +): + batch, seq_len, feature = x.shape + block_s = min(block_s, seq_len) + + # The input block index does not depend on the head-block axis, so Pallas + # fetches each `[block_s, H*D]` tile from HBM once per `(b, s)` tile and + # reduces `x_ref[0].astype(jnp.float32)` in VMEM without a separate HBM pass + # for `rsqrt`. + return pl.pallas_call( + functools.partial( + _fused_kernel, dim_head=dim_head, eps=eps, rope_accum=rope_accum, head_block=head_block, prescale=prescale + ), + grid=(batch, pl.cdiv(seq_len, block_s), heads // head_block), + in_specs=[ + pl.BlockSpec((1, block_s, feature), lambda b, s, h: (b, s, 0)), + pl.BlockSpec((head_block, 1, dim_head), lambda b, s, h: (h, 0, 0)), + pl.BlockSpec((1, block_s, dim_head), lambda b, s, h: (0, s, 0)), + pl.BlockSpec((1, block_s, dim_head), lambda b, s, h: (0, s, 0)), + ], + out_specs=pl.BlockSpec((1, head_block, block_s, dim_head), lambda b, s, h: (b, h, s, 0)), + out_shape=jax.ShapeDtypeStruct((batch, heads, seq_len, dim_head), x.dtype), + compiler_params=pltpu.CompilerParams( + dimension_semantics=("parallel", "parallel", "arbitrary"), vmem_limit_bytes=vmem_limit_bytes + ), + interpret=interpret, + )(x, scale.reshape(heads, 1, dim_head), cos_full, sin_signed) + + +def fused_rmsnorm_rope_pallas( + raw_q: jax.Array, + raw_k: jax.Array, + q_norm_scale: jax.Array, + k_norm_scale: jax.Array, + freqs_cis: jax.Array, + q_heads: int = 40, + kv_heads: int | None = None, + dim_head: int = 128, + eps: float = 1e-6, + k_eps: float | None = None, + heads: int | None = None, + norm_mode: str = "exact", + rope_accum: str = "dtype", + block_s: int | None = None, + head_block: int | None = None, + vmem_limit_bytes: int | None = None, + interpret: bool = False, + mesh=None, + q_prescale: float = 1.0, + k_prescale: float = 1.0, +) -> Tuple[jax.Array, jax.Array]: + """Fused FP32 RMSNorm + RoPE + head transposition on TPU via Pallas. + + Drop-in replacement for `fused_producers.fused_rmsnorm_rope`. + + Args: + raw_q: Raw query projection, `[B, Sq, q_heads * dim_head]`. + raw_k: Raw key projection, `[B, Sk, kv_heads * dim_head]`. + q_norm_scale: RMSNorm scale for the query, `[q_heads * dim_head]`. + k_norm_scale: RMSNorm scale for the key, `[kv_heads * dim_head]`. + freqs_cis: Complex rotary embedding, `[1, 1, S, dim_head // 2]`. + q_heads: Number of query heads. + kv_heads: Number of key/value heads (defaults to `q_heads` for MHA). + dim_head: Per-head dimension. Must be a multiple of 128 and even. + eps: RMSNorm epsilon for `raw_q` (and `raw_k` when `k_eps` is None). + k_eps: Optional separate RMSNorm epsilon for `raw_k`. + heads: Deprecated alias for `q_heads`. + norm_mode: `"exact"` leaves the FP32 feature-axis reduction to XLA, written + exactly as the reference writes it; `"fused"` folds it into the kernel, + reading the activation once but changing the summation order. Measured at + the Wan shape, `"fused"` stays within 2x pair-relative bf16 machine + epsilon (`2 * eps(bf16) * ||(q_{2i}, q_{2i+1})||_2`, altering ~0.001% of + elements by up to 3.125e-2), which is too coarse for a hash-equality bar. + rope_accum: Rounding of the RoPE combine. `"dtype"` rounds each product to + the activation dtype before summing; `"f32"` keeps both products and the + sum in FP32 and rounds once. The reference matches `"dtype"` when its + multiply-adds are left uncontracted and `"f32"` when XLA contracts them + into FMAs, which depends on the surrounding graph. Pick whichever + reproduces the build you must match. + block_s: Sequence tile size (must be a positive multiple of 8); defaults per + `norm_mode` (512 for `"exact"`, 256 for `"fused"`, and clamped to at most + 256 in `"fused"` mode). + head_block: Number of heads emitted per grid step. Must divide the head + count of each tensor. Defaults to all of them, which measured fastest: + one head per step leaves a 2960-step grid whose per-step overhead costs + ~45%. Lower it only if VMEM is tight, since the output tile is + `head_block * block_s * dim_head`. + vmem_limit_bytes: Scoped VMEM budget handed to Mosaic (resolved per TPU + generation when None, see `_resolve_vmem_limit_bytes`). + interpret: Run in Pallas interpret mode (for CPU tests). + mesh: Mesh the kernel runs under; its devices pick the VMEM budget. + + Returns: + `(q_out, k_out)` of shapes `[B, q_heads, Sq, dim_head]` and + `[B, kv_heads, Sk, dim_head]`. + """ + if heads is not None: + q_heads = heads + kv_heads = q_heads if kv_heads is None else kv_heads + effective_k_eps = eps if k_eps is None else k_eps + + if norm_mode not in NORM_MODES: + raise ValueError(f"norm_mode must be one of {NORM_MODES}, got {norm_mode!r}.") + if rope_accum not in ROPE_ACCUM_MODES: + raise ValueError(f"rope_accum must be one of {ROPE_ACCUM_MODES}, got {rope_accum!r}.") + if dim_head % NUM_LANES != 0: + raise ValueError(f"fused_rmsnorm_rope_pallas requires dim_head to be a multiple of {NUM_LANES}, got {dim_head}.") + if dim_head % 2 != 0: + raise ValueError(f"RoPE requires an even dim_head, got {dim_head}.") + if head_block is not None: + if head_block <= 0: + raise ValueError(f"head_block must be a positive integer, got {head_block}.") + for name, n in (("q_heads", q_heads), ("kv_heads", kv_heads)): + if n % head_block != 0: + raise ValueError(f"head_block ({head_block}) must divide {name} ({n}).") + if block_s is not None and (block_s <= 0 or block_s % 8 != 0): + raise ValueError(f"block_s must be a positive multiple of 8, got {block_s}.") + + _, seq_q, feature_q = raw_q.shape + _, seq_k, feature_k = raw_k.shape + if feature_q != q_heads * dim_head: + raise ValueError(f"raw_q feature dim ({feature_q}) must equal q_heads ({q_heads}) * dim_head ({dim_head})") + if feature_k != kv_heads * dim_head: + raise ValueError(f"raw_k feature dim ({feature_k}) must equal kv_heads ({kv_heads}) * dim_head ({dim_head})") + if freqs_cis.shape[-1] * 2 != dim_head: + raise ValueError(f"freqs_cis last dim ({freqs_cis.shape[-1]}) must be dim_head // 2 ({dim_head // 2}).") + if freqs_cis.shape[2] < max(seq_q, seq_k): + raise ValueError( + f"freqs_cis sequence dim ({freqs_cis.shape[2]}) must be at least max(seq_q, seq_k) ({max(seq_q, seq_k)})." + ) + + if block_s is None: + block_s = DEFAULT_BLOCK_S_EXACT if norm_mode == "exact" else DEFAULT_BLOCK_S_FUSED + elif norm_mode == "fused": + block_s = min(block_s, DEFAULT_BLOCK_S_FUSED) + resolved_vmem_limit_bytes = _resolve_vmem_limit_bytes(vmem_limit_bytes, mesh) + runner = _run_exact if norm_mode == "exact" else _run_fused + + cos_q, sin_q = _rope_tables(freqs_cis, seq_q, raw_q.dtype) + if seq_k == seq_q and raw_k.dtype == raw_q.dtype: + cos_k, sin_k = cos_q, sin_q + else: + cos_k, sin_k = _rope_tables(freqs_cis, seq_k, raw_k.dtype) + + common = { + "dim_head": dim_head, + "block_s": block_s, + "vmem_limit_bytes": resolved_vmem_limit_bytes, + "interpret": interpret, + "rope_accum": rope_accum, + } + q_out = runner( + raw_q, + q_norm_scale, + cos_q, + sin_q, + heads=q_heads, + head_block=q_heads if head_block is None else head_block, + eps=eps, + prescale=q_prescale, + **common, + ) + k_out = runner( + raw_k, + k_norm_scale, + cos_k, + sin_k, + heads=kv_heads, + head_block=kv_heads if head_block is None else head_block, + eps=effective_k_eps, + prescale=k_prescale, + **common, + ) + return q_out, k_out + + +def rope_pair_swap_reference(x: jax.Array) -> jax.Array: + """Pure-JAX twin of `_pair_swap`, used to pin the rotation identity in tests.""" + axis = x.ndim - 1 + roll_right = jnp.roll(x, 1, axis=axis) + roll_left = jnp.roll(x, -1, axis=axis) + lane = jax.lax.broadcasted_iota(jnp.int32, x.shape, axis) + return jnp.where(lane % 2 == 0, roll_left, roll_right) diff --git a/src/maxdiffusion/models/attention_flax.py b/src/maxdiffusion/models/attention_flax.py index 9d52e3f98..c6b1c7e08 100644 --- a/src/maxdiffusion/models/attention_flax.py +++ b/src/maxdiffusion/models/attention_flax.py @@ -24,11 +24,18 @@ from jax.experimental import shard_map from jax.experimental.pallas.ops.tpu.splash_attention import splash_attention_mask from jax.experimental.pallas.ops.tpu.splash_attention import splash_attention_kernel +from maxdiffusion import wan_runtime_options from maxdiffusion.kernels.splash_attention import splash_attention_mask as tokamax_splash_attention_mask from maxdiffusion.kernels.splash_attention import splash_attention_kernel as tokamax_splash_attention_kernel from maxdiffusion.kernels.splash_attention import ring_attention_kernel as tokamax_ring_attention_kernel from maxdiffusion.kernels.splash_attention import base as tokamax_splash_base from maxdiffusion.kernels.fused_producers import fused_rmsnorm_rope +from maxdiffusion.kernels.fused_rmsnorm_rope_pallas import ( + fused_rmsnorm_rope_pallas, + resolve_rope_accum, + rope_accum_is_measured, + with_xla_backward, +) from einops import rearrange from .. import common_types, max_logging from maxdiffusion.tpu_utils import get_tpu_type, TpuType @@ -1017,6 +1024,8 @@ def _ulysses_attention( ulysses_shards: int = -1, kernel_name: str = "ulysses_custom", use_k_centering: bool = True, + qk_prescaled: bool = False, + transpose_out: bool = False, ) -> jax.Array: """Ulysses sequence-parallel attention. @@ -1030,6 +1039,7 @@ def _ulysses_attention( """ if kv_heads is None: kv_heads = heads + transpose_out = bool(transpose_out) if not use_custom_kernel and kv_heads != heads: raise NotImplementedError(f"{kernel_name} does not support GQA (got heads={heads}, kv_heads={kv_heads}).") axis_name = CONTEXT @@ -1082,7 +1092,7 @@ def wrap_ulysses_attention(query, key, value, attention_mask): # so this is bit-identical. Done after the a2a it sat between the collective # and the kernel and XLA wrapped it in relayout copies; done before, it fuses # into the producer of Q and its 185MB round-trip disappears. - if use_custom_kernel and use_base2_exp: + if use_custom_kernel and use_base2_exp and not qk_prescaled: query = query * LOG2E # Swap sharding: each device gives up a slice of heads and gathers # a slice of sequence, so the local kernel sees the full sequence. @@ -1153,7 +1163,7 @@ def wrap_ulysses_attention(query, key, value, attention_mask): k_mean=k_mean, # Use the unpadded V: `all_fixed` gates the whole kernel through a # lax.cond, so anything feeding it sits on the critical path. Reading - # the padded copy chained a 193MB pad + reduction behind the V + # the padded copy chained a 193.5MB (per batch element) pad + reduction behind the V # all-to-all and left that collective fully exposed. The padding is # zeros and the check is a max of squares, so this is output-invariant. value=raw_value, @@ -1177,6 +1187,7 @@ def wrap_ulysses_attention(query, key, value, attention_mask): vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=True, uniform_fixed_m=True, + transpose_out=transpose_out, ) splash_kernel_hybrid = custom_splash.make_splash_mha( block_sizes=bsizes, @@ -1188,6 +1199,7 @@ def wrap_ulysses_attention(query, key, value, attention_mask): vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=True, uniform_fixed_m=False, + transpose_out=transpose_out, ) def _run_uniform(q, k, v, m, km): @@ -1216,19 +1228,32 @@ def _run_hybrid(q, k, v, m, km): use_experimental_scheduler=use_experimental_scheduler, vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=False, + transpose_out=transpose_out, ) vmapped_splash = jax.vmap(splash_kernel, in_axes=(0, 0, 0)) attention_output = vmapped_splash(query, key, value) - attention_output = attention_output[:, :, :kv_size, :context_q_seq_len].astype(query.dtype) - # Restore original layout: head-sharded/full-sequence -> sequence-sharded/full-heads. - # Sequence axis is at index 3, heads axis is at index 1. - attention_output = jax.lax.all_to_all( - attention_output, - axis_name=axis_name, - split_axis=3, - concat_axis=1, - tiled=True, - ) + if transpose_out: + attention_output = attention_output[:, :, :context_q_seq_len, :kv_size].astype(query.dtype) + # Restore original layout: head-sharded/full-sequence -> sequence-sharded/full-heads. + # Sequence axis is at index 2 (sublanes), heads axis is at index 1. + attention_output = jax.lax.all_to_all( + attention_output, + axis_name=axis_name, + split_axis=2, + concat_axis=1, + tiled=True, + ) + else: + attention_output = attention_output[:, :, :kv_size, :context_q_seq_len].astype(query.dtype) + # Restore original layout: head-sharded/full-sequence -> sequence-sharded/full-heads. + # Sequence axis is at index 3, heads axis is at index 1. + attention_output = jax.lax.all_to_all( + attention_output, + axis_name=axis_name, + split_axis=3, + concat_axis=1, + tiled=True, + ) return attention_output else: # Run the same local splash kernel as standard TPU flash attention, but now @@ -1318,9 +1343,9 @@ def _run_hybrid(q, k, v, m, km): effective_num_heads = num_heads out_q_axis_names = ( - jax.sharding.PartitionSpec(q_axis_names[0], q_axis_names[1], q_axis_names[3], q_axis_names[2]) - if use_custom_kernel - else q_axis_names + q_axis_names + if (not use_custom_kernel or transpose_out) + else jax.sharding.PartitionSpec(q_axis_names[0], q_axis_names[1], q_axis_names[3], q_axis_names[2]) ) if attention_mask is None: @@ -1357,7 +1382,7 @@ def run_ulysses_attention(q, k, v): run_ulysses_attention, ) - if use_custom_kernel: + if use_custom_kernel and not transpose_out: if fold_batch: x = x.reshape(batch, num_heads, *x.shape[2:]) x = x[:, :, :, :orig_q_seq_len] @@ -1862,10 +1887,13 @@ def _ulysses_ring_custom_attention( per_q_block: bool = True, kv_heads: int | None = None, use_k_centering: bool = False, + qk_prescaled: bool = False, + transpose_out: bool = False, ) -> jax.Array: """2D USP attention (Ulysses + Ring) using custom splash kernel with exact Fixed-m support.""" if kv_heads is None: kv_heads = heads + transpose_out = bool(transpose_out) if attention_mask is not None: raise NotImplementedError("ulysses_ring_custom does not support attention_mask.") @@ -1941,7 +1969,7 @@ def wrap_ulysses_ring_attention(query, key, value): # so this is bit-identical. Done after the a2a it sat between the collective # and the kernel and XLA wrapped it in relayout copies; done before, it fuses # into the producer of Q and its 185MB round-trip disappears. - if use_base2_exp: + if use_base2_exp and not qk_prescaled: query = query * LOG2E # (0) R>1 fixed-m reductions and global eligibility predicates, computed @@ -2034,6 +2062,7 @@ def wrap_ulysses_ring_attention(query, key, value): vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=True, uniform_fixed_m=True, + transpose_out=transpose_out, ) splash_kernel_hybrid = custom_splash.make_splash_mha( block_sizes=bsizes, @@ -2045,6 +2074,7 @@ def wrap_ulysses_ring_attention(query, key, value): vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=True, uniform_fixed_m=False, + transpose_out=transpose_out, ) def _run_uniform(q, k, v, m, km): @@ -2064,9 +2094,13 @@ def _run_hybrid(q, k, v, m, km): use_experimental_scheduler=use_experimental_scheduler, vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=False, + transpose_out=transpose_out, ) raw_out = jax.vmap(splash_kernel, in_axes=(0, 0, 0))(query, key, value) - attention_output = jnp.swapaxes(raw_out, 2, 3) + if transpose_out: + attention_output = raw_out + else: + attention_output = jnp.swapaxes(raw_out, 2, 3) # (2b) Ring: Cross-chip ppermute schedule with custom ring kernel else: @@ -2139,6 +2173,9 @@ def _apply_attention_dot( use_memory_efficient_attention: bool, attention_mask: Array = None, kv_heads: int | None = None, + qk_prescaled: bool = False, + k_prescaled: bool = False, + use_base2_exp: bool = False, ): """Apply Attention.""" effective_kv_heads = kv_heads if kv_heads is not None else heads @@ -2185,9 +2222,14 @@ def _to_bshd(x: Array, n_heads: int) -> Array: if float32_qk_product: query_states = query_states.astype(jnp.float32) key_states = key_states.astype(jnp.float32) - query_states = query_states.transpose(1, 0, 2) - key_states = key_states.transpose(1, 0, 2) - value_states = value_states.transpose(1, 0, 2) + if k_prescaled or qk_prescaled: + key_states = key_states / jnp.asarray(scale, dtype=key_states.dtype) + if qk_prescaled and use_base2_exp: + query_states = query_states / jnp.asarray(LOG2E, dtype=query_states.dtype) + if not split_head_dim: + query_states = query_states.transpose(1, 0, 2) + key_states = key_states.transpose(1, 0, 2) + value_states = value_states.transpose(1, 0, 2) # this if statement create a chunk size for each layer of the unet # the chunk size is equal to the query_length dimension of the deepest layer of the unet @@ -2210,7 +2252,13 @@ def _to_bshd(x: Array, n_heads: int) -> Array: key_chunk_size=4096 * 4, ) - hidden_states = hidden_states.transpose(1, 0, 2) + if split_head_dim: + b = hidden_states.shape[0] + hidden_states = jnp.reshape(hidden_states, (b, -1, heads * dim_head)) + else: + hidden_states = hidden_states.transpose(1, 0, 2) + hidden_states = _reshape_batch_dim_to_heads(hidden_states, heads) + hidden_states = hidden_states.astype(dtype) else: preferred_element_type = jnp.float32 if float32_qk_product else None if split_head_dim: @@ -2228,7 +2276,13 @@ def _to_bshd(x: Array, n_heads: int) -> Array: preferred_element_type=preferred_element_type, ) - attention_scores = attention_scores * scale + if qk_prescaled: + if use_base2_exp: + attention_scores = attention_scores / jnp.asarray(LOG2E, dtype=attention_scores.dtype) + elif k_prescaled: + pass + elif scale != 1.0: + attention_scores = attention_scores * scale if attention_mask is not None: attention_scores = attention_scores + attention_mask.astype(attention_scores.dtype) attention_probs = nn.softmax(attention_scores, axis=-1 if split_head_dim else 2) @@ -2296,14 +2350,18 @@ def dot_product_kernel(q, k, v, context): context["use_memory_efficient_attention"], context["attention_mask"], kv_heads=context.get("kv_heads", None), + qk_prescaled=context.get("qk_prescaled", False), + k_prescaled=context.get("k_prescaled", False), + use_base2_exp=context.get("use_base2_exp", False), ) @register_kernel("ulysses_custom") def ulysses_custom_kernel(q, k, v, context): + qk_prescaled = context.get("qk_prescaled", False) return _ulysses_attention( q, - k * context["scale"], + k if (qk_prescaled or context.get("k_prescaled", False)) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2323,14 +2381,17 @@ def ulysses_custom_kernel(q, k, v, context): kv_heads=context.get("kv_heads", None), ulysses_shards=context.get("ulysses_shards", -1), kernel_name="ulysses_custom", + qk_prescaled=qk_prescaled, + transpose_out=bool(context.get("transpose_out")), ) @register_kernel("ulysses_ring_custom") def ulysses_ring_custom_kernel(q, k, v, context): + qk_prescaled = context.get("qk_prescaled", False) return _ulysses_ring_custom_attention( q, - k * context["scale"], + k if (qk_prescaled or context.get("k_prescaled", False)) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2346,22 +2407,18 @@ def ulysses_ring_custom_kernel(q, k, v, context): use_experimental_scheduler=context.get("use_experimental_scheduler", False), ulysses_attention_chunks=context["ulysses_attention_chunks"], kv_heads=context.get("kv_heads", None), + qk_prescaled=qk_prescaled, + transpose_out=bool(context.get("transpose_out")), ) @register_kernel("ulysses_ring_custom_fixed_m") def ulysses_ring_custom_fixed_m_kernel(q, k, v, context): - """Fixed-m variant of ulysses_ring_custom. - - Reduces squared Q/K row norms and V-magnitude safety predicates before the - Ulysses all-to-all (overlapping with the collective) and gathers K-shard - norms across the ring axis. When all heads pass the global Cauchy-Schwarz - bound, ring hops use a uniform fixed shift `m` and merge by direct FP32 - accumulation; otherwise each hop gates independently and merges in LSE space. - """ + """Fixed-m variant of ulysses_ring_custom with monolithic per-head gating.""" + qk_prescaled = context.get("qk_prescaled", False) return _ulysses_ring_custom_attention( q, - k * context["scale"], + k if (qk_prescaled or context.get("k_prescaled", False)) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2380,15 +2437,18 @@ def ulysses_ring_custom_fixed_m_kernel(q, k, v, context): ulysses_attention_chunks=context.get("ulysses_attention_chunks", 1), kv_heads=context.get("kv_heads", None), use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=True), + qk_prescaled=qk_prescaled, + transpose_out=bool(context.get("transpose_out")), ) @register_kernel("ulysses_ring_custom_fixed_m_per_q_block") def ulysses_ring_custom_fixed_m_per_q_block_kernel(q, k, v, context): """fixed-m variant of ulysses_ring_custom with per-Q-block gating.""" + qk_prescaled = context.get("qk_prescaled", False) return _ulysses_ring_custom_attention( q, - k * context["scale"], + k if (qk_prescaled or context.get("k_prescaled", False)) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2407,6 +2467,8 @@ def ulysses_ring_custom_fixed_m_per_q_block_kernel(q, k, v, context): ulysses_attention_chunks=context.get("ulysses_attention_chunks", 1), kv_heads=context.get("kv_heads", None), use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=True), + qk_prescaled=qk_prescaled, + transpose_out=bool(context.get("transpose_out")), ) @@ -2415,9 +2477,10 @@ def ulysses_ring_custom_bidir_kernel(q, k, v, context): """Wrap-free (bidirectional) variant of ulysses_ring_custom: the ring streams K/V both directions one hop at a time, avoiding the diameter-length wrap hop on a non-wrapping ring axis. Same USP split as ulysses_ring_custom otherwise.""" + qk_prescaled = context.get("qk_prescaled", False) return _ulysses_ring_custom_attention( q, - k * context["scale"], + k if (qk_prescaled or context.get("k_prescaled", False)) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2434,14 +2497,17 @@ def ulysses_ring_custom_bidir_kernel(q, k, v, context): bidirectional=True, ulysses_attention_chunks=context["ulysses_attention_chunks"], kv_heads=context.get("kv_heads", None), + qk_prescaled=qk_prescaled, + transpose_out=bool(context.get("transpose_out")), ) @register_kernel("ulysses_custom_fixed_m") def ulysses_custom_fixed_m_kernel(q, k, v, context): + qk_prescaled = context.get("qk_prescaled", False) return _ulysses_attention( q, - k * context["scale"], + k if (qk_prescaled or context.get("k_prescaled", False)) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2462,14 +2528,17 @@ def ulysses_custom_fixed_m_kernel(q, k, v, context): ulysses_shards=context.get("ulysses_shards", -1), kernel_name="ulysses_custom_fixed_m", use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=False), + qk_prescaled=qk_prescaled, + transpose_out=bool(context.get("transpose_out")), ) @register_kernel("ulysses_custom_fixed_m_per_q_block") def ulysses_custom_fixed_m_per_q_block_kernel(q, k, v, context): + qk_prescaled = context.get("qk_prescaled", False) return _ulysses_attention( q, - k * context["scale"], + k if (qk_prescaled or context.get("k_prescaled", False)) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2490,6 +2559,8 @@ def ulysses_custom_fixed_m_per_q_block_kernel(q, k, v, context): ulysses_shards=context.get("ulysses_shards", -1), kernel_name="ulysses_custom_fixed_m_per_q_block", use_k_centering=resolve_k_centering(context.get("use_k_centering"), ring=False), + qk_prescaled=qk_prescaled, + transpose_out=bool(context.get("transpose_out")), ) @@ -2497,7 +2568,7 @@ def ulysses_custom_fixed_m_per_q_block_kernel(q, k, v, context): def ulysses_kernel(q, k, v, context): return _ulysses_attention( q, - k * context["scale"], + k if context.get("k_prescaled", False) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2520,7 +2591,7 @@ def ulysses_kernel(q, k, v, context): def ulysses_ring_kernel(q, k, v, context): return _ulysses_ring_attention( q, - k * context["scale"], + k if context.get("k_prescaled", False) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2544,7 +2615,7 @@ def ulysses_ring_kernel(q, k, v, context): def flash_kernel(q, k, v, context): return _tpu_flash_attention( q, - k * context["scale"], + k if context.get("k_prescaled", False) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2568,7 +2639,7 @@ def flash_kernel(q, k, v, context): def tokamax_flash_kernel(q, k, v, context): return _tpu_flash_attention( q, - k * context["scale"], + k if context.get("k_prescaled", False) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2594,7 +2665,7 @@ def tokamax_flash_kernel(q, k, v, context): def tokamax_ring_kernel(q, k, v, context): return _tpu_flash_attention( q, - k * context["scale"], + k if context.get("k_prescaled", False) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2618,7 +2689,7 @@ def tokamax_ring_kernel(q, k, v, context): def tokamax_ring_custom_kernel(q, k, v, context): return _tpu_flash_attention( q, - k * context["scale"], + k if context.get("k_prescaled", False) else k * context["scale"], v, context["heads"], context["mesh"], @@ -2638,9 +2709,23 @@ def tokamax_ring_custom_kernel(q, k, v, context): @register_kernel("cudnn_flash_te") def cudnn_flash_te_kernel(q, k, v, context): + if context.get("k_prescaled", False): + k = k / jnp.asarray(context["scale"], dtype=k.dtype) return _cudnn_flash_attention(q, k, v, context["heads"], context["mesh"], context["dpa_layer"]) +_QK_PRESCALE_SUPPORTED_KERNELS = { + "dot_product", + "ulysses_custom", + "ulysses_ring_custom", + "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", + "ulysses_ring_custom_bidir", + "ulysses_custom_fixed_m", + "ulysses_custom_fixed_m_per_q_block", +} + + def _apply_attention( query: Array, key: Array, @@ -2672,6 +2757,9 @@ def _apply_attention( spatiotemporal_shape: Optional[Tuple[int, int, int]] = None, kv_heads: Optional[int] = None, use_k_centering: bool | str = "auto", + qk_prescaled: bool = False, + k_prescaled: bool = False, + transpose_out: bool = False, ): """Routes to different attention kernels using a module-level registry.""" @@ -2704,6 +2792,12 @@ def _apply_attention( if attention_kernel == "dot_product" or use_memory_efficient_attention or not can_use_flash_attention: effective_attention_kernel = "dot_product" + if qk_prescaled and effective_attention_kernel not in _QK_PRESCALE_SUPPORTED_KERNELS: + raise ValueError( + f"qk_prescaled=True is only supported by {sorted(_QK_PRESCALE_SUPPORTED_KERNELS)}, " + f"got {effective_attention_kernel!r}." + ) + # Masks enter the dispatcher as canonical [B, K] keep masks. Adapt them # only after fallback selection because a configured flash kernel may use # dot-product attention for short sequences. @@ -2748,6 +2842,9 @@ def _apply_attention( "spatiotemporal_config": spatiotemporal_config, "spatiotemporal_shape": spatiotemporal_shape, "use_k_centering": use_k_centering, + "qk_prescaled": qk_prescaled, + "k_prescaled": k_prescaled, + "transpose_out": transpose_out, } if spatiotemporal_config and spatiotemporal_config.get("use_svg_attention"): @@ -2774,6 +2871,11 @@ def _apply_attention( def _head_local_svg_attention(query, key, value, context): from .wan.transformers import svg_attention, svg_head_local + if context.get("qk_prescaled", False) or context.get("k_prescaled", False): + raise ValueError( + "Head-local SVG does not support prescaled Q/K (`qk_prescaled` / `k_prescaled`) " + "because `svg_profile_temporal_heads` applies `scale` during routing." + ) cfg = context["spatiotemporal_config"] grid = context["spatiotemporal_shape"] mesh = context["mesh"] @@ -3102,6 +3204,7 @@ def __init__( ulysses_attention_chunks: int = 1, kv_heads: Optional[int] = None, use_k_centering: bool | str = "auto", + transpose_out: bool = False, ): self.dpa_layer = None self.use_base2_exp = use_base2_exp @@ -3109,6 +3212,7 @@ def __init__( self.ulysses_shards = ulysses_shards self.ulysses_attention_chunks = ulysses_attention_chunks self.use_k_centering = use_k_centering + self.transpose_out = bool(transpose_out) if attention_kernel == "cudnn_flash_te": from transformer_engine.jax.flax.transformer import DotProductAttention # pytype: disable=import-error @@ -3157,6 +3261,8 @@ def apply_attention( preserve_asymmetric_block_sizes: bool = False, spatiotemporal_shape: Optional[Tuple[int, int, int]] = None, sparse_config_override: Optional[dict] = None, + qk_prescaled: bool = False, + k_prescaled: bool = False, ): return _apply_attention( query=query, @@ -3188,6 +3294,9 @@ def apply_attention( spatiotemporal_shape=spatiotemporal_shape, kv_heads=self.kv_heads, use_k_centering=getattr(self, "use_k_centering", "auto"), + qk_prescaled=qk_prescaled, + k_prescaled=k_prescaled, + transpose_out=getattr(self, "transpose_out", False), ) @@ -3213,6 +3322,7 @@ class AttentionOp(nn.Module): is_causal: bool = False kv_heads: Optional[int] = None use_k_centering: bool | str = "auto" + transpose_out: bool = False def setup(self): self.dpa_layer = None @@ -3247,6 +3357,8 @@ def apply_attention( preserve_asymmetric_block_sizes: bool = False, spatiotemporal_shape: Optional[Tuple[int, int, int]] = None, sparse_config_override: Optional[dict] = None, + qk_prescaled: bool = False, + k_prescaled: bool = False, ): return _apply_attention( query=query, @@ -3277,7 +3389,70 @@ def apply_attention( spatiotemporal_shape=spatiotemporal_shape, kv_heads=self.kv_heads, use_k_centering=getattr(self, "use_k_centering", "auto"), + qk_prescaled=qk_prescaled, + k_prescaled=k_prescaled, + transpose_out=getattr(self, "transpose_out", False), + ) + + +@functools.lru_cache(maxsize=64) +def _build_sharded_fused_rope_producer( + kernel_fn: Callable, + mesh: jax.sharding.Mesh, + act_spec: jax.sharding.PartitionSpec, + replicated_spec: jax.sharding.PartitionSpec, + freqs_spec: jax.sharding.PartitionSpec, + out_spec: jax.sharding.PartitionSpec, + q_heads: int, + dim_head: int, + eps: float, + k_eps: Optional[float], + norm_mode: str, + rope_accum: str, + block_s: Optional[int], + head_block: Optional[int], + q_prescale: float, + k_prescale: float, +): + """Builds and caches the `with_xla_backward(jax.shard_map(...))` RoPE producer.""" + sharded = jax.shard_map( + functools.partial( + kernel_fn, + q_heads=q_heads, + dim_head=dim_head, + eps=eps, + k_eps=k_eps, + norm_mode=norm_mode, + # Under jit, XLA on v6e contracts the reference's RoPE multiply-adds + # into FP32 FMAs ("f32"), while XLA on tpu7x emits native bf16 + # multiply-adds ("dtype"). Matching the platform's contraction + # yields 0-ULP bit-identical output on both v6e and tpu7x. + rope_accum=rope_accum, + block_s=block_s, + head_block=head_block, + q_prescale=q_prescale, + k_prescale=k_prescale, + mesh=mesh, + ), + mesh=mesh, + in_specs=(act_spec, act_spec, replicated_spec, replicated_spec, freqs_spec), + out_specs=(out_spec, out_spec), + check_vma=False, + ) + + def xla_equivalent(q, k, q_scale, k_scale, freqs): + """The same producer in unfused primitives, prescaling included.""" + out_q, out_k = fused_rmsnorm_rope( + q, k, q_scale, k_scale, freqs, q_heads=q_heads, dim_head=dim_head, eps=eps, k_eps=k_eps ) + if q_prescale != 1.0: + out_q = out_q * jnp.asarray(q_prescale, out_q.dtype) + if k_prescale != 1.0: + out_k = out_k * jnp.asarray(k_prescale, out_k.dtype) + return out_q, out_k + + # `pallas_call` has no transpose rule; forward stays the kernel. + return with_xla_backward(sharded, xla_equivalent) class FlaxWanAttention(nnx.Module): @@ -3341,6 +3516,19 @@ def __init__( "svg_low_noise_density": -1.0, "svg_flash_block_sizes": None, "use_k_centering": "auto", + # Fused RMSNorm+RoPE+head-transpose Pallas producer. Off by default: + # 0-ULP against the separately-jitted XLA producer on measured platforms, + # but end-to-end output is equivalent, not identical (XLA fuses the + # unfused producer differently). Enabled only where + # rope_accum_is_measured() holds. + "use_fused_rope_kernel": False, + "fused_rope_block_s": 1024, + "fused_rope_head_block": None, + "wan_rope_norm_mode": None, + "wan_fuse_qk_prescale": None, + "wan_splash_transpose_out": None, + "wan_cross_attn_prescale_kv": None, + "wan_rope_accum": None, **(attention_config or {}), } @@ -3365,6 +3553,26 @@ def __init__( self.svg_low_noise_density = attention_config["svg_low_noise_density"] self.svg_flash_block_sizes = attention_config["svg_flash_block_sizes"] self.is_self_attention = is_self_attention + # Assigned before the check below, which references it: without this the + # validation raises AttributeError instead of the intended message. + self.mesh = mesh + + self.use_fused_rope_kernel = attention_config["use_fused_rope_kernel"] + self.fused_rope_block_s = attention_config["fused_rope_block_s"] + self.fused_rope_head_block = attention_config["fused_rope_head_block"] + + # Graph-changing Wan switches. Resolved here, at build time, so they live + # on the GraphDef (and in the AOT cache key); a missing entry means the + # built-in default, never process-global state read at trace time. + def _wan_opt(name): + value = attention_config[name] + return wan_runtime_options.default(name) if value is None else wan_runtime_options.coerce(name, value) + + self.wan_rope_norm_mode = _wan_opt("wan_rope_norm_mode") + self.wan_fuse_qk_prescale = _wan_opt("wan_fuse_qk_prescale") + self.wan_splash_transpose_out = _wan_opt("wan_splash_transpose_out") + self.wan_cross_attn_prescale_kv = _wan_opt("wan_cross_attn_prescale_kv") + self.wan_rope_accum = _wan_opt("wan_rope_accum") if attention_kernel in {"flash", "cudnn_flash_te"} and mesh is None: raise ValueError(f"The flash attention kernel requires a value for mesh, but mesh is {self.mesh}") @@ -3437,6 +3645,7 @@ def __init__( ulysses_shards=attention_config["ulysses_shards"], ulysses_attention_chunks=attention_config["ulysses_attention_chunks"], use_k_centering=attention_config["use_k_centering"], + transpose_out=self.wan_splash_transpose_out, ) # None axes corresponds to the stacked weights across all blocks # because of the use of nnx.vmap and nnx.scan. @@ -3568,6 +3777,12 @@ def __init__( ("norm",), ), ) + _prescale_unsupported_kernels = {"cudnn_flash_te"} + self.cross_attn_prescale_kv = bool( + self.wan_cross_attn_prescale_kv + and getattr(self.attention_op, "attention_kernel", "") not in _prescale_unsupported_kernels + and not getattr(self.attention_op, "use_memory_efficient_attention", False) + ) def _apply_rope(self, xq: jax.Array, xk: jax.Array, freqs_cis: jax.Array) -> Tuple[jax.Array, jax.Array]: # 1. Extract cos and sin, keeping them in native bfloat16 @@ -3595,6 +3810,174 @@ def _apply_rope(self, xq: jax.Array, xk: jax.Array, freqs_cis: jax.Array) -> Tup return xq_out, xk_out + def _fused_rope_producer(self): + """Selects the RMSNorm+RoPE+transpose producer, falling back when unsafe. + + The Pallas kernel is a pure fusion of `fused_rmsnorm_rope`: 1.9x-2.1x faster + in isolation on v6e at the Wan shard shape (40 heads, per-shard seq 9450 / + 18900), and 0 ULP against the *separately jitted* XLA + producer there. That equality does not survive inlining -- inside the + 40-layer graph XLA fuses the unfused producer with its neighbours and rounds + it differently (v6e 720p/81f, same seed: 53.8 dB PSNR after 1 denoise step, + 34.1 dB after 40). Equivalent, not identical; hence opt-in, on measured + hardware only. + + Always returns a callable returning `((q_out, k_out), qk_prescaled: bool)`. + Dispatches to the Pallas producer when every static precondition below + holds, or to `fused_rmsnorm_rope` with `qk_prescaled=False` otherwise: + * the mesh devices are TPUs (the Pallas kernel uses Mosaic TPU primitives) + and `rope_accum_is_measured(self.mesh)` holds; + * a mesh is available to wrap the custom call in `shard_map` -- a raw + `pallas_call` on sharded operands would make GSPMD all-gather them; + * the feature axis is unsharded. RMSNorm reduces across `heads * dim_head` + and RoPE pairs lanes within a head, so a sharded feature axis would + silently produce a per-shard norm instead of the true one; + * `dim_head` is a whole number of lanes, which the kernel requires; + * at call time inside the returned wrapper, the sequence dimension must + divide evenly by its mesh axis (otherwise it falls back to + `fused_rmsnorm_rope` with `qk_prescaled=False`); an indivisible batch + dimension is treated as replicated. + """ + + def _xla_fallback(q, k, q_scale, k_scale, freqs, *, q_heads, dim_head, eps, k_eps=None): + return ( + fused_rmsnorm_rope(q, k, q_scale, k_scale, freqs, q_heads=q_heads, dim_head=dim_head, eps=eps, k_eps=k_eps), + False, + ) + + is_tpu = self.mesh is not None and all(getattr(d, "platform", None) == "tpu" for d in self.mesh.devices.flat) + # Matching the XLA producer's rounding is a per-platform property: on v4 the + # kernel drifts by up to one bf16 ULP. Unverified hardware gets XLA. + rounding_verified = is_tpu and rope_accum_is_measured(self.mesh) + if ( + not self.use_fused_rope_kernel + or not is_tpu + or self.mesh is None + or self.dim_head % 128 != 0 + or not rounding_verified + ): + if self.use_fused_rope_kernel: + kinds = sorted({getattr(d, "device_kind", "?") for d in self.mesh.devices.flat}) if self.mesh is not None else [] + _warn_once( + "fused_rope_kernel_unusable", + f"fused RoPE kernel requested but unusable (is_tpu={is_tpu}, mesh={self.mesh is not None}, " + f"dim_head={self.dim_head}, rounding_verified={rounding_verified}, device_kind={kinds}); " + "using the XLA producer.", + ) + return _xla_fallback + + in_spec = nn.logical_to_mesh_axes((BATCH, LENGTH, HEAD)) + feature_axis = in_spec[2] + if feature_axis is not None: + axes = (feature_axis,) if isinstance(feature_axis, str) else tuple(feature_axis) + if any(self.mesh.shape[a] > 1 for a in axes): + _warn_once( + "fused_rope_kernel_sharded_feature", + f"fused RoPE kernel disabled: the feature axis is sharded over {axes}, which would turn the RMSNorm " + "reduction into a per-shard norm; using the XLA producer.", + ) + return _xla_fallback + + replicated = jax.sharding.PartitionSpec() + block_s = self.fused_rope_block_s + head_block = self.fused_rope_head_block + mesh = self.mesh + + def _shards_over(axis) -> int: + """Total mesh width a single PartitionSpec entry splits a dimension by.""" + if axis is None: + return 1 + names = (axis,) if isinstance(axis, str) else tuple(axis) + return math.prod(mesh.shape[n] for n in names) + + batch_axis, seq_axis = in_spec[0], in_spec[1] + batch_shards = _shards_over(batch_axis) + seq_shards = _shards_over(seq_axis) + + def producer(q, k, q_scale, k_scale, freqs, *, q_heads, dim_head, eps, k_eps=None): + # `shard_map` demands exact divisibility on every sharded dimension, + # whereas the surrounding GSPMD program pads uneven splits. Wan runs a + # global batch of 1 over a multi-way data axis, which GSPMD degenerates + # into replication anyway -- so declare it replicated here rather than + # giving up the kernel over a dimension no device actually splits. + local_batch_axis = batch_axis if q.shape[0] % batch_shards == 0 else None + divides = ( + q.shape[1] % seq_shards == 0 + and k.shape[1] % seq_shards == 0 + # Each shard rotates its own contiguous window of positions, which + # only lines up if q, k and the table are split the same way. + and freqs.shape[2] == q.shape[1] + and q.shape[1] == k.shape[1] + ) + if not divides: + _warn_once( + "fused_rope_kernel_indivisible", + f"fused RoPE kernel disabled: q={q.shape} k={k.shape} freqs={freqs.shape} are not all splittable " + f"{seq_shards} ways on the sequence axis; using the XLA producer.", + ) + return _xla_fallback(q, k, q_scale, k_scale, freqs, q_heads=q_heads, dim_head=dim_head, eps=eps, k_eps=k_eps) + + act_spec = jax.sharding.PartitionSpec(local_batch_axis, seq_axis, in_spec[2]) + # [B, S, H*D] -> [B, H, S, D]: the head axis inherits the feature axis' + # sharding (unsharded, per the guard above) and `dim_head` is never split. + out_spec = jax.sharding.PartitionSpec(local_batch_axis, in_spec[2], seq_axis, None) + # `freqs_cis` is `[1, 1, S, dim_head // 2]` and reaches this point + # replicated, but each shard owns a contiguous window of positions. + # Splitting it on the same axis as the activations is what GSPMD does + # implicitly for the reference's elementwise multiply; leaving it + # replicated would make every shard rotate by positions [0, S_local), + # which is right only on the first shard. + freqs_spec = jax.sharding.PartitionSpec(None, None, seq_axis, None) + + attn_kernel = getattr(self.attention_op, "attention_kernel", "dot_product") + min_seq_len = getattr(self.attention_op, "flash_min_seq_length", 4096) + use_mem_eff = getattr(self.attention_op, "use_memory_efficient_attention", False) + supported_custom_kernels = { + "ulysses_custom", + "ulysses_ring_custom", + "ulysses_ring_custom_fixed_m", + "ulysses_ring_custom_fixed_m_per_q_block", + "ulysses_ring_custom_bidir", + "ulysses_custom_fixed_m", + "ulysses_custom_fixed_m_per_q_block", + } + can_prescale = attn_kernel in supported_custom_kernels and not use_mem_eff and q.shape[1] >= min_seq_len + norm_mode_env = self.wan_rope_norm_mode + fuse_qk_prescale = bool(self.wan_fuse_qk_prescale and can_prescale) + q_prescale_val = float(LOG2E) if (fuse_qk_prescale and getattr(self.attention_op, "use_base2_exp", True)) else 1.0 + k_prescale_val = float(self.attention_op.scale) if fuse_qk_prescale else 1.0 + _warn_once( + "fused_rope_kernel_active", + f"fused RoPE Pallas kernel ACTIVE: q={q.shape}, per-shard seq={q.shape[1] // seq_shards}, " + f"batch_spec={local_batch_axis}, block_s={block_s}, head_block={head_block or q_heads}, " + f"norm_mode={norm_mode_env}, q_prescale={q_prescale_val}, k_prescale={k_prescale_val}.", + ) + # Resolved through the shared helper so this and the AOT cache + # fingerprint in generate_wan.py cannot disagree about which mode was + # compiled (see `resolve_rope_accum`). + rope_accum_env = resolve_rope_accum(mesh, self.wan_rope_accum) + differentiable = _build_sharded_fused_rope_producer( + fused_rmsnorm_rope_pallas, + mesh, + act_spec, + replicated, + freqs_spec, + out_spec, + q_heads, + dim_head, + eps, + k_eps, + norm_mode_env, + rope_accum_env, + block_s, + head_block, + q_prescale_val, + k_prescale_val, + ) + return differentiable(q, k, q_scale, k_scale, freqs), fuse_qk_prescale + + return producer + def conditional_named_scope(self, name: str): """Return a JAX named scope if enabled, otherwise a null context.""" return jax.named_scope(name) if self.enable_jax_named_scopes else contextlib.nullcontext() @@ -3641,6 +4024,7 @@ def __call__( with jax.named_scope("query_proj"): query_proj = self.query(hidden_states) + k_prescaled = False if is_self_attention: with jax.named_scope("key_proj"): key_proj = self.key(hidden_states) @@ -3648,29 +4032,48 @@ def __call__( value_proj = self.value(hidden_states) elif cached_kv is not None and "text" in cached_kv: key_proj, value_proj = cached_kv["text"] + k_prescaled = self.cross_attn_prescale_kv else: with jax.named_scope("key_proj"): key_proj = self.key(encoder_hidden_states) with jax.named_scope("value_proj"): value_proj = self.value(encoder_hidden_states) + qk_prescaled = False if rotary_emb is not None and self.qk_norm and is_self_attention: with self.conditional_named_scope("fused_rmsnorm_rope"): q_scale = self.norm_q.scale[...] k_scale = self.norm_k.scale[...] q_eps = getattr(self.norm_q, "epsilon", self.eps) k_eps = getattr(self.norm_k, "epsilon", self.eps) - query_proj, key_proj = fused_rmsnorm_rope( - query_proj, - key_proj, - q_scale, - k_scale, - rotary_emb, - q_heads=self.heads, - dim_head=self.dim_head, - eps=q_eps, - k_eps=k_eps, - ) + # The SVG sparse path does not understand prescaled Q/K, so SVG layers + # keep the XLA fused_rmsnorm_rope producer rather than the Pallas kernel. + if self.use_svg_attention: + query_proj, key_proj = fused_rmsnorm_rope( + query_proj, + key_proj, + q_scale, + k_scale, + rotary_emb, + q_heads=self.heads, + dim_head=self.dim_head, + eps=q_eps, + k_eps=k_eps, + ) + qk_prescaled = False + else: + producer = self._fused_rope_producer() + (query_proj, key_proj), qk_prescaled = producer( + query_proj, + key_proj, + q_scale, + k_scale, + rotary_emb, + q_heads=self.heads, + dim_head=self.dim_head, + eps=q_eps, + k_eps=k_eps, + ) value_proj = _unflatten_heads(value_proj, self.heads) else: if self.qk_norm: @@ -3715,6 +4118,8 @@ def run_dense(_): key_proj, value_proj, attention_mask=encoder_attention_mask, + qk_prescaled=qk_prescaled, + k_prescaled=k_prescaled, ) def run_sparse_svg(_): @@ -3758,6 +4163,8 @@ def run_sparse_svg(_): key_proj, value_proj, attention_mask=encoder_attention_mask, + qk_prescaled=qk_prescaled, + k_prescaled=k_prescaled, ) else: @@ -3799,6 +4206,7 @@ def run_sparse_svg(_): # Text K/V if cached_kv is not None and "text" in cached_kv: key_proj_text, value_proj_text = cached_kv["text"] + k_prescaled_text = self.cross_attn_prescale_kv else: with self.conditional_named_scope("proj_key"): key_proj_text = self.key(encoder_hidden_states_text) @@ -3807,11 +4215,13 @@ def run_sparse_svg(_): key_proj_text = self.norm_k(key_proj_text) with self.conditional_named_scope("proj_value"): value_proj_text = self.value(encoder_hidden_states_text) + k_prescaled_text = False # Image K/V (only if image embeddings are present) if encoder_hidden_states_img is not None: if cached_kv is not None and "image" in cached_kv: key_proj_img, value_proj_img = cached_kv["image"] + k_prescaled_img = self.cross_attn_prescale_kv else: with self.conditional_named_scope("add_proj_k"): key_proj_img = self.add_k_proj(encoder_hidden_states_img) @@ -3819,6 +4229,7 @@ def run_sparse_svg(_): key_proj_img = self.norm_added_k(key_proj_img) with self.conditional_named_scope("add_proj_v"): value_proj_img = self.add_v_proj(encoder_hidden_states_img) + k_prescaled_img = False query_proj_img = query_proj_raw # Check norm_added_k too # Checkpointing @@ -3831,7 +4242,9 @@ def run_sparse_svg(_): # Attention - tensors are (B, S, D) with self.conditional_named_scope("cross_attn_text_apply"): - attn_output_text = self.attention_op.apply_attention(query_proj_text, key_proj_text, value_proj_text) + attn_output_text = self.attention_op.apply_attention( + query_proj_text, key_proj_text, value_proj_text, k_prescaled=k_prescaled_text + ) with self.conditional_named_scope("cross_attn_img_apply"): # Pass encoder_attention_mask_img for image cross-attention to mask padded tokens attn_output_img = self.attention_op.apply_attention( @@ -3839,6 +4252,7 @@ def run_sparse_svg(_): key_proj_img, value_proj_img, attention_mask=encoder_attention_mask_img, + k_prescaled=k_prescaled_img, ) attn_output = attn_output_text + attn_output_img @@ -3849,7 +4263,9 @@ def run_sparse_svg(_): value_proj_text = checkpoint_name(value_proj_text, "value_proj_text") with self.conditional_named_scope("cross_attn_text_apply"): - attn_output = self.attention_op.apply_attention(query_proj_text, key_proj_text, value_proj_text) + attn_output = self.attention_op.apply_attention( + query_proj_text, key_proj_text, value_proj_text, k_prescaled=k_prescaled_text + ) attn_output = attn_output.astype(dtype=dtype) attn_output = checkpoint_name(attn_output, "attn_output") @@ -3876,6 +4292,8 @@ def compute_kv( if self.qk_norm: with self.conditional_named_scope("attn_k_norm"): key_proj = self.norm_k(key_proj) + if self.cross_attn_prescale_kv: + key_proj = key_proj * jnp.asarray(self.attention_op.scale, key_proj.dtype) return {"text": (key_proj, value_proj)} else: @@ -3899,6 +4317,8 @@ def compute_kv( if self.qk_norm: with self.conditional_named_scope("attn_k_norm"): key_proj_text = self.norm_k(key_proj_text) + if self.cross_attn_prescale_kv: + key_proj_text = key_proj_text * jnp.asarray(self.attention_op.scale, key_proj_text.dtype) with self.conditional_named_scope("proj_value"): value_proj_text = self.value(encoder_hidden_states_text) @@ -3908,6 +4328,8 @@ def compute_kv( key_proj_img = self.add_k_proj(encoder_hidden_states_img) with self.conditional_named_scope("norm_add_k"): key_proj_img = self.norm_added_k(key_proj_img) + if self.cross_attn_prescale_kv: + key_proj_img = key_proj_img * jnp.asarray(self.attention_op.scale, key_proj_img.dtype) with self.conditional_named_scope("add_proj_v"): value_proj_img = self.add_v_proj(encoder_hidden_states_img) diff --git a/src/maxdiffusion/models/wan/transformers/transformer_wan.py b/src/maxdiffusion/models/wan/transformers/transformer_wan.py index fccf45c6f..db083ae0c 100644 --- a/src/maxdiffusion/models/wan/transformers/transformer_wan.py +++ b/src/maxdiffusion/models/wan/transformers/transformer_wan.py @@ -24,6 +24,7 @@ import flax.linen as nn import numpy as np from .... import common_types +from .... import wan_runtime_options from ...modeling_flax_utils import FlaxModelMixin, get_activation from ....configuration_utils import ConfigMixin, register_to_config from ...embeddings_flax import ( @@ -354,6 +355,7 @@ def __init__( mask_padding_tokens: bool = True, enable_jax_named_scopes: bool = False, attention_config: Optional[dict] = None, + split_head_dim: bool = False, ): self.enable_jax_named_scopes = enable_jax_named_scopes attention_config = { @@ -375,6 +377,7 @@ def __init__( qk_norm=qk_norm, eps=eps, flash_min_seq_length=flash_min_seq_length, + split_head_dim=split_head_dim, flash_block_sizes=flash_block_sizes, mesh=mesh, dtype=dtype, @@ -400,6 +403,7 @@ def __init__( added_kv_proj_dim=added_kv_proj_dim, image_seq_len=image_seq_len, flash_min_seq_length=flash_min_seq_length, + split_head_dim=split_head_dim, flash_block_sizes=flash_block_sizes, mesh=mesh, dtype=dtype, @@ -583,12 +587,19 @@ def __init__( scan_layers: bool = True, enable_jax_named_scopes: bool = False, attention_config: Optional[dict] = None, + split_head_dim: bool = False, + wan_cfg_before_unpatchify: Optional[bool] = None, ): inner_dim = num_attention_heads * attention_head_dim out_channels = out_channels or in_channels self.num_layers = num_layers self.scan_layers = scan_layers self.enable_jax_named_scopes = enable_jax_named_scopes + self.wan_cfg_before_unpatchify = ( + wan_runtime_options.default("wan_cfg_before_unpatchify") + if wan_cfg_before_unpatchify is None + else wan_runtime_options.coerce("wan_cfg_before_unpatchify", wan_cfg_before_unpatchify) + ) attention_config = { "use_base2_exp": False, "use_experimental_scheduler": False, @@ -657,6 +668,7 @@ def init_block(rngs): added_kv_proj_dim=added_kv_proj_dim, image_seq_len=image_seq_len, attention_config=attention_config, + split_head_dim=split_head_dim, ) self.gradient_checkpoint = GradientCheckpointType.from_str(remat_policy) @@ -686,6 +698,7 @@ def init_block(rngs): attention=attention, enable_jax_named_scopes=enable_jax_named_scopes, attention_config=attention_config, + split_head_dim=split_head_dim, ) blocks.append(block) self.blocks = nnx.data(blocks) @@ -783,6 +796,7 @@ def __call__( rotary_emb: Optional[jax.Array] = None, encoder_attention_mask: Optional[jax.Array] = None, svg_step_index: Optional[int | jax.Array] = None, + unpatchify: bool = True, ) -> Union[jax.Array, Tuple[jax.Array, jax.Array], Dict[str, jax.Array]]: hidden_states = nn.with_logical_constraint(hidden_states, ("batch", None, None, None, None)) batch_size, _, num_frames, height, width = hidden_states.shape @@ -960,6 +974,24 @@ def layer_forward(hidden_states, l_kv): with jax.named_scope("proj_out"): hidden_states = self.proj_out(hidden_states) + if not unpatchify: + if return_residual: + return hidden_states, residual_x + return hidden_states + + hidden_states = self.unpatchify_tokens(hidden_states, num_frames, height, width) + + if return_residual: + return hidden_states, residual_x + return hidden_states + + def unpatchify_tokens(self, hidden_states: jax.Array, num_frames: int, height: int, width: int) -> jax.Array: + batch_size = hidden_states.shape[0] + p_t, p_h, p_w = self.config.patch_size + post_patch_num_frames = num_frames // p_t + post_patch_height = height // p_h + post_patch_width = width // p_w + if p_t == 1: # Lossless HLO optimization: collapse p_t=1 dimension to avoid 8D non-contiguous stride copies hidden_states = hidden_states.reshape( @@ -972,7 +1004,7 @@ def layer_forward(hidden_states, l_kv): -1, ) hidden_states = jnp.transpose(hidden_states, (0, 6, 1, 2, 4, 3, 5)) - hidden_states = hidden_states.reshape(batch_size, -1, num_frames, height, width) + return hidden_states.reshape(batch_size, -1, num_frames, height, width) else: hidden_states = hidden_states.reshape( batch_size, @@ -985,8 +1017,4 @@ def layer_forward(hidden_states, l_kv): -1, ) hidden_states = jnp.transpose(hidden_states, (0, 7, 1, 4, 2, 5, 3, 6)) - hidden_states = hidden_states.reshape(batch_size, -1, num_frames, height, width) - - if return_residual: - return hidden_states, residual_x - return hidden_states + return hidden_states.reshape(batch_size, -1, num_frames, height, width) diff --git a/src/maxdiffusion/pipelines/wan/wan_pipeline.py b/src/maxdiffusion/pipelines/wan/wan_pipeline.py index 7b7a1ce36..82d165cfe 100644 --- a/src/maxdiffusion/pipelines/wan/wan_pipeline.py +++ b/src/maxdiffusion/pipelines/wan/wan_pipeline.py @@ -15,6 +15,7 @@ from abc import abstractmethod from typing import Any, List, Union, Optional, Tuple from functools import partial +from maxdiffusion import wan_runtime_options from maxdiffusion.image_processor import PipelineImageInput import numpy as np import math @@ -369,6 +370,9 @@ def create_model(rngs: nnx.Rngs, wan_config: dict): use_svg_for_expert = bool(getattr(config, "use_svg_attention", False)) and (expert_density < 1.0) + wan_config["split_head_dim"] = getattr(config, "split_head_dim", False) + wan_config["wan_cfg_before_unpatchify"] = wan_runtime_options.resolve_from_config(config, "wan_cfg_before_unpatchify") + fused_rope_head_block = getattr(config, "fused_rope_head_block", -1) wan_config["attention_config"] = { "use_base2_exp": config.use_base2_exp, "use_experimental_scheduler": config.use_experimental_scheduler, @@ -395,6 +399,11 @@ def create_model(rngs: nnx.Rngs, wan_config: dict): "svg_low_noise_density": low_density, "svg_flash_block_sizes": getattr(config, "svg_flash_block_sizes", None) or None, "use_k_centering": getattr(config, "use_k_centering", "auto"), + "use_fused_rope_kernel": getattr(config, "use_fused_rope_kernel", False), + "fused_rope_block_s": getattr(config, "fused_rope_block_s", 1024), + # -1 is the YAML spelling of "let the kernel pick" (i.e. all heads). + "fused_rope_head_block": None if fused_rope_head_block in (None, -1) else fused_rope_head_block, + **wan_runtime_options.attention_config_entries(config), } # 2. eval_shape - will not use flops or create weights on device @@ -1503,31 +1512,69 @@ def transformer_forward_pass( if do_classifier_free_guidance and latents.shape[0] != prompt_embeds.shape[0]: latents = jnp.concatenate([latents, latents], axis=0) wan_transformer = nnx.merge(graphdef, sharded_state, rest_of_state) - outputs = wan_transformer( - hidden_states=latents, - timestep=timestep, - encoder_hidden_states=prompt_embeds, - encoder_hidden_states_image=encoder_hidden_states_image, - skip_blocks=skip_blocks, - cached_residual=cached_residual, - return_residual=return_residual, - kv_cache=kv_cache, - rotary_emb=rotary_emb, - encoder_attention_mask=encoder_attention_mask, - svg_step_index=svg_step_index, + # Set on the model at build time (WanModel resolves None to the default), so + # it is part of the GraphDef; never read from process-global state here. + wan_cfg_before_unpatchify = bool( + getattr( + wan_transformer, + "wan_cfg_before_unpatchify", + wan_runtime_options.default("wan_cfg_before_unpatchify"), + ) ) - if return_residual: - noise_pred, residual_x = outputs - else: - noise_pred = outputs + if do_classifier_free_guidance and wan_cfg_before_unpatchify: + outputs = wan_transformer( + hidden_states=latents, + timestep=timestep, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_image=encoder_hidden_states_image, + skip_blocks=skip_blocks, + cached_residual=cached_residual, + return_residual=return_residual, + kv_cache=kv_cache, + rotary_emb=rotary_emb, + encoder_attention_mask=encoder_attention_mask, + svg_step_index=svg_step_index, + unpatchify=False, + ) + if return_residual: + noise_pred, residual_x = outputs + else: + noise_pred = outputs - if do_classifier_free_guidance: bsz = latents.shape[0] // 2 - noise_cond = noise_pred[:bsz] # First half = conditional - noise_uncond = noise_pred[bsz:] # Second half = unconditional + noise_cond = noise_pred[:bsz] + noise_uncond = noise_pred[bsz:] noise_pred = noise_uncond + guidance_scale * (noise_cond - noise_uncond) + _, _, num_frames, height, width = latents.shape + noise_pred = wan_transformer.unpatchify_tokens(noise_pred, num_frames, height, width) + else: + outputs = wan_transformer( + hidden_states=latents, + timestep=timestep, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_image=encoder_hidden_states_image, + skip_blocks=skip_blocks, + cached_residual=cached_residual, + return_residual=return_residual, + kv_cache=kv_cache, + rotary_emb=rotary_emb, + encoder_attention_mask=encoder_attention_mask, + svg_step_index=svg_step_index, + ) + + if return_residual: + noise_pred, residual_x = outputs + else: + noise_pred = outputs + + if do_classifier_free_guidance: + bsz = latents.shape[0] // 2 + noise_cond = noise_pred[:bsz] # First half = conditional + noise_uncond = noise_pred[bsz:] # Second half = unconditional + noise_pred = noise_uncond + guidance_scale * (noise_cond - noise_uncond) + if return_residual: return noise_pred, residual_x return noise_pred diff --git a/src/maxdiffusion/pipelines/wan/wan_pipeline_2_1.py b/src/maxdiffusion/pipelines/wan/wan_pipeline_2_1.py index 0cd0d7f7b..d6ceb5df0 100644 --- a/src/maxdiffusion/pipelines/wan/wan_pipeline_2_1.py +++ b/src/maxdiffusion/pipelines/wan/wan_pipeline_2_1.py @@ -475,7 +475,7 @@ def scan_body(carry, t): svg_step_index=jnp.asarray(step, dtype=jnp.int32), ) - elif do_cfg: + elif do_cfg and use_cfg_cache: latents_doubled = jnp.concatenate([latents] * 2) timestep = jnp.broadcast_to(t, bsz * 2) ( @@ -496,6 +496,23 @@ def scan_body(carry, t): svg_step_index=jnp.asarray(step, dtype=jnp.int32), ) + elif do_cfg: + timestep = jnp.broadcast_to(t, bsz * 2) + noise_pred = transformer_forward_pass( + graphdef, + sharded_state, + rest_of_state, + latents, + timestep, + prompt_embeds_combined, + do_classifier_free_guidance=True, + guidance_scale=guidance_scale, + kv_cache=kv_cache, + rotary_emb=rotary_emb, + encoder_attention_mask=encoder_attention_mask, + svg_step_index=jnp.asarray(step, dtype=jnp.int32), + ) + else: timestep = jnp.broadcast_to(t, bsz) noise_pred = transformer_forward_pass( diff --git a/src/maxdiffusion/pipelines/wan/wan_pipeline_animate.py b/src/maxdiffusion/pipelines/wan/wan_pipeline_animate.py index 6a926f9ea..c917dd355 100644 --- a/src/maxdiffusion/pipelines/wan/wan_pipeline_animate.py +++ b/src/maxdiffusion/pipelines/wan/wan_pipeline_animate.py @@ -40,7 +40,7 @@ from flax import nnx from flax.linen import partitioning as nn_partitioning from jax.sharding import NamedSharding, PartitionSpec as P -from maxdiffusion import max_logging +from maxdiffusion import max_logging, wan_runtime_options from maxdiffusion.image_processor import PipelineImageInput, VaeImageProcessor from maxdiffusion.max_utils import get_flash_block_sizes, get_precision from maxdiffusion.video_processor import VideoProcessor @@ -96,6 +96,7 @@ def _create_model(rngs: nnx.Rngs, wan_config: dict): "use_experimental_scheduler": config.use_experimental_scheduler, "ulysses_shards": getattr(config, "ulysses_shards", -1), "ulysses_attention_chunks": getattr(config, "ulysses_attention_chunks", 1), + **wan_runtime_options.attention_config_entries(config), } # 2. eval_shape – creates the model structure without allocating HBM. diff --git a/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py b/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py index 80ab6dbb3..7cca3dfcd 100644 --- a/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py +++ b/src/maxdiffusion/pipelines/wan/wan_vace_pipeline_2_1.py @@ -24,6 +24,7 @@ from ...pyconfig import HyperParameters from ... import aot_cache from ... import max_logging +from ... import wan_runtime_options from ...image_processor import PipelineImageInput from ...max_utils import get_flash_block_sizes, get_precision from ...models.wan.wan_utils import load_wan_transformer @@ -91,6 +92,7 @@ def create_model(rngs: nnx.Rngs, wan_config: dict): "use_experimental_scheduler": config.use_experimental_scheduler, "ulysses_shards": getattr(config, "ulysses_shards", -1), "ulysses_attention_chunks": getattr(config, "ulysses_attention_chunks", 1), + **wan_runtime_options.attention_config_entries(config), } wan_config["scan_layers"] = False diff --git a/src/maxdiffusion/pyconfig.py b/src/maxdiffusion/pyconfig.py index c995e7f4e..1eb9662af 100644 --- a/src/maxdiffusion/pyconfig.py +++ b/src/maxdiffusion/pyconfig.py @@ -331,6 +331,15 @@ def user_init(raw_keys): if "wan_debug_cond_timers" not in raw_keys: raw_keys["wan_debug_cond_timers"] = False + from maxdiffusion import wan_runtime_options # pylint: disable=import-outside-toplevel + + # Validate/coerce the Wan switches (e.g. "false" from the command line). + # They are read from the config by the Wan pipelines at build time; there + # is no process-wide store to configure here. + for name in wan_runtime_options.names(): + if name in raw_keys and raw_keys[name] is not None: + raw_keys[name] = wan_runtime_options.coerce(name, raw_keys[name]) + def get_num_slices(raw_keys): if int(raw_keys["compile_topology_num_slices"]) > 0: diff --git a/src/maxdiffusion/tests/aot_cache_test.py b/src/maxdiffusion/tests/aot_cache_test.py index 09327c803..065185bd3 100644 --- a/src/maxdiffusion/tests/aot_cache_test.py +++ b/src/maxdiffusion/tests/aot_cache_test.py @@ -166,8 +166,8 @@ def test_signature_deterministic_across_processes(self): "os.environ['JAX_PLATFORMS'] = 'cpu'", "import jax", "import jax.numpy as jnp", - "from flax import nnx", "from maxdiffusion import aot_cache", + "from flax import nnx", "", "class T(nnx.Module):", " def __init__(self, rngs):", @@ -381,6 +381,131 @@ def fake_generate(config, pipe, prefix, writer, load_time, persistent_dir): self.assertEqual(calls, [("", ephemeral)]) self.assertEqual(out, None if fail else ["out.mp4"]) + def test_wan_aot_metadata_includes_runtime_switches(self): + """Runtime optimization switches (via config or env) must change the Wan AOT metadata fingerprint.""" + import os + import types + from unittest import mock + from maxdiffusion import generate_wan + + cfg = types.SimpleNamespace(attention="ulysses_ring_custom_fixed_m") + base_meta = generate_wan._build_wan_aot_metadata(cfg, self._mesh, "rev1") + base_fp = aot_cache._metadata_fingerprint(base_meta) + + switches = [ + ("WAN_ROPE_NORM_MODE", "wan_rope_norm_mode", "fused"), + ("WAN_FUSE_QK_PRESCALE", "wan_fuse_qk_prescale", "0"), + ("WAN_SPLASH_TRANSPOSE_OUT", "wan_splash_transpose_out", "1"), + ("WAN_CFG_BEFORE_UNPATCHIFY", "wan_cfg_before_unpatchify", "0"), + ("WAN_CROSS_ATTN_PRESCALE_KV", "wan_cross_attn_prescale_kv", "1"), + ] + for env_var, attr_name, new_val in switches: + with mock.patch.dict(os.environ, {env_var: new_val}): + switched_meta = generate_wan._build_wan_aot_metadata(cfg, self._mesh, "rev1") + switched_fp = aot_cache._metadata_fingerprint(switched_meta) + self.assertNotEqual( + base_fp, + switched_fp, + f"Changing {env_var} to {new_val} did not invalidate the Wan AOT fingerprint.", + ) + cfg_explicit = types.SimpleNamespace(attention="ulysses_ring_custom_fixed_m", **{attr_name: new_val}) + explicit_meta = generate_wan._build_wan_aot_metadata(cfg_explicit, self._mesh, "rev1") + self.assertNotEqual( + base_fp, + aot_cache._metadata_fingerprint(explicit_meta), + f"Setting config.{attr_name}={new_val} did not invalidate the Wan AOT fingerprint.", + ) + + def test_wan_runtime_options_on_graphdef_change_dynamic_signature(self): + """FlaxWanAttention stores wan_* runtime options on graphdef so _dynamic_signature changes per module config.""" + from flax import nnx + from maxdiffusion.models.attention_flax import FlaxWanAttention + + attn_exact = FlaxWanAttention( + rngs=nnx.Rngs(0), + query_dim=64, + heads=2, + dim_head=32, + attention_kernel="dot_product", + attention_config={"wan_rope_norm_mode": "exact", "wan_fuse_qk_prescale": True}, + ) + attn_fused = FlaxWanAttention( + rngs=nnx.Rngs(0), + query_dim=64, + heads=2, + dim_head=32, + attention_kernel="dot_product", + attention_config={"wan_rope_norm_mode": "fused", "wan_fuse_qk_prescale": True}, + ) + gd_exact, _ = nnx.split(attn_exact) + gd_fused, _ = nnx.split(attn_fused) + self.assertNotEqual( + aot_cache._dynamic_signature((gd_exact, self._a), {}), + aot_cache._dynamic_signature((gd_fused, self._a), {}), + ) + + def test_wan_aot_metadata_includes_resolved_rope_accum(self): + """WAN_ROPE_ACCUM changes the lowered graph's rounding and must key the executable. + + This is deliberately NOT folded into the generic switches list above. The + fingerprint records the *resolved* mode, whose default is + platform-dependent ("dtype" on tpu7x and CPU, "f32" on v6e and other TPUs). Asserting on a + fixed literal would be vacuous on whichever platform already defaults to + it -- e.g. WAN_ROPE_ACCUM="f32" is a no-op on v6e. So the override is + chosen to be the opposite of whatever this platform resolves to. + """ + import os + import types + from unittest import mock + from maxdiffusion import generate_wan + from maxdiffusion.kernels.fused_rmsnorm_rope_pallas import resolve_rope_accum + + cfg = types.SimpleNamespace(attention="ulysses_ring_custom_fixed_m") + + with mock.patch.dict(os.environ, {}, clear=False): + os.environ.pop("WAN_ROPE_ACCUM", None) + default_mode = resolve_rope_accum(self._mesh) + base_meta = generate_wan._build_wan_aot_metadata(cfg, self._mesh, "rev1") + + self.assertIn(default_mode, ("dtype", "f32")) + self.assertEqual(base_meta["wan_rope_accum"], default_mode) + + other_mode = "f32" if default_mode == "dtype" else "dtype" + with mock.patch.dict(os.environ, {"WAN_ROPE_ACCUM": other_mode}): + switched_meta = generate_wan._build_wan_aot_metadata(cfg, self._mesh, "rev1") + self.assertEqual(switched_meta["wan_rope_accum"], other_mode) + self.assertNotEqual( + aot_cache._metadata_fingerprint(base_meta), + aot_cache._metadata_fingerprint(switched_meta), + f"WAN_ROPE_ACCUM={other_mode} (platform default {default_mode}) did not invalidate the Wan AOT fingerprint.", + ) + + # An explicit override equal to the platform default describes the same + # executable and must NOT force a recompile. + with mock.patch.dict(os.environ, {"WAN_ROPE_ACCUM": default_mode}): + same_meta = generate_wan._build_wan_aot_metadata(cfg, self._mesh, "rev1") + self.assertEqual( + aot_cache._metadata_fingerprint(base_meta), + aot_cache._metadata_fingerprint(same_meta), + f"WAN_ROPE_ACCUM={default_mode} matches the platform default and must reuse the cache.", + ) + + def test_resolve_rope_accum_rejects_unknown_mode(self): + """An unrecognised rope-accum mode must fail loudly rather than silently defaulting.""" + import os + import types + from unittest import mock + from maxdiffusion import wan_runtime_options + from maxdiffusion.kernels.fused_rmsnorm_rope_pallas import resolve_rope_accum + + with mock.patch.dict(os.environ, {"WAN_ROPE_ACCUM": "fp32"}): + with self.assertRaises(ValueError): + wan_runtime_options.resolve_from_config(types.SimpleNamespace(), "wan_rope_accum") + with self.assertRaises(ValueError): + wan_runtime_options.resolve_from_config(types.SimpleNamespace(wan_rope_accum="fp32"), "wan_rope_accum") + with self.assertRaises(ValueError): + resolve_rope_accum(self._mesh, rope_accum="fp32") + def test_plan_rejects_gcs_and_uses_the_real_revision_resolver(self): """gs:// raises whatever the revision; only an explicitly dirty/unversioned revision downgrades.""" import types diff --git a/src/maxdiffusion/tests/dot_fallback_layout_test.py b/src/maxdiffusion/tests/dot_fallback_layout_test.py index c79e71a83..d4951ae27 100644 --- a/src/maxdiffusion/tests/dot_fallback_layout_test.py +++ b/src/maxdiffusion/tests/dot_fallback_layout_test.py @@ -28,6 +28,7 @@ """ import math +import os import unittest import jax @@ -36,6 +37,13 @@ from maxdiffusion.models.attention_flax import _apply_attention_dot +IN_GITHUB_ACTIONS = os.getenv("GITHUB_ACTIONS") == "true" +# Kernel numerics grids and multi-device tests are skipped in CI; see +# end_to_end/tpu/run_wan_stack_tests.sh. +_SKIP_IN_GITHUB_ACTIONS = unittest.skipIf( + IN_GITHUB_ACTIONS, "TPU kernel / multi-device test, skipped in GitHub Actions; run end_to_end/tpu/run_wan_stack_tests.sh" +) + def _call_dot(query, key, value, heads, dim_head, kv_heads=None): """Runs the fallback and returns [B, H, S, D]. @@ -71,6 +79,15 @@ def _reference_attention(query, key, value, scale): class DotFallbackLayoutTest(unittest.TestCase): """`_apply_attention_dot` must transpose, not reshape, 4-D inputs.""" + def setUp(self): + super().setUp() + self._matmul_precision_ctx = jax.default_matmul_precision("float32") + self._matmul_precision_ctx.__enter__() + + def tearDown(self): + self._matmul_precision_ctx.__exit__(None, None, None) + super().tearDown() + def test_zero_logits_return_per_head_token_means(self): """The reviewer's counterexample, reproduced exactly. @@ -145,6 +162,287 @@ def test_gqa_repeat_still_applies_on_4d(self): ) np.testing.assert_allclose(out, expected, rtol=1e-4, atol=1e-4) + def test_qk_prescaled_with_base2_matches_unscaled(self): + """When Q is scaled by log2(e) and K is scaled by scale, dot attention with + qk_prescaled=True and use_base2_exp=True must match to within 1e-5.""" + heads, seq, dim_head = 4, 8, 16 + shape = (2, heads, seq, dim_head) + query = jax.random.normal(jax.random.PRNGKey(10), shape, jnp.float32) + key = jax.random.normal(jax.random.PRNGKey(11), shape, jnp.float32) + value = jax.random.normal(jax.random.PRNGKey(12), shape, jnp.float32) + scale = 1.0 / math.sqrt(dim_head) + log2e = math.log2(math.e) + + with jax.default_matmul_precision("float32"): + out_unscaled = _apply_attention_dot( + query=query, + key=key, + value=value, + dtype=jnp.float32, + heads=heads, + dim_head=dim_head, + scale=scale, + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=False, + qk_prescaled=False, + use_base2_exp=False, + ) + + q_prescaled = query * log2e + k_prescaled = key * scale + out_prescaled = _apply_attention_dot( + query=q_prescaled, + key=k_prescaled, + value=value, + dtype=jnp.float32, + heads=heads, + dim_head=dim_head, + scale=scale, + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=False, + qk_prescaled=True, + use_base2_exp=True, + ) + np.testing.assert_allclose(np.asarray(out_prescaled), np.asarray(out_unscaled), rtol=1e-5, atol=1e-5) + + def test_k_prescaled_matches_unscaled(self): + """When only K is prescaled by scale, dot attention with k_prescaled=True + must match to within 1e-5 without double-scaling.""" + heads, seq, dim_head = 4, 8, 16 + shape = (2, heads, seq, dim_head) + query = jax.random.normal(jax.random.PRNGKey(13), shape, jnp.float32) + key = jax.random.normal(jax.random.PRNGKey(14), shape, jnp.float32) + value = jax.random.normal(jax.random.PRNGKey(15), shape, jnp.float32) + scale = 1.0 / math.sqrt(dim_head) + + out_unscaled = _apply_attention_dot( + query=query, + key=key, + value=value, + dtype=jnp.float32, + heads=heads, + dim_head=dim_head, + scale=scale, + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=False, + qk_prescaled=False, + k_prescaled=False, + ) + + k_prescaled_arr = key * scale + out_k_prescaled = _apply_attention_dot( + query=query, + key=k_prescaled_arr, + value=value, + dtype=jnp.float32, + heads=heads, + dim_head=dim_head, + scale=scale, + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=False, + qk_prescaled=False, + k_prescaled=True, + ) + np.testing.assert_allclose(np.asarray(out_k_prescaled), np.asarray(out_unscaled), rtol=1e-5, atol=1e-5) + + def test_uncached_unaffected_by_cross_attn_prescale_kv(self): + """wan_cross_attn_prescale_kv=True prescales cached cross-attn K in compute_kv, without affecting uncached or cached outputs.""" + from flax import nnx + from maxdiffusion.models.attention_flax import FlaxWanAttention + + heads, seq_q, seq_kv, dim_head = 4, 8, 6, 16 + dim = heads * dim_head + hidden_states = jax.random.normal(jax.random.PRNGKey(16), (2, seq_q, dim), jnp.float32) + encoder_hidden_states = jax.random.normal(jax.random.PRNGKey(17), (2, seq_kv, dim), jnp.float32) + scale = 1.0 / math.sqrt(dim_head) + + def build(prescale_kv): + return FlaxWanAttention( + rngs=nnx.Rngs(0), + query_dim=dim, + heads=heads, + dim_head=dim_head, + attention_kernel="dot_product", + is_self_attention=False, + dtype=jnp.float32, + weights_dtype=jnp.float32, + attention_config={"wan_cross_attn_prescale_kv": prescale_kv}, + ) + + attn0 = build(False) + self.assertFalse(attn0.cross_attn_prescale_kv) + out_uncached_0 = attn0(hidden_states, encoder_hidden_states) + kv0 = attn0.compute_kv(encoder_hidden_states) + out_cached_0 = attn0(hidden_states, cached_kv=kv0) + + attn1 = build(True) + self.assertTrue(attn1.cross_attn_prescale_kv) + out_uncached_1 = attn1(hidden_states, encoder_hidden_states) + kv1 = attn1.compute_kv(encoder_hidden_states) + out_cached_1 = attn1(hidden_states, cached_kv=kv1) + + # Cached K in kv1["text"][0] is prescaled by scale compared to kv0["text"][0] + np.testing.assert_allclose(np.asarray(kv1["text"][0]), np.asarray(kv0["text"][0] * scale), rtol=1e-6, atol=1e-6) + # Uncached outputs are identical, and cached outputs with prescaled K match unprescaled within float32 precision + np.testing.assert_allclose(np.asarray(out_uncached_1), np.asarray(out_uncached_0), rtol=1e-6, atol=1e-6) + np.testing.assert_allclose(np.asarray(out_cached_1), np.asarray(out_cached_0), rtol=1e-5, atol=1e-5) + np.testing.assert_allclose(np.asarray(out_cached_0), np.asarray(out_uncached_0), rtol=1e-5, atol=1e-5) + + def test_memory_efficient_attention_with_prescaled_k(self): + """When use_memory_efficient_attention=True, k_prescaled=True must not double-scale.""" + heads, seq, dim_head = 4, 8, 16 + shape = (2, heads, seq, dim_head) + query = jax.random.normal(jax.random.PRNGKey(19), shape, jnp.float32) + key = jax.random.normal(jax.random.PRNGKey(20), shape, jnp.float32) + value = jax.random.normal(jax.random.PRNGKey(21), shape, jnp.float32) + scale = 1.0 / math.sqrt(dim_head) + + out_unscaled = _apply_attention_dot( + query=query, + key=key, + value=value, + dtype=jnp.float32, + heads=heads, + dim_head=dim_head, + scale=scale, + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=True, + qk_prescaled=False, + k_prescaled=False, + ) + out_prescaled = _apply_attention_dot( + query=query, + key=key * scale, + value=value, + dtype=jnp.float32, + heads=heads, + dim_head=dim_head, + scale=scale, + split_head_dim=True, + float32_qk_product=True, + use_memory_efficient_attention=True, + qk_prescaled=False, + k_prescaled=True, + ) + np.testing.assert_allclose(np.asarray(out_prescaled), np.asarray(out_unscaled), rtol=1e-5, atol=1e-5) + + def test_cfg_before_unpatchify_is_bit_identical(self): + """Applying CFG on packed tokens before unpatchify_tokens is 0-ULP bit-identical to CFG after unpatchify_tokens.""" + import types + from maxdiffusion.models.wan.transformers.transformer_wan import WanModel + + batch, frames, height, width = 2, 3, 4, 4 + p_t, p_h, p_w = 1, 2, 2 + out_dim = 16 + tokens = (frames // p_t) * (height // p_h) * (width // p_w) + patch_dim = p_t * p_h * p_w * out_dim + + pred_cond = jax.random.normal(jax.random.PRNGKey(30), (batch, tokens, patch_dim), jnp.bfloat16) + pred_uncond = jax.random.normal(jax.random.PRNGKey(31), (batch, tokens, patch_dim), jnp.bfloat16) + guidance_scale = 4.0 + + dummy = types.SimpleNamespace(config=types.SimpleNamespace(patch_size=(p_t, p_h, p_w))) + + cfg_tokens = pred_uncond + guidance_scale * (pred_cond - pred_uncond) + before = WanModel.unpatchify_tokens(dummy, cfg_tokens, frames, height, width) + + cond_vol = WanModel.unpatchify_tokens(dummy, pred_cond, frames, height, width) + uncond_vol = WanModel.unpatchify_tokens(dummy, pred_uncond, frames, height, width) + after = uncond_vol + guidance_scale * (cond_vol - uncond_vol) + + np.testing.assert_array_equal(np.asarray(before, np.float32), np.asarray(after, np.float32)) + + +def _bf16_cfg_reroute_may_round_differently() -> bool: + """True on TPUs where the bf16 CFG reroute is not bit-identical (bit-identical on CPU and v6e).""" + kinds = [d.device_kind.lower() for d in jax.devices() if d.platform == "tpu"] + return bool(kinds) and not all("v6" in k for k in kinds) + + +@_SKIP_IN_GITHUB_ACTIONS +class CfgBeforeUnpatchifyCompiledParityTest(unittest.TestCase): + """wan_cfg_before_unpatchify and the Wan 2.1 no-cache CFG reroute, compiled through a real WanModel. + + The eager test above only checks unpatchify_tokens. Here both paths go + through `transformer_forward_pass` under jit (with the production mesh and + logical axis rules), so XLA is free to fuse the CFG combine with the + unpatchify transpose. The result must still be bit-identical, for p_t == 1 + (the collapsed reshape) and p_t == 2, in f32 everywhere and in bf16 on CPU + and v6e; in bf16 on tpu7x it is within one bf16 ULP. The reroute is also + checked against `transformer_forward_pass_full_cfg`, the path Wan 2.1 no-cache + CFG used before. + """ + + def test_compiled_paths_are_bit_identical(self): + import os + from flax import nnx + from flax.linen import partitioning as nn_partitioning + from jax.sharding import Mesh + from maxdiffusion import pyconfig + from maxdiffusion.max_utils import create_device_mesh + from maxdiffusion.models.wan.transformers.transformer_wan import WanModel + from maxdiffusion.pipelines.wan import wan_pipeline + + pyconfig.initialize( + [None, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "configs", "base_wan_14b.yml")], + unittest=True, + ) + config = pyconfig.config + mesh = Mesh(create_device_mesh(config), config.mesh_axes) + + def forward(patch_size, dtype, cfg_before_unpatchify): + model = WanModel( + rngs=nnx.Rngs(0), + patch_size=patch_size, + num_attention_heads=2, + attention_head_dim=16, + in_channels=4, + out_channels=4, + text_dim=32, + freq_dim=32, + ffn_dim=64, + num_layers=2, + mesh=mesh, + dtype=dtype, + weights_dtype=dtype, + scan_layers=False, + wan_cfg_before_unpatchify=cfg_before_unpatchify, + ) + graphdef, state, rest = nnx.split(model, nnx.Param, ...) + latents = jax.random.normal(jax.random.PRNGKey(1), (1, 4, 4, 8, 8), dtype) + embeds = jax.random.normal(jax.random.PRNGKey(2), (2, 8, 32), dtype) + timestep = jnp.full((2,), 500.0) + out = wan_pipeline.transformer_forward_pass( + graphdef, state, rest, latents, timestep, embeds, do_classifier_free_guidance=True, guidance_scale=5.0 + ) + full_cfg, _, _ = wan_pipeline.transformer_forward_pass_full_cfg( + graphdef, state, rest, jnp.concatenate([latents] * 2), timestep, embeds, guidance_scale=5.0 + ) + return np.asarray(out, np.float32), np.asarray(full_cfg, np.float32) + + with mesh, nn_partitioning.axis_rules(config.logical_axis_rules): + for patch_size in ((1, 2, 2), (2, 2, 2)): + for dtype in (jnp.float32, jnp.bfloat16): + with self.subTest(patch_size=patch_size, dtype=jnp.dtype(dtype).name): + before, full_cfg = forward(patch_size, dtype, True) + after, _ = forward(patch_size, dtype, False) + self.assertEqual(before.shape, (1, 4, 4, 8, 8)) + if dtype == jnp.float32 or not _bf16_cfg_reroute_may_round_differently(): + np.testing.assert_array_equal(before, after) + np.testing.assert_array_equal(before, full_cfg) + else: + # tpu7x fuses the bf16 CFG combine differently on the two paths, so + # it is rounded at a different point: ~30% of elements differ by + # one bf16 ULP of the output scale (measured on tpu7x-8). + atol = 2**-7 * float(np.max(np.abs(before))) + np.testing.assert_allclose(before, after, rtol=2**-7, atol=atol) + np.testing.assert_allclose(before, full_cfg, rtol=2**-7, atol=atol) + class DispatcherLayoutContractTest(unittest.TestCase): """The threshold check and the dot path must agree on where S lives.""" diff --git a/src/maxdiffusion/tests/fused_rmsnorm_rope_pallas_test.py b/src/maxdiffusion/tests/fused_rmsnorm_rope_pallas_test.py new file mode 100644 index 000000000..709b51c74 --- /dev/null +++ b/src/maxdiffusion/tests/fused_rmsnorm_rope_pallas_test.py @@ -0,0 +1,1679 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Correctness tests for the fused RMSNorm + RoPE + head-transpose Pallas kernel. + +The kernel is a *pure fusion* of `fused_producers.fused_rmsnorm_rope`, so every +test here pins numerics rather than performance. Tests that only exercise the +algebra run anywhere via Pallas interpret mode; tests that must exercise the +Mosaic lowering (`pltpu.roll`, the dynamic lane-tile slice) are skipped off TPU. +""" + +import functools +import os +import unittest +from unittest import mock + +import jax +import jax.numpy as jnp +import numpy as np +from absl.testing import parameterized + +from maxdiffusion.kernels.fused_producers import fused_rmsnorm_rope +from maxdiffusion.kernels.fused_rmsnorm_rope_pallas import ( + DEFAULT_VMEM_LIMIT_BYTES, + ROPE_ACCUM_MODES, + _pair_swap, + _resolve_vmem_limit_bytes, + _rope_tables, + fused_rmsnorm_rope_pallas, + resolve_rope_accum, + rope_accum_is_measured, + vmem_limit_is_validated, + rope_pair_swap_reference, + with_xla_backward, +) +from flax import nnx + +IN_GITHUB_ACTIONS = os.getenv("GITHUB_ACTIONS") == "true" +# Kernel numerics grids and multi-device tests are skipped in CI; see +# end_to_end/tpu/run_wan_stack_tests.sh. +_SKIP_IN_GITHUB_ACTIONS = unittest.skipIf( + IN_GITHUB_ACTIONS, "TPU kernel / multi-device test, skipped in GitHub Actions; run end_to_end/tpu/run_wan_stack_tests.sh" +) + +DIM_HEAD = 128 + +# Bit-exactness against the XLA producer is a per-platform property, not a +# universal one: it holds only where an accumulation mode has been measured to +# reproduce that hardware's rounding of the RoPE combine. That has been done on +# tpu7x-8 (rope_accum="dtype"), v6e-8 ("f32") and XLA:CPU ("dtype"), which +# covers both platforms Wan actually serves on. +# +# It is not true on the v4-8 CI runner, where *no* mode is bit-identical. The +# measured grid there, at seq=50, is: +# +# 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 (see +# `test_exact_mode_normalisation_alone_is_bit_identical`, which passes on v4), +# 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 simply rounds +# `a*cos + b*sin` in a way neither mode models. Asserting 0 ULP there would be +# asserting something untrue, so these tests skip instead; `_rounding_report` +# prints the grid above if a *measured* platform ever regresses. +_UNMEASURED_ROUNDING_SKIP = ( + "Bit-exact parity is only asserted where the platform's RoPE rounding has been measured " + "(tpu7x, v6e, CPU). See rope_accum_is_measured()." +) + + +def _on_tpu() -> bool: + return jax.devices()[0].platform == "tpu" + + +def _on_validated_tpu() -> bool: + """On a TPU generation the kernel's 64 MiB VMEM budget was validated on (v6e, tpu7x).""" + return _on_tpu() and vmem_limit_is_validated() + + +_UNVALIDATED_VMEM_SKIP = "Production-shape kernel runs need a TPU with a validated VMEM budget (v6e, tpu7x)." + + +def _assert_bit_identical(test, name, ref, got): + """Asserts `ref` and `got` are the same bit pattern, without a host FP32 copy. + + At the production shape each output is ~774M elements, so the usual + `np.asarray(x, np.float32)` idiom pulls ~3 GB per array onto the host. The + comparison is done on-device over the raw storage bits instead, which is both + cheaper and a more direct statement of "0 ULP": two bf16 values are equal iff + their bit patterns are. The expensive diagnostics are computed only on the + failure path. + """ + test.assertEqual(ref.dtype, got.dtype, f"{name}: dtype differs ({ref.dtype} vs {got.dtype}).") + test.assertTrue(bool(jnp.all(jnp.isfinite(ref))), f"{name}: reference output contains non-finite values.") + test.assertTrue(bool(jnp.all(jnp.isfinite(got))), f"{name}: kernel output contains non-finite values (NaN/Inf).") + + bits = jnp.uint16 if ref.dtype.itemsize == 2 else jnp.uint32 + mismatches = int(jnp.count_nonzero(jax.lax.bitcast_convert_type(ref, bits) != jax.lax.bitcast_convert_type(got, bits))) + if mismatches: + max_diff = float(jnp.max(jnp.abs(ref.astype(jnp.float32) - got.astype(jnp.float32)))) + test.fail( + f"{name}: {mismatches}/{ref.size} elements are not bit-identical " + f"({100.0 * mismatches / ref.size:.6f}%, max |diff| = {max_diff:.3e})." + ) + + +def _make_inputs(seed, batch, seq, q_heads, kv_heads, dim_head, dtype): + d_q = q_heads * dim_head + d_k = kv_heads * dim_head + k0, k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(seed), 5) + raw_q = jax.random.normal(k0, (batch, seq, d_q), jnp.float32).astype(dtype) + raw_k = jax.random.normal(k1, (batch, seq, d_k), jnp.float32).astype(dtype) + q_scale = jax.random.normal(k2, (d_q,), jnp.float32) + k_scale = jax.random.normal(k3, (d_k,), jnp.float32) + # A genuine unit-modulus rotation, matching how WanRotaryPosEmbed builds + # freqs_cis; a degenerate table would hide sign/ordering bugs. + angle = jax.random.uniform(k4, (1, 1, seq, dim_head // 2), jnp.float32, minval=-np.pi, maxval=np.pi) + freqs_cis = jnp.cos(angle) + 1j * jnp.sin(angle) + return raw_q, raw_k, q_scale, k_scale, freqs_cis + + +@functools.partial(jax.jit, static_argnames=("q_heads", "kv_heads", "dim_head")) +def _compiled_reference(raw_q, raw_k, q_scale, k_scale, freqs_cis, *, q_heads, kv_heads=None, dim_head): + """The XLA reference producer, compiled -- which is the only fair comparand. + + `fused_rmsnorm_rope` called eagerly and the same function under `jit` need + not be the same computation: eager dispatch evaluates each op separately, + while under `jit` XLA may contract the RoPE multiply-add into an FP32 FMA + that rounds once (as on `v6e`; `tpu7x` and `XLA:CPU` emit uncontracted ops). + The Pallas kernel is always compiled, so an eager reference would make every + parity test a measurement of Python dispatch. Production calls this producer + from inside a jitted graph, so the compiled form is also the one that + actually ships. + + `fused_rmsnorm_rope` is itself `jax.jit`-wrapped (so eager callers get the + fused HBM footprint), which makes the two rows of `_rounding_report` + coincide today; this wrapper is kept so the parity tests state the compiled + comparand explicitly rather than relying on that implementation detail. + """ + return fused_rmsnorm_rope(raw_q, raw_k, q_scale, k_scale, freqs_cis, q_heads=q_heads, kv_heads=kv_heads, dim_head=dim_head) + + +@functools.partial( + jax.jit, + static_argnames=("q_heads", "kv_heads", "dim_head", "norm_mode", "rope_accum", "block_s", "head_block", "interpret"), +) +def _compiled_kernel( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + *, + q_heads, + kv_heads=None, + dim_head, + norm_mode, + rope_accum, + block_s=None, + head_block=None, + interpret=False, +): + """The Pallas kernel, compiled, so both sides of a comparison are dispatched alike. + + `norm_mode="exact"` deliberately leaves the FP32 feature-axis reduction in + XLA rather than Mosaic, which means the kernel has an XLA prologue that is + itself subject to eager-vs-jit differences. Calling the kernel eagerly while + the reference is jitted therefore reintroduces the very mismatch + `_compiled_reference` exists to remove, just one op earlier. + + It is not hypothetical. At the 18,900 x 40 x 128 production shape the eager + and jitted reductions disagree on a few hundred of the 96,768,000 outputs + (665 on tpu7x, 503 on v6e, max |diff| 3.125e-02) -- rare enough to survive a + small-shape test and be caught only at full size. Compiling both sides makes + the test measure the kernel instead of the dispatch path, and matches how + production invokes it. + """ + return fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=q_heads, + kv_heads=kv_heads, + dim_head=dim_head, + norm_mode=norm_mode, + rope_accum=rope_accum, + block_s=block_s, + head_block=head_block, + interpret=interpret, + ) + + +def _rounding_report(raw_q, raw_k, q_scale, k_scale, freqs_cis, *, q_heads, kv_heads, dim_head, norm_mode, block_s): + """Cross-tabulates every reference dispatch against every kernel rounding mode. + + There are only four ways the kernel can round (eager or compiled, times + `dtype` or `f32`) and two ways the reference can (eager or compiled). If any + cell is 0, the kernel is algebraically correct and the only question is which + convention this platform wants -- a one-line change to `resolve_rope_accum`. + If no cell is 0, the kernel genuinely disagrees with the producer and the + arithmetic needs looking at. Printing the whole grid answers that question + from a CI log, without needing the hardware in hand. + """ + interpret = not _on_tpu() + common = {"q_heads": q_heads, "kv_heads": kv_heads, "dim_head": dim_head} + + refs = { + "ref:eager": fused_rmsnorm_rope(raw_q, raw_k, q_scale, k_scale, freqs_cis, **common), + "ref:jit": _compiled_reference(raw_q, raw_k, q_scale, k_scale, freqs_cis, **common), + } + kernels = {} + for accum in ROPE_ACCUM_MODES: + kernels[f"kernel:eager/{accum}"] = fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + norm_mode=norm_mode, + rope_accum=accum, + block_s=block_s, + interpret=interpret, + **common, + ) + kernels[f"kernel:jit/{accum}"] = _compiled_kernel( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + norm_mode=norm_mode, + rope_accum=accum, + block_s=block_s, + interpret=interpret, + **common, + ) + + device = jax.devices()[0] + lines = [ + f" rounding report [{device.platform}/{getattr(device, 'device_kind', '?')}, " + f"norm_mode={norm_mode!r}, resolve_rope_accum()={resolve_rope_accum()!r}]:" + ] + for ref_name, (ref_q, _) in refs.items(): + for kernel_name, (got_q, _) in kernels.items(): + a = np.asarray(ref_q, np.float32) + b = np.asarray(got_q, np.float32) + n = int(np.count_nonzero(a != b)) + verdict = " <-- bit-identical" if n == 0 else "" + lines.append(f" {ref_name:10s} vs {kernel_name:20s}: {n:6d}/{a.size} ({100.0 * n / a.size:6.2f}%){verdict}") + return "\n".join(lines) + + +@_SKIP_IN_GITHUB_ACTIONS +class PairSwapIdentityTest(unittest.TestCase): + """The lane-rotation trick must reproduce the strided pair swap exactly.""" + + @staticmethod + def _pallas_pair_swap(x: jax.Array) -> jax.Array: + from jax.experimental import pallas as pl + + return pl.pallas_call( + lambda x_ref, o_ref: o_ref.__setitem__(Ellipsis, _pair_swap(x_ref[...])), + out_shape=jax.ShapeDtypeStruct(x.shape, x.dtype), + interpret=True, + )(x) + + def test_pair_swap_matches_strided_gather(self): + x = jax.random.normal(jax.random.PRNGKey(0), (3, 5, DIM_HEAD), jnp.float32) + got_ref = rope_pair_swap_reference(x) + got_fn = self._pallas_pair_swap(x) + + pairs = x.reshape(3, 5, DIM_HEAD // 2, 2) + want = jnp.stack([pairs[..., 1], pairs[..., 0]], axis=-1).reshape(x.shape) + + np.testing.assert_array_equal(np.asarray(got_ref), np.asarray(want)) + np.testing.assert_array_equal(np.asarray(got_fn), np.asarray(want)) + + def test_pair_swap_wraps_correctly_at_both_ends(self): + """Circular rotation is only valid because dim_head is even; pin that.""" + x = jnp.arange(8, dtype=jnp.float32).reshape(1, 8) + expected = np.array([[1, 0, 3, 2, 5, 4, 7, 6]], dtype=np.float32) + np.testing.assert_array_equal(np.asarray(rope_pair_swap_reference(x)), expected) + np.testing.assert_array_equal(np.asarray(self._pallas_pair_swap(x)), expected) + + +@_SKIP_IN_GITHUB_ACTIONS +class FusedRmsNormRopePallasParityTest(parameterized.TestCase): + """Parity against the unfused reference producer, for both norm modes. + + Every rotating comparison here is dispatched the same way on both sides: + the reference goes through `_compiled_reference` and the kernel through + `_compiled_kernel`, with the accumulation mode from `resolve_rope_accum()` + (the sole exception, and why, is documented on + `test_exact_mode_normalisation_alone_is_bit_identical`). + + That symmetry is load-bearing, because eager and compiled XLA can produce two + different correct answers and comparing across them measures the dispatch + path rather than the kernel. RoPE's `a*cos + b*sin` may be contracted into a + fused multiply-add (which rounds once instead of twice) depending on the + platform: on `v6e` XLA contracts under `jit` (while `tpu7x` and `XLA:CPU` + emit separate rounded ops; see `resolve_rope_accum`). 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 and with + nothing in between. `resolve_rope_accum()` returns the mode matching the + compiled reference on the current platform, which is also the mode production + compiles with -- so these tests pin the shipped configuration. + + With that pairing fixed, the admissible deviation differs per mode, and both + bounds are deliberately tight enough to catch a real algebra bug: + + `norm_mode="exact"` keeps the FP32 feature-axis reduction in XLA, so the + normalisation is bit-identical by construction (verified directly: under an + identity rotation this path matches the reference exactly, in both bf16 and + fp32). What remains is the RoPE combine, and the resolved accumulation mode + makes it round the same way the reference does, so bf16 must be + *bit-identical on TPU* (off-TPU interpret mode tolerates <=2 half-way ties at + 1 ULP; see `_assert_parity`). In float32 the kernel still evaluates in + float32 either way, so a single rounding of drift is allowed for the + reduction order. + + `norm_mode="fused"` additionally moves the reduction into Mosaic, whose + summation tree may differ from XLA's, which can flip the final rounding. + + Drift is measured relative to the norm of the *rotated pair*, not of the + individual component. RoPE rotates each `(x[2i], x[2i+1])` 2-vector, so that + pair norm is the rotation invariant, and the rounding error of either output + component is bounded by eps times it. An individual component, by contrast, + is free to be arbitrarily close to zero, which would make a + component-relative bound meaningless. + """ + + # Pair-relative machine-epsilon slack allowed on top of the exactness each + # mode guarantees (tol = budget * eps * pair_norm). + _PAIR_EPS_BUDGET = { + ("exact", jnp.bfloat16): 0, # bit-identical: kernel and reference round identically + ("exact", jnp.float32): 1, # FP32 reduction order in the RoPE add + ("fused", jnp.bfloat16): 2, # + Mosaic's VMEM reduction tree + RoPE rounding + ("fused", jnp.float32): 3, + } + _ULP_BUDGET = _PAIR_EPS_BUDGET + + @staticmethod + def _pair_norm(a): + """Norm of each RoPE 2-vector, broadcast back over both of its lanes.""" + pairs = a.reshape(*a.shape[:-1], a.shape[-1] // 2, 2) + norm = np.sqrt(np.sum(np.square(pairs.astype(np.float64)), axis=-1, keepdims=True)) + return np.repeat(norm, 2, axis=-1).reshape(a.shape).astype(np.float32) + + @staticmethod + def _ulp_distance(ref, got): + """Computes per-element integer ULP distance between two arrays of identical dtype.""" + ref_arr = np.asarray(ref) + got_arr = np.asarray(got) + if ref_arr.dtype == jnp.bfloat16: + u_ref = ref_arr.view(np.uint16).astype(np.int32) + u_got = got_arr.view(np.uint16).astype(np.int32) + i_ref = np.where(u_ref < 0x8000, u_ref, 0x8000 - u_ref) + i_got = np.where(u_got < 0x8000, u_got, 0x8000 - u_got) + return np.abs(i_ref - i_got) + u_ref = ref_arr.astype(np.float32).view(np.uint32).astype(np.int64) + u_got = got_arr.astype(np.float32).view(np.uint32).astype(np.int64) + i_ref = np.where(u_ref < 0x80000000, u_ref, 0x80000000 - u_ref) + i_got = np.where(u_got < 0x80000000, u_got, 0x80000000 - u_got) + return np.abs(i_ref - i_got) + + def _assert_parity(self, name, ref, got, dtype, norm_mode, diagnose=None): + """Asserts parity, and on failure says *which* rounding convention would have matched. + + `diagnose`, when supplied, is a zero-argument callable returning the + dispatch x accumulation table for this case. A bare "31% of elements + differ" cannot distinguish a real algebra bug from the kernel being + compiled against the wrong rounding convention for the platform, and the + two want opposite fixes. The table separates them, which matters most on + hardware the author cannot reach interactively -- a CI runner on a TPU + generation nobody has locally is exactly where that happens. + """ + self.assertEqual(ref.shape, got.shape, f"{name} shape mismatch") + self.assertEqual(ref.dtype, got.dtype, f"{name} dtype mismatch") + ref_np = np.asarray(ref, np.float32) + got_np = np.asarray(got, np.float32) + self.assertTrue(np.all(np.isfinite(ref_np)), f"{name}: reference output contains non-finite values.") + self.assertTrue(np.all(np.isfinite(got_np)), f"{name}: kernel output contains non-finite values (NaN/Inf).") + + report = f"\n{diagnose()}" if diagnose is not None else "" + budget = self._PAIR_EPS_BUDGET[(norm_mode, jnp.dtype(dtype).type)] + if budget == 0 and not rope_accum_is_measured(): + self.skipTest(_UNMEASURED_ROUNDING_SKIP) + if budget == 0: + mismatches = int(np.count_nonzero(ref_np != got_np)) + if mismatches: + # On CPU `interpret=True`, `pallas_call` lowers the grid via `lax.scan`, + # introducing an XLA:CPU fusion boundary between `rsqrt` and the consumer + # that can shift `rsqrt` by 1 FP32 ULP and flip a single exact 0x8000 + # bfloat16 half-way rounding tie by 1 bfloat16 ULP. On TPU (`_on_tpu()`), + # strict 0-ULP bit-identity (`mismatches == 0`) is always required. + if not _on_tpu() and mismatches <= 2 and int(np.max(self._ulp_distance(ref, got))) <= 1: + return + max_diff = float(np.max(np.abs(ref_np - got_np))) + self.fail( + f"{name}: norm_mode={norm_mode!r} in {jnp.dtype(dtype).name} must be bit-identical, but " + f"{mismatches}/{ref_np.size} elements differ ({100.0 * mismatches / ref_np.size:.2f}%, " + f"max |diff| = {max_diff:.3e}).{report}" + ) + return + + tol = budget * float(jnp.finfo(dtype).eps) * self._pair_norm(ref_np) + drift = np.abs(ref_np - got_np) + n_bad = int(np.count_nonzero(drift > tol)) + self.assertEqual( + n_bad, + 0, + f"{name}: norm_mode={norm_mode!r} in {jnp.dtype(dtype).name} exceeded its " + f"{budget}x pair-relative eps budget on {n_bad}/{ref_np.size} elements " + f"(worst excess ratio {float(np.max(drift / np.maximum(tol, np.finfo(np.float32).tiny))):.2f}x).{report}", + ) + + @parameterized.named_parameters( + {"testcase_name": f"_{tag}_{mode}", "q_heads": q, "kv_heads": kv, "dtype": dt, "norm_mode": mode} + for tag, q, kv, dt in ( + ("mha_bf16", 4, 4, jnp.bfloat16), + ("gqa_bf16", 8, 2, jnp.bfloat16), + ("mha_f32", 2, 2, jnp.float32), + ) + for mode in ("exact", "fused") + ) + def test_matches_reference(self, q_heads, kv_heads, dtype, norm_mode): + seq = 96 + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(0, 1, seq, q_heads, kv_heads, DIM_HEAD, dtype) + + ref_q, ref_k = _compiled_reference( + raw_q, raw_k, q_scale, k_scale, freqs, q_heads=q_heads, kv_heads=kv_heads, dim_head=DIM_HEAD + ) + got_q, got_k = _compiled_kernel( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + kv_heads=kv_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + rope_accum=resolve_rope_accum(), + block_s=32, + interpret=not _on_tpu(), + ) + + diagnose = functools.partial( + _rounding_report, + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + kv_heads=kv_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + block_s=32, + ) + for name, ref, got in (("q", ref_q, got_q), ("k", ref_k, got_k)): + self._assert_parity(name, ref, got, dtype, norm_mode, diagnose=diagnose) + + @parameterized.named_parameters(("_exact", "exact"), ("_fused", "fused")) + def test_handles_sequence_not_divisible_by_block(self, norm_mode): + """Wan's 18,900-token shard is not a multiple of any power-of-two tile.""" + seq, q_heads = 50, 3 + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(1, 1, seq, q_heads, q_heads, DIM_HEAD, jnp.bfloat16) + + ref_q, ref_k = _compiled_reference(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=q_heads, dim_head=DIM_HEAD) + got_q, got_k = _compiled_kernel( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + rope_accum=resolve_rope_accum(), + block_s=16, + interpret=not _on_tpu(), + ) + + diagnose = functools.partial( + _rounding_report, + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + kv_heads=None, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + block_s=16, + ) + self._assert_parity("q", ref_q, got_q, jnp.bfloat16, norm_mode, diagnose=diagnose) + self._assert_parity("k", ref_k, got_k, jnp.bfloat16, norm_mode, diagnose=diagnose) + + @parameterized.named_parameters(("_exact", "exact"), ("_fused", "fused")) + def test_result_is_independent_of_block_size(self, norm_mode): + seq, q_heads = 64, 2 + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(2, 1, seq, q_heads, q_heads, DIM_HEAD, jnp.bfloat16) + + outs = [] + for block_s in (16, 32, 64): + q, _ = fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + block_s=block_s, + interpret=not _on_tpu(), + ) + outs.append(np.asarray(q, np.float32)) + for other in outs[1:]: + np.testing.assert_array_equal(outs[0], other, err_msg="Tiling must not change the result.") + + @parameterized.named_parameters(("_exact", "exact"), ("_fused", "fused")) + def test_rmsnorm_component_matches_flax(self, norm_mode): + """Under an identity rotation the kernel must reproduce nnx.RMSNorm to within 1e-6.""" + seq, q_heads = 32, 2 + d_model = q_heads * DIM_HEAD + k0, k1 = jax.random.split(jax.random.PRNGKey(3)) + raw_q = jax.random.normal(k0, (1, seq, d_model), jnp.float32) + q_scale = jax.random.normal(k1, (d_model,), jnp.float32) + freqs = jnp.ones((1, 1, seq, DIM_HEAD // 2), jnp.complex64) # cos=1, sin=0 + + layer = nnx.RMSNorm(d_model, epsilon=1e-6, dtype=jnp.float32, param_dtype=jnp.float32, rngs=nnx.Rngs(0)) + layer.scale.value = q_scale + want = layer(raw_q).reshape(1, seq, q_heads, DIM_HEAD).transpose(0, 2, 1, 3) + + got, _ = fused_rmsnorm_rope_pallas( + raw_q, + raw_q, + q_scale, + q_scale, + freqs, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + block_s=16, + interpret=not _on_tpu(), + ) + np.testing.assert_allclose(np.asarray(got, np.float32), np.asarray(want, np.float32), rtol=1e-6, atol=1e-6) + + @parameterized.named_parameters(("_bf16", jnp.bfloat16), ("_f32", jnp.float32)) + def test_exact_mode_normalisation_alone_is_bit_identical(self, dtype): + """With the rotation switched off, `exact` must match bit-for-bit in *both* dtypes. + + This localises the float32 slack granted in `_ULP_BUDGET`. Under an + identity rotation the kernel degenerates to `x * (rsqrt * scale)` with no + `a*cos + b*sin` anywhere, so no multiply-add contraction is possible. If + this test ever fails, the normalisation itself has drifted and the float32 + budget is masking a real bug rather than a benign contraction. + + This is the one parity test that deliberately keeps the *eager* reference + and the kernel's default accumulation mode. Everywhere else that pairing + is meaningless, but here it is the point: with the rotation switched off + the two accumulation modes and the two dispatch modes all denote the same + arithmetic, so a mismatch cannot be blamed on rounding conventions. + """ + seq, q_heads = 32, 2 + raw_q, raw_k, q_scale, k_scale, _ = _make_inputs(8, 1, seq, q_heads, q_heads, DIM_HEAD, dtype) + identity = jnp.ones((1, 1, seq, DIM_HEAD // 2), jnp.complex64) # cos = 1, sin = 0 + + ref_q, ref_k = fused_rmsnorm_rope(raw_q, raw_k, q_scale, k_scale, identity, q_heads=q_heads, dim_head=DIM_HEAD) + got_q, got_k = fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + identity, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode="exact", + block_s=16, + interpret=not _on_tpu(), + ) + + for name, ref, got in (("q", ref_q, got_q), ("k", ref_k, got_k)): + np.testing.assert_array_equal( + np.asarray(ref, np.float32), + np.asarray(got, np.float32), + err_msg=f"{name}: exact-mode RMSNorm must be bit-identical in {jnp.dtype(dtype).name}.", + ) + + @parameterized.named_parameters(("_exact", "exact"), ("_fused", "fused")) + def test_rope_preserves_per_head_norm(self, norm_mode): + """RoPE is a rotation, so it must not change the per-position head norm.""" + seq, q_heads = 32, 2 + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(4, 1, seq, q_heads, q_heads, DIM_HEAD, jnp.float32) + + got_q, _ = fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + block_s=16, + interpret=not _on_tpu(), + ) + pre = jnp.linalg.norm( + (raw_q * (jax.lax.rsqrt(jnp.mean(raw_q**2, -1, keepdims=True) + 1e-6) * q_scale)) + .reshape(1, seq, q_heads, DIM_HEAD) + .transpose(0, 2, 1, 3), + axis=-1, + ) + post = jnp.linalg.norm(got_q.astype(jnp.float32), axis=-1) + np.testing.assert_allclose(np.asarray(pre), np.asarray(post), rtol=2e-5, atol=2e-5) + + @parameterized.named_parameters( + {"testcase_name": f"_batch_{b}_{mode}", "batch": b, "norm_mode": mode} for b in (2, 4) for mode in ("exact", "fused") + ) + def test_multi_batch(self, batch, norm_mode): + """Pallas kernel must process all batch elements, not just element 0.""" + seq, q_heads = 48, 2 + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(7, batch, seq, q_heads, q_heads, DIM_HEAD, jnp.bfloat16) + + ref_q, ref_k = _compiled_reference(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=q_heads, dim_head=DIM_HEAD) + got_q, got_k = _compiled_kernel( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + rope_accum=resolve_rope_accum(), + block_s=16, + interpret=not _on_tpu(), + ) + for b in range(batch): + self._assert_parity(f"q_b{b}", ref_q[b : b + 1], got_q[b : b + 1], jnp.bfloat16, norm_mode) + self._assert_parity(f"k_b{b}", ref_k[b : b + 1], got_k[b : b + 1], jnp.bfloat16, norm_mode) + + @parameterized.named_parameters(("_exact", "exact"), ("_fused", "fused")) + def test_rope_accum_f32_and_prescale(self, norm_mode): + """Exercises rope_accum='f32' and non-unit Q/K prescaling.""" + seq, q_heads = 32, 2 + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(9, 2, seq, q_heads, q_heads, DIM_HEAD, jnp.bfloat16) + q_prescale = 1.4426950408889634 # LOG2E + k_prescale = 1.0 / np.sqrt(DIM_HEAD) + + # Reference under FMA (rope_accum='f32') evaluates RoPE rotation in FP32 with lane-swapped partner + cos_full, sin_signed = _rope_tables(freqs, seq, jnp.bfloat16) + + def _build_exact_f32_ref(raw, scale, prescale, num_heads): + x_fp32 = raw.astype(jnp.float32) + rms = jax.lax.rsqrt(jnp.mean(jnp.square(x_fp32), axis=-1, keepdims=True) + 1e-6) + outs = [] + for h in range(num_heads): + x_head = raw[:, :, h * DIM_HEAD : (h + 1) * DIM_HEAD] + scale_head = scale[h * DIM_HEAD : (h + 1) * DIM_HEAD].reshape(1, 1, DIM_HEAD) + mul = rms * scale_head + normed = (x_head.astype(jnp.float32) * mul).astype(jnp.bfloat16) + wide = normed.astype(jnp.float32) + out_head = ( + wide * cos_full[0].astype(jnp.float32) + rope_pair_swap_reference(wide) * sin_signed[0].astype(jnp.float32) + ).astype(jnp.bfloat16) + out_head = out_head * jnp.asarray(prescale, jnp.bfloat16) + outs.append(out_head) + return jnp.stack(outs, axis=1) + + ref_q = _build_exact_f32_ref(raw_q, q_scale, q_prescale, q_heads) + ref_k = _build_exact_f32_ref(raw_k, k_scale, k_prescale, q_heads) + + got_q, got_k = fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + rope_accum="f32", + block_s=16, + q_prescale=q_prescale, + k_prescale=k_prescale, + interpret=not _on_tpu(), + ) + for b in range(2): + self._assert_parity(f"q_b{b}", ref_q[b : b + 1], got_q[b : b + 1], jnp.bfloat16, norm_mode) + self._assert_parity(f"k_b{b}", ref_k[b : b + 1], got_k[b : b + 1], jnp.bfloat16, norm_mode) + + @parameterized.named_parameters(("_dtype", "dtype"), ("_f32", "f32")) + def test_explicit_rope_accum_algebra_on_all_platforms(self, rope_accum): + """Tests both rope_accum modes against an explicit reference matching that mode without skipping on v4 CI.""" + seq, q_heads = 32, 2 + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(10, 1, seq, q_heads, q_heads, DIM_HEAD, jnp.bfloat16) + cos_full, sin_signed = _rope_tables(freqs, seq, jnp.bfloat16) + + def _explicit_ref(raw, scale): + x_fp32 = raw.astype(jnp.float32) + rms = jax.lax.rsqrt(jnp.mean(jnp.square(x_fp32), axis=-1, keepdims=True) + 1e-6) + outs = [] + for h in range(q_heads): + x_head = raw[:, :, h * DIM_HEAD : (h + 1) * DIM_HEAD] + scale_head = scale[h * DIM_HEAD : (h + 1) * DIM_HEAD].reshape(1, 1, DIM_HEAD) + normed = (x_head.astype(jnp.float32) * (rms * scale_head)).astype(jnp.bfloat16) + if rope_accum == "f32": + wide = normed.astype(jnp.float32) + out_head = ( + wide * cos_full[0].astype(jnp.float32) + rope_pair_swap_reference(wide) * sin_signed[0].astype(jnp.float32) + ).astype(jnp.bfloat16) + else: + out_head = normed * cos_full[0] + rope_pair_swap_reference(normed) * sin_signed[0] + outs.append(out_head) + return jnp.stack(outs, axis=1) + + ref_q = _explicit_ref(raw_q, q_scale) + ref_k = _explicit_ref(raw_k, k_scale) + got_q, got_k = fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=q_heads, + dim_head=DIM_HEAD, + norm_mode="exact", + rope_accum=rope_accum, + block_s=16, + interpret=not _on_tpu(), + ) + # Within at most 1 ULP on any platform (including v4 CI where XLA FMA rules differ) + for name, ref, got in (("q", ref_q, got_q), ("k", ref_k, got_k)): + max_ulp = int(np.max(self._ulp_distance(ref, got))) + self.assertLessEqual(max_ulp, 1, f"{name}[{rope_accum}]: exceeded 1 bf16 ULP (max ULP = {max_ulp}).") + + +class FusedRmsNormRopePallasGuardTest(unittest.TestCase): + """Misconfiguration must fail loudly rather than silently produce garbage.""" + + def _inputs(self, dim_head, q_heads=2, seq=16): + return _make_inputs(5, 1, seq, q_heads, q_heads, dim_head, jnp.bfloat16) + + def test_rejects_unaligned_dim_head(self): + raw_q, raw_k, q_scale, k_scale, freqs = self._inputs(64) + with self.assertRaisesRegex(ValueError, "multiple of 128"): + fused_rmsnorm_rope_pallas(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=2, dim_head=64, interpret=not _on_tpu()) + + def test_rejects_unaligned_block_s(self): + raw_q, raw_k, q_scale, k_scale, freqs = self._inputs(DIM_HEAD) + with self.assertRaisesRegex(ValueError, "multiple of 8"): + fused_rmsnorm_rope_pallas( + raw_q, raw_k, q_scale, k_scale, freqs, q_heads=2, dim_head=DIM_HEAD, block_s=12, interpret=not _on_tpu() + ) + + def test_rejects_feature_dim_mismatch(self): + raw_q, raw_k, q_scale, k_scale, freqs = self._inputs(DIM_HEAD) + with self.assertRaisesRegex(ValueError, "raw_q feature dim"): + fused_rmsnorm_rope_pallas(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=3, dim_head=DIM_HEAD, interpret=not _on_tpu()) + + def test_rejects_freqs_cis_mismatch(self): + raw_q, raw_k, q_scale, k_scale, _ = self._inputs(DIM_HEAD) + bad_freqs = jnp.ones((1, 1, 16, 8), jnp.complex64) + with self.assertRaisesRegex(ValueError, "freqs_cis last dim"): + fused_rmsnorm_rope_pallas( + raw_q, raw_k, q_scale, k_scale, bad_freqs, q_heads=2, dim_head=DIM_HEAD, interpret=not _on_tpu() + ) + + def test_nan_output_fails_parity_assertion(self): + ref = jnp.ones((1, 2, 16, DIM_HEAD), dtype=jnp.bfloat16) + got_nan = ref.at[0, 0, 0, 0].set(jnp.nan) + helper = FusedRmsNormRopePallasParityTest() + with self.assertRaisesRegex(AssertionError, "non-finite"): + helper._assert_parity("q", ref, got_nan, jnp.bfloat16, "fused") + + def test_non_tpu_platform_falls_back_to_xla_producer(self): + from types import SimpleNamespace + from maxdiffusion.models.attention_flax import FlaxWanAttention + + fake_gpu_device = SimpleNamespace(platform="gpu") + fake_mesh = SimpleNamespace(devices=np.array([fake_gpu_device]), shape={"data": 1}) + dummy_self = SimpleNamespace( + use_fused_rope_kernel=True, + mesh=fake_mesh, + dim_head=128, + heads=2, + kv_heads=2, + scale=1.0 / np.sqrt(128.0), + use_base2_exp=True, + attention_kernel="dot_product", + wan_rope_norm_mode="exact", + wan_fuse_qk_prescale=True, + wan_rope_accum="auto", + fused_rope_block_s=512, + fused_rope_head_block=None, + ) + producer = FlaxWanAttention._fused_rope_producer(dummy_self) + raw_q, raw_k, q_scale, k_scale, freqs = self._inputs(DIM_HEAD) + (q_out, k_out), qk_prescaled = producer(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=2, dim_head=DIM_HEAD, eps=1e-6) + self.assertFalse(qk_prescaled) + ref_q, ref_k = fused_rmsnorm_rope(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=2, kv_heads=2, dim_head=DIM_HEAD) + np.testing.assert_array_equal(np.asarray(q_out, np.float32), np.asarray(ref_q, np.float32)) + np.testing.assert_array_equal(np.asarray(k_out, np.float32), np.asarray(ref_k, np.float32)) + + def test_svg_attention_uses_xla_fused_producer_and_rejects_prescaled_qk(self): + from unittest import mock + from flax import nnx + from maxdiffusion.models import attention_flax + from maxdiffusion.models.attention_flax import FlaxWanAttention + + attn = FlaxWanAttention( + rngs=nnx.Rngs(0), + query_dim=256, + heads=2, + dim_head=128, + attention_kernel="dot_product", + attention_config={"use_svg_attention": True, "use_fused_rope_kernel": True}, + dtype=jnp.bfloat16, + weights_dtype=jnp.bfloat16, + ) + hidden_states = jnp.ones((1, 16, 256), dtype=jnp.bfloat16) + freqs = jnp.ones((1, 1, 16, 64), dtype=jnp.complex64) + with mock.patch.object( + FlaxWanAttention, + "_fused_rope_producer", + side_effect=AssertionError("_fused_rope_producer must not be called when use_svg_attention=True"), + ): + with mock.patch.object(attention_flax, "fused_rmsnorm_rope", wraps=fused_rmsnorm_rope) as mock_fused: + with mock.patch.object( + attn.attention_op, + "apply_attention", + side_effect=lambda q, k, v, **kwargs: jnp.transpose(q, (0, 2, 1, 3)).reshape(q.shape[0], q.shape[2], -1), + ): + _ = attn(hidden_states, hidden_states, rotary_emb=freqs, spatiotemporal_shape=(1, 4, 4)) + self.assertEqual(mock_fused.call_count, 1) + + q_4d = jnp.ones((1, 2, 16, 128), dtype=jnp.bfloat16) + with self.assertRaisesRegex(ValueError, "does not support prescaled Q/K"): + attention_flax._head_local_svg_attention(q_4d, q_4d, q_4d, {"qk_prescaled": True}) + with self.assertRaisesRegex(ValueError, "does not support prescaled Q/K"): + attention_flax._head_local_svg_attention(q_4d, q_4d, q_4d, {"k_prescaled": True}) + + def test_cudnn_flash_te_unscales_prescaled_k(self): + from unittest import mock + from maxdiffusion.models import attention_flax + + q = jnp.ones((1, 16, 2, 64), dtype=jnp.float32) + k_raw = jnp.full((1, 16, 2, 64), 4.0, dtype=jnp.float32) + scale = 0.125 + k_prescaled = k_raw * scale + v = jnp.ones((1, 16, 2, 64), dtype=jnp.float32) + captured = {} + + def fake_cudnn(q_in, k_in, v_in, heads, mesh, dpa_layer): + captured["k"] = np.asarray(k_in) + return q_in + + with mock.patch.object(attention_flax, "_cudnn_flash_attention", side_effect=fake_cudnn): + attention_flax.cudnn_flash_te_kernel( + q, + k_prescaled, + v, + {"heads": 2, "mesh": None, "dpa_layer": None, "scale": scale, "k_prescaled": True}, + ) + np.testing.assert_allclose(captured["k"], np.asarray(k_raw), rtol=1e-6, atol=1e-6) + + +class ResolveVmemLimitBytesTest(unittest.TestCase): + """The VMEM budget follows the mesh's TPU generation; unvalidated TPUs get Mosaic's default.""" + + @staticmethod + def _mesh(*kinds, platform="tpu"): + from types import SimpleNamespace + + return SimpleNamespace(devices=np.array([SimpleNamespace(platform=platform, device_kind=k) for k in kinds])) + + def test_explicit_value_wins(self): + self.assertEqual(_resolve_vmem_limit_bytes(123, self._mesh("TPU v4")), 123) + + def test_validated_generations_get_the_default(self): + for kinds in (("TPU v6 lite",) * 4, ("TPU7x",) * 4): + with self.subTest(kinds=kinds[0]): + self.assertEqual(_resolve_vmem_limit_bytes(None, self._mesh(*kinds)), DEFAULT_VMEM_LIMIT_BYTES) + + def test_unvalidated_generations_defer_to_mosaic(self): + for kind in ("TPU v4", "TPU v5 lite", "TPU v5p", "TPU v5e"): + with self.subTest(kind=kind): + self.assertIsNone(_resolve_vmem_limit_bytes(None, self._mesh(kind, kind))) + # A mixed mesh is only as good as its weakest member. + self.assertIsNone(_resolve_vmem_limit_bytes(None, self._mesh("TPU v6 lite", "TPU v4"))) + + def test_mesh_is_used_instead_of_global_devices(self): + with mock.patch.object(jax, "devices", side_effect=AssertionError("must use the mesh")): + self.assertIsNone(_resolve_vmem_limit_bytes(None, self._mesh("TPU v4"))) + + def test_off_tpu_returns_default(self): + self.assertEqual(_resolve_vmem_limit_bytes(None, self._mesh("cpu", platform="cpu")), DEFAULT_VMEM_LIMIT_BYTES) + + +@_SKIP_IN_GITHUB_ACTIONS +class FusedRmsNormRopeBackwardTest(unittest.TestCase): + """`with_xla_backward` must buy differentiability without moving the forward pass. + + Interpret mode, so the property is covered off-TPU too: the wrapper is pure + JAX plumbing and nothing here is platform-specific. + """ + + Q_HEADS = 2 + SEQ = 16 + + def _inputs(self): + return _make_inputs(11, 1, self.SEQ, self.Q_HEADS, self.Q_HEADS, DIM_HEAD, jnp.float32) + + def _fns(self): + kernel = functools.partial( + fused_rmsnorm_rope_pallas, + q_heads=self.Q_HEADS, + dim_head=DIM_HEAD, + norm_mode="exact", + interpret=not _on_tpu(), + ) + xla = functools.partial(fused_rmsnorm_rope, q_heads=self.Q_HEADS, dim_head=DIM_HEAD) + return kernel, xla + + @staticmethod + def _loss(fn, *args): + out_q, out_k = fn(*args) + return jnp.sum(out_q * 2.0) + jnp.sum(out_k * 3.0) + + def test_forward_is_bit_identical_to_the_unwrapped_kernel(self): + """The wrapper must be invisible to inference -- 0 ULP, not merely close.""" + kernel, xla = self._fns() + args = self._inputs() + bare_q, bare_k = jax.jit(kernel)(*args) + wrapped_q, wrapped_k = jax.jit(with_xla_backward(kernel, xla))(*args) + _assert_bit_identical(self, "q", bare_q, wrapped_q) + _assert_bit_identical(self, "k", bare_k, wrapped_k) + + def test_gradient_matches_the_xla_producer(self): + """Training must work, and differentiate the function the kernel computes.""" + kernel, xla = self._fns() + args = self._inputs() + wrapped = with_xla_backward(kernel, xla) + + got = jax.jit(jax.grad(lambda *a: self._loss(wrapped, *a), argnums=(0, 1, 2, 3)))(*args) + want = jax.jit(jax.grad(lambda *a: self._loss(xla, *a), argnums=(0, 1, 2, 3)))(*args) + + for name, g, w in zip(("raw_q", "raw_k", "q_scale", "k_scale"), got, want): + g_np, w_np = np.asarray(g, np.float32), np.asarray(w, np.float32) + self.assertTrue(np.all(np.isfinite(g_np)), f"d/d{name}: non-finite gradient.") + self.assertGreater(float(np.max(np.abs(w_np))), 0.0, f"d/d{name}: reference gradient is all zeros; test is vacuous.") + np.testing.assert_allclose(g_np, w_np, rtol=1e-6, atol=1e-6, err_msg=f"d/d{name} disagrees with the XLA producer.") + + def test_bare_kernel_still_has_no_transpose_rule(self): + """Tripwire: if Pallas grows a transpose rule the wrapper may be removable.""" + kernel, xla = self._fns() + args = self._inputs() + grad_fn = jax.grad(lambda *a: self._loss(kernel, *a), argnums=(0, 1, 2, 3)) + try: + got = jax.jit(grad_fn)(*args) + except Exception: # pylint: disable=broad-except + return # Expected: `pallas_call` is not transposable. + want = jax.jit(jax.grad(lambda *a: self._loss(xla, *a), argnums=(0, 1, 2, 3)))(*args) + for name, g, w in zip(("raw_q", "raw_k", "q_scale", "k_scale"), got, want): + np.testing.assert_allclose( + np.asarray(g, np.float32), + np.asarray(w, np.float32), + rtol=1e-6, + atol=1e-6, + err_msg=f"pallas_call became differentiable but d/d{name} disagrees with XLA.", + ) + + +@_SKIP_IN_GITHUB_ACTIONS +class FusedRmsNormRopePallasProductionShapeTest(unittest.TestCase): + """Exercises the real Wan 2.2 per-shard shape on TPU. + + 18,900 tokens x 40 heads x 128 is the v6e-8 / tpu7x-8 per-context-shard + self-attention projection. It is deliberately not a multiple of 8, so it also + covers the ragged trailing sequence tile at full width. + """ + + SEQ = 18900 + HEADS = 40 + + # The tile production actually compiles with: `fused_rope_block_s` / + # `fused_rope_head_block` from base_wan_27b.yml, read by FlaxWanAttention + # (attention_flax.py). The same tile is used on v6e and tpu7x; the kernel + # only clamps `block_s` in norm_mode="fused" (to 256). + PRODUCTION_BLOCK_S = 1024 + PRODUCTION_HEAD_BLOCK = None + + def _run(self, norm_mode): + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(6, 1, self.SEQ, self.HEADS, self.HEADS, DIM_HEAD, jnp.bfloat16) + ref = _compiled_reference(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=self.HEADS, dim_head=DIM_HEAD) + got = _compiled_kernel( + raw_q, + raw_k, + q_scale, + k_scale, + freqs, + q_heads=self.HEADS, + dim_head=DIM_HEAD, + norm_mode=norm_mode, + rope_accum=resolve_rope_accum(), + block_s=self.PRODUCTION_BLOCK_S, + head_block=self.PRODUCTION_HEAD_BLOCK, + ) + return ref, got + + @unittest.skipUnless(_on_tpu(), "Production-shape parity requires a TPU.") + @unittest.skipUnless(rope_accum_is_measured(), _UNMEASURED_ROUNDING_SKIP) + def test_exact_mode_is_bit_identical(self): + (ref_q, ref_k), (got_q, got_k) = self._run("exact") + + for name, ref, got in (("q", ref_q, got_q), ("k", ref_k, got_k)): + ref_np = np.asarray(ref, np.float32) + got_np = np.asarray(got, np.float32) + self.assertTrue(np.all(np.isfinite(ref_np)), f"{name}: reference output contains non-finite values.") + self.assertTrue(np.all(np.isfinite(got_np)), f"{name}: kernel output contains non-finite values (NaN/Inf).") + mismatches = int(np.count_nonzero(ref_np != got_np)) + self.assertEqual( + mismatches, + 0, + f"{name}: {mismatches}/{ref_np.size} elements differ " + f"(max |diff| = {float(np.max(np.abs(ref_np - got_np))):.3e}).", + ) + + def _production_shard_map_setup(self): + """Builds the jit + shard_map harness at the production per-shard shape. + + Shared so the reference-parity and prescaling assertions can be gated + independently; only the former depends on this platform's XLA rounding. + """ + import functools + from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + + devices = jax.devices() + global_seq = self.SEQ * len(devices) + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(7, 1, global_seq, self.HEADS, self.HEADS, DIM_HEAD, jnp.bfloat16) + + mesh = Mesh(np.array(devices), ("context",)) + act_spec = P(None, "context", None) + freqs_spec = P(None, None, "context", None) + rep_spec = P() + out_spec = P(None, None, "context", None) + + args = ( + jax.device_put(raw_q, NamedSharding(mesh, act_spec)), + jax.device_put(raw_k, NamedSharding(mesh, act_spec)), + jax.device_put(q_scale, NamedSharding(mesh, rep_spec)), + jax.device_put(k_scale, NamedSharding(mesh, rep_spec)), + jax.device_put(freqs, NamedSharding(mesh, freqs_spec)), + ) + + def run_kernel(**kwargs): + return jax.jit( + jax.shard_map( + functools.partial( + fused_rmsnorm_rope_pallas, + q_heads=self.HEADS, + dim_head=DIM_HEAD, + norm_mode="exact", + **kwargs, + ), + mesh=mesh, + in_specs=(act_spec, act_spec, rep_spec, rep_spec, freqs_spec), + out_specs=(out_spec, out_spec), + check_vma=False, + ) + ) + + return mesh, args, run_kernel + + @unittest.skipUnless(_on_tpu(), "Production-shape parity requires a TPU.") + @unittest.skipUnless(rope_accum_is_measured(), _UNMEASURED_ROUNDING_SKIP) + def test_production_jit_shard_map_matches_xla_reference(self): + """Tests jax.jit + shard_map + prescaling against the XLA producer at the kernel's default block_s.""" + import functools + + mesh, args, run_kernel = self._production_shard_map_setup() + q_prescale = float(np.log2(np.e)) # LOG2E, matching use_base2_exp. + k_prescale = float(1.0 / np.sqrt(DIM_HEAD)) # The attention softmax scale. + + # Ask the shared resolver rather than re-deriving the platform rule here. An + # inline copy is how this file previously claimed "dtype on tpu7x, f32 + # everywhere else", which silently became wrong on v4 -- the same + # duplicated-rule failure that motivated `resolve_rope_accum` in the first + # place. + exact_rope_accum = resolve_rope_accum(mesh) + + run_ref = jax.jit(functools.partial(fused_rmsnorm_rope, q_heads=self.HEADS, dim_head=DIM_HEAD)) + ref_q0, ref_k0 = run_ref(*args) + got_exact_q0, got_exact_k0 = run_kernel(rope_accum=exact_rope_accum)(*args) + got_q0, got_k0 = run_kernel(rope_accum="f32")(*args) + + # 1. The resolved rope_accum must be 0-ULP against the jitted XLA producer, + # and rope_accum="f32" must stay within 1x pair-eps of it. + for name, ref, got_exact, got_f32 in ( + ("q_unscaled", ref_q0, got_exact_q0, got_q0), + ("k_unscaled", ref_k0, got_exact_k0, got_k0), + ): + ref_np = np.asarray(ref, np.float32) + got_exact_np = np.asarray(got_exact, np.float32) + got_f32_np = np.asarray(got_f32, np.float32) + self.assertTrue(np.all(np.isfinite(ref_np)), f"{name}: reference output contains non-finite values.") + self.assertTrue(np.all(np.isfinite(got_exact_np)), f"{name}: exact kernel output contains non-finite values.") + self.assertTrue(np.all(np.isfinite(got_f32_np)), f"{name}: f32 kernel output contains non-finite values.") + self.assertEqual( + int(np.count_nonzero(ref_np != got_exact_np)), + 0, + f"{name}: expected 0-ULP bit-identical output with rope_accum={exact_rope_accum!r}.", + ) + drift = np.abs(ref_np - got_f32_np) + tol = float(jnp.finfo(jnp.bfloat16).eps) * FusedRmsNormRopePallasParityTest._pair_norm(ref_np) + self.assertEqual(int(np.count_nonzero(drift > tol)), 0, f"{name}: rope_accum='f32' exceeded 1x pair-eps.") + + # 2. The prescaled kernel must stay within 2x pair-eps of the prescaled XLA + # reference. (That prescaling is *exactly* a post-hoc multiply is asserted + # separately, without a platform gate.) + got_q, got_k = run_kernel(rope_accum="f32", q_prescale=q_prescale, k_prescale=k_prescale)(*args) + for name, ref, got in ( + ("q_prescaled", ref_q0 * jnp.asarray(q_prescale, jnp.bfloat16), got_q), + ("k_prescaled", ref_k0 * jnp.asarray(k_prescale, jnp.bfloat16), got_k), + ): + ref_np = np.asarray(ref, np.float32) + got_np = np.asarray(got, np.float32) + self.assertTrue(np.all(np.isfinite(ref_np)), f"{name}: reference output contains non-finite values.") + self.assertTrue(np.all(np.isfinite(got_np)), f"{name}: kernel output contains non-finite values (NaN/Inf).") + drift = np.abs(ref_np - got_np) + tol = 2.0 * float(jnp.finfo(jnp.bfloat16).eps) * FusedRmsNormRopePallasParityTest._pair_norm(ref_np) + over_pair_eps = int(np.count_nonzero(drift > tol)) + self.assertEqual( + over_pair_eps, + 0, + f"{name}: {over_pair_eps}/{ref_np.size} elements exceed 2x pair-eps.", + ) + + @unittest.skipUnless(_on_tpu(), "Production-shape parity requires a TPU.") + @unittest.skipUnless(_on_validated_tpu(), _UNVALIDATED_VMEM_SKIP) + def test_in_register_prescaling_equals_post_hoc_multiply(self): + """In-register Q/K prescaling vs. a post-hoc XLA multiply on the unscaled kernel output. + + On validated platforms (v6e, tpu7x) the in-kernel multiply is 0-ULP bit-identical + to a separate post-hoc XLA multiply. + """ + _, args, run_kernel = self._production_shard_map_setup() + q_prescale = float(np.log2(np.e)) + k_prescale = float(1.0 / np.sqrt(DIM_HEAD)) + + base_q, base_k = run_kernel(rope_accum="f32")(*args) + got_q, got_k = run_kernel(rope_accum="f32", q_prescale=q_prescale, k_prescale=k_prescale)(*args) + + for name, base, got, scale in ( + ("q_prescaled", base_q, got_q, q_prescale), + ("k_prescaled", base_k, got_k, k_prescale), + ): + want = base * jnp.asarray(scale, base.dtype) + want_np = np.asarray(want, np.float32) + got_np = np.asarray(got, np.float32) + self.assertTrue(np.all(np.isfinite(want_np)), f"{name}: post-hoc reference contains non-finite values.") + self.assertTrue(np.all(np.isfinite(got_np)), f"{name}: kernel output contains non-finite values (NaN/Inf).") + mismatches = int(np.count_nonzero(want_np != got_np)) + self.assertEqual( + mismatches, + 0, + f"{name}: in-register prescaling differs from a post-hoc multiply on {mismatches}/{want_np.size} elements.", + ) + + @unittest.skipUnless(_on_tpu(), "Production-shape parity requires a TPU.") + @unittest.skipUnless(rope_accum_is_measured(), _UNMEASURED_ROUNDING_SKIP) + def test_shipped_config_is_bit_identical_to_compiled_xla_reference(self): + """Pins the exact configuration that ships, end to end, at 0 ULP. + + `test_production_jit_shard_map_matches_xla_reference` checks the prescaled + kernel with `rope_accum="f32"` within 2x pair-eps at the default `block_s=512` + tile, and `test_in_register_prescaling_equals_post_hoc_multiply` checks + in-register prescaling against `base_q * prescale` (another Pallas output) + with `rope_accum="f32"`. Neither covers what production runs on tpu7x, where + `resolve_rope_accum` selects `"dtype"` with prescale ON and + `PRODUCTION_BLOCK_S = 1024`. + + Three things are therefore aligned with production rather than with the + kernel's defaults: + + * the accumulation mode comes from `resolve_rope_accum(mesh)`, the same + resolver `FlaxWanAttention` and the AOT fingerprint call, so this test + follows the shipped default onto new hardware instead of hard-coding a + platform's answer; + * the tile is `fused_rope_block_s` / `fused_rope_head_block`, not + `DEFAULT_BLOCK_S_EXACT`; + * the reference is a single jitted graph that applies the prescale + *inside* the jit. A host-side multiply on a materialised bf16 array + leaves XLA free to fuse the multiply into the RoPE combine in production + but not in the reference, which is exactly the kind of rounding + difference this assertion exists to catch. Measured on both v6e-8 and + tpu7x-8, the fused, barrier-staged and post-hoc references all agree + bit-for-bit with the kernel, so 0 ULP is the correct bound and not an + optimistic one. + """ + import functools + from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + + devices = jax.devices() + global_seq = self.SEQ * len(devices) + q_prescale = float(np.log2(np.e)) # LOG2E, matching use_base2_exp. + k_prescale = float(1.0 / np.sqrt(DIM_HEAD)) # The attention softmax scale. + + raw_q, raw_k, q_scale, k_scale, freqs = _make_inputs(11, 1, global_seq, self.HEADS, self.HEADS, DIM_HEAD, jnp.bfloat16) + mesh = Mesh(np.array(devices), ("context",)) + act_spec = P(None, "context", None) + freqs_spec = P(None, None, "context", None) + rep_spec = P() + out_spec = P(None, None, "context", None) + + raw_q = jax.device_put(raw_q, NamedSharding(mesh, act_spec)) + raw_k = jax.device_put(raw_k, NamedSharding(mesh, act_spec)) + q_scale = jax.device_put(q_scale, NamedSharding(mesh, rep_spec)) + k_scale = jax.device_put(k_scale, NamedSharding(mesh, rep_spec)) + freqs = jax.device_put(freqs, NamedSharding(mesh, freqs_spec)) + + rope_accum = resolve_rope_accum(mesh) + self.assertIn(rope_accum, ROPE_ACCUM_MODES) + + def reference(raw_q, raw_k, q_scale, k_scale, freqs): + q_out, k_out = fused_rmsnorm_rope(raw_q, raw_k, q_scale, k_scale, freqs, q_heads=self.HEADS, dim_head=DIM_HEAD) + return ( + q_out * jnp.asarray(q_prescale, q_out.dtype), + k_out * jnp.asarray(k_prescale, k_out.dtype), + ) + + run_ref = jax.jit(reference) + run_got = jax.jit( + jax.shard_map( + functools.partial( + fused_rmsnorm_rope_pallas, + q_heads=self.HEADS, + dim_head=DIM_HEAD, + norm_mode="exact", + rope_accum=rope_accum, + block_s=self.PRODUCTION_BLOCK_S, + head_block=self.PRODUCTION_HEAD_BLOCK, + q_prescale=q_prescale, + k_prescale=k_prescale, + ), + mesh=mesh, + in_specs=(act_spec, act_spec, rep_spec, rep_spec, freqs_spec), + out_specs=(out_spec, out_spec), + check_vma=False, + ) + ) + + args = (raw_q, raw_k, q_scale, k_scale, freqs) + ref_q, ref_k = run_ref(*args) + got_q, got_k = run_got(*args) + + for name, ref, got in ((f"q[{rope_accum}]", ref_q, got_q), (f"k[{rope_accum}]", ref_k, got_k)): + _assert_bit_identical(self, name, ref, got) + + @unittest.skipUnless(_on_tpu(), "Production-shape parity requires a TPU.") + @unittest.skipUnless(_on_validated_tpu(), _UNVALIDATED_VMEM_SKIP) + def test_fused_mode_drift_is_within_two_pair_eps(self): + """Quantifies the Mosaic reduction-tree difference and verifies it stays within 2x pair-relative bf16 eps. + + This is the measurement that justifies `exact` being the default: the + fused reduction is not bit-identical to XLA's reduction tree, and Wan's 40 + autoregressive denoise steps amplify even small bf16 rounding differences. + """ + (ref_q, ref_k), (got_q, got_k) = self._run("fused") + + for name, ref, got in (("q", ref_q, got_q), ("k", ref_k, got_k)): + ref_np = np.asarray(ref, np.float32) + got_np = np.asarray(got, np.float32) + self.assertTrue(np.all(np.isfinite(ref_np)), f"{name}: reference output contains non-finite values.") + self.assertTrue(np.all(np.isfinite(got_np)), f"{name}: kernel output contains non-finite values (NaN/Inf).") + drift = np.abs(ref_np - got_np) + tol = 2 * float(jnp.finfo(jnp.bfloat16).eps) * FusedRmsNormRopePallasParityTest._pair_norm(ref_np) + mismatches = int(np.count_nonzero(ref_np != got_np)) + over_pair_eps = int(np.count_nonzero(drift > tol)) + max_ulp = int(np.max(FusedRmsNormRopePallasParityTest._ulp_distance(ref, got))) + print( + f"[fused/{name}] {mismatches}/{ref_np.size} elements differ " + f"({100.0 * mismatches / ref_np.size:.6f}%), max |diff| = {float(drift.max()):.3e}, max ULP = {max_ulp}" + ) + self.assertEqual( + over_pair_eps, + 0, + f"{name}: {over_pair_eps} elements exceeded 2x pair-relative bf16 eps tolerance (max ULP = {max_ulp}).", + ) + + def test_fused_mode_multi_head_block_interpret(self): + """Verifies norm_mode='fused' produces accurate outputs across all head blocks when head_block < heads.""" + raw_q, raw_k, q_scale, k_scale, freqs_cis = _make_inputs( + seed=0, batch=1, seq=16, q_heads=4, kv_heads=4, dim_head=128, dtype=jnp.bfloat16 + ) + ref_q, ref_k = fused_rmsnorm_rope(raw_q, raw_k, q_scale, k_scale, freqs_cis, q_heads=4, kv_heads=4, dim_head=128) + got_q, got_k = fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=4, + kv_heads=4, + dim_head=128, + norm_mode="fused", + head_block=2, + block_s=16, + interpret=True, + ) + np.testing.assert_allclose(np.asarray(got_q, np.float32), np.asarray(ref_q, np.float32), rtol=2e-2, atol=2e-2) + np.testing.assert_allclose(np.asarray(got_k, np.float32), np.asarray(ref_k, np.float32), rtol=2e-2, atol=2e-2) + + def test_invalid_inputs_raise(self): + raw_q, raw_k, q_scale, k_scale, freqs_cis = _make_inputs(1, 1, 16, 4, 4, 128, jnp.bfloat16) + with self.assertRaises(ValueError): + fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=4, + kv_heads=4, + dim_head=128, + head_block=0, + interpret=True, + ) + with self.assertRaises(ValueError): + fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=4, + kv_heads=4, + dim_head=128, + block_s=0, + interpret=True, + ) + with self.assertRaises(ValueError): + fused_rmsnorm_rope_pallas( + raw_q, + raw_k, + q_scale, + k_scale, + freqs_cis, + q_heads=4, + kv_heads=4, + dim_head=128, + norm_mode="invalid", + interpret=True, + ) + + +@_SKIP_IN_GITHUB_ACTIONS +class FlaxWanAttentionFusedRopeMeshTest(unittest.TestCase): + """Exercises FlaxWanAttention._fused_rope_producer end-to-end on a sharded Mesh.""" + + def test_sharded_flax_wan_attention_fused_rope_producer_matches_xla(self): + """Sequence-sharded over `context` via real logical axis rules (bit-exact on CPU, close on TPU). + + Uses the dot_product attention kernel, so the Q/K prescale fold is not + exercised here (it is only enabled for the custom Ulysses kernels). + """ + from unittest import mock + from flax.linen import partitioning as nn_partitioning + from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + from maxdiffusion.models import attention_flax + from maxdiffusion.models.attention_flax import FlaxWanAttention + + devices = jax.devices() + mesh = Mesh(np.array(devices).reshape(1, len(devices)), ("data", "context")) + heads = 4 + seq = 32 * len(devices) + dim = heads * DIM_HEAD + + k0, k1 = jax.random.split(jax.random.PRNGKey(42)) + hidden_states = jax.random.normal(k0, (1, seq, dim), jnp.float32).astype(jnp.bfloat16) + angle = jax.random.uniform(k1, (1, 1, seq, DIM_HEAD // 2), jnp.float32, minval=-np.pi, maxval=np.pi) + freqs = jnp.cos(angle) + 1j * jnp.sin(angle) + + hidden_states = jax.device_put(hidden_states, NamedSharding(mesh, P("data", "context", None))) + freqs = jax.device_put(freqs, NamedSharding(mesh, P(None, None, "context", None))) + + attn_xla = FlaxWanAttention( + rngs=nnx.Rngs(0), + query_dim=dim, + heads=heads, + dim_head=DIM_HEAD, + mesh=mesh, + attention_kernel="dot_product", + attention_config={ + "use_fused_rope_kernel": False, + "use_base2_exp": True, + "wan_fuse_qk_prescale": True, + "wan_rope_norm_mode": "exact", + "wan_rope_accum": "auto", + "fused_rope_block_s": 16, + }, + dtype=jnp.bfloat16, + weights_dtype=jnp.bfloat16, + ) + attn_pallas = FlaxWanAttention( + rngs=nnx.Rngs(0), + query_dim=dim, + heads=heads, + dim_head=DIM_HEAD, + mesh=mesh, + attention_kernel="dot_product", + attention_config={ + "use_fused_rope_kernel": True, + "use_base2_exp": True, + "wan_fuse_qk_prescale": True, + "wan_rope_norm_mode": "exact", + "wan_rope_accum": "auto", + "fused_rope_block_s": 16, + }, + dtype=jnp.bfloat16, + weights_dtype=jnp.bfloat16, + ) + + axis_rules = ( + (attention_flax.BATCH, "data"), + (attention_flax.LENGTH, "context"), + (attention_flax.HEAD, None), + ) + built_specs = [] + orig_build = attention_flax._build_sharded_fused_rope_producer + + def recording_build(*args): + built_specs.append(args[2]) # act_spec + return orig_build(*args) + + def run(m): + return jax.jit(lambda m, h, f: m(h, h, rotary_emb=f))(m, hidden_states, freqs) + + with nn_partitioning.axis_rules(axis_rules): + out_xla = run(attn_xla) + with mock.patch.object(attention_flax, "_build_sharded_fused_rope_producer", side_effect=recording_build): + if _on_tpu(): + if not (rope_accum_is_measured() and _on_validated_tpu()): + self.skipTest(_UNMEASURED_ROUNDING_SKIP + " (the producer falls back to XLA here)") + out_pallas = run(attn_pallas) + else: + # On CPU, exercise the exact shard_map + Pallas path via interpret=True and a hashable fake-TPU mesh proxy + from types import SimpleNamespace + + class _HashableFakeMesh: + + def __init__(self, real_mesh): + self.devices = np.array( + [SimpleNamespace(platform="tpu", device_kind="TPU v7x") for _ in range(real_mesh.devices.size)] + ).reshape(real_mesh.devices.shape) + self.shape = real_mesh.shape + + orig_shard_map = jax.shard_map + pallas_interpret = functools.partial(fused_rmsnorm_rope_pallas, interpret=True) + attn_pallas.mesh = _HashableFakeMesh(mesh) + with ( + mock.patch.object( + jax, + "shard_map", + side_effect=lambda f, mesh, **kw: orig_shard_map(f, mesh=real_mesh, **kw), + ), + mock.patch.object(attention_flax, "fused_rmsnorm_rope_pallas", side_effect=pallas_interpret), + ): + real_mesh = mesh + out_pallas = run(attn_pallas) + + # The kernel really ran, sequence-sharded on `context` (not replicated). + self.assertEqual(len(built_specs), 1) + self.assertEqual(built_specs[0][1], "context") + if _on_tpu(): + # On TPU the XLA producer is inlined into the jitted attention graph and + # fused with its neighbours, which changes its rounding (the documented + # "equivalent, not identical" behaviour), so only closeness holds here. + np.testing.assert_allclose(np.asarray(out_pallas, np.float32), np.asarray(out_xla, np.float32), rtol=2e-2, atol=2e-2) + else: + # XLA:CPU does not re-round the inlined producer, so with "exact" norm + # mode and the platform's rope_accum the outputs are bit-identical. + np.testing.assert_array_equal(np.asarray(out_pallas, np.float32), np.asarray(out_xla, np.float32)) + + +@_SKIP_IN_GITHUB_ACTIONS +class CustomSplashTransposeOutTest(parameterized.TestCase): + """Verifies transpose_out=True in custom_splash_attention and _ulysses_attention.""" + + @parameterized.named_parameters( + ("_standard", False, False, False, 1), + ("_fixed_m_uniform", True, True, False, 1), + ("_fixed_m_hybrid", True, False, False, 1), + # Production uses virtual K-centering: k_mean goes into the kernel. + ("_fixed_m_uniform_k_centered", True, True, True, 1), + ("_fixed_m_hybrid_k_centered", True, False, True, 1), + ("_standard_mhpt", False, False, False, 2), + ) + def test_splash_transpose_out_matches_transposed_default(self, use_fixed_m, uniform_fixed_m, k_centered, heads_per_tile): + from maxdiffusion.kernels import custom_splash_attention + + heads, q_seq, kv_seq, dim_head = 2, 256, 256, 128 + k0, k1, k2 = jax.random.split(jax.random.PRNGKey(99), 3) + q = jax.random.normal(k0, (heads, q_seq, dim_head), jnp.float32).astype(jnp.bfloat16) + k = jax.random.normal(k1, (heads, kv_seq, dim_head), jnp.float32).astype(jnp.bfloat16) + v = jax.random.normal(k2, (heads, kv_seq, dim_head), jnp.float32).astype(jnp.bfloat16) + + block_sizes = custom_splash_attention._BlockSizes( + block_q=128, + block_kv=128, + block_kv_compute=128, + block_kv_compute_in=128, + ) + num_q_blocks = q_seq // 128 + mk = None + if use_fixed_m: + m_b = jnp.full((1, heads, num_q_blocks), 16.0, dtype=jnp.float32) + elig = jnp.ones((1, heads, num_q_blocks), dtype=jnp.float32) + mk = jnp.concatenate([m_b, elig], axis=0) + + kernel_bhsd = custom_splash_attention.make_splash_mha( + block_sizes=block_sizes, + orig_q_seq_len=q_seq, + orig_kv_seq_len=kv_seq, + use_base2_exp=True, + heads_per_tile=heads_per_tile, + use_fixed_m=use_fixed_m, + uniform_fixed_m=uniform_fixed_m, + interpret=not _on_tpu(), + transpose_out=False, + ) + kernel_bshd = custom_splash_attention.make_splash_mha( + block_sizes=block_sizes, + orig_q_seq_len=q_seq, + orig_kv_seq_len=kv_seq, + use_base2_exp=True, + heads_per_tile=heads_per_tile, + use_fixed_m=use_fixed_m, + uniform_fixed_m=uniform_fixed_m, + interpret=not _on_tpu(), + transpose_out=True, + ) + + if heads_per_tile > 1: + if not _on_tpu(): + self.skipTest("the heads_per_tile > 1 kernel has no interpret mode") + out_hds = kernel_bhsd(q, k, v) + out_hsd = kernel_bshd(q, k, v) + else: + k_mean = jnp.mean(k.astype(jnp.float32), axis=1) if k_centered else None + out_hds = kernel_bhsd(q, k, v, mk=mk, k_mean=k_mean) + out_hsd = kernel_bshd(q, k, v, mk=mk, k_mean=k_mean) + self.assertEqual(out_hds.shape, (heads, dim_head, q_seq)) + self.assertEqual(out_hsd.shape, (heads, q_seq, dim_head)) + np.testing.assert_array_equal( + np.asarray(jnp.swapaxes(out_hds, 1, 2), np.float32), + np.asarray(out_hsd, np.float32), + ) + + @parameterized.named_parameters( + ("_standard", "ulysses_custom", False), + ("_fixed_m_k_centered", "ulysses_custom_fixed_m", True), + ) + def test_ulysses_wrapper_transpose_out_is_bit_identical(self, kernel_name, use_fixed_m): + """`_ulysses_attention(transpose_out=True)` (the Wan production wrapper) matches the default layout.""" + self._check_wrapper_transpose_out(kernel_name, use_fixed_m) + + @parameterized.named_parameters(("_standard", False), ("_fixed_m_k_centered", True)) + def test_ulysses_ring_wrapper_transpose_out_is_bit_identical_at_r1(self, use_fixed_m): + """`_ulysses_ring_custom_attention(transpose_out=True)` at U=2, R=1 matches the default layout.""" + self._check_wrapper_transpose_out("ulysses_ring_custom", use_fixed_m) + + def _check_wrapper_transpose_out(self, kernel_name, use_fixed_m): + from flax.linen import partitioning as nn_partitioning + from jax.sharding import Mesh + from maxdiffusion.kernels import custom_splash_attention as custom_splash + from maxdiffusion.models import attention_flax + + if len(jax.devices()) < 2: + self.skipTest("needs >= 2 devices for a 2-way context axis") + mesh = Mesh(np.array(jax.devices()[:2]).reshape(1, 1, 2, 1), ("data", "fsdp", "context", "tensor")) + axis_rules = ( + (attention_flax.BATCH, "data"), + (attention_flax.SELF_ATTN_HEAD, None), + (attention_flax.SELF_ATTN_Q_LENGTH, "context"), + (attention_flax.SELF_ATTN_KV_LENGTH, "context"), + (attention_flax.D_KV, None), + ) + heads, seq, dim_head = 4, 256, 128 + k0, k1, k2 = jax.random.split(jax.random.PRNGKey(7), 3) + q = jax.random.normal(k0, (1, seq, heads * dim_head), jnp.float32).astype(jnp.bfloat16) + k = ((jax.random.normal(k1, (1, seq, heads * dim_head), jnp.float32) + 0.5) / np.sqrt(dim_head)).astype(jnp.bfloat16) + v = jax.random.normal(k2, (1, seq, heads * dim_head), jnp.float32).astype(jnp.bfloat16) + block_sizes = attention_flax.BlockSizes(block_q=128, block_kv_compute=128, block_kv=128) + names_q = (attention_flax.BATCH, attention_flax.SELF_ATTN_HEAD, attention_flax.SELF_ATTN_Q_LENGTH, attention_flax.D_KV) + names_kv = (attention_flax.BATCH, attention_flax.SELF_ATTN_HEAD, attention_flax.SELF_ATTN_KV_LENGTH, attention_flax.D_KV) + + def run(transpose_out): + if kernel_name == "ulysses_ring_custom": + return attention_flax._ulysses_ring_custom_attention( + q, + k, + v, + heads=heads, + mesh=mesh, + axis_names_q=names_q, + axis_names_kv=names_kv, + flash_block_sizes=block_sizes, + dtype=jnp.bfloat16, + ulysses_shards=2, + use_base2_exp=True, + use_fixed_m=use_fixed_m, + use_k_centering=use_fixed_m, + transpose_out=transpose_out, + ) + return attention_flax._ulysses_attention( + q, + k, + v, + heads=heads, + mesh=mesh, + axis_names_q=names_q, + axis_names_kv=names_kv, + flash_block_sizes=block_sizes, + dtype=jnp.bfloat16, + use_custom_kernel=True, + use_base2_exp=True, + use_fixed_m=use_fixed_m, + kernel_name=kernel_name, + use_k_centering=use_fixed_m, + transpose_out=transpose_out, + ) + + with mesh, nn_partitioning.axis_rules(axis_rules): + if use_fixed_m: + q_bhsd = (q.reshape(1, seq, heads, dim_head).transpose(0, 2, 1, 3) * attention_flax.LOG2E).astype(jnp.bfloat16) + k_bhsd = k.reshape(1, seq, heads, dim_head).transpose(0, 2, 1, 3) + v_bhsd = v.reshape(1, seq, heads, dim_head).transpose(0, 2, 1, 3) + recenter, safe_bound = custom_splash.get_fixed_m_constants(seq) + k_mean = jnp.mean(k_bhsd.astype(jnp.float32), axis=2) + _, all_fixed = attention_flax._compute_fixed_m_metadata( + q_bhsd, + k_bhsd, + block_q=128, + safe_bound=safe_bound, + recenter=recenter, + per_q_block=False, + k_mean=k_mean, + value=v_bhsd, + ) + self.assertTrue(bool(jnp.all(all_fixed)), "Expected uniform fixed-m branch (all_fixed == True).") + default = jax.jit(lambda: run(False))() + transposed = jax.jit(lambda: run(True))() + self.assertEqual(default.shape, transposed.shape) + np.testing.assert_array_equal(np.asarray(default, np.float32), np.asarray(transposed, np.float32)) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/tests/wan_runtime_options_test.py b/src/maxdiffusion/tests/wan_runtime_options_test.py new file mode 100644 index 000000000..a479a1a41 --- /dev/null +++ b/src/maxdiffusion/tests/wan_runtime_options_test.py @@ -0,0 +1,166 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +import glob +import os +import types +import unittest +from unittest import mock + +import yaml + +from maxdiffusion import wan_runtime_options as opts + +_CONFIG_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "configs") + + +class WanRuntimeOptionsTest(unittest.TestCase): + + def test_every_wan_yaml_declares_every_switch_with_the_module_default(self): + """The YAML is the source of truth, so it must list every switch -- and agree with the code default.""" + paths = sorted(glob.glob(os.path.join(_CONFIG_DIR, "base_wan*.yml"))) + self.assertTrue(paths) + for path in paths: + with open(path, encoding="utf-8") as f: + cfg = yaml.safe_load(f) + for name in opts.names(): + with self.subTest(config=os.path.basename(path), option=name): + self.assertIn(name, cfg) + self.assertEqual(opts.coerce(name, cfg[name]), opts.default(name)) + + def test_config_is_authoritative_over_legacy_env(self): + with mock.patch.dict(os.environ, {"WAN_SPLASH_TRANSPOSE_OUT": "1", "WAN_ROPE_NORM_MODE": "fused"}): + cfg = types.SimpleNamespace(wan_splash_transpose_out=False, wan_rope_norm_mode="exact") + self.assertFalse(opts.resolve_from_config(cfg, "wan_splash_transpose_out")) + self.assertEqual(opts.resolve_from_config(cfg, "wan_rope_norm_mode"), "exact") + + def test_legacy_env_used_only_without_config_key(self): + with mock.patch.dict(os.environ, {"WAN_CFG_BEFORE_UNPATCHIFY": "0"}): + self.assertFalse(opts.resolve_from_config(types.SimpleNamespace(), "wan_cfg_before_unpatchify")) + self.assertFalse(opts.resolve_from_config(None, "wan_cfg_before_unpatchify")) + with mock.patch.dict(os.environ, {}, clear=False): + os.environ.pop("WAN_CFG_BEFORE_UNPATCHIFY", None) + self.assertEqual( + opts.resolve_from_config(None, "wan_cfg_before_unpatchify"), opts.default("wan_cfg_before_unpatchify") + ) + + def test_command_line_strings_are_coerced(self): + cfg = types.SimpleNamespace(wan_fuse_qk_prescale="false", wan_splash_transpose_out="true") + self.assertFalse(opts.resolve_from_config(cfg, "wan_fuse_qk_prescale")) + self.assertTrue(opts.resolve_from_config(cfg, "wan_splash_transpose_out")) + with self.assertRaises(ValueError): + opts.resolve_from_config(types.SimpleNamespace(wan_fuse_qk_prescale="maybe"), "wan_fuse_qk_prescale") + + def test_invalid_norm_mode_raises(self): + with self.assertRaises(ValueError): + opts.resolve_from_config(types.SimpleNamespace(wan_rope_norm_mode="invalid"), "wan_rope_norm_mode") + + def test_rope_accum_validation_and_snapshot_from_config(self): + for mode in ("auto", "dtype", "f32"): + self.assertEqual(opts.resolve_from_config(types.SimpleNamespace(wan_rope_accum=mode), "wan_rope_accum"), mode) + with self.assertRaises(ValueError): + opts.resolve_from_config(types.SimpleNamespace(wan_rope_accum="fp32"), "wan_rope_accum") + + snap_default = opts.snapshot_from_config(types.SimpleNamespace()) + snap_custom = opts.snapshot_from_config(types.SimpleNamespace(wan_splash_transpose_out=True)) + self.assertNotEqual(snap_default, snap_custom) + + def test_no_process_global_store(self): + """Resolving one config must not change what another (config-less) caller sees.""" + self.assertFalse(hasattr(opts, "get")) + self.assertFalse(hasattr(opts, "configure_from_config")) + with mock.patch.dict(os.environ, {}, clear=False): + os.environ.pop("WAN_SPLASH_TRANSPOSE_OUT", None) + opts.attention_config_entries(types.SimpleNamespace(wan_splash_transpose_out=True)) + self.assertEqual(opts.resolve_from_config(None, "wan_splash_transpose_out"), opts.default("wan_splash_transpose_out")) + + def test_pyconfig_initialize_coerces_wan_switches(self): + from maxdiffusion import pyconfig + + pyconfig.initialize([ + None, + os.path.join(_CONFIG_DIR, "base_wan_14b.yml"), + "run_name=test_wan_opts", + "wan_splash_transpose_out=true", + "wan_rope_accum=f32", + ]) + self.assertIs(pyconfig.config.wan_splash_transpose_out, True) + self.assertEqual(pyconfig.config.wan_rope_accum, "f32") + entries = opts.attention_config_entries(pyconfig.config) + self.assertTrue(entries["wan_splash_transpose_out"]) + self.assertEqual(entries["wan_rope_accum"], "f32") + + +class WanModuleOptionsTest(unittest.TestCase): + """Wan modules carry every switch on the GraphDef; None is resolved at build time.""" + + def _attn(self, **attention_config): + from flax import nnx + from maxdiffusion.models.attention_flax import FlaxWanAttention + + return FlaxWanAttention( + rngs=nnx.Rngs(0), + query_dim=64, + heads=2, + dim_head=32, + attention_kernel="dot_product", + attention_config=attention_config or None, + ) + + def test_unset_options_resolve_to_defaults_not_env(self): + env = { + opts.env_var(name): "1" if opts.default(name) is False else "0" + for name in opts.names() + if isinstance(opts.default(name), bool) + } + env.update({"WAN_ROPE_NORM_MODE": "fused", "WAN_ROPE_ACCUM": "f32"}) + with mock.patch.dict(os.environ, env): + attn = self._attn() + for name in opts.ATTENTION_OPTIONS: + with self.subTest(option=name): + self.assertEqual(getattr(attn, name), opts.default(name)) + + def test_explicit_options_are_coerced_onto_the_module(self): + attn = self._attn(wan_splash_transpose_out="true", wan_rope_norm_mode="fused", wan_cross_attn_prescale_kv=True) + self.assertIs(attn.wan_splash_transpose_out, True) + self.assertEqual(attn.wan_rope_norm_mode, "fused") + self.assertTrue(attn.cross_attn_prescale_kv) + self.assertIs(attn.attention_op.transpose_out, True) + + def test_wan_model_resolves_cfg_before_unpatchify(self): + from flax import nnx + from maxdiffusion.models.wan.transformers.transformer_wan import WanModel + + kwargs = { + "num_attention_heads": 2, + "attention_head_dim": 16, + "in_channels": 4, + "out_channels": 4, + "text_dim": 32, + "freq_dim": 32, + "ffn_dim": 64, + "num_layers": 1, + "scan_layers": False, + } + with mock.patch.dict(os.environ, {"WAN_CFG_BEFORE_UNPATCHIFY": "0"}): + model = nnx.eval_shape(lambda: WanModel(rngs=nnx.Rngs(0), **kwargs)) + self.assertIs(model.wan_cfg_before_unpatchify, opts.default("wan_cfg_before_unpatchify")) + model = nnx.eval_shape(lambda: WanModel(rngs=nnx.Rngs(0), **kwargs, wan_cfg_before_unpatchify="false")) + self.assertIs(model.wan_cfg_before_unpatchify, False) + + +if __name__ == "__main__": + unittest.main() diff --git a/src/maxdiffusion/utils/wan_block_benchmark.py b/src/maxdiffusion/utils/wan_block_benchmark.py index 348beb6af..a5e6e345d 100644 --- a/src/maxdiffusion/utils/wan_block_benchmark.py +++ b/src/maxdiffusion/utils/wan_block_benchmark.py @@ -60,7 +60,7 @@ from flax import nnx from flax.linen import partitioning as nn_partitioning -from maxdiffusion import max_logging, max_utils, pyconfig +from maxdiffusion import max_logging, max_utils, pyconfig, wan_runtime_options from maxdiffusion.models.wan.transformers.transformer_wan import WanModel from maxdiffusion.utils.tile_size_grid_search import ( BenchResult, @@ -217,6 +217,7 @@ def _flash_block_sizes(self, bq, bkv, cmp): def _build_model(self, bq, bkv, cmp): c = self._config wan_config = dict(self._hf_cfg) + fused_rope_head_block = getattr(c, "fused_rope_head_block", -1) wan_config.update( mesh=self._mesh, dtype=c.activations_dtype, @@ -237,7 +238,13 @@ def _build_model(self, bq, bkv, cmp): "use_base2_exp": c.use_base2_exp, "use_experimental_scheduler": c.use_experimental_scheduler, "ulysses_shards": c.ulysses_shards, + "use_fused_rope_kernel": getattr(c, "use_fused_rope_kernel", False), + "fused_rope_block_s": getattr(c, "fused_rope_block_s", 1024), + "fused_rope_head_block": None if fused_rope_head_block in (None, -1) else fused_rope_head_block, + # Same graph-changing switches the pipelines pass (wan_runtime_options). + **wan_runtime_options.attention_config_entries(c), }, + wan_cfg_before_unpatchify=wan_runtime_options.resolve_from_config(c, "wan_cfg_before_unpatchify"), ) model = WanModel(**wan_config, rngs=nnx.Rngs(params=0)) gd, state, rest = nnx.split(model, nnx.Param, ...) diff --git a/src/maxdiffusion/wan_runtime_options.py b/src/maxdiffusion/wan_runtime_options.py new file mode 100644 index 000000000..fd6facdc0 --- /dev/null +++ b/src/maxdiffusion/wan_runtime_options.py @@ -0,0 +1,145 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Wan inference switches that change the compiled graph. + +These were previously read ad hoc from `WAN_*` environment variables at trace +time. They now come from the YAML config (`wan_rope_norm_mode`, ...). The +pipelines resolve them once at model-build time and store them on the modules +(`attention_config` / `WanModel.wan_cfg_before_unpatchify`), so the values are +part of the GraphDef and of the AOT cache fingerprint. There is no process-wide +store: nothing reads these switches at trace time, and a module built without +them uses the built-in default (`default(name)`), never a value left behind by +an earlier run in the same process. + +Resolution order for `resolve_from_config(config, name)`, evaluated at build +time only: + 1. the value in `config` (the YAML / command line), if the key is present + and not None; + 2. otherwise the legacy `WAN_*` environment variable; + 3. otherwise the built-in default, which matches the shipped YAML. + +The shipped YAMLs set every key, so on a normal run the env var only triggers a +one-time "ignoring" log when it disagrees. The env var takes effect only when +`resolve_from_config` is called with a config that lacks the key (or sets it to +None). Modules built without going through `resolve_from_config` (e.g. a bare +`FlaxWanAttention(...)` or `WanModel(...)` with no such `attention_config` +entry) use `default(name)` and ignore the environment. +""" + +import os +from typing import Any + +from maxdiffusion import max_logging + +# config key -> (legacy env var, default, kind) +_OPTIONS: dict[str, tuple[str, Any, str]] = { + "wan_rope_norm_mode": ("WAN_ROPE_NORM_MODE", "exact", "str"), + "wan_fuse_qk_prescale": ("WAN_FUSE_QK_PRESCALE", True, "bool"), + "wan_splash_transpose_out": ("WAN_SPLASH_TRANSPOSE_OUT", False, "bool"), + "wan_cfg_before_unpatchify": ("WAN_CFG_BEFORE_UNPATCHIFY", True, "bool"), + "wan_cross_attn_prescale_kv": ("WAN_CROSS_ATTN_PRESCALE_KV", False, "bool"), + "wan_rope_accum": ("WAN_ROPE_ACCUM", "auto", "str"), +} + + +def _coerce(name: str, value: Any) -> Any: + kind = _OPTIONS[name][2] + if kind == "bool": + if isinstance(value, bool): + return value + v = str(value).strip().lower() + if v in ("1", "true", "yes", "on"): + return True + if v in ("0", "false", "no", "off"): + return False + raise ValueError(f"{name} must be a boolean, got {value!r}.") + val = str(value) + if name == "wan_rope_norm_mode" and val not in ("exact", "fused"): + raise ValueError(f"wan_rope_norm_mode must be 'exact' or 'fused', got {value!r}.") + if name == "wan_rope_accum" and val not in ("auto", "dtype", "f32"): + raise ValueError(f"wan_rope_accum must be 'auto', 'dtype', or 'f32', got {value!r}.") + return val + + +coerce = _coerce + + +def names() -> tuple[str, ...]: + return tuple(_OPTIONS) + + +def env_var(name: str) -> str: + return _OPTIONS[name][0] + + +def _extract_keys(config: Any) -> dict[str, Any]: + if config is None: + return {} + if hasattr(config, "get_keys"): + return config.get_keys() + if isinstance(config, dict): + return config + return vars(config) + + +def default(name: str) -> Any: + """Built-in default of `name` (matches the shipped YAML).""" + return _OPTIONS[name][1] + + +_WARNED_ENV: set[str] = set() + + +def resolve_from_config(config: Any, name: str) -> Any: + """Coerced value of `name`: `config` first, then the legacy env var, then the default.""" + env, dflt, _ = _OPTIONS[name] + keys = _extract_keys(config) + legacy = os.environ.get(env) + if name in keys and keys[name] is not None: + value = _coerce(name, keys[name]) + if legacy is not None and env not in _WARNED_ENV: + try: + legacy_coerced = _coerce(name, legacy) + except ValueError: + legacy_coerced = legacy + if legacy_coerced != value: + _WARNED_ENV.add(env) + max_logging.log( + f"[wan_runtime_options] Ignoring legacy env {env}={legacy!r}: config sets {name}={value!r}. " + f"Set {name} in the YAML or on the command line instead." + ) + return value + return _coerce(name, legacy) if legacy is not None else dflt + + +ATTENTION_OPTIONS = ( + "wan_rope_norm_mode", + "wan_fuse_qk_prescale", + "wan_splash_transpose_out", + "wan_cross_attn_prescale_kv", + "wan_rope_accum", +) + + +def attention_config_entries(config: Any) -> dict[str, Any]: + """The `attention_config` entries every Wan pipeline passes to its attention modules.""" + return {name: resolve_from_config(config, name) for name in ATTENTION_OPTIONS} + + +def snapshot_from_config(config: Any = None) -> dict[str, str]: + """Coerced string snapshot of every switch, preferring explicit `config` fields.""" + return {name: str(resolve_from_config(config, name)) for name in _OPTIONS}