Repository navigation
Fix: reject unsupported paged-attention shapes before hardware submission - #1972
Merged
ChaoZheng109 merged 1 commit intoAug 24, 2026
Merged
Conversation
|
Warning Review limit reachedNext included review available in 15 minutes. View limit detailsLimit details: You’ve used the included review currently available. You've used all free OSS reviews for now. Wait for the free limit to reset to keep reviewing this public repository. Review configuration: ⚙️ Run configurationConfiguration used: Organization UI Review profile: CHILL Plan: Pro Plus Run ID: 📒 Files selected for processing (11)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
ChaoZheng109
force-pushed
the
feat/pa-shape-assert
branch
from
August 24, 2026 03:32
5743de5 to
2a9dd93
Compare
…sion The paged-attention QK/PV/online-update kernels dispatch on runtime tensor shapes, but their fallback arm hardcodes (q_tile=64, head_dim=128, block_size=64) and the 16-row arm never consults head_dim. A case with an unshaped tuple — e.g. head_dim=192/512, or head_dim=256 with block_size=128, or the (16, 256) corner — computes a wrong result with no diagnostic. That silent-wrong-answer mode is how hw-native-sys#1832 surfaced months late as a bare max_diff ~= 1.0 in a Daily run (hw-native-sys#1925 fixed the Case3 dispatch but kept the property). generate_inputs() in the shared goldens module now validates the (q_tile, head_dim, block_size) tuple against the kernels' dispatch table and raises ValueError naming the tuple and the fix direction, before any device work. q_tile derives from num_heads exactly as the orchestrations do (min(num_heads, 128)). The unroll-family scenes (whose kernels lack the 16x16 small-shape arm) pass variant="paged_attention_unroll"; all other variants use the default. All 102 existing case params across the 24 kernel-using scene files pass; verified end-to-end that a mutated head_dim=192 case fails in generate_args with the message, and that batch_paged_attention and multi_round_paged_attention still pass on a2a3sim. Part of hw-native-sys#1832.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
assert_supported_shape()to the shared paged-attention goldens module (simpler_setup/goldens/paged_attention.py) and call it fromgenerate_inputs(), so every scene that uses the shared generator validates its shape before any device work.variant="paged_attention_unroll"fromgenerate_args; all other variants use the default.Problem
Follow-up to #1925 (part of #1832). The QK/PV/online-update kernels dispatch on runtime tensor shapes, but:
(q_tile=64, head_dim=128, block_size=64)— it assumes, not checks;q_tile_size == 16arm never consultshead_dim.A case with an unsupported tuple —
head_dim=192/512,head_dim=256, block_size=128, or the(16, 256)corner — compiles fine, runs fine, and silently computes a wrong result. That failure mode is exactly how #1832 surfaced months late as a baremax_diff ~= 1.0in a Daily run. #1925 fixed theCase3dispatch but deliberately kept this property (out of scope per review discussion); this PR closes the gap on the Python side, where the failure becomes a loudValueErrornaming the tuple and the fix direction, before any hardware submission.Shape table
q_tilederives fromnum_headsexactly as the orchestrations do (min(num_heads, 128)).(q_tile, head_dim, block_size)paged_attention(default)(16, ≤16, ≤16),(16, 128, 128),(64, 128, 64),(64, 256, 64)paged_attention_unroll(16, 128, 128),(64, 128, 64),(64, 256, 64)The table matches the actual kernel dispatch arms verified across all 24 scene files that reference the dispatch kernels (local copies and
TMR_CASE/_PA_KERNELSreferences resolved individually). The SPMD variants use a different mix-kernel dispatch (fixedHEAD_DIM=128, no runtime fallback) and are out of scope.Verification
(16, 256)andhead_dim=192/512) raise.Case1tohead_dim=192→ pytest fails ingenerate_argswithunsupported paged-attention shape (q_tile=16, head_dim=192, block_size=128) ... extend the kernel dispatch (and this table) first— no device touched.a2a3sim:batch_paged_attentionandmulti_round_paged_attentionpass with--manual include.Part of #1832.
🤖 Generated with Claude Code