Skip to content

Govern the weight-transpose cache under the memory plan (#1056 item 2) - #1079

Merged
justinchuby merged 4 commits into
mainfrom
squad/1056-transpose-cache-governed
Aug 16, 2026
Merged

justinchuby merged 4 commits into
mainfrom
squad/1056-transpose-cache-governed

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 16, 2026 •

Copy link
Copy Markdown
Owner

Bring the weight-transpose cache under the memory plan (#1056, item 2)

The process-global weight-transpose cache holds one full K x N f32/f16 copy per transposed constant weight for the session (crates/onnx-runtime-ep-cpu/src/kernels/weight_transpose.rs). #1056's rule: any allocation that outlives a single kernel call and scales with weight size must be declared to the plan before allocation, in the bytes actually allocated, and be declinable. Reporting landed in 1ae696c9; this PR finishes the governing half, mirroring #1051 (MLAS packed buffer).

The exact rule for when the kernels transpose (file:line citations)

I read the MatMul and Gemm kernels to find every call into cached_transpose_f32/f16. The populating call sites are:

call site platform gate dtype cached condition
gemm.rs:119 transposed_b(&b, n, k) all platforms f32 (N*K*4) transB != 0 and constant B; B is first widened to dense f32 by MatMulPrepack::dense. Added by #1035.
matmul.rs:1434,1457 transposed_b #[cfg(any(macos, ios))] f32 (N*K*4) Accelerate GEMV / thin-M, constant f32 B
matmul.rs:735 transposed_b_f16 #[cfg(any(macos, ios))] f16 (N*K*2) Accelerate GEMV, contiguous Float16 constant B
fused_matmul_bias.rs:71 transposed_b_f16 #[cfg(any(macos, ios))] f16 (N*K*2) as above

Key consequence: on x86/Windows (non-Apple, no Accelerate) the only path that populates the cache is Gemm with transB and a constant B. Every MatMul transpose call site is compiled out. The predictor's MatMul/FusedMatMulBias arm is therefore #[cfg(any(macos, ios))]-gated too — a binary predicts exactly what its own kernels allocate.

Effect (b) from #1051 — measured, does NOT apply here

The shape-keyed KernelCache instantiates a node once for prefill (m>1) and once for decode (m=1). But every instance keys the global cache on (weight address, K, N), so the decode instance hits the prefill entry and allocates nothing extra. The transpose is held once per weight, not once per instantiation — no per-copy multiplier (unlike the MLAS packed buffer, which retains an owned copy per instance). This is asserted directly by the exactness test, which runs the real kernel at m=4 then m=1 and observes a single copy.

Predicted-vs-actual (ratio 1.00)

Synthetic exactness test predicted_transpose_bytes_equal_actual_after_gemm_execution (gemm.rs) — runs the real GemmKernel through the Kernel trait at prefill (m=4) and decode (m=1) on a constant transB weight, then asserts the predictor equals the bytes the process-global cache actually holds for that weight:

N K predicted actual held ratio
37 91 13,468 B 13,468 B 1.00

This asserts predictor-vs-kernel, not predictor-vs-formula, so a drift between the two would fail it.

Real models (predicted measured on this host via the loader + predictor):

model Gemm nodes MatMul nodes predicted transpose actual (x86 native EP) ratio
qwen2.5-0.5b (f32, 1.98 GB data) 0 169 0 B 0 B 1.00
qwen2.5-0.5b-q4 (int4, 873 MB data) 0 0 0 B 0 B 1.00

Neither exported model contains a single Gemm node, and their MatMul transposes are Apple-only, so the cache is genuinely empty on x86 — predicted 0 matches actual 0. The non-trivial ratio-1.00 proof is the synthetic test above.

RSS measurements (this machine, CPU-only, no CUDA runtime)

Peak working set is per-process and reliable; timing is process CPU time (TotalProcessorTime), not wall clock.

model peak RSS proc CPU time resident_f32_cache_bytes (plan) notes
qwen2.5-0.5b 3,007.7 MB 256.2 s 0 (incl. transpose predictor) fell back to ORT for decode¹
qwen2.5-0.5b-q4 ~90 MB (partial) 1.6 s 0 (incl. transpose predictor) native decoder load failed¹

The memory-strategy plan log confirms the wiring end-to-end on the native path: scope="single_model_native" resident_f32_cache_bytes=0 f32_weight_cache_admitted=true — the transpose predictor (0) is folded into resident_f32_cache_bytes and set_weight_transpose_cache_enabled(true) is applied from the same verdict.

¹ Both exported models failed to load on the native decode backend on this host with model.io.position_ids_input declares port 'position_ids', but the graph exposes [...] when I first measured (pre-rebase). That was a model/IO-metadata incompatibility related to the same area c5385b2f addressed on origin/main; I have not re-measured a full generate post-rebase (the CLI release build is ~50 min on this shared host). Regardless, neither exported model contains a Gemm node, so the transpose cache is empty on x86 and the predictor's 0 is exact for them either way; the authoritative ratio-1.00 proof is the synthetic executor test above.

What happens on decline

set_weight_transpose_cache_enabled(false) (called by the plan when the folded resident total does not fit):

  • MatMulPrepack::transposed_b / transposed_b_f16 return None → the Gemm/MatMul kernels recompute a transient transpose per call (Cow::Owned(transpose_row_major(...)), freed at the end of the call), so no per-kernel-instance Arc is memoized and nothing session-lifetime accrues.
  • cached_transpose_f32/f16 compute and return the transpose without inserting it into the global map, so weight_transpose_cache_bytes() stays put.

Test declined_transpose_cache_retains_nothing_and_is_byte_identical proves both: after a declined run the global byte total is unchanged (retains nothing), and the declined output is byte-identical to the admitted output — declining is a pure performance tradeoff, never a numerical one (both paths call the same transpose_row_major).

Case I could not exercise exactly

The half-typed Gemm case: when both operands are f16/bf16 the node takes try_half_gemm and does not transpose, but whether A is half is not a graph-static property, so the predictor counts every constant transB Gemm at N*K*4. This over-predicts (never under — the #1056-mandated safe direction) for a half Gemm; it is exact for the f32 case. Documented in the predictor's doc comment. No such node exists in the available models, so it did not affect any measurement.

Gates (run locally on this machine, unfiltered, after rebasing on origin/main @ 8635e4db)

  • cargo test -p onnx-runtime-ep-cpu --lib → 1254 passed, 0 failed, 11 ignored (includes the 2 new tests)
  • cargo clippy -p onnx-runtime-ep-cpu --lib -- -D warnings → clean
  • cargo test -p onnx-genai-engine --features native-backend --lib → 496 passed, 0 failed, 1 ignored

Test isolation (review follow-up): the admission flag was originally a bare process-global AtomicBool that a decline test toggled while the matmul::transposed_b tests read it concurrently — passes-alone/fails-in-company. Fixed by making cache_enabled() consult a thread-local override first (only production writes the global) plus a #[cfg(test)] RAII CacheEnabledScope that restores the previous value on drop (even on panic). The decline test now probes the exact (addr, n, k) cache key via a test-only peek rather than a global byte delta. See the PR comment for detail.

Closes part 2 of #1056.

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

@codecov

codecov Bot commented Aug 16, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 79.70297% with 41 lines in your changes missing coverage. Please review.
✅ Project coverage is 80.36%. Comparing base (8635e4d) to head (7bc4698).
⚠️ Report is 4 commits behind head on main.

Files with missing lines Patch % Lines
...nnx-runtime-ep-cpu/src/kernels/weight_transpose.rs 60.00% 19 Missing and 1 partial ⚠️
crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs 62.74% 16 Missing and 3 partials ⚠️
crates/onnx-runtime-ep-cpu/src/kernels/gemm.rs 98.01% 0 Missing and 2 partials ⚠️
Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1079      +/-   ##
==========================================
+ Coverage   79.70%   80.36%   +0.65%     
==========================================
  Files         369      369              
  Lines      160559   161491     +932     
  Branches   160559   161491     +932     
==========================================
+ Hits       127975   129778    +1803     
+ Misses      27857    26970     -887     
- Partials     4727     4743      +16     
Flag Coverage Δ
cli-ort-linux 83.79% <ø> (+0.09%) ⬆️
cli-ort-windows 83.40% <ø> (+0.18%) ⬆️
mlas 83.24% <ø> (-0.16%) ⬇️
offline 80.19% <79.70%> (+0.68%) ⬆️

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

Files with missing lines Coverage Δ
crates/onnx-runtime-ep-cpu/src/kernels/gemm.rs 87.90% <98.01%> (+2.25%) ⬆️
crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs 88.22% <62.74%> (+3.66%) ⬆️
...nnx-runtime-ep-cpu/src/kernels/weight_transpose.rs 91.89% <60.00%> (-2.87%) ⬇️

... and 16 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@github-actions

github-actions Bot commented Aug 16, 2026 •

Copy link
Copy Markdown

⚠️ Benchmark Change Detected

Comparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).

ℹ️ Absolute times are informational only — they vary with runner load. The % change column is the reliable signal because both sides ran under identical conditions.

Status Scenario Base PR Change
⚠️ sampling_latency/top_k_per_token 52.02 µs 64.25 µs +23.5%
⚠️ kv_cache/alloc_dealloc_pages 40.45 µs 49.54 µs +22.5%
⚠️ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.66 ms 6.81 ms +20.4%
⚠️ grammar_masking/llguidance_compute_mask/32 76.67 µs 91.22 µs +19.0%
⚠️ sampling_latency/top_p_per_token 383.84 µs 449.04 µs +17.0%
⚠️ tokenization/decode_tokens_per_second 6.15 ms 7.14 ms +16.1%
⚠️ logit_processing/seven_processor_chain_per_step 327.66 µs 377.64 µs +15.3%
⚠️ qwen3_sampling_processors/top_k_full_sort_baseline 2.14 ms 2.47 ms +15.1%
✅ sampling_latency/min_p_per_token 209.56 µs 240.99 µs +15.0%
✅ qwen3_sampling_processors/top_k_partial_selection 139.69 µs 159.25 µs +14.0%
✅ tokenization/encode_tokens_per_second 395.73 µs 447.95 µs +13.2%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.52 ms 3.98 ms +13.1%
✅ sampling_latency/greedy_per_token 3.19 µs 3.60 µs +12.6%
✅ qwen3_sampling_processors/top_k_top_p_fast 655.96 µs 723.59 µs +10.3%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 531.29 µs 581.20 µs +9.4%
✅ add/large_f32_threads=1-internal/4194304 631.02 µs 651.24 µs +3.2%
✅ gather/small_f16_threads=1-internal/4096 478.3 ns 491.5 ns +2.8%
✅ reduce_mean/small_f32_threads=1-internal/4096 14.85 µs 15.14 µs +1.9%
✅ add/medium_bf16_threads=1-internal/262144 102.12 µs 103.97 µs +1.8%
✅ gather/small_f32_threads=1-internal/4096 667.8 ns 665.7 ns -0.3%
✅ add/large_bf16_threads=1-internal/4194304 1.67 ms 1.66 ms -0.7%
✅ add/medium_f32_threads=1-internal/262144 24.95 µs 24.73 µs -0.9%
✅ add/large_f16_threads=1-internal/4194304 1.65 ms 1.64 ms -1.1%
✅ reduce_mean/medium_f32_threads=1-internal/65536 245.15 µs 242.49 µs -1.1%
✅ gather/small_bf16_threads=1-internal/4096 478.9 ns 473.3 ns -1.2%
✅ add/medium_f16_threads=1-internal/262144 105.10 µs 103.47 µs -1.5%
✅ reduce_mean/large_f32_threads=1-internal/262144 1.01 ms 990.15 µs -1.8%
✅ add/small_f32_threads=1-internal/1024 217.2 ns 211.5 ns -2.6%
✅ matmul/medium_generic_f16_threads=1/32x512x512 32.80 µs 30.63 µs -6.6%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 511.32 µs 476.31 µs -6.8%
✅ add/small_f16_threads=1-internal/1024 507.3 ns 451.6 ns -11.0%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 192.54 µs 168.52 µs -12.5%
✅ add/small_bf16_threads=1-internal/1024 512.1 ns 446.4 ns -12.8%
✅ matmul/small_generic_f32_threads=8/1x256x256 38.48 µs 33.23 µs -13.6%
🟢 matmul/small_generic_f32_threads=1/1x256x256 42.96 µs 36.50 µs -15.0%
🟢 block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 60.16 µs 49.78 µs -17.2%
🟢 matmul/small_generic_f16_threads=1/1x256x256 38.28 µs 31.11 µs -18.7%
🟢 gather/medium_f16_threads=1-internal/32768 2.95 µs 2.39 µs -18.9%
🟢 gather/medium_f32_threads=1-internal/32768 4.68 µs 3.75 µs -19.9%
🟢 gather/medium_bf16_threads=1-internal/32768 3.02 µs 2.38 µs -21.0%
🟢 matmul/large_generic_f16_threads=1/32x1024x1024 100.98 µs 79.60 µs -21.2%
🟢 gather/large_f32_threads=1-internal/131072 37.09 µs 27.41 µs -26.1%
🟢 gather/large_bf16_threads=1-internal/131072 14.58 µs 10.56 µs -27.6%
🟢 matmul/small_generic_bf16_threads=1/1x256x256 44.11 µs 31.71 µs -28.1%
🟢 matmul/large_generic_f32_threads=8/32x1024x1024 5.37 ms 3.82 ms -28.9%
🟢 gather/large_f16_threads=1-internal/131072 14.96 µs 10.42 µs -30.4%
🟢 matmul/large_generic_bf16_threads=8/32x1024x1024 2.09 ms 1.37 ms -34.4%
🟢 block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 922.23 µs 585.69 µs -36.5%
🟢 matmul/medium_generic_f32_threads=1/32x512x512 3.74 ms 2.33 ms -37.6%
🟢 block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 74.39 µs 46.34 µs -37.7%
🟢 matmul/large_generic_bf16_threads=1/32x1024x1024 3.16 ms 1.96 ms -38.0%
🟢 matmul/medium_generic_bf16_threads=1/32x512x512 876.76 µs 527.09 µs -39.9%
🟢 matmul/medium_generic_f32_threads=8/32x512x512 1.59 ms 912.62 µs -42.4%
🟢 matmul/large_generic_f32_threads=1/32x1024x1024 16.49 ms 9.40 ms -43.0%
🟢 matmul/medium_generic_f16_threads=8/32x512x512 67.78 µs 30.46 µs -55.1%
🟢 matmul/small_generic_f16_threads=8/1x256x256 89.99 µs 31.16 µs -65.4%
🟢 matmul/large_generic_f16_threads=8/32x1024x1024 273.29 µs 85.11 µs -68.9%
🟢 matmul/small_generic_bf16_threads=8/1x256x256 127.22 µs 31.86 µs -75.0%
🟢 matmul/medium_generic_bf16_threads=8/32x512x512 1.58 ms 370.00 µs -76.6%

Visual flags: ⚠️ ≥ 15% slower, 🔴 ≥ 30% slower — calibrated against measured runner noise (~27% worst-case on multi-threaded matmul)

Host info
CPU: Apple M1 (Virtual)
Cores: 3
OS: Darwin 25.5.0 arm64
Rust: rustc 1.97.1 (8bab26f4f 2026-07-14)
Load avg: { 3.75 3.45 4.40 }
What this cannot catch
  • Regressions in code paths not covered by these benchmarks (e.g., end-to-end decode with a real model)
  • Sub-threshold regressions that compound over multiple PRs
  • Performance changes that only manifest under GPU execution
  • Latency changes in the ORT integration path (these benchmarks exercise the native Rust kernels)

@justinchuby

Copy link
Copy Markdown
Owner Author

The design is right and the rule you derived is correct -- I verified it independently. But the PR introduces a process-global mutable flag that tests flip while other tests run, and it is already breaking the suite. Requesting changes.

The failure

cargo test -p onnx-runtime-ep-cpu --lib transpose on your branch merged with current origin/main:

kernels::matmul::tests::transposed_b_at_a_reused_address_never_serves_another_tensors_transpose ... FAILED kernels::matmul::tests::transposed_b_handles_zero_sized_weights ... FAILED

Both pass in isolation:

cargo test -p onnx-runtime-ep-cpu --lib \ kernels::matmul::tests::transposed_b_at_a_reused_address_never_serves_another_tensors_transpose -- --exact test result: ok. 1 passed; 0 failed

Passing alone and failing in the suite is the signature of shared process state, not of a wrong assertion. The mechanism is visible in your own diff: gemm.rs:821 calls set_weight_transpose_cache_enabled(false) inside a test, and matmul.rs's tests run concurrently in the same process against the same AtomicBool. Your declined_transpose_cache_retains_nothing_and_is_byte_identical turns the cache off; whichever transposed_b test is mid-flight then takes the new if !weight_transpose::cache_enabled() { return None; } early return and fails.

Your own gate reported this to you: the summary says "engine 494 passed (2 pre-existing unrelated failures)" -- those two were mine, fixed on main in c5385b2f, so rebasing will clear them. But the ep-cpu suite failing was not pre-existing and needs fixing here.

This is the third time today that process-global state has produced a test that passes alone and fails in company (#983's ORT handle cache, #1033's fp16 GEMV). Worth treating as a known hazard in this repo rather than an accident.

What I want

  1. Serialize or scope the flag. Either take a test-local mutex around every test that flips it (and restore the previous value on the way out, including on panic -- a guard type, not a bare set/reset pair), or give the decline path a scoped form so a test never mutates process state at all. The second is better if it is cheap: a global that only production sets is much harder to misuse than one tests toggle.
  2. Re-run the full cargo test -p onnx-runtime-ep-cpu --lib and report the count. Your 1216-passed figure predates this interaction or was taken with a filter that excluded one side of it.
  3. Rebase on current main to clear the two engine failures.

What I verified and agree with

The rule is right. I checked every call site rather than taking the citation on trust:

  • gemm.rs:772 -- the only unconditional populator.
  • matmul.rs:859, matmul.rs:1558, fused_matmul_bias.rs:71 -- all inside #[cfg(any(target_os = "macos", target_os = "ios"))].
  • matmul.rs:285/312 -- Apple-only prewarm helpers.
  • transposed_b / transposed_b_f16 carry allow(dead_code, reason = "consumed only by the Apple Accelerate paths") on non-Apple, which matches.

So on x86/Windows the only populator is Gemm + transB + constant B, and the qwen models having zero Gemm nodes explains predicted = actual = 0 there.

The ratio-1.00 evidence is the right shape, and the finding that the cache keys on (addr, K, N) -- so prefill and decode share one copy and there is no shape-keyed multiplier here, unlike MLAS in #1051 -- is exactly the thing I asked you to measure rather than assume. Good that you checked it in both directions.

The over-prediction case is handled correctly: half-typed Gemm skipping the transpose is unpredictable from graph-static dtypes, so over-predicting at N*K*4 is the safe direction, and you said so instead of quietly approximating.

Fix the test isolation and I will merge.

Copilot AI added 3 commits August 16, 2026 09:56
…1056)

The process-global weight-transpose cache holds one full K x N f32/f16 copy
per transposed constant weight for the session. #1056 requires every
session-lifetime, weight-scaled allocation to be declared to the memory plan
before it is allocated, in the bytes actually allocated, and to be declinable.
Reporting (TransposeCache::bytes / weight_transpose_cache_bytes) landed in
1ae696c; this finishes the governing half.

- Predictor: weight_transpose_cache_predicted_bytes(&Graph) mirrors the exact
  kernel decision. On all platforms it counts constant-weight Gemm nodes with
  transB != 0 (gemm.rs transposed_b, added in #1035) at N*K*4 f32 bytes; the
  MatMul/FusedMatMulBias transpose call sites are cfg(macos/ios)-only, so the
  predictor's MatMul arm is likewise cfg-gated -- a binary predicts exactly
  what its own kernels allocate. The shape-keyed kernel cache instantiates a
  node once per activation shape, but every instance keys the global cache on
  (weight address, K, N), so the second instantiation hits the first entry:
  the transpose is held once per weight, not once per instantiation (no #1051
  per-copy multiplier here). Verified by a test that runs the real Gemm kernel
  at prefill (m=4) and decode (m=1) and asserts predicted == bytes held, 1.00.

- Admission gate: set_weight_transpose_cache_enabled(bool). When declined,
  MatMulPrepack::transposed_b / transposed_b_f16 return None (kernels recompute
  a transient transpose per call, freed each call) and cached_transpose_f32/f16
  compute without inserting -- nothing session-lifetime accrues. Declining is a
  pure performance tradeoff: the transpose is byte-identical either way.

- Wiring: engine load.rs folds the predicted bytes into resident_f32_cache_bytes
  and calls set_weight_transpose_cache_enabled beside the resident dequant f32
  cache (#987) and MLAS SQNBit packed buffer (#1051) gates, so one plan verdict
  governs all three.

Gates: cargo test -p onnx-runtime-ep-cpu --lib (1216 passed, 0 failed, 11
ignored); cargo clippy -p onnx-runtime-ep-cpu --lib -D warnings clean; engine
--features native-backend --lib green except two pre-existing native_decode IO
failures unrelated to this change (confirmed failing on origin/main).

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
…l harness

The #1056 admission flag was a bare process-global AtomicBool that the gemm
decline test toggled with set_weight_transpose_cache_enabled(false) while the
matmul transposed_b tests read it concurrently in the same process, so those
tests took the new declined early-return and failed in company (passing alone).
This is the same process-global-races-the-harness trap as #983 and #1033.

Make the decline decision consult a thread-local override first; only
production writes the global. Add a #[cfg(test)] RAII CacheEnabledScope that
sets the override on the current thread and restores the previous value on drop
(including on panic), so a test's decline can never leak to another thread.
Rework the decline test to probe the exact (addr, k, n) key via a test-only
f32_cache_contains peek instead of a global byte delta, so a concurrent test
caching an unrelated weight cannot mask a leak.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
…s it

The decline test probed f32_cache_contains(b_ptr, k, n), but the Gemm kernel
installs the entry via transposed_b(&b, n, k), so the resident key is
(b_ptr, n, k). Match that order so the admitted-run assertion sees the entry.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
@justinchuby
justinchuby force-pushed the squad/1056-transpose-cache-governed branch from e4d2a27 to 38af02e Compare August 16, 2026 17:33
@justinchuby

Copy link
Copy Markdown
Owner Author

Addressed — thanks for the precise diagnosis. The mechanism was exactly as you described: the admission flag was a bare process-global AtomicBool, and the gemm decline test flipped it while the matmul::tests::transposed_b_* tests read it concurrently, so whichever transposed_b test was mid-flight took the new declined early-return and failed. Passes-alone-fails-in-company, same shape as #983 / #1033.

Fix (your preferred option 1 — scope so tests never mutate process state):

  • cache_enabled() now consults a thread-local override first, falling back to the process-global. Only production (set_cache_enabled, called once at load from the plan) writes the global; the global is now something tests cannot reach.
  • Added a #[cfg(test)] RAII CacheEnabledScope that sets the override on the current thread and restores the previous value on drop, including on panic — so a failing test cannot leave the flag flipped for anyone. It's thread-local, so a decline on this thread does not race the transposed_b tests on other threads.
  • Reworked the decline test to probe the exact (addr, n, k) key the Gemm kernel installs (via a new test-only f32_cache_contains peek) instead of a global byte-total delta, so a concurrent test caching an unrelated weight cannot mask a leak. (The kernel calls transposed_b(&b, n, k), so the resident key is (addr, n, k) — noted in the test.)

No production code path changed; production still reads the global on every worker thread.

Rebased on current origin/main (8635e4db). The two native_decode failures were indeed yours, fixed by c5385b2f — after the rebase the engine suite is clean.

Unfiltered gate counts, re-run on this machine after the fix:

gate result
cargo test -p onnx-runtime-ep-cpu --lib (full, unfiltered) 1254 passed; 0 failed; 11 ignored
cargo clippy -p onnx-runtime-ep-cpu --lib -- -D warnings clean
cargo test -p onnx-genai-engine --features native-backend --lib 496 passed; 0 failed; 1 ignored

I also confirmed cargo test -p onnx-runtime-ep-cpu --lib transpose (the filter that exercises both sides together) is 53 passed / 0 failed, and that both new tests still pass with the two transposed_b tests in company.

Carrying your note forward: any process-global mutable state in this repo should assume the parallel harness will race it, and the test access should be thread-local/RAII-scoped from the start.

@justinchuby

Copy link
Copy Markdown
Owner Author

Update: the isolation fix works -- the two transposed_b tests no longer fail in company. But the new decline test is itself flaky, and I can reproduce it failing on an unmodified tree.

Same commit 38af02ed, clean working tree, no patches applied:

run result
full suite, first time 1254 passed; 0 failed
full suite, later 1253 passed; 1 failed
that test in isolation (-- --exact) 0 passed; 1 failed

Failing assertion is the first one, gemm.rs:832 -- "a declined transpose cache must retain nothing for this weight". Under decline, f32_cache_contains(b_ptr, n, k) returns true. My earlier green run was luck.

Before blaming the test I falsified it, because a test that cannot fail is worse than no test. Removing the decline gate in MatMulPrepack::transposed_b alone: still passes. Removing the gate in cached_transpose_f32 alone: still passes. Removing both: fails with exactly that assertion.

That is a good finding about the design -- the two gates are genuine defence in depth, either one sufficient -- and it establishes the test can detect a broken decline. So something is really populating the global cache for that key under decline, at least sometimes.

Two candidate mechanisms, which need different fixes:

  1. The transpose runs on a different thread from the scope. The override is thread-local; any part of the Gemm path executing off the calling thread reads the still-true global and caches. That would make the failure depend on whether a worker pool happened to be warm.
  2. Address reuse. The cache keys on (addr, K, N) and weight_transpose.rs already documents this hazard -- clear_weight_transpose_caches exists for it. A fresh Vec can land on a recycled address of a previously cached weight with the same dimensions.

The distinction matters more than the test going green: if a worker thread can bypass the thread-local override in tests, something analogous may bypass the global verdict in production -- for instance a transpose performed before the plan's set_cache_enabled is applied at load. That would mean the decline is incomplete, not merely the test.

Sent back for root-cause with evidence, and for five consecutive full-suite runs rather than one. A single green run does not establish stability for a test now known to be order- and timing-sensitive; I made exactly that mistake an hour ago, which is why I re-ran.

The decline/exactness tests peek the process-global cache by (addr, K, N).
An address only names a weight while it is live; once freed, the allocator can
recycle it for an unrelated buffer of the same dims, leaving a stale entry that
makes f32_cache_contains a false positive for a weight that never went through
the cache. That is the intermittent 'a declined transpose cache must retain
nothing' failure under the parallel harness.

Instrumented the gate sites (thread id, override, global, key) and confirmed the
production decline is complete: during decline cached_transpose_f32 is never
reached and before/after are both false; every insert observed carries
override=Some(true), i.e. the admitted phase. So this is a test-only false
positive, not a decline gap.

Add a #[cfg(test)] TransposeCache::remove and a f32_cache_evict helper, and evict
this weight's exact key at the start of both tests so the 'before' state is
deterministically empty. Safe under the harness: no live concurrent allocation
can share the address, so the only entry evicted is a stale, unused one.

Full suite run five times back to back: 1254 passed / 0 failed each time.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
@justinchuby

Copy link
Copy Markdown
Owner Author

Root-caused with instrumentation, not reasoning — thanks for insisting on that, it settled the important question (production vs. test).

What I measured

I instrumented both gate sites to print, on every cached_transpose_f32 insert, the source pointer, key (k, n), thread id, thread-local override, and the global flag; and the decline assertion to print the probed b_ptr, thread id, and f32_cache_contains before and after the declined run. Then I ran the isolated decline test 40× in a loop. Representative line pair (identical every iteration):

decline assert: b_ptr=0x1f9d97fc100 n=29 k=53 thread=ThreadId(2) before=false after=false
cached_transpose_f32 INSERT-PATH src_ptr=0x1f9d97fc100 k=29 n=53 thread=ThreadId(2) override=Some(true) global=true

Two facts fall straight out:

  1. before=false after=false during decline — the declined run inserts nothing. The only insert in the whole test carries override=Some(true), i.e. it is the admitted phase, and it is on the same thread (ThreadId(2)) that set the scope. No insert ever happens with override=Some(false), and none happens on a worker thread.

So mechanism 1 (a worker thread bypassing the thread-local override) is ruled out, and — the distinction you flagged as mattering most — the production decline is complete: under decline, cached_transpose_f32 is never even reached (transposed_b returns None on the calling thread and the kernel takes the transient Cow::Owned path). Nothing analogous can bypass the global verdict in production either, because the same calling-thread check gates the only populating call site.

The real cause: mechanism 2 (address reuse), and it is test-only

f32_cache_contains(b_ptr, n, k) keys on (addr, K, N). An address only names a weight while that weight is live; once a weight is freed, the allocator can hand b_ptr to an unrelated buffer of the same dims. If an earlier (now-freed) weight had been cached under (that same addr, 29, 53), the probe answers true for a buffer that never went through the cache — exactly the intermittent a declined transpose cache must retain nothing failure. weight_transpose.rs already documents this recycling hazard; my byte-total-free "probe one specific key" rewrite traded the concurrent-growth problem for the reuse problem.

Your "remove both gates → fails, either gate alone → passes" falsification is fully consistent with this: the assertion is genuinely capable of catching a broken decline; it is the starting state that was non-deterministic.

Fix (deterministic, not a retry or a relaxed assertion)

  • Added #[cfg(test)] TransposeCache::remove and a f32_cache_evict(ptr, k, n) helper.
  • Both the decline and exactness tests now evict this weight's exact key before probing, so the "before" state is deterministically empty. This is safe under the parallel harness precisely because no live concurrent allocation can share b_ptr — the only entry eviction can touch is a stale one nobody is using. The decline test also asserts the precondition (key starts absent) so the eviction can't silently mask a regression.
  • Removed the instrumentation.

Production code is unchanged; this commit is test-only.

Stability — five consecutive full runs

cargo test -p onnx-runtime-ep-cpu --lib, back to back on this machine:

run result wall
1 1254 passed; 0 failed; 11 ignored 214.8s
2 1254 passed; 0 failed; 11 ignored 421.6s
3 1254 passed; 0 failed; 11 ignored 529.8s
4 1254 passed; 0 failed; 11 ignored 363.0s
5 1254 passed; 0 failed; 11 ignored 449.1s

Plus the two transpose tests looped 120× in isolation: 0 failures. cargo clippy -p onnx-runtime-ep-cpu --lib -- -D warnings clean. (The wall-clock spread is the shared-host noise you noted; the pass/fail and counts are invariant.)

Pushed as 7bc4698c on the same branch.

@github-actions

Copy link
Copy Markdown

🔴 Benchmark Regression Detected

Comparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).

ℹ️ Absolute times are informational only — they vary with runner load. The % change column is the reliable signal because both sides ran under identical conditions.

Status Scenario Base PR Change
🔴 matmul/small_generic_f32_threads=8/1x256x256 32.42 µs 57.24 µs +76.6%
🔴 matmul/small_generic_f16_threads=1/1x256x256 27.91 µs 45.70 µs +63.8%
🔴 matmul/small_generic_f16_threads=8/1x256x256 29.69 µs 48.57 µs +63.6%
🔴 gather/large_f32_threads=1-internal/131072 24.33 µs 38.52 µs +58.3%
🔴 gather/large_f16_threads=1-internal/131072 12.26 µs 18.91 µs +54.2%
🔴 matmul/medium_generic_f32_threads=8/32x512x512 900.72 µs 1.36 ms +50.8%
🔴 gather/large_bf16_threads=1-internal/131072 10.70 µs 15.11 µs +41.2%
🔴 matmul/small_generic_bf16_threads=8/1x256x256 28.90 µs 39.28 µs +35.9%
🔴 gather/medium_f16_threads=1-internal/32768 2.24 µs 2.95 µs +31.5%
⚠️ matmul/medium_generic_bf16_threads=8/32x512x512 362.95 µs 467.23 µs +28.7%
⚠️ gather/medium_f32_threads=1-internal/32768 3.41 µs 4.38 µs +28.6%
⚠️ matmul/large_generic_f32_threads=8/32x1024x1024 4.14 ms 5.07 ms +22.4%
⚠️ matmul/medium_generic_bf16_threads=1/32x512x512 482.62 µs 585.65 µs +21.3%
⚠️ matmul/small_generic_f32_threads=1/1x256x256 35.07 µs 42.29 µs +20.6%
⚠️ matmul/small_generic_bf16_threads=1/1x256x256 29.37 µs 35.27 µs +20.1%
⚠️ gather/medium_bf16_threads=1-internal/32768 2.20 µs 2.64 µs +19.9%
✅ add/large_f32_threads=1-internal/4194304 644.56 µs 730.71 µs +13.4%
✅ reduce_mean/small_f32_threads=1-internal/4096 13.81 µs 15.64 µs +13.2%
✅ gather/small_f32_threads=1-internal/4096 628.7 ns 700.8 ns +11.5%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.13 ms 2.37 ms +11.2%
✅ gather/small_f16_threads=1-internal/4096 457.0 ns 507.1 ns +11.0%
✅ sampling_latency/min_p_per_token 192.95 µs 212.44 µs +10.1%
✅ add/small_f16_threads=1-internal/1024 419.5 ns 460.3 ns +9.7%
✅ gather/small_bf16_threads=1-internal/4096 469.4 ns 513.6 ns +9.4%
✅ reduce_mean/large_f32_threads=1-internal/262144 964.11 µs 1.05 ms +9.0%
✅ qwen3_sampling_processors/top_k_top_p_fast 618.49 µs 671.11 µs +8.5%
✅ logit_processing/seven_processor_chain_per_step 313.20 µs 336.88 µs +7.6%
✅ add/small_bf16_threads=1-internal/1024 419.4 ns 450.5 ns +7.4%
✅ add/medium_f32_threads=1-internal/262144 23.70 µs 25.42 µs +7.3%
✅ add/large_bf16_threads=1-internal/4194304 1.58 ms 1.68 ms +6.5%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 8.76 ms 9.31 ms +6.3%
✅ kv_cache/alloc_dealloc_pages 37.14 µs 39.44 µs +6.2%
✅ add/medium_bf16_threads=1-internal/262144 95.60 µs 101.42 µs +6.1%
✅ reduce_mean/medium_f32_threads=1-internal/65536 245.13 µs 258.68 µs +5.5%
✅ tokenization/encode_tokens_per_second 352.26 µs 371.44 µs +5.4%
✅ sampling_latency/top_p_per_token 353.54 µs 366.54 µs +3.7%
✅ sampling_latency/top_k_per_token 48.57 µs 50.32 µs +3.6%
✅ add/medium_f16_threads=1-internal/262144 104.67 µs 106.97 µs +2.2%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 80.06 µs 81.35 µs +1.6%
✅ qwen3_sampling_processors/top_k_partial_selection 132.26 µs 134.16 µs +1.4%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 1.95 ms 1.98 ms +1.2%
✅ matmul/medium_generic_f16_threads=1/32x512x512 30.64 µs 30.79 µs +0.5%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.29 ms 3.30 ms +0.4%
✅ block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 522.27 µs 524.06 µs +0.3%
✅ add/small_f32_threads=1-internal/1024 199.3 ns 200.0 ns +0.3%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 457.93 µs 457.16 µs -0.2%
✅ sampling_latency/greedy_per_token 3.03 µs 3.03 µs -0.2%
✅ tokenization/decode_tokens_per_second 5.69 ms 5.67 ms -0.2%
✅ matmul/large_generic_bf16_threads=8/32x1024x1024 1.26 ms 1.26 ms -0.3%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 484.58 µs 480.77 µs -0.8%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 161.08 µs 157.43 µs -2.3%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 83.42 µs 80.84 µs -3.1%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.39 ms 5.20 ms -3.5%
✅ block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 43.96 µs 42.40 µs -3.6%
✅ matmul/medium_generic_f16_threads=8/32x512x512 29.89 µs 28.76 µs -3.8%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 1.97 ms 1.84 ms -6.5%
✅ add/large_f16_threads=1-internal/4194304 1.88 ms 1.73 ms -7.9%
✅ grammar_masking/llguidance_compute_mask/32 88.78 µs 78.89 µs -11.1%
✅ block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 43.30 µs 38.38 µs -11.4%

Visual flags: ⚠️ ≥ 15% slower, 🔴 ≥ 30% slower — calibrated against measured runner noise (~27% worst-case on multi-threaded matmul)

Host info
CPU: Apple M1 (Virtual)
Cores: 3
OS: Darwin 25.5.0 arm64
Rust: rustc 1.97.1 (8bab26f4f 2026-07-14)
Load avg: { 5.88 3.85 4.72 }
What this cannot catch
  • Regressions in code paths not covered by these benchmarks (e.g., end-to-end decode with a real model)
  • Sub-threshold regressions that compound over multiple PRs
  • Performance changes that only manifest under GPU execution
  • Latency changes in the ORT integration path (these benchmarks exercise the native Rust kernels)

@justinchuby

Copy link
Copy Markdown
Owner Author

Root cause accepted, independently re-verified, merging.

Three consecutive full runs on 7bc4698c, my machine, unfiltered:

run result
1 1254 passed; 0 failed; 11 ignored
2 1254 passed; 0 failed; 11 ignored
3 1254 passed; 0 failed; 11 ignored

Together with your five, that is eight clean runs of a test that previously failed roughly every other time. Good enough.

The root cause is the right one and the evidence supports it. The instrumentation showing before=false after=false during decline, with the only insert carrying override=Some(true) on the same thread that set the scope, rules out the worker-thread hypothesis directly rather than by argument. That was the answer I most wanted, because it settles the question I actually cared about: the production decline is complete. Under decline cached_transpose_f32 is never reached at all, so nothing bypasses the plan's verdict at runtime either.

So the defect was a false positive in the probe, not a leak: (addr, K, N) keying plus allocator address reuse meant a stale entry could answer "resident" for a weight this run never cached. Evicting the exact key before probing makes the starting state deterministic without weakening the assertion — and that is the important part. A retry or a relaxed assertion would have made the symptom go away while leaving the test unable to distinguish "we cached nothing" from "we happened to look at a recycled address".

Worth noting how consistent this is with the rest of the day: clear_weight_transpose_caches already exists in this file because address reuse across model lifetimes is a known hazard here, documented in its own doc comment. The hazard was known; the test just did not account for it.

Summary of what lands

That closes item 2 of #1056. The remaining item is the audit for a fourth such buffer.

@justinchuby
justinchuby merged commit cefbcc5 into main Aug 16, 2026
11 of 17 checks passed
@justinchuby
justinchuby deleted the squad/1056-transpose-cache-governed branch August 16, 2026 19:43
justinchuby added a commit that referenced this pull request Aug 16, 2026
## What

`cargo fmt --all -- --check` currently **fails on `origin/main`**
(`bce03cabb`) in five
places across `crates/onnx-runtime-ep-cpu/src/kernels/gemm.rs` and
`crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs`. This is the `cargo
fmt --all` output
and nothing else.

## Why it happened

Nobody wrote badly-formatted code. #1073, #1079 and #1080 each touched
these two files and
each was fmt-clean against its own base. Squash-merging them produced a
combined text that
rustfmt formats differently:

- `TRANSPOSE_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner())` is 61
characters, which
exceeds rustfmt's default `chain_width` of 60 once it sits at test-body
indentation (two
  sites).
- The `onnx_runtime_ir` import list grew past `max_width` and now wants
braces on their own
  lines.
- One `Gemm`/`transB` attribute chain became short enough to fit on a
single line after a
  neighbouring edit.
- One stray double blank line at end of a `mod tests`.

`main` is unprotected, so no required check re-ran fmt on the merge
result and the breakage
landed silently.

## How it was found

It blocked #1086: that PR's `Fast (Linux x86_64)` and `Rust quality`
jobs failed on the
PR **merge ref** with diffs in files #1086 does not touch. Reproduced
independently by
checking out `origin/main` into a clean worktree and running `cargo fmt
--all -- --check`.

## Verification

- `cargo fmt --all -- --check` -> clean (was: 5 diffs).
- Diff is whitespace/line-breaking only; `git diff -w` on the two files
is empty apart from
  the import-brace move. No logic, no behaviour, no test changes.

## Risk

None. Mechanical formatter output.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 17, 2026
#1100)

## #1056: bring `MatMulPrepack::dense` under the memory plan

Fourth resident, weight-scaled buffer brought under the memory-strategy
plan after the resident dequant f32 cache (#987), the MLAS SQNBit packed
buffer (#1051), and the weight-transpose cache (#1079). Follows #1056's
rule: *any allocation that outlives a single kernel call and scales with
weight size must be declared to the plan before it is allocated, in the
bytes actually allocated, and must be declinable.*

### The exact condition under which the kernel caches (verified)

`MatMulPrepack::dense(index, view)` in
`crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs` caches a
session-lifetime `Vec<f32>` iff **both**:

1. the operand is a constant initializer (`constant_inputs[index]`), and
2. `to_dense_f32_widen("MatMul", view)` returns `Cow::Owned` — i.e. the
operand is **not** already a contiguous f32 view.

`to_dense_f32_widen` (`crates/onnx-runtime-ep-cpu/src/dtype.rs:806`)
borrows contiguous f32 zero-copy (`Cow::Borrowed`, no cache) and
allocates an owned `4*numel` f32 copy for every other float case:
f16/bf16/f64, or a strided/column-major f32. So a contiguous f32
constant costs **nothing** (this is why the buffer stayed invisible on
the int4 and f32-contiguous models we exercise most), while an
f16/bf16/f64 or non-contiguous constant costs a permanent `4*K*N` per
kernel instance.

### The shape-keyed instantiation multiplier (×2) — corrected after
review

The executor's kernel cache is **shape-keyed** (`KernelKey { node,
resolved_input_shapes }`,
`crates/onnx-runtime-session/src/executor/kernel_cache.rs`). A decoder
instantiates each `MatMul` node at **two** activation shapes — prefill
(`m>1`) and decode (`m==1`) — as **separate** `MatMulKernel`s, each with
its own `MatMulPrepack::dense`. Unlike the weight-transpose cache
(process-global, keyed on the weight address, so a second instantiation
reuses the first), `dense` is per-instance.

The only case that populates `dense` is a constant non-f32 `B` paired
with an f32 `A` (a same-half `B`+`A` takes `try_matmul_half`, and both
the `m==1` GEMV and the MLAS half-prefill fast paths require `A` to be
half, so they are skipped when `A` is f32). In that case **both** the
prefill instance and the decode instance take the generic/direct-f32
GEMM that widens `B`, so **each retains its own `4*K*N` copy** for the
session. The resident footprint is therefore `2 × 4·K·N`, not `4·K·N`.

The first cut of this PR counted **one** copy — an **under**-prediction
by 2×, the exact under-reporting defect #1051 corrected for the MLAS
packed buffer (it reported 247 MB where steady state grew to ~592 MB,
the missing factor being this same prefill/decode shape-keyed doubling).
Fixed by adding `MATMUL_DENSE_DECODE_INSTANTIATIONS = 2` (mirroring
#1051's `MLAS_PACKED_DECODE_INSTANTIATIONS`) and multiplying the
per-node prediction by it. **Measured, not reasoned**: the ratio test
below instantiates prefill + decode and counts two live copies.

### What was built

- **Predictor** `matmul_dense_cache_predicted_bytes(&Graph)` — mirrors
the kernel condition per node (`MatMul` + `FusedMatMulBias`, both
operand indices). A graph initializer is contiguous (`WeightRef` carries
only dtype + dims, no strides), so from the graph the condition reduces
to *"a constant operand whose dtype is a non-f32 float"* → `4*numel`,
**× `MATMUL_DENSE_DECODE_INSTANTIATIONS`** for the prefill+decode
instances; a contiguous f32 constant → 0.
- **`GovernedWeightCache<f32>`** — `dense: [OnceLock<Vec<f32>>; 2]` →
`[GovernedWeightCache<f32>; 2]`. Declined → `dense()` widens transiently
per call and frees it. The only extension needed was a read-only
`filled()` accessor (reuse an existing session copy without ever running
a builder); the "array of two" and the `f32` element type were already
expressible, so nothing else changed in that type.
- **Plan wiring** in `crates/onnx-genai-engine/src/engine/load.rs`,
beside `set_resident_dequant_f32_cache_enabled` /
`set_mlas_sqnbit_packing_enabled` /
`set_weight_transpose_cache_enabled`, folding the prediction into
`resident_f32_cache_bytes` and gating on the same
`f32_weight_cache_admitted` verdict — one decision governs all four
buffers.
- **Decline path** — process-global `AtomicBool` written only by
production (`set_matmul_dense_cache_enabled`), with a test-only
thread-local RAII `DenseCacheEnabledScope` that restores on drop
including on panic. No process-global mutable state tests mutate
(#983/#1033/#1079).

### Predicted vs actual (ratio 1.00) — now across both instantiations

Proven in one run by
`predicted_dense_bytes_equal_actual_after_matmul_execution`: it computes
the plan's graph prediction, then — mirroring the shape-keyed kernel
cache — builds a **separate** `MatMulKernel` for the prefill shape
(`m=4`) and the decode shape (`m=1`), executes each with a constant
**f16** `B` and an **f32** `A`, asserts each retains a copy, and
compares the **summed** `live_bytes()` against the prediction.

| case | K×N | instantiations retaining a copy | per-copy | predicted
(×2) | summed actual | ratio |
|---|---|---|---|---|---|---|
| f16 constant B, f32 A | 48×33 | 2 (prefill m=4 + decode m=1) | 6336 B
| 12672 B | 12672 B | **1.00** |

The test asserts `summed_actual == predicted` and `instances ==
MATMUL_DENSE_DECODE_INSTANTIATIONS`, so a dropped multiplier would fail
it — the single-instantiation version could not.

### What happens on decline

`declined_dense_cache_retains_nothing_and_is_byte_identical` runs one
kernel instance admitted vs declined:

| arm | bytes held after run | output |
|---|---|---|
| admitted | `40*24*4` = 3840 B (one copy) | reference |
| declined | **0 B** | **byte-identical** |

Declining is a pure performance tradeoff: the widened f32 is
byte-identical whether cached or recomputed. Because the transient widen
is freed at the end of each call, no session-lifetime footprint accrues
when declined — and with the ×2 multiplier, admitting an f16-weight
model now costs the plan the full `2×` it will actually hold.

### Peak RSS before/after — no available model populates this cache

Checked every model in `C:\Users\justinchu\dev\models` for plain
`MatMul` nodes with a constant operand:

- `qwen2.5-14b-f32`, `qwen2.5-14b-onnx`, and all `qwen*`/`qwen05b*`
variants: **0 plain MatMul nodes** — every weight routes through
`MatMulNBits` (+ `GatherBlockQuantized`).
- `gemma-3-27b-onnx`: exactly **1 plain MatMul**, and **both its
operands are activations** (there is a `Transpose` feeding it), so
`constant_inputs = [false, false]` and it caches nothing.

So there is no local model that genuinely populates `dense`, and a CLI
A/B would truthfully report predicted = actual = 0 — which proves
nothing. Per #1056 I therefore prove predicted==actual and
byte-identical-on-decline on a **synthetic graph** (the tests above)
rather than reporting a meaningless zero. An f16 peak-RSS A/B awaits an
export that uses plain `MatMul` with f16/non-contiguous constant
weights.

### Anything not predicted exactly (over-predict, documented)

The predictor **over-predicts, never under-predicts**, with documented
directions:

1. **Same-half packed path**: when both operands are the same half dtype
and contiguous, the node takes `try_matmul_half` and never calls
`dense`, so the true cost is 0 while the predictor still counts the
constant operand. Whether the *other* operand is half is not a
graph-static property of the constant one, so counting is the safe
direction.
2. **Prefill-only / multi-shape workloads**: a prefill-only run
instantiates one copy, so the ×2 over-estimates there (safe — the gate
declines sooner). Conversely a session presenting *more* than two
distinct activation shapes to a node (several prompt lengths across
`generate` calls) could instantiate more than two; the multiplier
follows #1051's convention of bounding the autoregressive decode
workload at prefill + decode, and that residual is the same documented
class as #1051's.
3. **Residual gap — non-contiguous f32 constant** (a column-major
weight): the kernel *would* cache it, but `WeightRef` exposes no
strides, so from the graph it is indistinguishable from a contiguous f32
constant and is **not** counted — a potential under-prediction. No such
operand occurs in the models this repo exercises: the one documented
column-major weight (the lm_head projection, per `contiguous_b_f16`'s
doc) is **f16**, so it *is* counted via the dtype rule. Documented in
the predictor rather than over-counting every f32 weight (which would
inflate the budget on the overwhelmingly common contiguous-f32 case). If
f32 column-major MatMul weights ever appear, the loader must expose
their layout for the predictor to see them.

### Gates

- `cargo test -p onnx-runtime-ep-cpu --lib` — **five consecutive runs**,
all `1299 passed; 0 failed; 11 ignored` (46.0 / 49.9 / 52.6 / 52.7 /
49.6 s). Includes the 4 new tests (3 in `matmul`, 1 `filled()` in
`governed_weight_cache`).
- `cargo test -p onnx-genai-engine --features native-backend --lib` —
`496 passed; 0 failed; 1 ignored`.
- `cargo clippy -p onnx-runtime-ep-cpu --lib -- -D warnings` — clean.

Closes #1056.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624

---------

Co-authored-by: justinchuby <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
justinchuby added a commit that referenced this pull request Aug 17, 2026
## Share one MLAS SQNBit packed buffer per weight (#1056)

Refs #1027 (MLAS SQNBit route), #1051 (packed-buffer accounting), #1056
(this dedup).

### The problem

`MatMulNBits` int4 `accuracy_level=0` nodes route to MLAS SQNBit
CompFp32 (#1027). The
executor's `KernelCache` is **shape-keyed**, so each node compiles two
kernel instances —
prefill (`m > 1`) and decode (`m == 1`) — and before this change **each
instance packed its
own full copy of the same constant weight**. The resident packed
footprint was therefore `2x`
the single-copy cost, held for the whole session.

### The fix

A process-global, weight-identity-keyed store (`MlasPackedCaches`) keyed
on
`(address, N, K, bits, block_size, has_zero_points, compute_type)`. The
first kernel instance to
reach a weight packs it once; the sibling instance takes the same `Arc`.
The session now holds
**one** packed copy per weight.

* Keyed on the mmap **address** plus every pack-determining shape/param,
so a same-address
different-shape weight (allocator recycling a freed address — the
#845/#1079 hazard) misses
  rather than serving the wrong bytes.
* `clear_mlas_packed_caches()` runs on `Executor` drop — the **same
lifetime boundary** as
`weight_transpose::clear_all` — closing the
same-address/same-shape/across-lifetimes window.
* Accounting is updated **in the same commit**:
`MLAS_PACKED_DECODE_INSTANTIATIONS` goes `2 -> 1`,
so `resident_dequant_f32_cache_bytes` (the plan's prediction) equals
`mlas_sqnbit_packed_live_bytes`
(the actual allocation). The existing accounting test additionally
asserts **pointer identity**
  of the shared `Arc` across the prefill and decode instances.
* **No new process-global mutable state that tests mutate**
(#983/#1033/#1079): production reads
only the global store; tests use a `cfg(test)` **thread-local** store
(each libtest thread gets a
private cache, so no cross-test recycled-address contamination, while a
single test's prefill+decode
  still share).

### Acceptance criteria

**1. Predicted bytes == actual bytes, ratio 1.00.**
`resident_f32_cache_bytes` (plan prediction) vs
`live_total` (profiler actual `SQNBIT_PACKED_LIVE_BYTES`), measured
on-model:

| model / arm | predicted (bytes) | actual live (bytes) | ratio |
|---|---:|---:|---:|
| qwen05b after, admitted, multi-token | 316,443,904 | 316,443,904 |
**1.00** |
| qwen05b before, admitted, multi-token | 632,887,808 | 632,887,808 |
1.00 |
| qwen14b after, admitted (18GiB ceiling) | 8,962,744,320 |
8,962,744,320 | **1.00** |

The dedup test
(`int4_acc0_mlas_packed_accounting_equals_actual_allocated`) that ties
the plan's
prediction to the profiler's actual bytes stays green and now also
asserts the shared `Arc`.

**2. Pack count halves.** `ONNX_GENAI_PROFILE_MM=1`, `[mm_prepack]
calls=` on `qwen05b-symzp`
(169 weight boundaries):

| run | before (origin/main) | after (this branch) |
|---|---:|---:|
| 1-token (single activation shape) | 169 | 169 |
| multi-token (prefill + decode shapes) | **338** | **169** |

Before, the multi-token run packed twice as many buffers as the 1-token
run; after, they pack the
**same** count. (The 1-token run already packed 169 before, but the
pre-dedup predictor still
accounted `2x = 632,887,808` for it — over-report, the safe direction;
after, both the pack count
and the accounting are single-copy.)

**3. Peak RSS + accounted, with ratios, both models, before/after.**
Every number measured on this
host (Windows, 68,535,443,456 B RAM, CPU-only, AVX2/FMA/F16C/AVX-VNNI).
Peak RSS = polled
`PeakWorkingSet64` while running; CPU time = `TotalProcessorTime`.

**qwen05b-symzp** (weights 366,846,066 B), multi-token autoregressive
run:

| arm | packs | accounted | live | peak RSS | admitted |
|---|---:|---:|---:|---:|:--:|
| route OFF (`QNBIT=0`) | – | 0 | – | 485.9 MB | – |
| route ON **before** | 338 | 632,887,808 | 632,887,808 | 1135.9 MB |
yes |
| route ON **after** | 169 | 316,443,904 | 316,443,904 | **818.5 MB** |
yes |

The packed accounting halved (632,887,808 -> 316,443,904) and peak RSS
dropped **317.4 MB** — almost
exactly the one deduplicated packed buffer (316,443,904 B).

**qwen14b-symzp** (weights 8,549,241,669 B). Default residency ceiling =
`0.25 x RAM` =
**17,133,860,864 B**. Admission tests the **expanded** footprint
(`on-disk weights + packed cache`):

| arm | accounted (predicted) | live | expanded footprint | peak RSS |
verdict |
|---|---:|---:|---:|---:|:--:|
| **before**, default ceiling | 17,925,488,640 | 0 (declined) |
26,474,730,309 | 8643.4 MB | **declined** |
| **after**, default ceiling | 8,962,744,320 | 0 (declined) |
17,511,985,989 | 8669.3 MB | **declined** |
| **after**, ceiling 18 GiB | 8,962,744,320 | 8,962,744,320 |
17,511,985,989 | 17,606.9 MB | **admitted** (ratio 1.00) |

**Does the 14B flip declined -> admitted at the default ceiling? No —
but only just, and the reason
is precise.** The dedup halved the predicted packed cache
(17,925,488,640 -> 8,962,744,320) and
shrank the expanded footprint from 26,474,730,309 to 17,511,985,989. But
admission compares that
**expanded** footprint (weights **+** cache), not the cache alone,
against the `0.25 x RAM` ceiling
of 17,133,860,864 B. After the dedup the expanded footprint is
**17,511,985,989 B — still 378,125,125 B
(2.2%) over** the default ceiling, so it stays declined and runs the
borrowed zero-copy path
(peak ~8.6 GB, unchanged from before).

What the dedup *does* change is the admission threshold: admitting the
14B previously required a
ceiling >= 26,474,730,309 B = **0.386 of RAM**; it now requires >=
17,511,985,989 B = **0.256 of RAM**
— i.e. barely above the 0.25 default. `--host-ram-limit 18GiB` (0.263)
now admits it at peak
17,606.9 MB, comfortably inside 68.5 GB. So the dedup moves the 14B from
"unreachable without
allowing a 26.5 GB expansion" to "one notch above the default," but does
**not** cross the 0.25 line
on its own on this box. (The original prediction that it would flip
rested on comparing the ~8.4 GB
cache to the ceiling; the gate actually tests the 17.5 GB expanded
footprint.)

**4. Byte-identical generated text** (SHA-256 of generated text), greedy
decode:

| prompt | before | after (declined) | after (admitted) |
|---|---|---|---|
| qwen05b, "…relativity…" 48 tok | `EF7CA14F…` | `EF7CA14F…` | – |
| qwen05b, raw "a" 8 tok | `FD4972FF…` | `FD4972FF…` | – |
| qwen14b, "…relativity…" 16 tok | `EB88829D…` | `EB88829D…` |
`EB88829D…` |

Identical across before/after and across admitted/declined (the borrowed
zero-copy path and the MLAS
packed path produce the same tokens). qwen05b route ON and route OFF
also match (`EF7CA14F…`).

**5. No new process-global mutable state that tests mutate.** Production
writes only the
`LazyLock` global; tests use a `cfg(test)` thread-local store restored
automatically by each
libtest thread ending. No env/RAII toggles were added to production.

### Gates

* `cargo test -p onnx-runtime-ep-cpu --features mlas --lib matmul_nbits`
— **110 passed, 0 failed, 6 ignored**.
* `cargo test -p onnx-runtime-ep-cpu --lib` — **five consecutive runs**:
`1269/0/11`, `1269/0/11`,
  `1269/0/11`, `1269/0/11`, `1269/0/11` (passed / failed / ignored).
* `cargo clippy -p onnx-runtime-ep-cpu --lib -- -D warnings` — clean
(both default and `--features mlas`).

### Cases not deduplicated

None on the constant-weight route. Every MLAS SQNBit route
(`weight_prepacked`, static shards, and
the `NO_SHARD` A/B) goes through the shared store. The only non-shared
fallback is a weight with no
stable contiguous host address to key on — which never occurs on the
constant-weight (`can_prepack`)
route this touches, since the initializer is a contiguous mmap slice. A
non-constant weight rebuilds
a transient pack per call and retains nothing, so there is nothing to
share.

---

_Note: `Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624`._

Co-authored-by: justinchuby <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
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