Repository navigation
cpu: copy the KV cache by plane instead of scattering it element by element - #1756
Conversation
`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>
Adversarial review:
|
| 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
msft_attention::concat_cachehad no dim-compatibility check, unlike its
two siblings. Unreachable today (the caller validates atmsft: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.- 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. - The bit-identity grid never exercised the
plane == 0early return or
thepast == Noneclone 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.
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ 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
Flags with carried forward coverage won't be shown. Click here to find out more.
🚀 New features to boost your workflow:
|
concat_cacheinattention.rsandmsft_attention.rscopied the KV cache oneelement at a time with the head-dimension loop outermost, so every store landed
dimfloats past the previous one — a fresh cache line per store on decodeshapes.
multi_head_attention.rswas already converted to plane-sizedcopy_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 afixed
bh = b*heads + h:pastsource[bh*past_plane, bh*past_plane + past_plane)— contiguouscursource[bh*cur_plane, bh*cur_plane + cur_plane)— contiguous[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:
3.4 GB/s → 32.0 GB/s on the
past=1023cell. That cell was also run on agenuinely 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_cacheruns 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_scatterin each file keeps the oldscatter as an independent oracle and compares element-for-element over seven
shapes, including
dim = 1, an empty past, a single past step, and multi-batchmulti-head — the ragged cases where a plane-stride mistake still yields a
plausible-looking buffer.
Verified non-vacuous both ways:
attention.rs)b=1 h=2 past=3 cur=1 d=4msft_attention.rs)b=1 h=1 past=1 cur=1 d=1The oracle indexes the backing
Vecdirectly rather than calling the productionaccessor, so it shares no index arithmetic with the code under test.
Two deliberate choices
par_chunks_mut. MHA has one aboveMIN_PARALLEL_TRANSPOSE_ELEMENTSthat 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 forconcat_cache—data.len()there spans the whole past, so a long-contextdecode step does take its parallel path. Whether parallelising this copy
pays is a separate measured question, not one to settle by porting silently.
at()deleted from both structs. It existed only to feed the scatter andis 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 failedoffline-linuxpackage set — exit 0--check; clippy--all-targets -D warnings;--features mlas;--no-default-features;aarch64-unknown-linux-gnucross-clippy — all cleanCloses #1755