Repository navigation
Prefetch each key tile from its own KV cache offset in flash attention - #81
Open
Shreyas8612 wants to merge 1 commit into
Open
Shreyas8612 wants to merge 1 commit into
Shreyas8612 wants to merge 1 commit into
Conversation
flash_attn_asm iterated the key tiles with a discarded loop index, and qkt_multiply and computing_pv_code built the HBM prefetch offset from the KV-head index alone. For a cache longer than MLEN every iteration therefore re-read the first MLEN keys and values, so attention over a multi-tile cache only ever covered the first tile. Enumerate the key-tile loop and pass the tile's element offset into both prefetches. A key tile is MLEN rows of the cache, whose row width is the one reset_kv_prefetch already programs into the stride register, and the K offset is loaded through the large-immediate helper because the tile term outgrows the 18-bit immediate a few tiles in. asm_templates/tests/test_flashattn_kv_tile_addressing.py lowers a four-tile cache and checks that the K and V prefetches use the four distinct tile bases; it fails on the previous emitter. The asm_templates, aten and generator suites and the clm-60m and SmolVLM2 codegen-plus- assemble smoke give the same results as before.
There was a problem hiding this comment.
🟡 Changes recommended
Unresolved critical and moderate findings affect correctness, addressing, and API compatibility.
Get a fresh assessment by requesting another Copilot review.
Pull request overview
This pull request updates Flash Attention to prefetch each K/V cache tile from its own offset.
Changes:
- Computes and propagates per-tile K/V offsets.
- Adds large-immediate handling for K offsets.
- Adds multi-tile addressing tests.
File summaries
| File | Review summary |
|---|---|
asm_templates/tests/test_flashattn_kv_tile_addressing.py |
Adds tile-address tests. Nit (1 vote): does not test offsets beyond 2**18 or verify the relevant K/V source-register encoding. |
asm_templates/flashattn/qkt.py |
Applies offsets to K prefetches. Moderate (3 votes): the new parameter can break existing positional callers. |
asm_templates/flashattn/pv.py |
Applies offsets to V prefetches. |
asm_templates/flashattn/overall.py |
Critical (1 vote): per-tile resets discard online-softmax state. Moderate (1 vote): tile stride is not batch-aware. Moderate (1 vote): prefetch addressing does not match non-contiguous cache layouts. |
Review details
Suppressed comments (3)
asm_templates/flashattn/overall.py:67
- For
batch > 1andhkv * d >= mlen,reset_kv_prefetchprograms the row stride ashkv * d * batch(flashattn/reset.py:121-124), but this tile step uses onlyhkv * d. Consequently tile 1 starts before the next batch-strided cache row. Derive this value from the same batch-aware stride (and add a batch>1 assertion) so the tile offset matches the configured prefetch layout.
kv_row_elements = mlen if hkv * d < mlen else hkv * d
asm_templates/flashattn/overall.py:153
- This offset calculation assumes the cache is row-major with
kv_row_elementselements between consecutive cache rows, but both K and V prefetches still useH_PREFETCH_M ... rstride=0inqkt.py:73andpv.py:66, which selects contiguous source addressing. Whenhkv * d > mlen(for example GQA with 8 KV heads and 128-wide heads), an MLEN-row tile is not contiguous, so advancing bymlen * kv_row_elementsonly moves to the right nominal tile while each prefetch still reads the wrong rows. Use the configured row stride for these prefetches (with the appropriate precision units), or otherwise match the actual packed cache layout, and cover this branch in the test.
kv_tile_offset = k_tile_index * mlen * kv_row_elements
asm_templates/tests/test_flashattn_kv_tile_addressing.py:74
- These four tiles reach only
3 * 64 * 64 = 12,288, below the2**18immediate limit. The new LUI path is therefore untested here; reverting to the old single-ADDI emission would still yield the same parsed offsets and pass these assertions. Add a targeted case whose tile offset exceeds2**18and verify the encoding for the K/V source register rather than merely checking for an unrelated LUI elsewhere.
tiles = 4
asm = self._lower(kv_len=tiles * self.MLEN)
- Files reviewed: 4/4 changed files
- Comments generated: 2
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
|
|
||
| # loop over per kv head kv_len // MLEN | ||
| for _ in range(k_seq_iteration_number): | ||
| for k_tile_index in range(k_seq_iteration_number): |
Comment on lines
+20
to
22
| k_tile_offset: int = 0, | ||
| use_batched: bool = True, | ||
| blen: int = 4, |
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.
flash_attn_asmiterates the key tiles of the KV cache with a discarded loop index, and bothqkt_multiplyandcomputing_pv_codebuild the H_PREFETCH_M offset from the KV-head index alone. For any cache longer than MLEN every iteration re-reads the first MLEN keys and values, so attention over a multi-tile cache only ever sees tile 0.This enumerates the key-tile loop and passes the tile's element offset into the K and V prefetches. A tile is MLEN cache rows; the row width is the one
reset_kv_prefetchalready programs into the stride register. The K offset goes through the large-immediate helper because it exceeds the 18-bit immediate a few tiles in. For a single-tile cache the emitted code is unchanged.Test:
asm_templates/tests/test_flashattn_kv_tile_addressing.pylowers a four-tile cache and checks the K and V prefetch bases are the four distinct tile offsets; it fails on main. The asm_templates, aten and generator suites and the clm-60m / SmolVLM2 codegen+assemble smoke are unchanged. For a single-tile cache the generated program is byte-identical to main's.Note for reviewers: I tried to confirm this end to end with PLENA_Simulator's direct_emit flashattn_prefill testbench, but that testbench currently panics on the emulator's matrix-SRAM alignment assertion (per-head PV loop) with and without this change, so it could not serve as an oracle.