Skip to content

feat(attention): fixed-m splash attention kernel with dynamic bounds and safety fallbacks - #477

Open
Perseus14 wants to merge 1 commit into
mainfrom
feat/fixed-m-kernel
Open

Perseus14 wants to merge 1 commit into
mainfrom
feat/fixed-m-kernel

Conversation

@Perseus14

@Perseus14 Perseus14 commented Sep 13, 2026 •

Copy link
Copy Markdown
Collaborator

PR #477 fixed-m kernel: summary of changes

Summary

This PR reworks main's existing fixed-m splash-attention path, on both the single-device/Ulysses side and the ring side. The fixed Cauchy-Schwarz constants are replaced by bounds computed from the KV length. Metadata is now per (head, Q-block), and there are explicit safety fallbacks. Eligible tiles subtract a precomputed shift instead of tracking an online running max, so numerator and denominator accumulate directly in FP32.

Changes

Dynamic bounds (custom_splash_attention.py)

  • get_fixed_m_constants(kv_seq_len, v_max_bound=256) returns C(N) and the gate:

    $$C(N) = 127 - \lceil\log_2 N\rceil - \lceil\log_2 V_{\max}\rceil,\quad W(N) = C(N) + 125,\quad \text{gate} = \lfloor W(N)/2 \rfloor$$

  • For a Q-block B, let $U_B = \max_{i\in B}\lVert q_i\rVert\cdot\max_j\lVert k_j\rVert$ and $m_B = \lceil U_B\rceil - C(N)$. Over both logit signs, $z - m_B \ge C(N) - (U_B + \lceil U_B\rceil) \ge -125$, so no term flushes to subnormal and no FP32 accumulator overflows.

  • Examples: N = 4096 gives C = 107, gate = 116. N = 75,600 gives C = 102, gate = 113.

  • The legacy constants _FIXED_M_RECENTER, _FIXED_M_SAFE_BOUND, _FIXED_M_RING_SAFE_BOUND and FP32_OUTPUT_HEADROOM_BITS are removed.

Why the gate is ⌊W/2⌋ and not W

One tempting argument says that centering K makes the realized row max ≥ 0, so only one side of the window needs covering. It doesn't hold, because the shift is built from U, not from the realized max: the worst case is still C − (U + ⌈U⌉), centered or not.

  • Example at N = 4096 (W = 232): q = [231, 2], and K is half [0, 1] and half [0, −1] (exactly centered).
  • U ≈ 231.01 would pass a W-sized gate. The negative-logit terms, which carry 1/17 ≈ 5.9% of the softmax mass, would then flush to zero.
  • With V = ±1, the output becomes 1.0 instead of 15/17 ≈ 0.882, an absolute error of 0.118.
  • test_adversarial_centered_keys_softmax_mass_loss checks that this input is rejected and falls back correctly.

Behaviour changes vs main

This PR is not purely additive. It changes main's fixed-m path:

Metadata (mk)

  • mk changes from (2, num_heads), which held max‖k‖ and a flag, to (2, num_heads, num_q_blocks), which holds the precomputed shift m_B and a flag.
  • The kernel rejects legacy 2D mk arrays of shape (2, num_heads) with a ValueError.
  • _compute_fixed_m_metadata also raises when q_len % block_q != 0.

Ulysses and R = 1 paths (_ulysses_attention, R = 1 branch of _ulysses_ring_custom_attention)

  • Gate: 213 becomes ⌊W(N)/2⌋.
  • K is now centered virtually inside the kernel (k_mean), instead of key - mean being written out before the kernel. Centering is always on for these paths.
  • A V-magnitude / dtype safety check is added (|V| ≤ 256, and the dtype must be able to hold 2^C(N)).
  • A lax.cond chooses between a uniform-fixed kernel and a hybrid fixed/online kernel.
  • ulysses_custom_fixed_m now passes per_q_block=False. A new ulysses_custom_fixed_m_per_q_block kernel is registered.

Ring path, R > 1 (make_custom_ring_attention)

  • Global gate: 106.5 becomes ⌊W(N_total)/2⌋. The per-hop gate is ⌊W(N_local)/2⌋.
  • Norms are passed squared (fixed_m_norms_squared=True by default) and gated in squared space.
  • A v_ok predicate is now required when use_fixed_m=True. The attention_flax caller computes it and reduces it over both internal mesh axes.
  • Virtual K-centering is used only when the caller passes k_mean. The in-tree R > 1 caller passes none, so it stays uncentered.
  • New parameters:
    • per_q_block, default True. The fixed-m caller passes False.
    • uniform_fixed_m: True forces the accumulate path and skips the gates and v_ok.
    • pregathered_mk.
    • k_mean.
  • The keyword-only *, delimiter has been removed from make_custom_ring_attention.

Ragged tails

  • The last KV slice is rounded up to a multiple of NUM_SUBLANES, and the padding rows are masked, on both the fixed and online paths.
  • In _ulysses_attention, KV is left unpadded when N % 8 == 0 (otherwise it is still padded to block_kv).

Head-dim padding of k_mean

  • The ring kernel pads k_mean to q.shape[-1].
  • The Ulysses and R = 1 wrappers pad it to 128.

Measured cost of the strict gate (earlier revision)

This was measured on an earlier revision of the stack: a full Wan 2.2 720p generation on v6e-8, N = 75,600, 1280 metadata calls per gate.

gate derivation entries falling back calls with any fallback
113 ⌊W/2⌋ (shipped) 0.73% 112 / 1280
201 relative mass loss ≤ 2⁻¹⁰ incl. the N multiplicity 0.45% 56 / 1280
218 relative mass loss ≤ 2⁻¹⁰ per term 0.44% 56 / 1280
227 W 0.43% 54 / 1280

Testing

Results at this head, run on CPU. Pallas runs in interpret mode.

custom_splash_fixed_m_test.py: 25 passed

  • CustomSplashFixedMTest (16): online and fixed-m vs an f32 reference, sink and per-Q-block fallback, batched per-sample mk via _compute_fixed_m_metadata, the gate-crossing sweep, invariant bounds, non-divisible sequences, anti-aligned uncentered keys with heavily negative logits, virtual K-centering.
  • FixedMDtypeSafetyTest (3).
  • FixedMMetadataSafetyTest (4): includes the adversarial centered-key regression.
  • FixedMAttentionFlaxIntegrationTest (2): end to end through _ulysses_attention (per_q_block True and False) and the R = 1 branch of _ulysses_ring_custom_attention, compared with fp32 softmax. max|err| ≈ 5e-3.

ring_fixed_m_test.py: 16 passed with XLA_FLAGS=--xla_force_host_platform_device_count=8

  • RingFixedMTest (7), RingFixedMContractTest (6), RingRawKeyBoundUnsoundTest (3).
  • On a single CPU device, the 7 RingFixedMTest tests skip.

Lint: pyink --pyink-indentation=2 --line-length=125 and ruff check are clean.

Known limitations

  • Multi-device ring tests skip in single-device CPU CI.
  • all_fixed is reduced over the whole batch, so one outlier sample sends the whole batch to the hybrid kernel.
  • Near the gate, individual p·v products can flush even though the denominator bound holds. The effect is harmless in practice.
  • Fixed-m on the bidirectional ring is not supported.
  • Base: origin/main 49cc98b4.

@github-actions

Copy link
Copy Markdown

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request introduces dynamic fixed-m constants and safe bounds calculation based on KV sequence length, adds support for virtual K-centering, and enhances input validation and dtype safety checks in the Pallas flash attention kernel. It also expands the test suite to cover various edge cases, including per-Q-block fallback, batched isolation, and virtual K-centering. The review feedback focuses on optimizing the TPU kernel performance by replacing expensive dynamic integer division with optimized BlockSpec mapping and introducing a static use_k_centering flag to conditionally compile the centering logic at trace time.

Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
@Perseus14
Perseus14 force-pushed the feat/fixed-m-kernel branch 10 times, most recently from 3b9d7aa to cb1e4d1 Compare September 15, 2026 06:37
@syhuang22
syhuang22 self-requested a review September 15, 2026 20:04

@syhuang22 syhuang22 left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A few things before this can go in though:

  • If this lands on its own, fixed-m gets turned off for R>1 (use_fixed_m = False in _ulysses_ring_custom_attention, it only comes back in #478). That breaks the recipe we ship today. Can we split the stack so each PR is safe to merge by itself?
  • Cutting the Ulysses gate from 213 to 113 feels more conservative than we need. With centered keys, the mass you can lose is bounded by 2^(ceil(U)-C-126), so a gate around C+116 already keeps the loss under 0.1%. Could you share fallback rates so we can pick the number?
  • mk[0] means something different now, but a 2D mk still gets silently broadcast. Can we just raise on that?
  • The V check never fired in my runs on WAN 2.2 (every call passed), and it only exists because C moved from 88 to 102. Is it worth the extra complexity?
  • Please put back the comments that explain the Mosaic cliff and k-smoothing. They save the next person a lot of pain.

@syhuang22 syhuang22 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Some line-level notes to go with my comment above.

Comment thread src/maxdiffusion/models/attention_flax.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/models/attention_flax.py Outdated
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py
Comment thread src/maxdiffusion/generate_ltx_video.py
@Perseus14
Perseus14 force-pushed the feat/fixed-m-kernel branch 4 times, most recently from 44dfe63 to fc91941 Compare September 17, 2026 12:31
syhuang22
syhuang22 previously approved these changes Sep 17, 2026
@Perseus14
Perseus14 force-pushed the feat/fixed-m-kernel branch 3 times, most recently from d57775a to 678e806 Compare September 17, 2026 18:57
@Perseus14
Perseus14 requested a review from eltsai September 17, 2026 19:09
@Perseus14 Perseus14 self-assigned this Sep 17, 2026
@Perseus14
Perseus14 added this pull request to stack #486 September 17, 2026 19:10
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py Outdated
def _run_hybrid(q, k, v, m, km):
return jax.vmap(splash_kernel_hybrid, in_axes=(0, 0, 0, 0, 0))(q, k, v, m, km)

raw_out = jax.lax.cond(all_fixed, _run_uniform, _run_hybrid, query, key, value, mk_arr, k_mean)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I wonder what is the cost of the lax.cond here (maybe it's worth separating a cond-free variant)

…and safety fallbacks

Implements exact fixed-m splash attention in Pallas on TPU:
- Dynamic C(N) headroom constants guaranteeing FP32 accumulator safety
- fixed_m_dtype_is_safe checks rejecting FP16/FP8 exponent overflow
- Value bound validation (|V| <= 256) with safe online softmax fallback
- Unit tests covering all boundary conditions, dtypes, and scale factors

Explicit metadata contracts on the fixed-m ring path
----------------------------------------------------
Three implicit contracts are made explicit. Each failed silently rather than
loudly, and one of them produced non-finite output on TPU.

1. Norm representation is declared, not inferred. The gate previously guessed
   whether `fixed_m_norms` were squared with
   `(qn.max() * mk.max()) < 1000.0`. Magnitude cannot answer that question:
   legacy unsquared norms of (1000, 2) have a true bound of 2000, but read as
   already-squared they yield sqrt(2000) ~= 44.7 -- a ~45x under-estimate that
   admits fixed-m where it must fall back, and overflows. Replaced by
   `fixed_m_norms_squared` (default True, matching every in-tree caller); the
   test harness is migrated to squared norms.

2. The V-safety predicate is required, not assumed. Unlike the Cauchy-Schwarz
   norm bounds, the V-magnitude and dtype verdict is not re-derivable from a
   single hop's Q/K, so the kernel cannot reconstruct it. Omission previously
   meant "safe", which let fixed-m run on inputs it cannot represent: float16
   with Q=K=0 and V=1 returns inf instead of 1.0. The ring path now raises
   unless `v_ok` is passed, mirroring the existing `fixed_m_recenter` rule, and
   `_ulysses_ring_custom_attention` computes it -- dtype safety plus
   |V| <= DEFAULT_MAX_V_BOUND, reduced with pmin over BOTH internal axes, since
   after the all-to-all neither axis alone observes the whole V and the fixed-m
   branch must be taken uniformly by every ppermute participant.

3. Norm shape is validated against per_q_block. This is the defect behind the
   `test_sink_head_falls_back_everywhere` TPU failure. Both gates compute
   `qn * mk[:, None]`, so a (num_heads,) array supplied while per_q_block=True
   does not raise -- it broadcasts to (num_heads, num_heads), pairing head j's
   query norm with head h's key norm. A sink head then inherits a small bound
   from an unrelated head, is wrongly marked eligible, and the kernel evaluates
   exp2(large_logit - small_m) -> inf. The kernel now rejects the mismatch, and
   the test declares per_q_block=False to match the per-head norms it supplies,
   as the production ring caller already did.

Regression coverage: six backend-independent contract tests (both omissions
raise, v_ok=False is accepted, mis-shaped norms are rejected, correctly shaped
per-Q-block norms are accepted, fp16 is rejected while bf16/fp32 pass) plus two
TPU tests (the two declared norm representations must agree, and an explicit
unsafe verdict must force a finite fallback).

Verified on v6e-8: ring_fixed_m_test 13 passed; attention_test,
custom_splash_fixed_m_test, attention_block_sizes_test and ring_fixed_m_test
together 58 passed.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants