Conversation
There was a problem hiding this comment.
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.
3b9d7aa to
cb1e4d1
Compare
There was a problem hiding this comment.
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 = Falsein_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 2Dmkstill 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
left a comment
There was a problem hiding this comment.
Some line-level notes to go with my comment above.
44dfe63 to
fc91941
Compare
d57775a to
678e806
Compare
678e806 to
07d7e72
Compare
| 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) |
There was a problem hiding this comment.
I wonder what is the cost of the lax.cond here (maybe it's worth separating a cond-free variant)
b426abd to
73ab07f
Compare
73ab07f to
17679b5
Compare
…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.
17679b5 to
b65a99a
Compare
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: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_BOUNDandFP32_OUTPUT_HEADROOM_BITSare 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.
test_adversarial_centered_keys_softmax_mass_losschecks that this input is rejected and falls back correctly.Behaviour changes vs
mainThis PR is not purely additive. It changes
main's fixed-m path:Metadata (
mk)mkchanges from(2, num_heads), which heldmax‖k‖and a flag, to(2, num_heads, num_q_blocks), which holds the precomputed shift m_B and a flag.mkarrays of shape(2, num_heads)with aValueError._compute_fixed_m_metadataalso raises whenq_len % block_q != 0.Ulysses and R = 1 paths (
_ulysses_attention, R = 1 branch of_ulysses_ring_custom_attention)k_mean), instead ofkey - meanbeing written out before the kernel. Centering is always on for these paths.lax.condchooses between a uniform-fixed kernel and a hybrid fixed/online kernel.ulysses_custom_fixed_mnow passesper_q_block=False. A newulysses_custom_fixed_m_per_q_blockkernel is registered.Ring path, R > 1 (
make_custom_ring_attention)fixed_m_norms_squared=Trueby default) and gated in squared space.v_okpredicate is now required whenuse_fixed_m=True. The attention_flax caller computes it and reduces it over both internal mesh axes.k_mean. The in-tree R > 1 caller passes none, so it stays uncentered.per_q_block, defaultTrue. The fixed-m caller passesFalse.uniform_fixed_m:Trueforces the accumulate path and skips the gates andv_ok.pregathered_mk.k_mean.*,delimiter has been removed frommake_custom_ring_attention.Ragged tails
NUM_SUBLANES, and the padding rows are masked, on both the fixed and online paths._ulysses_attention, KV is left unpadded when N % 8 == 0 (otherwise it is still padded toblock_kv).Head-dim padding of
k_meank_meantoq.shape[-1].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.
Testing
Results at this head, run on CPU. Pallas runs in interpret mode.
custom_splash_fixed_m_test.py: 25 passedCustomSplashFixedMTest(16): online and fixed-m vs an f32 reference, sink and per-Q-block fallback, batched per-samplemkvia_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_blockTrue 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 withXLA_FLAGS=--xla_force_host_platform_device_count=8RingFixedMTest(7),RingFixedMContractTest(6),RingRawKeyBoundUnsoundTest(3).RingFixedMTesttests skip.Lint:
pyink --pyink-indentation=2 --line-length=125andruff checkare clean.Known limitations
all_fixedis reduced over the whole batch, so one outlier sample sends the whole batch to the hybrid kernel.p·vproducts can flush even though the denominator bound holds. The effect is harmless in practice.origin/main49cc98b4.