Skip to content

Fix: reject unsupported paged-attention shapes before hardware submission - #1972

Merged
ChaoZheng109 merged 1 commit into
hw-native-sys:mainfrom
ChaoZheng109:feat/pa-shape-assert
Aug 24, 2026
Merged

ChaoZheng109 merged 1 commit into
hw-native-sys:mainfrom
ChaoZheng109:feat/pa-shape-assert

Conversation

@ChaoZheng109

Copy link
Copy Markdown
Collaborator

Summary

  • Add assert_supported_shape() to the shared paged-attention goldens module (simpler_setup/goldens/paged_attention.py) and call it from generate_inputs(), so every scene that uses the shared generator validates its shape before any device work.
  • The unroll-family scenes (kernels without the 16×16 small-shape arm) pass variant="paged_attention_unroll" from generate_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:

  • the fallback arm hardcodes (q_tile=64, head_dim=128, block_size=64) — it assumes, not checks;
  • the q_tile_size == 16 arm never consults head_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 bare max_diff ~= 1.0 in a Daily run. #1925 fixed the Case3 dispatch 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 loud ValueError naming the tuple and the fix direction, before any hardware submission.

Shape table

q_tile derives from num_heads exactly as the orchestrations do (min(num_heads, 128)).

Variant Supported (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_KERNELS references resolved individually). The SPMD variants use a different mix-kernel dispatch (fixed HEAD_DIM=128, no runtime fallback) and are out of scope.

Verification

  • Unit-checked the helper: 7 supported tuples pass, 6 unsupported tuples (including (16, 256) and head_dim=192/512) raise.
  • Dry-ran the assertion against all 102 existing case params across the 24 scene files — zero false rejections.
  • End-to-end negative test: mutated Case1 to head_dim=192 → pytest fails in generate_args with unsupported paged-attention shape (q_tile=16, head_dim=192, block_size=128) ... extend the kernel dispatch (and this table) first — no device touched.
  • Positive smoke on a2a3sim: batch_paged_attention and multi_round_paged_attention pass with --manual include.
  • Pre-commit hooks all pass.

Part of #1832.

🤖 Generated with Claude Code

@coderabbitai

coderabbitai Bot commented Aug 24, 2026 •

Copy link
Copy Markdown

Warning

Review limit reached

Next included review available in 15 minutes.

View limit details

Limit 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.

Learn how review limits work.

Review configuration:

⚙️ Run configuration

Configuration used: Organization UI

Review profile: CHILL

Plan: Pro Plus

Run ID: fbb95780-7aef-4922-8fe6-0dcd9b51a72c

📥 Commits

Reviewing files that changed from the base of the PR and between 969c6a6 and 2a9dd93.

📒 Files selected for processing (11)
  • examples/a2a3/host_build_graph/paged_attention_unroll_manual_scope/test_paged_attention_unroll_manual_scope.py
  • examples/a2a3/tensormap_and_ringbuffer/paged_attention_unroll_manual_scope/test_paged_attention_unroll_manual_scope.py
  • examples/a5/host_build_graph/paged_attention_unroll_manual_scope/test_paged_attention_unroll_manual_scope.py
  • examples/a5/tensormap_and_ringbuffer/paged_attention_unroll_manual_scope/test_paged_attention_unroll_manual_scope.py
  • simpler_setup/goldens/paged_attention.py
  • tests/st/a2a3/host_build_graph/paged_attention_unroll/test_paged_attention_unroll.py
  • tests/st/a2a3/tensormap_and_ringbuffer/paged_attention_unroll/test_paged_attention_unroll.py
  • tests/st/a2a3/tensormap_and_ringbuffer/paged_attention_unroll_4dims/test_paged_attention_unroll_4dims.py
  • tests/st/a5/host_build_graph/paged_attention_unroll/test_paged_attention_unroll.py
  • tests/st/a5/tensormap_and_ringbuffer/paged_attention_unroll/test_paged_attention_unroll.py
  • tests/st/a5/tensormap_and_ringbuffer/paged_attention_unroll_4dims/test_paged_attention_unroll_4dims.py

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.

❤️ Share

Comment @coderabbitai help to get the list of available commands.

…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.
@ChaoZheng109
ChaoZheng109 merged commit 97a4217 into hw-native-sys:main Aug 24, 2026
19 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant