Skip to content

cpu: copy the KV cache by plane instead of scattering it element by element - #1756

Merged
justinchuby merged 2 commits into
mainfrom
seb/concat-cache-memcpy
Aug 22, 2026
Merged

justinchuby merged 2 commits into
mainfrom
seb/concat-cache-memcpy

Conversation

@justinchuby

Copy link
Copy Markdown
Owner

concat_cache in attention.rs and msft_attention.rs copied the KV cache one
element at a time with the head-dimension loop outermost, so every store landed
dim floats past the previous one — a fresh cache line per store on decode
shapes. multi_head_attention.rs was already converted to plane-sized
copy_from_slice; these two never were.

Why it is safe

The accessor is at(b,h,s,d) = data[((b*heads+h)*seq + s)*dim + d], so for a
fixed bh = b*heads + h:

range
past source [bh*past_plane, bh*past_plane + past_plane) — contiguous
cur source [bh*cur_plane, bh*cur_plane + cur_plane) — contiguous
destination back to back inside [bh*plane, bh*plane + plane)

So the triple nest was always two memcpys per plane. Nothing is transposed
and no element changes position — this is pure data movement, bit-exact by
construction, no tolerance involved.

Measured

Interleaved A/B in one process, medians over 20–60 reps, both arms asserted
bit-identical before timing:

case MiB scalar scatter memcpy/plane speedup
llama decode past=1023, h=32, d=128 16.0 8.973 ms 0.986 ms 9.10x
llama decode past=255, h=32, d=128 4.0 2.378 ms 0.185 ms 12.87x
llama decode past=4095, h=8, d=128 16.0 13.628 ms 0.989 ms 13.78x
gqa decode past=1023, h=8, d=128 4.0 2.527 ms 0.183 ms 13.80x
prefill past=0, h=32, d=128 8.0 5.167 ms 0.476 ms 10.86x
prefill+past 512, h=32, d=128 16.0 11.342 ms 1.818 ms 6.24x
batch=4 decode past=1023, h=32 64.0 55.109 ms 6.534 ms 8.43x

3.4 GB/s → 32.0 GB/s on the past=1023 cell. That cell was also run on a
genuinely quiet host (loadavg 1.20) at 9.36x versus 9.10x here; the A/B is
interleaved, so steady contention loads both arms equally.

concat_cache runs twice per attention op per token (key and value), so this is
~32 MiB per layer per decode step at 1023 past tokens.

No end-to-end claim. These are kernel-level numbers. I have not attributed a
tokens/sec delta on a full model, and the shared host is not currently quiet
enough for me to stand behind one.

Falsifier

concat_cache_is_bit_identical_to_the_scalar_scatter in each file keeps the old
scatter as an independent oracle and compares element-for-element over seven
shapes, including dim = 1, an empty past, a single past step, and multi-batch
multi-head — the ragged cases where a plane-stride mistake still yields a
plausible-looking buffer.

Verified non-vacuous both ways:

mutation result
reverse the copied past plane (attention.rs) fails at b=1 h=2 past=3 cur=1 d=4
corrupt one element after the copy (msft_attention.rs) fails at b=1 h=1 past=1 cur=1 d=1

The oracle indexes the backing Vec directly rather than calling the production
accessor, so it shares no index arithmetic with the code under test.

Two deliberate choices

  1. No par_chunks_mut. MHA has one above MIN_PARALLEL_TRANSPOSE_ELEMENTS
    that enters the global Rayon pool. Not porting it: 32 GB/s is already a
    fair share of one socket, and two more global-Rayon entry points cut against
    the thread-count reduction work. Worth noting MHA's threshold comment
    ("Decode steps (seq == 1) never reach it") does not actually hold for
    concat_cache — data.len() there spans the whole past, so a long-context
    decode step does take its parallel path. Whether parallelising this copy
    pays is a separate measured question, not one to settle by porting silently.
  2. at() deleted from both structs. It existed only to feed the scatter and
    is dead afterwards; clippy fails the build on it. MHA dropped its own
    accessor for the same reason when it was converted, so this matches the
    existing precedent rather than inventing one.

Validation

  • cargo test -p onnx-runtime-ep-cpu --lib — 1612 passed, 0 failed
  • full CI offline-linux package set — exit 0
  • fmt --check; clippy --all-targets -D warnings; --features mlas;
    --no-default-features; aarch64-unknown-linux-gnu cross-clippy — all clean
  • all seven Rust-quality scripts — pass

Closes #1755

justinchuby and others added 2 commits August 22, 2026 16:19
`concat_cache` in attention.rs and msft_attention.rs walked the head
dimension outermost, so every store was `dim` floats past the last one --
a fresh cache line per store on decode shapes. The accessor is
`((b*heads+h)*seq + s)*dim + d`, so for a fixed `b*heads+h` both sources
are contiguous runs that land back to back in the destination plane: the
whole nest is two memcpys per plane. Nothing is transposed and no element
moves, so the traversal order was the only thing making it slow.

multi_head_attention.rs was already fixed this way; port it to the other
two. 6.2-13.8x on interleaved A/B across decode, prefill and batched
shapes -- 3.4 GB/s to 32.0 GB/s on llama decode past=1023.

Deliberately not porting MHA's `par_chunks_mut` path: it enters the
global Rayon pool, and whether parallelising a 32 GB/s memcpy pays is a
separate measured question.

`at()` existed only to serve the scatter and is now dead, so it is
removed from both structs -- matching MHA, which dropped its own when it
was converted. The test oracle indexes explicitly instead, which also
makes it independent of the production accessor.

Closes #1755

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
…nches

Adversarial review NITs on #1756.

msft_attention's concat_cache never validated past-vs-current dims, unlike
its two sibling kernels. That was harmless while the copy was an
element-at-a-time scatter -- a mismatch quietly read the wrong elements --
but the plane-wise copy would surface it as an opaque slice-length panic
instead. The caller already validates, so this is unreachable; assert it
anyway so the unreachable case names itself.

Also cover the two branches the bit-identity grid stepped over: an absent
past cache, and degenerate shapes that take the plane == 0 early return.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby

Copy link
Copy Markdown
Owner Author

Adversarial review: claude-opus-4.8

Brief asked for falsification of six named claims, with instructions to report a
finding even at moderate confidence and to say plainly if nothing was found.

Verdict: no BLOCKING findings, no SHOULD-FIX findings. Three NITs, two of
which are now fixed in 73af488.

What the reviewer checked independently rather than taking on trust

  • Re-derived the index algebra symbolically instead of accepting the PR's
    table, confirming bh = b*heads + h enumerates 0..batch*heads exactly once
    ascending, that data.len() is always an exact multiple of plane (so
    chunks_mut yields exactly batch*heads full chunks with no ragged tail),
    and that past_plane + cur_plane == plane makes the two copies exactly fill
    each destination plane.
  • Chased the over-allocated-buffer question to the source. past.data
    comes from to_dense_f32_widen (dtype.rs:828), which returns exactly
    numel, and msft's make closure (msft_attention.rs:566) does
    dense[off..off+chunk].to_vec(). Also noted that even if a buffer were
    over-allocated the new code reads only the first batch*heads*past.seq*dim
    elements — precisely the range the old at() touched — so equivalence holds
    regardless.
  • Enumerated every new panic site and matched each against the old code's
    requirements: chunks_mut(plane) needs plane != 0 (guarded);
    split_at_mut(past_plane) is always in range because
    past_plane <= total*dim; both copy_from_slice lengths match by
    construction. The old scatter's maximum source index imposed the same
    bound, so the new code panics on exactly the same short-buffer inputs — never
    a superset.

The test-strength result is the one worth reading

The reviewer built its own mutation harness rather than trusting my two
falsifiers, and specifically hunted for a wrong implementation that all seven
shape tuples would miss:

mutant survives all 7 shapes? killed by
correct (control) yes —
reverse past plane no shapes 3–7
swap past/cur no all non-trivial
transpose (batch, heads) no only (2,3,5,2,7), (3,2,1,9,5), (2,2,4,4,1)
plane off-by-one no shapes 3–7

The bh ↔ (b,h) transpose mutant survives every shape with batch == 1 or
heads == 1, and is killed only by the tuples having batch > 1 and
heads > 1. So the grid's load-bearing case is (2,3,5,2,7) — the one that
also has past.seq != cur.seq. Good to know which row is actually doing the
work. The reviewer could not construct a wrong implementation passing all seven.

On the performance methodology

Confirmed bench2.rs faithfully reimplements both the real old and new code,
then looked for inflation and found the confounds run mostly the other way:
both arms pay vec![0.0; N] allocation, zero-fill and page faults every call,
which is shared cost that shrinks the ratio. The one possible inflation — LLVM
eliding the zero-init memset in the new arm, which cannot happen in the opaque
old arm — is worth a few percent of the old arm's time and cannot manufacture a
9–13x gap. The reviewer attributed the gap to write amplification from
stride-512B stores, which is the mechanism claimed.

NITs

  1. msft_attention::concat_cache had no dim-compatibility check, unlike its
    two siblings. Unreachable today (the caller validates at msft:559-563), but
    the reviewer made the sharp observation that my change alters the failure
    mode
    : the old scatter would silently read wrong elements, the new copy
    panics on slice length. Fixed in 73af488 with an assert that names the
    condition instead of failing as an opaque length mismatch.
  2. Serial here vs parallel in MHA for the same operation. Deliberate and
    documented; flagged only so nobody assumes the three kernels are identical.
    Left as-is — see the PR body.
  3. The bit-identity grid never exercised the plane == 0 early return or
    the past == None clone branch. Fixed in 73af488
    (concat_cache_handles_the_absent_past_and_degenerate_shapes).

The reviewer also independently confirmed CLAIM 5 by reading
multi_head_attention.rs:149-152 and :385: MHA's threshold comment ("Decode
steps (seq == 1) never reach it") was written for the transpose helpers, where
dst.len() = batch*heads*seq*dim = 4096 at decode. For concat_cache the
buffer spans the whole past, so at batch=1, heads=32, dim=128 it crosses
MIN_PARALLEL_TRANSPOSE_ELEMENTS once past.seq >= 15. Long-context decode
does take MHA's global-Rayon path.

@justinchuby
justinchuby marked this pull request as ready for review August 22, 2026 16:38
@justinchuby
justinchuby enabled auto-merge (squash) August 22, 2026 16:38
@justinchuby
justinchuby merged commit c977098 into main Aug 22, 2026
17 of 18 checks passed
@justinchuby
justinchuby deleted the seb/concat-cache-memcpy branch August 22, 2026 16:52
@codecov

codecov Bot commented Aug 22, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 80.46%. Comparing base (4431c23) to head (73af488).
⚠️ Report is 7 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1756      +/-   ##
==========================================
+ Coverage   80.40%   80.46%   +0.05%     
==========================================
  Files         410      410              
  Lines      198269   198346      +77     
  Branches   198269   198346      +77     
==========================================
+ Hits       159416   159596     +180     
+ Misses      33467    33360     -107     
- Partials     5386     5390       +4     
Flag Coverage Δ
cli-ort-linux 72.47% <ø> (ø)
cli-ort-windows 72.06% <ø> (+0.09%) ⬆️
mlas 85.33% <ø> (+0.09%) ⬆️
offline 80.60% <100.00%> (+0.06%) ⬆️

Flags with carried forward coverage won't be shown. Click here to find out more.

Files with missing lines Coverage Δ
...rates/onnx-runtime-ep-cpu/src/kernels/attention.rs 88.31% <100.00%> (+0.37%) ⬆️
.../onnx-runtime-ep-cpu/src/kernels/msft_attention.rs 82.41% <100.00%> (+2.10%) ⬆️

... and 12 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.
  • 📦 JS Bundle Analysis: Save yourself from yourself by tracking and limiting bundle sizes in JS merges.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

1 participant