Skip to content

perf(cpu-ep): probe the host pool with an empty dispatch, not the caller's work - #1154

Merged
justinchuby merged 2 commits into
mainfrom
deckard/probe-cost
Aug 19, 2026
Merged

justinchuby merged 2 commits into
mainfrom
deckard/probe-cost

Conversation

@justinchuby

Copy link
Copy Markdown
Owner

What

Follow-up to #1143, from its review. The host-pool probe no longer answers
"does this session's ORT pool have workers?" by sending the caller's real
work
to the pool being tested. It sends an empty eight-index dispatch
instead.

Root cause

#1143 landed with prefer_host() returning true on a probe dispatch, so the
slice in front of it went to the host pool. On a session that never latches —
intra_op_num_threads = 1, where ORT runs KernelContext_ParallelFor inline —
that means ORT running 1 Mi on one thread, where our own pool is ~5x
faster. It self-corrects after the opening burst, but the first PROBE_BURST
= 32 dispatches of every fused node paid for the question, and one dispatch in
a thousand kept paying.

Measured on this machine (AMD EPYC 9V74, taskset -c 0-15, ORT 1.28.0,
intra_op = 1, RAYON_NUM_THREADS = 16, 1 Mi f32, p50 over 4 interleaved
rounds) against a control build with probing compiled out:

op #1143 probing removed main
QuickGelu 406.9 µs 363.5 µs 371.3 µs
Swish 357.8 299.1 305.4
Erf 379.2 334.0 338.9
Tanh 357.2 327.9 346.0
Sigmoid 341.2 312.9 332.9
Clip 204.1 183.1 200.6

5-17% at 1 Mi across seven ops, and the intra_op = 1 benchmark in #1143's
body did not catch it because it was run with RAYON_NUM_THREADS = 1, where
"host, inline" and "our pool, one thread" are the same thing.

Fix

Probe with an empty dispatch. What is being measured is scheduling
behaviour, not arithmetic, so the probe does not need to carry anything — and
its cost stops scaling with the caller's slice.

That also makes the question much easier to answer, because nothing drains the
indices ahead of the workers. Instrumented over 15 sixteen-thread sessions:

probes needed to latch 1 2 3 never
sessions 9 5 1 0

against ~2.5 for the old work-carrying probe. So:

  • PROBE_BURST 32 → 8 — still ~3x the worst case observed.
  • PROBE_MIN 16 → 64 — the burst is what answers the question; this is
    only the recovery path for a session whose pool was busy through all eight.
  • PROBE_STALL is now a deadline, not a duration: hold the first index
    until another thread is seen, up to 400 µs. A pool with workers pays its
    wake-up latency and no more; only a pool with nobody to wake pays it in full.

The deadline had to grow because 8 × 100 µs was not decisive on a loaded
machine: at load 5-10 several sessions never latched and kept their work on the
wrong pool (Erf and Gelu at 65 Ki, Clip at 256 Ki showed no improvement
over main at all in that run). Early exit is what makes 400 µs affordable.

Benchmarks

Same method as #1143: .work/mt2.py alternates the two .sos round by round
in one process, ORT's CPU EP measured in the same process on the same inputs,
EP assignment asserted through ORT's profiler on every cell (anomalies=0),
taskset -c 0-15, f32, p50.

The regression this PR fixes — intra_op = 1, RAYON_NUM_THREADS = 16

µs at 1 Mi, 4 rounds. control is this branch with probing compiled out, which
bounds what the measurement noise on this shared box looks like:

op this PR control main (#1143)
Tanh 323.6 382.6 322.9
Erf 344.9 401.3 334.5
Swish 305.0 345.6 299.4
Sigmoid 338.8 374.2 328.3
QuickGelu 383.9 406.5 396.0
Gelu 501.5 499.0 465.3
Clip 204.4 181.9 191.2
Relu 176.2 170.2 197.9
Sqrt 265.3 259.2 260.5
FastGelu 450.4 414.9 446.8

The systematic 5-17% is gone: what is left is inside the control's own spread.

The win this PR must not break — intra_op = 16, rayon 16

6 rounds, µs, with this build's speedup against ORT in the same process. Every
op at 65 Ki and above still improves 2-10x over the pre-#1143 baseline:

op n before #1143 this PR ORT ORT/PR
Clip 1 Mi 791.2 87.7 45.8 0.52
Erf 65 Ki 73.0 53.2 20.1 0.38
1 Mi 1502.7 158.7 146.7 0.92
FastGelu 1 Mi 2050.4 251.0 136.5 0.54
Gelu 65 Ki 116.7 34.3 24.1 0.70
1 Mi 1916.8 247.6 192.9 0.78
QuickGelu 1 Mi 1696.7 228.6 145.0 0.63
Relu 65 Ki 38.8 16.4 17.9 1.09
1 Mi 762.4 79.3 40.2 0.51
Sigmoid 1 Mi 1515.0 185.3 76.8 0.41
Sqrt 65 Ki 56.0 22.6 26.3 1.16
1 Mi 1158.0 117.3 87.1 0.74
Swish 65 Ki 74.9 23.3 60.0 2.57
1 Mi 1543.6 165.3 150.3 0.91
Tanh 65 Ki 86.1 23.8 34.2 1.43
1 Mi 1545.4 194.5 72.2 0.37

That run was taken at load 6.5, i.e. under exactly the conditions where the
100 µs stall failed to latch.

intra_op = 1, rayon 1 — unchanged, as it must be

3 rounds, µs at 1 Mi: Relu 235.7 vs 312.7 on main, Sigmoid 779.5 vs 855.4,
Erf 894.2 vs 892.3, Gelu 1302.7 vs 1301.2, Sqrt 666.5 vs 666.0,
Tanh 809.9 vs 814.6, Clip 313.6 vs 313.7.

Also from the review of #1143

  • The bit-identity note on run_on_host claimed every chunk is a multiple of
    eight lanes and at least SIMD_MIN_LEN. The final chunk is neither —
    host_chunk_len(65537) ends with a chunk of one. The result is still
    bit-identical, but for the real reason: chunk starts are 8-aligned, every
    chunk runs the same masked-tail kernel, and the scalar-vs-vector decision is
    taken once on the whole slice. The old wording would have led someone to
    believe a per-chunk scalar fallback was safe.
  • a_nested_split_stays_serial used a slice shorter than PAR_MIN_LEN, so it
    passed whether or not the in_host_task guard existed, and it only asserted
    on the host counter. It now uses PAR_MIN_LEN + 4099 and asserts the rayon
    counter too — the guard's actual job.
  • try_host bumped PARALLEL_DISPATCHES (documented as rayon dispatches) on
    a host dispatch, while try_host_rows did not. Removed.

Correctness

cargo test -p onnx-runtime-ep-cpu --features mlas --lib → 1353 passed,
-p onnx-runtime-ep-api → 57, -p onnx-runtime-ep-plugin → 246. cargo fmt
clean for the files this PR touches; clippy clean.

the_probe_and_the_latch_agree is rewritten for the new semantics: it drives
prefer_host through the real ort_parallel_for over a threaded stand-in
(latches on the first probe, and stays latched) and over an inline one, where
it asserts nothing is ever handed to the pool across 4096 dispatches while
the cell keeps asking at the PROBE_MAX cap.

Nothing is handed to ORT's CPU EP

Threading only, as in #1143. Every node our EP claims is still computed by our
kernels; no op is declined, no capability filter changes, no fallback added.

Limitations

  • The remaining 16-thread gaps against ORT are untouched by this PR and stay
    open: Tanh 0.37x, Sigmoid 0.41x, Relu 0.51x, Clip 0.52x, FastGelu
    0.54x at 1 Mi. Those are per-op kernel and memory-bandwidth problems.
  • n ≤ 4096 is still below HOST_MIN_LEN and dominated by per-node plugin
    overhead.
  • PROBE_BURST = 8 and the 400 µs deadline are tuned on this one EPYC 9V74.
    A missed latch costs performance, never correctness, and the geometric
    back-off keeps re-asking.

@codecov

codecov Bot commented Aug 18, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 97.22222% with 2 lines in your changes missing coverage. Please review.
✅ Project coverage is 80.14%. Comparing base (06c62e0) to head (bf30828).
⚠️ Report is 78 commits behind head on main.

Files with missing lines Patch % Lines
crates/onnx-runtime-ep-api/src/host_parallel.rs 97.05% 1 Missing ⚠️
crates/onnx-runtime-ep-plugin/src/host_pool.rs 96.77% 1 Missing ⚠️
Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1154      +/-   ##
==========================================
- Coverage   80.88%   80.14%   -0.75%     
==========================================
  Files         364      364              
  Lines      160729   160978     +249     
  Branches   160729   160978     +249     
==========================================
- Hits       130005   129012     -993     
- Misses      26069    27310    +1241     
- Partials     4655     4656       +1     
Flag Coverage Δ
mlas 85.05% <ø> (-0.17%) ⬇️
offline 80.04% <97.22%> (-0.76%) ⬇️

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

Files with missing lines Coverage Δ
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 90.86% <100.00%> (ø)
crates/onnx-runtime-ep-api/src/host_parallel.rs 96.35% <97.05%> (+0.02%) ⬆️
crates/onnx-runtime-ep-plugin/src/host_pool.rs 98.18% <96.77%> (-0.56%) ⬇️

... and 9 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 18, 2026 •

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_bf16_threads=8/1x256x256 29.05 µs 52.50 µs +80.7%
🔴 matmul/medium_generic_f32_threads=8/32x512x512 1.15 ms 1.99 ms +72.9%
🔴 matmul/large_generic_f32_threads=8/32x1024x1024 5.66 ms 9.01 ms +59.1%
🔴 matmul/medium_generic_f16_threads=8/32x512x512 36.45 µs 57.70 µs +58.3%
🔴 qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.81 ms 5.81 ms +52.6%
🔴 matmul/medium_generic_f16_threads=1/32x512x512 31.27 µs 46.65 µs +49.2%
🔴 matmul/small_generic_f16_threads=1/1x256x256 31.97 µs 46.61 µs +45.8%
🔴 matmul/medium_generic_bf16_threads=8/32x512x512 585.61 µs 813.37 µs +38.9%
🔴 matmul/small_generic_f32_threads=8/1x256x256 38.13 µs 51.58 µs +35.3%
🔴 matmul/medium_generic_f32_threads=1/32x512x512 2.40 ms 3.18 ms +32.7%
🔴 matmul/small_generic_f32_threads=1/1x256x256 39.04 µs 51.60 µs +32.1%
⚠️ matmul/small_generic_bf16_threads=1/1x256x256 34.63 µs 45.02 µs +30.0%
⚠️ qwen3_sampling_processors/top_k_partial_selection 171.33 µs 221.66 µs +29.4%
⚠️ sampling_latency/min_p_per_token 213.72 µs 255.76 µs +19.7%
⚠️ logit_processing/seven_processor_chain_per_step 342.88 µs 408.72 µs +19.2%
⚠️ grammar_masking/llguidance_compute_mask/32 80.38 µs 93.96 µs +16.9%
⚠️ matmul/large_generic_f16_threads=8/32x1024x1024 114.57 µs 133.57 µs +16.6%
⚠️ block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 124.29 µs 143.36 µs +15.3%
⚠️ matmul/large_generic_bf16_threads=8/32x1024x1024 2.29 ms 2.64 ms +15.3%
✅ gather/medium_f16_threads=1-internal/32768 2.54 µs 2.90 µs +14.1%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 606.77 µs 688.63 µs +13.5%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 11.61 ms 13.16 ms +13.4%
✅ gather/medium_bf16_threads=1-internal/32768 2.49 µs 2.81 µs +12.5%
✅ kv_cache/alloc_dealloc_pages 41.28 µs 46.28 µs +12.1%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 2.73 ms 3.05 ms +11.7%
✅ tokenization/encode_tokens_per_second 407.43 µs 453.85 µs +11.4%
✅ add/small_bf16_threads=1-internal/1024 405.8 ns 450.0 ns +10.9%
✅ add/small_f16_threads=1-internal/1024 414.3 ns 459.1 ns +10.8%
✅ tokenization/decode_tokens_per_second 6.59 ms 7.29 ms +10.7%
✅ block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 1.31 ms 1.44 ms +9.6%
✅ sampling_latency/greedy_per_token 3.36 µs 3.65 µs +8.6%
✅ add/medium_f32_threads=1-internal/262144 24.58 µs 26.41 µs +7.4%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 2.32 ms 2.48 ms +7.3%
✅ gather/small_f32_threads=1-internal/4096 650.4 ns 697.5 ns +7.2%
✅ gather/large_f32_threads=1-internal/131072 39.93 µs 42.73 µs +7.0%
✅ matmul/small_generic_f16_threads=8/1x256x256 36.22 µs 38.65 µs +6.7%
✅ qwen3_sampling_processors/top_k_top_p_fast 716.01 µs 759.97 µs +6.1%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 593.17 µs 624.44 µs +5.3%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 558.55 µs 585.83 µs +4.9%
✅ gather/small_f16_threads=1-internal/4096 456.8 ns 478.2 ns +4.7%
✅ gather/medium_f32_threads=1-internal/32768 4.04 µs 4.18 µs +3.5%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 132.31 µs 136.48 µs +3.2%
✅ add/small_f32_threads=1-internal/1024 201.6 ns 207.0 ns +2.7%
✅ block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 100.98 µs 102.96 µs +2.0%
✅ gather/small_bf16_threads=1-internal/4096 490.7 ns 495.0 ns +0.9%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 111.07 µs 112.02 µs +0.9%
✅ reduce_mean/small_f32_threads=1-internal/4096 15.88 µs 15.92 µs +0.2%
✅ add/medium_f16_threads=1-internal/262144 109.81 µs 109.29 µs -0.5%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 6.42 ms 6.39 ms -0.6%
✅ reduce_mean/large_f32_threads=1-internal/262144 1.04 ms 1.02 ms -2.3%
✅ sampling_latency/top_k_per_token 57.63 µs 54.18 µs -6.0%
✅ gather/large_f16_threads=1-internal/131072 18.30 µs 17.16 µs -6.2%
✅ sampling_latency/top_p_per_token 407.46 µs 374.95 µs -8.0%
✅ reduce_mean/medium_f32_threads=1-internal/65536 281.90 µs 251.51 µs -10.8%
✅ add/large_bf16_threads=1-internal/4194304 2.09 ms 1.80 ms -13.5%
🟢 add/large_f32_threads=1-internal/4194304 1.17 ms 850.87 µs -27.3%
🟢 gather/large_bf16_threads=1-internal/131072 21.20 µs 15.33 µs -27.7%
🟢 add/medium_bf16_threads=1-internal/262144 153.07 µs 109.85 µs -28.2%
🟢 add/large_f16_threads=1-internal/4194304 2.84 ms 1.83 ms -35.8%

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: { 6.53 4.95 6.01 }
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
justinchuby marked this pull request as ready for review August 18, 2026 23:49
@justinchuby

Copy link
Copy Markdown
Owner Author

Rebased onto main (c55a3fa), re-validated, two test defects fixed — Sebastian

Picking this up to converge it. Rebased off its stale base (2291caa) onto current main, which matters here: #1201 landed the CPU task runtime after this branch was cut, and task_runtime::for_each_range is a new caller of prefer_host(). Since this PR makes prefer_host dispatch rather than just read an atomic, that interaction is the thing to check, and it is safe: for_each_range returns early on host_parallel::in_host_task() || in_task() before it asks, so a probe can never nest inside a host task or a rayon task. The elementwise call sites (simd_activations.rs:555/659/694) carry the same guard ahead of try_host.

Probe cost after the rebase is still bounded the same way: the cell lives on ExportedComputeInfo — one per compiled fused node, for the life of the session — so the opening burst is 8 probes per fused node total, not per Run. On an intra_op = 1 pool that is 8 × 400 µs once, then one probe in 1024 dispatches (~0.4 µs amortised). RoPE/Softmax/Transpose now share that same budget rather than adding to it.

Two defects fixed (commit 0a71f78c6)

Both are test-quality, and both let a broken schedule pass:

  1. a_probe_does_not_carry_the_callers_work built an AtomicUsize it never read and stored 0 into it on the last line — dead state from an earlier draft.
  2. the_probe_and_the_latch_agree asserted settled >= PROBE_MIN immediately after assert_eq!(settled, PROBE_MAX). That is 1024 >= 64 between two constants: it cannot fail, whatever the back-off does. It had replaced the count-based assertion the pre-rewrite test had, so the rewrite lost the property it was meant to keep — that a serial-looking session keeps re-asking often enough to recover from an unlucky burst.

The property is now measured where it is observable: counted_inline_parallel_for counts the dispatches ORT's stand-in is actually handed, and the test bounds that count on both sides. Measured 15 probes in 4096 dispatches (8 burst + 7 back-off). Falsified — raising the lower bound to 100000 turns it RED with probing stopped after the opening burst (15 probes).

Validation on this rebase

suite result
cargo test -p onnx-runtime-ep-api --lib 57 passed / 0 failed
cargo test -p onnx-runtime-ep-plugin --lib 246 passed / 0 failed
cargo test -p onnx-runtime-ep-cpu --lib 1417 passed / 0 failed / 17 ignored
cargo clippy --all-targets (all three) clean
rustfmt --check (touched files) clean

The Fast/Rust quality failures on the previous run were rustfmt drift in mlas-sys, governed_accumulator_budget.rs and qlinear_matmul.rs — files this PR does not touch, inherited from the old base and since repaired on main. They are gone on the rebase.

I did not re-run the ORT A/B sweep; the µs tables in the body are the author's, on the EPYC 9V74.

justinchuby and others added 2 commits August 19, 2026 00:28
Review caught that the unknown state was being paid for with the caller's
own work. `prefer_host` returned true on a probe dispatch, so the slice in
front of it went to the host pool -- and on an `intra_op = 1` session,
which never latches, that is ORT running 1 Mi on one thread where our own
pool is 5x faster. It recovered after the opening burst, but the first 32
dispatches of every fused node paid for the question. Measured against a
build with probing compiled out, at `intra_op = 1` with a 16-thread rayon
pool: 5-17% at 1 Mi across seven ops.

Ask with an empty eight-index dispatch instead. It is the scheduling
behaviour we are measuring, not the arithmetic, so the probe does not need
to carry anything, and its cost no longer scales with the caller's slice.

That also makes the question far easier to answer, because nothing drains
the indices ahead of the workers: instrumented over 15 sixteen-thread
sessions, all latched, taking 1 probe (nine), 2 (five) or 3 (one), against
~2.5 for the old work-carrying probe. So the burst drops 32 -> 8 and the
recovery period widens 16 -> 64.

A burst of 8 x 100 us turned out to be too little on a loaded machine --
several sessions never latched at load 5-10 and kept their work on the
wrong pool -- so the stall is now a *deadline*: hold the first index until
another thread is seen, up to 400 us. A pool with workers pays its wake-up
and no more; only a pool with nobody to wake pays it in full.

At intra_op=1/rayon=16, 1 Mi, this branch vs main vs probing-compiled-out
(us): Tanh 324/323/383, Erf 345/335/401, Swish 305/299/346, Gelu
502/465/499 -- i.e. within the control's own spread. The 16-thread win is
unchanged: every op at 65 Ki and above still improves 2-10x over main.

Also from review: the bit-identity note on `run_on_host` claimed every
chunk is a multiple of eight lanes and at least SIMD_MIN_LEN, which the
*final* chunk is not (65537 ends with a chunk of one). Restated for the
real reason -- chunk starts are 8-aligned and every chunk runs the same
masked-tail kernel. `a_nested_split_stays_serial` now uses a slice long
enough for the rayon path to fire and asserts on the rayon counter, so it
can actually fail; and `try_host` no longer bumps the rayon dispatch
counter on a host dispatch.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
(cherry picked from commit 0af4789)
Two review defects in the probe tests, both of which let a broken
schedule pass.

`a_probe_does_not_carry_the_callers_work` built an `AtomicUsize` it
never read and stored zero into it at the end of the test — dead state
left over from an earlier draft.

`the_probe_and_the_latch_agree` asserted `settled >= PROBE_MIN` one
line after `assert_eq!(settled, PROBE_MAX)`, which is `1024 >= 64`
between two constants: it holds no matter what the back-off does, and it
replaced the count-based assertion the previous revision of this test
had. The property it was meant to state — a serial-looking session keeps
re-asking, so an unlucky opening burst is recoverable — is now measured
where it is observable, on the pool: `counted_inline_parallel_for`
counts the dispatches ORT's stand-in is actually handed, and the test
bounds that count on both sides. 15 probes in 4096 dispatches (8 burst +
7 back-off). Falsified: raising the lower bound to 100000 turns it RED
with `probing stopped after the opening burst (15 probes)`.

`cargo test -p onnx-runtime-ep-api --lib` 57 passed,
`-p onnx-runtime-ep-plugin --lib` 246 passed,
`-p onnx-runtime-ep-cpu --lib` 1417 passed / 0 failed. Clippy and
rustfmt clean on the touched files.

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

Copy link
Copy Markdown
Owner Author

Convergence report — rebased onto current main, validated, merging

Rebase

Rebased from c55a3fab3 onto dcea14c74 (after #1346, #1352, #1232 and #1238
landed). Both commits cherry-picked clean. Final diff: 3 files, +201/−67.

The interaction that mattered

#1201's task_runtime::for_each_range is a new caller of prefer_host()
that did not exist when this branch was written, and #1238 has since routed the
MLAS prefill tiling through it too. That is safe: for_each_range returns early
on host_parallel::in_host_task() || in_task() before asking, so a probe is
never issued from inside a region. simd_activations.rs guards the same way
ahead of try_host/try_host_rows. The one accepted cost is that
for_each_range can spend a probe on a fan-out it then decides to run serially
— 8 dispatches of an empty body per fused node, once.

Worth stating plainly because it is not obvious from the diff: the probe cell
lives on ExportedComputeInfo::host_pool_probe, i.e. one per compiled fused
node for the life of the session
, not one per Run. The opening burst is
therefore paid once per node, not once per token.

Two test defects found and fixed

  1. host_parallel::a_probe_does_not_carry_the_callers_work carried a dead
    let indices = AtomicUsize::new(0) and a trailing no-op indices.store(0, …)
    that asserted nothing. Removed.
  2. host_pool::the_probe_and_the_latch_agree contained a vacuous assertion:
    assert!(settled >= PROBE_MIN) sitting one line after
    assert_eq!(settled, PROBE_MAX) — i.e. 1024 >= 64, a comparison between
    two constants that no code path can influence. Replaced with a real
    measurement: counted_inline_parallel_for + an INLINE_DISPATCHES counter
    now observe the actual number of probes, and the test bounds it from both
    sides.

Falsified, not assumed: the observed count is 15 probes in 4096 dispatches
(8 burst + 7 back-off). Raising the lower bound to 100000 turns the test RED
with "probing stopped after the opening burst (15 probes)", so the assertion can
fail. Restored to the real bounds afterwards.

Validation on the rebased head (rustc 1.97.1, the toolchain CI resolves)

check result
cargo test -p onnx-runtime-ep-api 57 passed, 0 failed
cargo test -p onnx-runtime-ep-plugin 246 passed, 0 failed
cargo test -p onnx-runtime-ep-cpu 1433 passed, 0 failed, 17 ignored
cargo clippy … --all-targets -- -D warnings (all three) clean
cargo fmt --all -- --check clean

CI

The Actions queue is saturated (every recent run queued, nothing
in_progress), so Fast (Linux x86_64) and Rust quality cannot report. Both
were reproduced locally, step for step, out of .github/workflows/ci.yml. The
fmt-drift failures this PR was showing earlier were in files it does not touch;
they were main's, and #1346/#1352 fixed them.

Working as sebastian (CPU perf).

@justinchuby
justinchuby merged commit 2a8994e into main Aug 19, 2026
9 checks passed
@justinchuby
justinchuby deleted the deckard/probe-cost branch August 19, 2026 00:30
justinchuby added a commit that referenced this pull request Aug 19, 2026
#1244)

## What this is

The per-`Run` cost of dispatching a **single-node fused subgraph**
through the plugin path, with the kernel removed from the picture.

On a 4 Ki f32 elementwise node our arm spends ~3.5 µs per `Run` where
plain ORT spends ~2.6 µs, and the kernels themselves account for well
under a microsecond of that — the same kernel on a 1 Mi tensor runs at
0.51x–1.03x of ORT. The difference is fixed dispatch overhead, and it is
why every cheap elementwise op sits at 1.1x–1.5x of ORT at 4 Ki while
the identical kernel wins at 1 Mi.

This PR removes four pieces of that overhead. No kernel is touched, no
numerics change, and nothing about assignment or execution ownership
changes: the CPU EP still executes every node it claims, locally, with
no ORT CPU fallback anywhere.

## The four cuts

**1. `device_mem_info` is no longer resolved eagerly (6 ORT FFI calls
per `Run`).**

`compute_execute` resolved `scratch_mem_info` at the top of *every*
call:

```rust
let scratch_mem_info =
    unsafe { device_mem_info(api_ref, kernel_context, exported.device_staging.as_ref()) };
```

`device_mem_info` calls `KernelContext_GetInputCount`, then for each
input `KernelContext_GetInput` + `GetTensorMemoryInfo` +
`GetMemoryInfoDeviceType`, then falls through to two more calls for
input 0. For a one-input node that is **six FFI calls**.

Who reads it? Two consumers. `intermediate_scratch` — routed multi-node
path only. And `PlacementSources::subgraph_fallback`, which
`operand_mem_info` reads **only when the node binds no ORT operands at
all**, and which `prepare_workspace` only reaches past its zero-byte and
lifetime gates. An elementwise node has operands and needs no workspace,
so it reached neither.

Now the routed path resolves it once per `Run` (unchanged) and passes
`SubgraphFallback::Resolved`; the single-node path passes
`SubgraphFallback::Deferred(staging)` and `operand_mem_info` calls
`device_mem_info` itself, with the same arguments, if it ever actually
needs the answer. Deferring is safe: the function only *reads* input
memory info, so running it after outputs are allocated cannot see a
different answer.

**2. `allocate_output` takes `want_mem_info` (1 ORT FFI call per output
per `Run`).**

`OwnedOutput::mem_info` has exactly one consumer in the tree —
`stage_host_boundary_inputs`, called under `if let Some(staging) =
exported.device_staging.as_ref()`. A host EP has no staging context, so
every `Run` made a `GetTensorMemoryInfo` call per output and dropped the
result. Both call sites now pass `exported.device_staging.is_some()`, so
a device EP is bit-for-bit unchanged and a host EP stops making the
call.

**3. `staging_log(&format!(..))` → `staging_log!(..)` (one `String` per
`Run`).**

`staging_log` checks `ONNX_GENAI_PLUGIN_TRANSFER_TRACE` *inside* the
function, so the `format!` was evaluated whether or not the trace was on
— and one of these sites is at the top of `compute_execute`, formatting
three fields into a heap `String` on every dispatch. The macro checks
the gate first. All 12 sites converted, so the footgun is gone rather
than papered over at one site.

**4. Absent-output bookkeeping is built only when there are absent
outputs (5 allocations per `Run`).**

```rust
let absent_shapes: Vec<Vec<usize>> = output_shapes.clone();
let absent_strides_storage: Vec<Vec<i64>> =
    absent_shapes.iter().map(|s| contiguous_strides(s)).collect();
let mut ort_views: Vec<TensorMut<'_>> =
    owned_outputs.iter_mut().map(|o| o.view_mut()).collect();
let mut ort_view_iter = ort_views.drain(..);
```

That storage exists solely to back the `TensorMut`s of *absent* output
slots. A node with no absent outputs — every elementwise op — cloned
every output shape and built a stride vector per output to lend them to
nobody. It is now conditional on `!absent_bufs.is_empty()`. The
`ort_views` `Vec` was collected and immediately drained; the iterator is
taken directly from `owned_outputs.iter_mut()`.

Related: placement operands are now described by an `OrtOperands` enum,
so the single-node path lends `entry.input_slots` (`&[Option<usize>]`,
flattened lazily) instead of collecting a fresh `Vec<usize>` per call
for a consumer that almost never runs. The routed path passes its
already-resolved slice.

Per `Run` on a 1-in/1-out elementwise node this is **7 fewer ORT FFI
calls** (16 → 9) and **7 fewer heap allocations**.

## A/B, one thread

`taskset -c 8-15`, this commit vs its parent, **interleaved**
(A,B,A,B,A,B with a rebuild before each arm, so drift hits both arms
equally), 400 iterations after 50 warmup, three rounds. Ratio is
**ours/ORT, lower is better**; the p50 column is the median of the three
rounds' p50 ratios, p90 likewise.

| case | before p50 | after p50 | before p90 | after p90 |
|---|---|---|---|---|
| `sqrt_f32_4k` | 1.132 | **1.026** | 1.153 | **1.033** |
| `sigmoid_f32_4k` | 1.484 | **1.346** | 1.517 | **1.362** |
| `tanh_f32_4k` | 1.534 | **1.406** | 1.595 | **1.434** |
| `erf_f32_4k` | 1.760 | **1.510** | 1.779 | **1.512** |

Best-of-three (the contention-robust statistic on this shared box)
agrees: sqrt 1.130 → 1.023, sigmoid 1.481 → 1.344, tanh 1.503 → 1.387,
erf 1.570 → 1.480. 12 of 12 arm-pairs favour the change; there is no
round in which any case regressed.

In absolute terms `tanh_f32_4k` goes 0.0046 ms → 0.0042 ms, i.e. **~0.4
µs off a ~0.9 µs gap**.

## A/B, threaded

Same protocol, `taskset -c 0-15`, two rounds, `NXRT_MM_BENCH_THREADS` =
`ONNX_GENAI_MLAS_THREADPOOL_THREADS` = `RAYON_NUM_THREADS`.

| case | 4t before | 4t after | 16t before | 16t after |
|---|---|---|---|---|
| `sqrt_f32_4k` | 1.257 | **1.129** | 1.262 | **1.156** |
| `sigmoid_f32_4k` | 1.520 | **1.355** | 1.499 | **1.340** |
| `tanh_f32_4k` | 1.551 | **1.418** | 1.548 | **1.455** |
| `erf_f32_4k` | 1.776 | **1.671** | 1.777 | **1.624** |

The overhead is per call, not per element or per worker, so the gain is
the same absolute number of microseconds at every thread count.

## Drift control: the 1 Mi grid is unchanged

Same protocol, `_f32_1m`, three interleaved rounds, one thread. A large
tensor amortises the per-call cost away, so these must *not* move — and
they don't:

| case | before | after |
|---|---|---|
| `relu_f32_1m` | 1.032 | 1.034 |
| `exp_f32_1m` | 1.021 | 1.022 |
| `sigmoid_f32_1m` | 1.066 | 1.063 |
| `tanh_f32_1m` | 1.117 | 1.115 |
| `gelu_tanh_f32_1m` | 1.243 | 1.241 |
| `gelu_exact_f32_1m` | 1.424 | 1.420 |
| `fastgelu_f32_1m` | 1.239 | 1.224 |
| `erf_f32_1m` | 1.461 | 1.439 |
| `quickgelu_f32_1m` | 0.809 | 0.811 |
| `sqrt_f32_1m` | 0.514 | 0.511 |

Ten of ten within ±0.5 %, which is this box's noise floor. That is the
shape of a per-call fix.

## Correctness

**New test, with a verified falsifier.**
`output_memory_info_is_queried_only_when_the_caller_asked_for_it` drives
`allocate_output` against a hand-built `OrtApi` whose
`GetTensorMemoryInfo` counts its calls, and asserts 0 calls for
`want_mem_info: false` and 1 for `true`. Falsifier: making
`allocate_output` ignore the flag fails it with `left: 1, right: 0` —
checked by breaking the code, not by inspection.

**Behaviour preserved, argued per cut.** (1) `device_mem_info` is called
with identical arguments, only later and only when read; it reads
inputs, which output allocation cannot change. (2) The gate on the
memory-info query is *the same condition* as the gate on its only
consumer. (3) The macro's only difference is when the `format!` runs.
(4) The absent storage is only ever indexed for absent slots.

**Suites.** `-p onnx-runtime-ep-plugin`: 247 unit tests pass (246
before, +1 new). `-p onnx-runtime-ep-cpu-plugin` with
`NXRT_REQUIRE_ORT_TESTS=1`: the full e2e suite passes, including all 54
`plugin_ort_e2e` cases — the routed multi-node fixtures
(`conformance_chain_add_mul`,
`..._repeated_runs_do_not_leak_stale_intermediates`,
`conformance_mixed_partition`) exercise the `SubgraphFallback::Resolved`
arm, and `every_assigned_node_is_also_executed_by_this_ep` passes, so
every node this EP claims is still executed here with ORT CPU fallback
disabled. `cargo fmt --all` and `cargo clippy --release --all-targets -p
onnx-runtime-ep-cpu -p onnx-runtime-ep-cpu-plugin -p
onnx-runtime-ep-plugin` are clean.

**Build identity.** Pure native CPU EP: no MLAS at runtime, no ORT CPU
EP fallback, no new dependency. AVX2/FMA host (`avx2 fma f16c`, no
AVX-512), so ORT/MLAS and we are on the same instruction footing.

## What is left

The remaining ~0.5 µs is, as far as I can attribute it without CPU
counters (`perf_event_paranoid` is 4 on this box, so `perf record` is
not available and everything here is A/B attribution):

* **~9 ORT FFI calls that are genuinely needed.** `read_inputs` costs 7
per input — `KernelContext_GetInput`, `GetTensorTypeAndShape` (which
allocates an ORT object we then release), `GetTensorElementType`,
`GetDimensionsCount`, `GetDimensions`, `ReleaseTensorTypeAndShapeInfo`,
`GetTensorData` — and there is no cheaper spelling in the stable C API.
ORT's own CPU kernels reach the same data through `OpKernelContext` with
no FFI at all, which is a structural part of what a plugin EP pays.
* **~7 remaining allocations**: `OwnedInput`'s shape and strides per
input, `kernel_inputs`, `infer_shapes`'s `Vec<Vec<usize>>`, `slot_map`,
`output_views`, `prepare_workspace`'s metadata, `allocate_output`'s
dims, and the `Box<HostPool>` in `host_pool::install`. Each is worth
~25–35 ns. Removing them needs either an inline-capacity vector type or
per-session caching of the parts that cannot change between `Run`s; both
are worth doing and neither belongs in this PR.

I deliberately did **not** touch `host_pool::install` — the per-call
`Box` is one allocation, and that file is @sebastian's 16-thread
scheduling work; a change there should come from him or after his PRs
land.


---

## Refreshed against `main` (2026-08-18)

The branch was behind `main` and its whole red CI wall came from that,
not from
this change: `crates/onnx-runtime-session/src/executor/mod.rs:175`
failed
`-D dead-code` on current stable, which `ca32b3adf` ("fix(ci): unbreak
the Rust
quality lane on current stable", #1239) fixed on `main` after this
branch forked.
`origin/main` (`c55a3fab3`) is merged in — no rebase, no force-push —
and the
diff this PR owns is unchanged at 2 files, +285/-79.

Revalidated on the merge commit, AVX2/FMA host, no AVX-512:

* `cargo test --release -p onnx-runtime-ep-plugin` — **247 passed, 0
failed**.
* `NXRT_REQUIRE_ORT_TESTS=1 cargo test --release -p
onnx-runtime-ep-cpu-plugin`
— every suite green, including all **55** `plugin_ort_e2e` cases. The
ones
  that matter to this change all pass on the merge:
  `every_assigned_node_is_also_executed_by_this_ep`,
  `no_supported_node_is_ever_left_to_the_ort_cpu_ep`,
  `no_matmul_family_node_escapes_to_the_ort_cpu_ep`, and
`every_fixture_loads_with_cpu_fallback_disabled` — so assigned still
equals
  executed with ORT CPU fallback off.
* `cargo fmt` clean for the crates this PR touches. The one `cargo fmt
--all`
hunk on this tree is in
`onnx-runtime-ep-cuda/src/kernels/standard_attention.rs`,
which arrived from `main` untouched by this PR and is a local
rustfmt-version
  difference, not a branch defect.

## Independent review

Reviewed by **Claude Opus 4.8**, read-only, with the four cuts and the
absent-slot history stated as the priority list. Verdict **APPROVE**, no
blockers. It independently confirmed:

* `SlotKind::Absent(idx)` is pushed into `slot_map` only in the same
branch that
pushes into `absent_bufs`, so `has_absent == false` implies `slot_map`
holds no
`Absent` and the skipped storage is never indexed; and the surviving
indices
  are the same full-slot indices as before.
* `ort_view_iter` from `owned_outputs.iter_mut()` has the identical
borrow
structure as the old `collect()` + `drain(..)`, so nothing borrows a
temporary.
* `OrtOperands::Slots(..).indices()` yields the same elements in the
same order
  as the removed `.iter().flatten().copied().collect()`, `None` skipped.
* `operand_mem_info` is the only reader of the deferred value, only
reachable
when the node binds **zero** ORT inputs, which makes the timing of the
deferred
  `device_mem_info` moot rather than merely argued.
* Both `allocate_output` call sites gate on the same predicate as the
only
  reader of `OwnedOutput::mem_info`, so a device EP is unchanged.

Its one correction is applied above: the prose said 14 `staging_log`
sites; the
real count in `origin/main` is **12**, and all 12 are converted.

---

## Refreshed against `main` @ `6a855d5e0`, and measured as a stack

`origin/main` moved a long way while this sat in the CI queue (#1346,
#1352 and
#1361 on the quality lane; #1154, #1232, #1238 on the CPU side). Merged
in
normally — no rebase — and re-measured from scratch against the new
baseline.

Production pure-native A/B, plain ORT as the control arm. No MLAS, no
ORT CPU
fallback, no deferral. `taskset -c 8-15`, one thread, 400 iterations,
five
interleaved rounds out of two worktrees, started only once cores 8-15
were
>=93% idle. **Ratio is ours/ORT, lower is better.** `before` is `main`
at
`6a855d5e0`; `after` is #1244 + #1246 together, since #1246 is stacked
on #1244
and the pair is what a user gets.

| case | ratio p50 main | ratio p50 stack | Δ | ratio p90 main | ratio
p90 stack | ours us | ORT drift | rounds won |
|---|---|---|---|---|---|---|---|---|
| `thresholdedrelu_f32_4k` | 1.512 | **1.216** | -19.6% | 1.519 |
**1.216** | 3.6 → **2.8** | -4.2% | 5/5 |
| `tanh_f32_4k` | 1.511 | **1.279** | -15.4% | 1.513 | **1.284** | 4.6 →
**3.9** | +0.0% | 5/5 |
| `sigmoid_f32_4k` | 1.475 | **1.248** | -15.4% | 1.484 | **1.256** |
4.7 → **4.0** | +0.0% | 5/5 |
| `erf_f32_4k` | 1.470 | **1.363** | -7.3% | 1.489 | **1.364** | 7.5 →
**7.0** | +0.0% | 5/5 |
| `hardsigmoid_f32_4k` | 1.416 | **1.127** | -20.4% | 1.426 | **1.133**
| 3.5 → **2.8** | +0.0% | 5/5 |
| `leakyrelu_f32_4k` | 1.361 | **1.094** | -19.6% | 1.371 | **1.104** |
3.5 → **2.8** | +0.0% | 5/5 |
| `sqrt_f32_4k` | 1.141 | **0.946** | -17.1% | 1.150 | **0.955** | 4.0 →
**3.4** | -2.8% | 5/5 |
| `log_f32_4k` | 0.767 | **0.698** | -9.0% | 0.776 | **0.703** | 7.9 →
**7.2** | +0.0% | 5/5 |
| `selu_f32_4k` | 0.505 | **0.440** | -12.9% | 0.512 | **0.444** | 5.6 →
**4.9** | +0.0% | 5/5 |
| `elu_f32_4k` | 0.480 | **0.416** | -13.3% | 0.488 | **0.421** | 5.3 →
**4.6** | +0.0% | 5/5 |
| `celu_f32_4k` | 0.470 | **0.411** | -12.6% | 0.478 | **0.418** | 5.7 →
**5.0** | +0.0% | 5/5 |
| `mish_f32_4k` | 0.276 | **0.264** | -4.3% | 0.278 | **0.268** | 17.3 →
**16.6** | +0.0% | 5/5 |

**Every case, every round.** The two rows with a moving control (`sqrt`
-2.8%,
`thresholdedrelu` -4.2%) are reported rather than dropped; both won 5/5
anyway
and their absolute time fell by the same ~0.7 us as everything else.

That constant ~0.7 us is the point. It is not proportional to tensor
size — the
same absolute amount comes off `hardsigmoid` (3.5 -> 2.8 us) as off
`mish`
(17.3 -> 16.6 us) — which is what a fixed per-`Run` cost looks like when
you
remove some of it. It moves the cheap ops the most because they had the
least
to hide it behind, and `sqrt` crosses from 1.141 to **0.946**, from a
loss to a
win.

### Where the remaining time goes

Measured directly, by instrumenting `compute_execute` segment by segment
on top
of this stack (temporary probe, not committed; `perf` is unavailable on
this
host — `perf_event_paranoid=4`). Per `Run`, one-in/one-out elementwise
node,
4096 `f32`, microseconds:

| segment | us | note |
|---|---|---|
| `KernelContext_GetOutput` | 0.35 | ORT's own API — ours to call, not
to optimise |
| `read_inputs` | 0.15 | 4 ORT FFI calls, already one shape call after
#1246 |
| rest of `allocate_output` | 0.13 | `GetTensorMutableData` + strides |
| `prepare_workspace` | 0.09 | metadata vector + plan-cache lookup, for
a kernel needing 0 bytes |
| `host_pool::install` | 0.05 | @sebastian's, not touched |
| `infer_shapes` | 0.05 | |
| `kernel_inputs` | 0.04 | |
| `output_views` | 0.04 | |

Non-kernel node cost is **~1.25 us and near-constant across all twelve
operators** (0.28 to 15.2 us of kernel time), which is the direct
confirmation
that small-node ratios on this EP are dispatch-bound rather than
kernel-bound.
There is no single large item left — the biggest,
`KernelContext_GetOutput`, is
ORT's. The rest is a long tail of 0.04-0.15 us items, which is what
#1358
(`InlineVec`) starts on.

### And nothing breaks at 1 Mi

Same harness, 1048576 elements, 120 iterations, 3 rounds. A fixed
per-`Run`
cost should be invisible here, and it is:

| case | ratio p50 main | ratio p50 stack | ours us | ORT drift |
|---|---|---|---|---|
| `celu_f32_1m` | 0.141 | 0.140 | 382.1 → 377.9 | -0.0% |
| `elu_f32_1m` | 0.138 | 0.136 | 347.7 → 343.6 | +0.1% |
| `erf_f32_1m` | 0.671 | 0.670 | 595.2 → 594.6 | +0.1% |
| `exp_f32_1m` | 0.606 | 0.590 | 247.0 → 240.5 | +0.1% |
| `fastgelu_f32_1m` | 0.643 | 0.648 | 415.3 → 413.2 | -1.7% |
| `gelu_exact_f32_1m` | 0.572 | 0.592 | 706.1 → 710.6 | -2.8% ⚠ |
| `gelu_tanh_f32_1m` | 0.657 | 0.651 | 414.8 → 410.5 | +0.3% |
| `hardsigmoid_f32_1m` | 0.376 | 0.372 | 88.0 → 69.6 | -0.9% |
| `leakyrelu_f32_1m` | 0.421 | 0.430 | 90.1 → 87.3 | -2.2% ⚠ |
| `log_f32_1m` | 0.270 | 0.274 | 616.4 → 612.4 | +1.7% |
| `mish_f32_1m` | 0.105 | 0.105 | 1626.0 → 1626.0 | -0.2% |
| `quickgelu_f32_1m` | 0.455 | 0.453 | 321.5 → 320.0 | -0.5% |
| `relu_f32_1m` | 1.034 | 1.022 | 131.7 → 130.2 | +0.0% |
| `selu_f32_1m` | 0.147 | 0.148 | 374.1 → 371.8 | -1.5% |
| `sigmoid_f32_1m` | 0.471 | 0.606 | 231.6 → 229.0 | -23.4% ⚠ |
| `sqrt_f32_1m` | 0.314 | 0.302 | 148.1 → 143.9 | +1.1% |
| `tanh_f32_1m` | 0.644 | 0.631 | 226.2 → 221.9 | -0.2% |
| `thresholdedrelu_f32_1m` | 0.500 | 0.486 | 70.9 → 68.6 | -0.3% |

Flat, as predicted — 0.7 us against 70-1626 us of work. Absolute time is
equal
or better in 16 of 18 cases. The two ⚠ rows had the control move more
than the
effect: `sigmoid` is unusable (ORT itself moved -23.4%; our own absolute
went
231.6 -> 229.0 us), and `gelu_exact`'s +0.6% absolute sits inside its
-2.8%
control. Reported rather than dropped.

This is the coverage claim for the change: it buys ~0.7 us at every
size, which
is 20% of a small node and nothing at all of a large one, and it costs
nothing
anywhere.

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

## What this is

`read_inputs` runs **once per input per `Run`**, and it asked ORT for
the element
type and the dimensions through the classic five-call sequence:

```
GetTensorTypeAndShape  →  GetTensorElementType
                       →  GetDimensionsCount
                       →  GetDimensions
                       →  ReleaseTensorTypeAndShapeInfo
```

The first of those makes ORT **heap-allocate** an
`OrtTensorTypeAndShapeInfo` and
copy the shape into it; the last frees it. Five FFI crossings and one
allocation
on ORT's side of the boundary, per input, per `Run`, to read data the
`OrtValue`
already owns.

`GetTensorElementTypeAndShapeDataReference` (ORT C API, since 1.24)
returns the
element type **and a reference to the value's own shape array** in one
call,
allocating nothing. This PR routes `read_inputs` through it.

Stacked on #1244 (since merged; this branch now targets `main`
directly).
No kernel is touched, no numerics change, and node assignment
and execution ownership are unchanged: the CPU EP still executes every
node it
claims, locally, with ORT CPU fallback off.

## The three cuts

**1. Five ORT calls and one ORT-side allocation per input → one call, no
allocation.**

The plugin already fails closed below API 27, so the hook is always
present in
practice. The five-call sequence is kept as a fallback for a host that
leaves it
null, and a host offering **neither** now fails closed with a message
naming
both — it does not silently read garbage shapes.

The borrowed pointer is used only to build the owned `shape`/`strides`
of the
`OwnedInput`; nothing retains it. ORT spells a scalar as a **null**
pointer with
count 0, and `slice::from_raw_parts` is UB on null, so that case is
guarded
explicitly — see *Correctness* below for how that guard is actually
enforced.

**2. `validate_dims` takes a lazily-formatted label, not `&str`.**

Every call site passed `&format!("input {i}")`, i.e. a heap `String`
built on the
success path purely to label an error that almost never happens. Callers
now pass
`format_args!(..)` and the string is only materialised inside the error
branches.
One `String` per input per `Run` gone. All call sites converted,
including
`transfer.rs` and the unit tests; no message text changed.

(The merge with `main` took main's `impl std::fmt::Display` signature
here, which
is a superset of this branch's original `Arguments<'_>` and keeps both
callers
working unchanged.)

**3. `allocate_output` converts the output shape to `i64` on the
stack.**

`let dims: Vec<i64> = shape.iter().map(|&d| d as i64).collect()` was a
heap
allocation per output per `Run` whose entire lifetime was the
`KernelContext_GetOutput`
call below it. Rank ≤ 8 now goes to a stack array; a taller tensor still
works,
through the `Vec`.

For a 1-in/1-out elementwise node that is **5 fewer ORT FFI calls**, one
fewer
ORT-side allocation, and **2 fewer heap allocations** per `Run`, on top
of #1244.

## A/B, one thread

`taskset -c 8-15`, this branch vs its merge parent (`origin/main` +
#1244),
**interleaved** (A,B,A,B,… so drift hits both arms equally), 400
iterations after
50 warmup, **five rounds**. Ratio is **ours/ORT, lower is better**; each
column is
the median of the five rounds. "ORT drift" is the control: ORT's own p50
measured
in the two arms, which must not move for the comparison to mean
anything.

| case | before p50 | after p50 | before p90 | after p90 | ORT drift |
rounds won |
|---|---|---|---|---|---|---|
| `sqrt_f32_4k` | 1.023 | **0.952** | 1.031 | **0.964** | +2.9 % | 4/5 |
| `leakyrelu_f32_4k` | 1.202 | **1.097** | 1.210 | **1.104** | 0.0 % |
5/5 |
| `hardsigmoid_f32_4k` | 1.213 | **1.144** | 1.226 | **1.149** | 0.0 % |
5/5 |
| `thresholdedrelu_f32_4k` | 1.316 | **1.209** | 1.314 | **1.220** | 0.0
% | 5/5 |
| `sigmoid_f32_4k` | 1.328 | **1.252** | 1.333 | **1.258** | 0.0 % | 5/5
|
| `tanh_f32_4k` | 1.368 | **1.287** | 1.369 | **1.291** | 0.0 % | 5/5 |
| `erf_f32_4k` | 1.383 | **1.327** | 1.387 | **1.329** | 0.0 % | 5/5 |
| `celu_f32_4k` | 0.430 | **0.411** | 0.435 | **0.418** | 0.0 % | 4/5 |
| `elu_f32_4k` | 0.437 | **0.416** | 0.443 | **0.422** | −1.8 % | 4/5 |
| `selu_f32_4k` | 0.467 | **0.441** | 0.471 | **0.446** | 0.0 % | 5/5 |
| `log_f32_4k` | 0.727 | **0.700** | 0.728 | **0.706** | 0.0 % | 5/5 |
| `mish_f32_4k` | 0.268 | **0.264** | 0.271 | **0.266** | 0.0 % | 4/5 |

12 of 12 cases improve at p50 and p90; 56 of 60 individual arm-pairs
favour the
change, and no case regressed in its median. In absolute terms it is
**0.2–0.3 µs
off every `Run`**: `tanh_f32_4k` 4.20 µs → 3.90 µs against ORT's 3.10
µs,
`thresholdedrelu_f32_4k` 3.10 µs → 2.80 µs against ORT's 2.30 µs. That
is the
shape of a fixed per-call cost being removed, which is what it is.

`sqrt_f32_4k` crosses below 1.00 — the plugin path now dispatches that
op faster
than plain ORT does.

## A/B, four threads

Same protocol, `NXRT_MM_BENCH_THREADS` =
`ONNX_GENAI_MLAS_THREADPOOL_THREADS` =
`RAYON_NUM_THREADS` = 4, three rounds.

| case | before p50 | after p50 | before p90 | after p90 |
|---|---|---|---|---|
| `leakyrelu_f32_4k` | 1.069 | **0.966** | 1.075 | **0.976** |
| `hardsigmoid_f32_4k` | 1.152 | **1.050** | 1.164 | **1.060** |
| `thresholdedrelu_f32_4k` | 1.212 | **1.091** | 1.228 | **1.106** |
| `sqrt_f32_4k` | 1.121 | **1.036** | 1.130 | **1.045** |
| `sigmoid_f32_4k` | 1.346 | **1.264** | 1.359 | **1.263** |
| `tanh_f32_4k` | 1.434 | **1.319** | 1.453 | **1.325** |
| `erf_f32_4k` | 1.581 | **1.502** | 1.588 | **1.510** |
| `celu_f32_4k` | 0.460 | **0.437** | 0.476 | **0.445** |
| `elu_f32_4k` | 0.481 | **0.448** | 0.486 | **0.468** |
| `selu_f32_4k` | 0.500 | **0.458** | 0.482 | **0.464** |
| `log_f32_4k` | 0.766 | **0.731** | 0.781 | **0.734** |
| `mish_f32_4k` | 0.332 | **0.322** | 0.345 | **0.339** |

36 of 36 arm-pairs favour the change. The cost is per call, not per
element or
per worker, so the gain is the same absolute microseconds at every
thread count.

## Drift control: the 1 Mi grid does not move

Same protocol, `_f32_1m`, 200 iterations after 30 warmup, five
interleaved rounds,
one thread. A 1 Mi tensor amortises a fixed per-call cost away, so
**these must
not move** — and they don't:

| case | before | after | | case | before | after |
|---|---|---|---|---|---|---|
| `relu_f32_1m` | 1.024 | 1.028 | | `sqrt_f32_1m` | 0.313 | 0.309 |
| `exp_f32_1m` | 0.580 | 0.576 | | `log_f32_1m` | 0.264 | 0.258 |
| `sigmoid_f32_1m` | 0.588 | 0.584 | | `mish_f32_1m` | 0.100 | 0.099 |
| `tanh_f32_1m` | 0.618 | 0.620 | | `celu_f32_1m` | 0.127 | 0.129 |
| `erf_f32_1m` | 0.651 | 0.668 | | `elu_f32_1m` | 0.129 | 0.129 |
| `gelu_exact_f32_1m` | 0.585 | 0.584 | | `selu_f32_1m` | 0.136 | 0.138
|
| `gelu_tanh_f32_1m` | 0.602 | 0.646 | | `hardsigmoid_f32_1m` | 0.455 |
0.453 |
| `fastgelu_f32_1m` | 0.608 | 0.639 | | `leakyrelu_f32_1m` | 0.409 |
0.417 |
| `quickgelu_f32_1m` | 0.429 | 0.438 | | `thresholdedrelu_f32_1m` |
0.567 | 0.573 |

The per-case round-win counts here are 0/5–4/5 with a median of 2/5,
i.e. a coin
flip — the signature of noise, not of an effect. The two largest movers,
`gelu_tanh` (+7 %) and `fastgelu` (+5 %), are the two cases whose
**ORT** side also
drifted most in the same rounds (−4.4 % and −2.8 %), so the ratio moved
because
the denominator did. This host's floor is roughly ±5 % on the 1 Mi grid
and I am
not claiming anything below it.

## Correctness

**Five new tests, three with a verified falsifier**, driving
`read_inputs` and
`allocate_output` against a hand-built `OrtApi` whose hooks count their
calls.

| test | what it pins | falsifier |
|---|---|---|
| `input_shapes_come_from_one_call_when_ort_offers_the_reference_hook` |
1 reference call, **0** legacy calls, and the shape/strides/dtype that
come out | forcing the legacy route fails it with `left: 0, right: 1`
(run) |
| `the_five_call_fallback_produces_the_same_input_as_the_reference_hook`
| the fallback is exercised and its `OwnedInput` matches the reference
path field for field | — (it *is* the parity check) |
| `a_borrowed_scalar_shape_never_dereferences_null` | ORT's documented
scalar spelling — null pointer, count 0 — takes the borrowed route and
yields rank 0 | deleting the null guard makes **Miri** report
`out-of-bounds pointer use: null pointer is a dangling pointer` at
`kernel_ctx.rs:280` (run) |
| `read_inputs_fails_closed_when_no_shape_route_exists` | the error
names **both** routes; no silent garbage shapes | — |
| `output_dims_are_identical_on_the_inline_and_heap_ranks` | ranks
0/1/8/9/12 all reach ORT with every dimension intact | dropping the heap
arm delivers rank 9 truncated to 8 dims (run) |

**The null guard is now actually enforced, not just asserted.**
Natively,
`from_raw_parts(null, 0)` returns an empty slice, so a test asserting
the
resulting shape passes with or without the guard — the guarantee only
exists if
Miri sees it. `onnx-runtime-ep-plugin` was not in the Miri matrix, so
this PR adds
`kernel_ctx::` as a lane in `.github/workflows/miri.yml` (the crate's
other
modules dlopen ORT and are not Miri-tractable, which is why the lane is
scoped to
the module rather than the crate). 25 tests pass under Miri in 4.15 s;
with the
guard deleted the lane fails. That pairing is what makes the test
non-vacuous.

**Suites.** `-p onnx-runtime-ep-plugin`: **252** unit tests pass (247
before,
+5 new). `-p onnx-runtime-ep-cpu-plugin` with
`NXRT_REQUIRE_ORT_TESTS=1`: every
suite green, including all **55** `plugin_ort_e2e` cases —
`every_assigned_node_is_also_executed_by_this_ep`,
`no_supported_node_is_ever_left_to_the_ort_cpu_ep`,
`no_matmul_family_node_escapes_to_the_ort_cpu_ep` and
`every_fixture_loads_with_cpu_fallback_disabled` all pass, so **assigned
still
equals executed** with ORT CPU fallback disabled. `cargo clippy
--release
--all-targets -p onnx-runtime-ep-plugin` and `cargo fmt` are clean.

**Build identity.** Pure native CPU EP: no MLAS at runtime, no ORT CPU
EP
fallback, no new dependency. AVX2/FMA host (`avx2 fma f16c`, no
AVX-512), so ORT
and we are on the same instruction footing.

## Independent review

Reviewed by **Claude Opus 4.8**, read-only, briefed with the exact ORT
contract
for the borrowed pointer and asked specifically to hunt UB and vacuous
tests.
Verdict **APPROVE**, no blockers. It confirmed the null/scalar guard is
unreachable-by-construction for `from_raw_parts`, that the borrow cannot
outlive
the `OrtValue` (its only consumer copies), that
`ReleaseTensorTypeAndShapeInfo`
still runs on every legacy error path — and noted that moving
`DataType::from_onnx`
after the match incidentally closes a **pre-existing** leak of
`type_shape` on the
unsupported-dtype path.

It raised two test-quality defects, both real and both **fixed** in
`6027f6e16`:

1. The three counting tests shared two process-wide `AtomicUsize`es
while
asserting exact counts, and libtest runs them in parallel — one test's
reset
   could land inside another's assertion window. They now take a shared
   `SHAPE_COUNTER_LOCK` that resets both counters under the guard.
2. `a_borrowed_scalar_shape_never_dereferences_null` was **vacuous**
with respect
to its name, for exactly the reason above, and the crate was not under
Miri.
Hence the Miri lane, and the test now also asserts the scalar went
through the
   borrowed route so it cannot pass on the fallback.

It also flagged that `OrtStatus` is not released on error paths in this
file.
That is pre-existing and repo-wide in `kernel_ctx.rs` — the old
five-call code
leaked identically — so it is not touched here; it belongs in its own
change.

## What is left

After #1244 and this PR, a 1-in/1-out elementwise `Run` is down to
roughly
`KernelContext_GetInputCount` + `GetInput` + the one shape reference +
`GetTensorData` + `GetOutput` + `GetTensorMutableData`. What remains:

* **`OwnedInput`'s `shape` and `strides` `Vec`s**, `kernel_inputs`,
`infer_shapes`'s `Vec<Vec<usize>>`, `slot_map`, `output_views`, and the
`Box<HostPool>` in `host_pool::install` — each worth ~25–35 ns. Removing
them
needs an inline-capacity vector type or per-session caching of the parts
that
cannot change between `Run`s. Both are worth doing; neither belongs
here.
* **The `HostPool` box** specifically is @sebastian's 16-thread
scheduling file
  and should come from him or after his PRs land.

---

## Refreshed against `main` @ `6a855d5e0`, and measured as a stack

`origin/main` moved a long way while this sat in the CI queue (#1346,
#1352 and
#1361 on the quality lane; #1154, #1232, #1238 on the CPU side). Merged
in
normally — no rebase — and re-measured from scratch against the new
baseline.

Production pure-native A/B, plain ORT as the control arm. No MLAS, no
ORT CPU
fallback, no deferral. `taskset -c 8-15`, one thread, 400 iterations,
five
interleaved rounds out of two worktrees, started only once cores 8-15
were
>=93% idle. **Ratio is ours/ORT, lower is better.** `before` is `main`
at
`6a855d5e0`; `after` is #1244 + #1246 together, since #1246 is stacked
on #1244
and the pair is what a user gets.

| case | ratio p50 main | ratio p50 stack | Δ | ratio p90 main | ratio
p90 stack | ours us | ORT drift | rounds won |
|---|---|---|---|---|---|---|---|---|
| `thresholdedrelu_f32_4k` | 1.512 | **1.216** | -19.6% | 1.519 |
**1.216** | 3.6 → **2.8** | -4.2% | 5/5 |
| `tanh_f32_4k` | 1.511 | **1.279** | -15.4% | 1.513 | **1.284** | 4.6 →
**3.9** | +0.0% | 5/5 |
| `sigmoid_f32_4k` | 1.475 | **1.248** | -15.4% | 1.484 | **1.256** |
4.7 → **4.0** | +0.0% | 5/5 |
| `erf_f32_4k` | 1.470 | **1.363** | -7.3% | 1.489 | **1.364** | 7.5 →
**7.0** | +0.0% | 5/5 |
| `hardsigmoid_f32_4k` | 1.416 | **1.127** | -20.4% | 1.426 | **1.133**
| 3.5 → **2.8** | +0.0% | 5/5 |
| `leakyrelu_f32_4k` | 1.361 | **1.094** | -19.6% | 1.371 | **1.104** |
3.5 → **2.8** | +0.0% | 5/5 |
| `sqrt_f32_4k` | 1.141 | **0.946** | -17.1% | 1.150 | **0.955** | 4.0 →
**3.4** | -2.8% | 5/5 |
| `log_f32_4k` | 0.767 | **0.698** | -9.0% | 0.776 | **0.703** | 7.9 →
**7.2** | +0.0% | 5/5 |
| `selu_f32_4k` | 0.505 | **0.440** | -12.9% | 0.512 | **0.444** | 5.6 →
**4.9** | +0.0% | 5/5 |
| `elu_f32_4k` | 0.480 | **0.416** | -13.3% | 0.488 | **0.421** | 5.3 →
**4.6** | +0.0% | 5/5 |
| `celu_f32_4k` | 0.470 | **0.411** | -12.6% | 0.478 | **0.418** | 5.7 →
**5.0** | +0.0% | 5/5 |
| `mish_f32_4k` | 0.276 | **0.264** | -4.3% | 0.278 | **0.268** | 17.3 →
**16.6** | +0.0% | 5/5 |

**Every case, every round.** The two rows with a moving control (`sqrt`
-2.8%,
`thresholdedrelu` -4.2%) are reported rather than dropped; both won 5/5
anyway
and their absolute time fell by the same ~0.7 us as everything else.

That constant ~0.7 us is the point. It is not proportional to tensor
size — the
same absolute amount comes off `hardsigmoid` (3.5 -> 2.8 us) as off
`mish`
(17.3 -> 16.6 us) — which is what a fixed per-`Run` cost looks like when
you
remove some of it. It moves the cheap ops the most because they had the
least
to hide it behind, and `sqrt` crosses from 1.141 to **0.946**, from a
loss to a
win.

### Where the remaining time goes

Measured directly, by instrumenting `compute_execute` segment by segment
on top
of this stack (temporary probe, not committed; `perf` is unavailable on
this
host — `perf_event_paranoid=4`). Per `Run`, one-in/one-out elementwise
node,
4096 `f32`, microseconds:

| segment | us | note |
|---|---|---|
| `KernelContext_GetOutput` | 0.35 | ORT's own API — ours to call, not
to optimise |
| `read_inputs` | 0.15 | 4 ORT FFI calls, already one shape call after
#1246 |
| rest of `allocate_output` | 0.13 | `GetTensorMutableData` + strides |
| `prepare_workspace` | 0.09 | metadata vector + plan-cache lookup, for
a kernel needing 0 bytes |
| `host_pool::install` | 0.05 | @sebastian's, not touched |
| `infer_shapes` | 0.05 | |
| `kernel_inputs` | 0.04 | |
| `output_views` | 0.04 | |

Non-kernel node cost is **~1.25 us and near-constant across all twelve
operators** (0.28 to 15.2 us of kernel time), which is the direct
confirmation
that small-node ratios on this EP are dispatch-bound rather than
kernel-bound.
There is no single large item left — the biggest,
`KernelContext_GetOutput`, is
ORT's. The rest is a long tail of 0.04-0.15 us items, which is what
#1358
(`InlineVec`) starts on.

### And nothing breaks at 1 Mi

Same harness, 1048576 elements, 120 iterations, 3 rounds. A fixed
per-`Run`
cost should be invisible here, and it is:

| case | ratio p50 main | ratio p50 stack | ours us | ORT drift |
|---|---|---|---|---|
| `celu_f32_1m` | 0.141 | 0.140 | 382.1 → 377.9 | -0.0% |
| `elu_f32_1m` | 0.138 | 0.136 | 347.7 → 343.6 | +0.1% |
| `erf_f32_1m` | 0.671 | 0.670 | 595.2 → 594.6 | +0.1% |
| `exp_f32_1m` | 0.606 | 0.590 | 247.0 → 240.5 | +0.1% |
| `fastgelu_f32_1m` | 0.643 | 0.648 | 415.3 → 413.2 | -1.7% |
| `gelu_exact_f32_1m` | 0.572 | 0.592 | 706.1 → 710.6 | -2.8% ⚠ |
| `gelu_tanh_f32_1m` | 0.657 | 0.651 | 414.8 → 410.5 | +0.3% |
| `hardsigmoid_f32_1m` | 0.376 | 0.372 | 88.0 → 69.6 | -0.9% |
| `leakyrelu_f32_1m` | 0.421 | 0.430 | 90.1 → 87.3 | -2.2% ⚠ |
| `log_f32_1m` | 0.270 | 0.274 | 616.4 → 612.4 | +1.7% |
| `mish_f32_1m` | 0.105 | 0.105 | 1626.0 → 1626.0 | -0.2% |
| `quickgelu_f32_1m` | 0.455 | 0.453 | 321.5 → 320.0 | -0.5% |
| `relu_f32_1m` | 1.034 | 1.022 | 131.7 → 130.2 | +0.0% |
| `selu_f32_1m` | 0.147 | 0.148 | 374.1 → 371.8 | -1.5% |
| `sigmoid_f32_1m` | 0.471 | 0.606 | 231.6 → 229.0 | -23.4% ⚠ |
| `sqrt_f32_1m` | 0.314 | 0.302 | 148.1 → 143.9 | +1.1% |
| `tanh_f32_1m` | 0.644 | 0.631 | 226.2 → 221.9 | -0.2% |
| `thresholdedrelu_f32_1m` | 0.500 | 0.486 | 70.9 → 68.6 | -0.3% |

Flat, as predicted — 0.7 us against 70-1626 us of work. Absolute time is
equal
or better in 16 of 18 cases. The two ⚠ rows had the control move more
than the
effect: `sigmoid` is unusable (ORT itself moved -23.4%; our own absolute
went
231.6 -> 229.0 us), and `gelu_exact`'s +0.6% absolute sits inside its
-2.8%
control. Reported rather than dropped.

This is the coverage claim for the change: it buys ~0.7 us at every
size, which
is 20% of a small node and nothing at all of a large one, and it costs
nothing
anywhere.

---

## Status after merging latest `main` (2026-08-19)

Validated on latest `main` after the six-PR #1077 stack landed
(#1387, #1409, #1412, #1430, #1433, #1472). Full evidence in the PR
comment
below; the measured summary:

**Deterministic counters, `relu_1_tiny`** — `OrtFfiCall` **10 → 6 per
`Run`**,
`DispatchAlloc` **24 → 20**. Four fewer round trips for one input, i.e.
the
predicted 7 → 3 per input.

**Timing A/B**, production build, 3 clean reps (a 4th discarded for
contention —
it read 0.973, which would have flattered us):

| case | main | this PR |
|---|---|---|
| `relu_1_tiny` | 1.397 | **1.271** |
| `relu_10_tiny` | 1.250 | **1.173** |
| `relu_100_tiny` | 1.163 | **1.110** |

The gain shrinks as depth grows — the signature of a fixed-cost fix,
since this
removes FFI calls per `Run`, not per node. Fixed per-`Run` overhead
**1.38 →
1.23** against ORT; per-node slope unchanged.

**Test coverage gap this merge exposed and closed:** the pinned costs
(`1 + 7` calls, `1 + 3` allocations) run against a zeroed `fake_api()`,
where the
reference hook is null — so they describe the *legacy fallback*, not the
fast
path this PR exists to add. Added
`the_reference_hook_path_costs_exactly_three_ort_calls_per_input` (pins
`1 + 3`)
and `the_reference_hook_path_allocates_less_than_the_legacy_path`.
`dispatch_probe`'s FFI-coverage table updated to 12 members / 12
`ort_call()`
sites, which is the guard that flagged the new API member in the first
place.

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Co-authored-by: Resch <resch@squad.local>
justinchuby added a commit that referenced this pull request Aug 19, 2026
## The build we ship had no integer GEMM at all

`QLinearMatMul` in the default build did this per call:

1. `read_quantized` widened operand `A` to a `Vec<i32>`, then did it
again for operand `B`. For a
2048x2048 `B` that is a 16 MiB allocation and fill on every single call,
thrown away at the end
   of it.
2. A scalar rank-1 update walked `A` row by row, and each row
re-streamed the whole of `B`. At
   `m = 128` that is 512 MiB of traffic for 1 GFLOP of work.

The result was 11.8x ORT at `m = 1` and 12.1x at `m = 128` — the largest
single loss on the
x86-64 CPU EP. The performance doc's `QLinearMatMul` rows never
described this build: they were
taken with `--features mlas`, which is a research build we do not ship.
That is now called out in
the doc.

This adds `kernels/qgemm_native.rs`, a native byte-operand integer GEMM,
and points
`qlinear_matmul.rs` at it. Nothing defers, and nothing falls back.

## Two kernels, chosen by `m`

| shape | kernel | why |
| --- | --- | --- |
| `m <= 4` (decode) | pack-free fused | one pass over `A` means a packed
panel of `B` is never reused, so packing is pure cost. Accumulators stay
in registers across a 256-row `k` block. |
| `m > 4` (prefill) | packed | `KC/2` pairs of `NC` columns (`KC = 512`,
`NC = 256`) is 256 KiB of `B`, which stays in L2 while every row of `A`
sweeps it. |

The inner tile is `vpmaddwd` over `NR = 16` columns and `MR = 4` rows,
`k` consumed two rows at a
time.

## Why `vpmaddwd` and not `vpmaddubsw`

MLAS gets 32 MACs from two instructions using `vpmaddubsw`, which
**saturates**: it needs a
sign-domain translation of `B` and its intermediate is only nominally
exact. `vpmaddwd` needs four
instructions for the same 32 MACs, but with centred `a` in `[-255, 255]`
and raw `b` in
`[-128, 255]` a product is at most 65025 and a pair sum at most 130050,
so it cannot saturate and
cannot overflow. No sign-domain flip, no reasoning about clamped
intermediates.

That instruction-count difference is the whole of the residual gap at `m
= 1`. Closing it means
giving up exact integer arithmetic, which is not a trade I am willing to
make for a quantized
kernel whose entire value is that it is exact.

## Determinism is structural, not tested-in

The kernel computes `sum_k (a - za)(b - zb)` as `sum_k (a - za) * b - zb
* sum_k (a - za)`, with
every accumulation a **wrapping** `i32` add. Wrapping addition is
arithmetic mod 2^32, which is
associative and commutative, so *any* blocking, tiling, column split,
row split or thread count
gives bit-identical output — including on overflow, where the wrap
itself is reproducible.
`wrapping_overflow_is_reordering_invariant` and
`the_thread_count_cannot_change_the_result` assert
exactly that, and the SIMD path is checked bit-for-bit against a
portable scalar oracle
(`the_simd_kernel_is_bit_identical_to_the_portable_loop`, and separately
for the fused path).

## Numbers

Session A/B against plain ORT, `K = N = 2048`, u8 x u8, ratio is `ours /
ORT`, **lower is better**,
p50 of 61 iterations. ORT's own timings moved under 1.5% between the two
arms at 1 and 4 threads,
which is the control that makes the comparison mean anything.

| M | threads | before | after | ours before | ours after |
| ---: | ---: | ---: | ---: | ---: | ---: |
| 1 | 1 | 12.20x | **2.17x** | 1.402 ms | 0.226 ms |
| 128 | 1 | 11.90x | **1.20x** | 99.84 ms | 9.99 ms |
| 1 | 4 | 37.11x | **4.03x** | 1.379 ms | 0.170 ms |
| 128 | 4 | 14.44x | **1.47x** | 31.03 ms | 3.14 ms |
| 1 | 16 | 83.12x | 35.63x | 2.366 ms | 1.336 ms |
| 128 | 16 | 42.28x | 15.05x | 51.05 ms | 16.55 ms |

`i8_m1` goes 0.206 ms to **0.049 ms** at one thread.

Kernel-level scaling (`bench_qgemm_ab`, `taskset -c 0-15`), with the
portable scalar arm as the
control:

| shape | 1t | 2t | 4t | 8t | 16t | portable 1t |
| --- | ---: | ---: | ---: | ---: | ---: | ---: |
| 1x2048x2048 | 0.229 ms | 0.136 | 0.090 | 0.098 | 0.166 | 4.92 ms (21x)
|
| 4x2048x2048 | 0.565 ms | 0.311 | 0.199 | 0.237 | 0.345 | 4.58 ms
(8.1x) |
| 128x2048x2048 | 8.911 ms | 4.773 | 2.755 | 2.780 | 1.991 | — |
| 128x5120x5120 | 53.56 ms | 27.06 | 14.35 | 8.98 | 11.51 | — |

The task grid splits rows as well as columns. Columns alone gave only `n
/ NC` tasks — eight for
`n = 2048` — so a sixteen-worker pool left half of itself spinning;
`128x2048x2048` was 2.69 ms at
sixteen threads against 1.62 ms at eight. Splitting columns further
would shrink the panel and
re-walk `B`; splitting rows duplicates only the pack, about a percent of
the GEMM it feeds.

## Things I measured and rejected

- **Software prefetch** of the next `B` rows (`PREFETCH_ROWS = 8`): a
consistent **8% regression**
with a stable `m = 128` control. The hardware prefetcher already has the
sequential stream.
- **Permuting inside the fused inner loop**: replaced by accumulators
held in the permuted order
with a single `vperm2i128` fixup per `k`-block flush. Saves eight
instructions per 32 MACs.

## Left open, deliberately

- **Constant-`B` packed cache.** The pack is repeated per call. Caching
it would remove it from
  prefill entirely, but any new weight-derived cache has to go through
`kernels/governed_weight_cache.rs` to satisfy the "New weight-derived
caches must be governed"
gate. That is a separate PR with its own eviction story, not a rider on
this one.
- **The session-level threading gap.** At four threads the session takes
0.170 ms while the kernel
alone does 0.090 ms, and past eight threads both arms get worse. That is
the pre-existing
oversubscription item — it is present before and after this change, so
it is not a regression
  here, and it is the next thing I am working on.

## Validation

- `cargo test --release -p onnx-runtime-ep-cpu --lib` — 1340 passed, 0
failed.
- Every `onnx-runtime-ep-cpu-plugin` suite with
`NXRT_REQUIRE_ORT_TESTS=1`, including the 53-test
  `plugin_ort_e2e` ORT conformance suite with CPU fallback disabled.
- `cargo clippy -p onnx-runtime-ep-cpu --all-targets` clean, `cargo fmt
--all --check` clean.
- `cargo check -p onnx-runtime-ep-cpu --lib --features mlas` — the
research build still compiles.
- Reviewed by Claude Opus 4.8 against the memory-safety, lane-semantics,
determinism and
edge-extent claims above; no blockers, two documentation fixes applied.

---

## Refreshed against `main` (2026-08-18)

The branch was behind `main` and its red CI wall came from that, not
from this
change: `crates/onnx-runtime-session/src/executor/mod.rs:175` failed `-D
dead-code`
on current stable, fixed on `main` by `ca32b3adf` (#1239) after this
branch forked.
`origin/main` (`c55a3fab3`) is merged in — no rebase, no force-push.

One conflict, in `docs/performance/CPU_MATMUL_ASSIGNMENT.md`, resolved
as a
**union**: this branch's `#### 3b` (the native integer GEMM) and
`main`'s
`### 4` (the f32 `M = 1` GEMV becoming the default, #1091) were both new
sections
appended after 3a. Both are kept, in that order. Taking either side
would have
silently deleted the other's record.

Revalidated on the merge commit, AVX2/FMA host, no AVX-512:

* `cargo test --release -p onnx-runtime-ep-cpu --lib` — **1424 passed, 0
failed**,
18 ignored, including
`qgemm_i32_matches_the_integer_oracle_for_every_signedness`,
  `the_simd_kernel_is_bit_identical_to_the_portable_loop`,
  `wrapping_overflow_is_reordering_invariant` and
  `the_thread_count_cannot_change_the_result`.

The measurements in this PR were taken before the merge; nothing in the
merged
range touches `qgemm_native.rs`, `qlinear_matmul.rs`, or the CPU
threadpool, so
they stand as recorded. The `main` change that did land in this range
(#1091's
f32 `M = 1` GEMV default) is on a different kernel family and is
documented in
the section-4 text kept above.

---

## Refreshed again against `main` @ `6a855d5e0`, and a real branch bug
found

`main` moved again while this was queued (#1346/#1352/#1361 on the
quality lane,
#1154/#1232/#1238 on the CPU side). Merged in normally — no rebase — and
revalidated.

The revalidation caught something the earlier ones had not. Running
`-p onnx-runtime-ep-cpu --lib` in a **debug** profile rather than
`--release` fails:

```
kernels::qgemm_native::tests::degenerate_extents_do_nothing
  assertion `left == right` failed
  left: 0
 right: 4
```

`degenerate_extents_do_nothing` called `qgemm` with an empty
`b_zero_points`
and `n == 4`. `qgemm` opens with `debug_assert_eq!(b_zero_points.len(),
n)`,
so that call is not one the function accepts — the test was exercising
the
`m == 0` early return through an argument list the contract forbids. It
passed
every previous run here only because `debug_assert` compiles out under
`--release`, which is how I had been validating this branch locally. A
debug
test profile fails it, and this is branch-caused: `qgemm_native.rs` is
new in
this PR.

Fixed in `9ca99e538` by sizing the test's zero points to `m` and `n`,
not by
weakening the assertion — the assertion states the contract the kernel's
indexing depends on, and a caller whose `m` is zero still has `n`
columns and
still knows their zero points.

**`-p onnx-runtime-ep-cpu --lib`, debug profile: 1440 passed, 0 failed**
(was
1439 passed, 1 failed).

This is the second time on this stack that the profile a test runs under
decided whether it caught anything. Worth remembering: `--release`
silently
disables every `debug_assert` in the crate under test, so a local
`cargo test --release` is not a substitute for what CI runs.

---

## Re-validated on latest `main` (`e0aedd0fa`), 2026-08-19

Latest `main` merged in normally (no rebase). Full re-measurement, 1
thread
pinned, `K = N = 2048`, 61 iters / 10 warmup, 2 reps, `ours_p50 /
ort_p50`:

| case | `main` ours | `main` ratio | this PR ours | this PR ratio |
speedup |
|---|---|---|---|---|---|
| `bench_qlinear_u8_m1` | 1.418 / 1.435 ms | 11.83x / 12.51x | **0.121 /
0.123 ms** | **1.16x / 1.18x** | **11.7x** |
| `bench_qlinear_u8_m128` | 29.55 / 29.56 ms | 3.57x / 3.57x | **3.055 /
3.065 ms** | **0.372x / 0.373x** | **9.6x** |
| `bench_qlinear_i8_m1` | 1.516 / 1.497 ms | 0.215x / 0.212x | **0.209 /
0.212 ms** | **0.030x / 0.030x** | **7.2x** |

ORT-side drift between the two arms was 0.7% at `m = 128` and 0.0% on
`i8`,
which is the control that makes the comparison mean anything.

**At `m = 128` we are now 2.7x faster than ORT outright**, and `m = 1`
closes
from 11.8x to 1.16x. These are better than the numbers originally posted
above
because the dispatch work in #1077 landed in between.

## Review fixes (`987aa0c5c`)

An independent review found no blockers but two things worth fixing:

1. **aarch64 built with 5 warnings** — `NR`/`MR`/`NC`/`KC`/`FUSED_KC`
are read
only by the x86 kernels, so every non-x86 target warned on all five. CI
builds with `-D warnings`, so this was a branch-caused CI failure
waiting to
happen; the local x86 clippy run could never have caught it. Now
`#[cfg]`-gated
alongside the code that uses them: **0 warnings on both x86-64 and
aarch64**.
2. **The fused-parallel path had no end-to-end coverage.** Every `m <=
4` shape
   in `qlinear_matmul_reordered_accumulation_is_bit_identical` sat below
`PARALLEL_MIN_WORK`, so the pack-free kernel's column split was only
ever
   checked at the kernel level, never through `requantize_rows`. Added
   `(4, 1029, 1100)`, which forks both.

The review independently re-derived the register-shuffle math in numpy
(`cvtep*_epi16`, `permute4x64_epi64(0xD8)`, `unpacklo/hi_epi16`,
`madd_epi16`,
`permute2x128`) against a plain per-column dot product over 2000 tiles
with
extreme values — 0 mismatches — and confirmed the `vpmaddwd`
non-saturation
bound for all four operand combos, the wrapping-add determinism claim,
and the
absence of out-of-bounds access in every tail path.

## Validation on the merged base

- `cargo fmt` clean; `cargo clippy --all-targets -D warnings` clean
- **1553 `onnx-runtime-ep-cpu` tests**, debug profile (so
`debug_assert`s are live)
- **55 plugin conformance tests** (`NXRT_REQUIRE_ORT_TESTS=1`, release)
- `every_assigned_node_is_also_executed_by_this_ep` and
`every_fixture_loads_with_cpu_fallback_disabled` green — nothing defers,
  nothing falls back to the ORT CPU EP
- **aarch64-unknown-linux-gnu** cross-check clean, 0 warnings

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Co-authored-by: Resch <resch@squad.local>
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.

1 participant