Skip to content

Prefetch each key tile from its own KV cache offset in flash attention - #81

Open
Shreyas8612 wants to merge 1 commit into
mainfrom
fix/flashattn-kv-tile-addressing
Open

Shreyas8612 wants to merge 1 commit into
mainfrom
fix/flashattn-kv-tile-addressing

Conversation

@Shreyas8612

Copy link
Copy Markdown
Collaborator

flash_attn_asm iterates the key tiles of the KV cache with a discarded loop index, and both qkt_multiply and computing_pv_code build 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_prefetch already 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.py lowers 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.

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.
Copilot AI lite review requested due to automatic review settings September 17, 2026 01:48

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 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 > 1 and hkv * d >= mlen, reset_kv_prefetch programs the row stride as hkv * d * batch (flashattn/reset.py:121-124), but this tile step uses only hkv * 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_elements elements between consecutive cache rows, but both K and V prefetches still use H_PREFETCH_M ... rstride=0 in qkt.py:73 and pv.py:66, which selects contiguous source addressing. When hkv * d > mlen (for example GQA with 8 KV heads and 128-wide heads), an MLEN-row tile is not contiguous, so advancing by mlen * kv_row_elements only 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 the 2**18 immediate 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 exceeds 2**18 and 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,
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.

2 participants