Repository navigation
Conversation
Contributor
There was a problem hiding this comment.
🟢 Approval recommended
The change is narrowly scoped to the failing head-dimension case, preserves existing power-of-two behavior, and adds targeted tests for the new configuration.
Pull request overview
Fixes a wave32 compilation failure in the fused RoPE + KV-cache kernel for the real-world configuration head_dim=96 by selecting a supported vector width and correctly pairing NeoX halves when the half-wave span is not a power of two.
Changes:
- Update
VEC_WIDTHselection to choose the smallest supported copy width in{1,2,4,8}that divides both the full head and each NeoX half while fitting in a single wave. - Generalize rotary pair-lane mapping: keep XOR pairing for power-of-two half spans, otherwise use explicit +/- half-span pairing (needed for D=96 on wave32).
- Add correctness coverage for
head_dim=96across bf16/f16 and flash/non-flash layouts.
File summaries
| File | Description |
|---|---|
| tests/kernels/test_fused_rope_cache.py | Adds new parametrized tests covering head_dim=96 for bf16/f16 and flash/non-flash KV cache layouts. |
| kernels/attention/fused_rope_cache_kernel.py | Fixes wave32 head_dim=96 by enforcing native vector widths and correcting the NeoX pair-lane mapping for non-power-of-two half spans. |
Review details
- Files reviewed: 2/2 changed files
- Comments generated: 0
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
Author
|
Hi @zhiding512 @coderfeli, could you review this when you get a chance? Thanks! |
This branch has not been deployed
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.
Motivation
The fused RoPE + KV cache kernel currently fails to compile for
head_dim=96on wave32.The previous vector-width calculation selects
VEC_WIDTH=3. Theds_bpermute_b32path handles the vector as 32-bit chunks and reconstructs only two BF16/FP16 elements, which causes:This is a real model configuration used by the Phi-3 Mini family. The existing comments also describe D=96 as a supported configuration.
Technical Details
1/2/4/8that divides both the complete head and each NeoX half while fitting the head into one wave.VEC_WIDTH=4and 24 active lanes.vecs_per_halfis a power of two.The D=64, D=128, and D=256 paths retain their existing vector widths and XOR mapping.
Test Plan
Run the fused RoPE + KV cache test suite:
The change was tested on an AMD Radeon 8060S using native
gfx1151wave32 execution.In addition to correctness testing, compare the generated ISA and CUDA Graph performance of the existing D=64/128/256 paths against
origin/main.Test Result
All four newly added D=96 combinations passed, and Q, K, key cache, and value cache matched the PyTorch reference.
The final gfx1151 ISA for D=64, D=128, and D=256 is byte-for-byte identical to
origin/main. Interleaved CUDA Graph measurements showed no performance regression in those existing paths.Representative D=96 CUDA Graph replay results for BF16,
QH=KH=32:The generated D=96 kernel uses 35 VGPRs, 26 SGPRs, no LDS, no scratch, and no register spills.
A direct D=96 comparison is unavailable because AITER's public Triton fused kernel currently rejects non-power-of-two head dimensions.
Submission Checklist