Skip to content

perf(cpu-ep): run the elementwise split on ORT's pool, not a second one - #1143

Merged
justinchuby merged 7 commits into
mainfrom
deckard/host-parallel
Aug 18, 2026
Merged

justinchuby merged 7 commits into
mainfrom
deckard/host-parallel

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 17, 2026 •

Copy link
Copy Markdown
Owner

What

When our CPU EP runs inside an ORT session, our elementwise activation kernels
stop starting a second thread pool. They borrow ORT's intra-op pool through
OrtApi::KernelContext_ParallelFor instead.

Root cause

simd_activations::run_chunked splits any f32 slice of ≥ 1 Mi across a rayon
pool. That is the right thing for the native executor, which owns the machine.
Inside an ORT session it is not: ORT already has an intra-op pool of its own,
and its workers spin. Sixteen rayon workers next to sixteen spinning ORT
workers puts 32 runnable threads on 16 cores.

Measured on this machine (AMD EPYC 9V74, taskset -c 0-15, ORT 1.28,
intra_op_num_threads = 16, 1 Mi f32, p50, same .so, RAYON_NUM_THREADS
1 vs 16):

op serial our rayon split
Sqrt 252 µs 777 µs 3.08× slower
Relu 345 µs 607 µs 1.76× slower
Clip 399 µs 858 µs 2.15× slower
Sigmoid 521 µs 993 µs 1.91× slower
Tanh 569 µs 944 µs 1.66× slower
QuickGelu 790 µs 1320 µs 1.67× slower
Erf 1020 µs 1388 µs 1.36× slower
FastGelu 1040 µs 1479 µs 1.42× slower
Gelu 1570 µs 1808 µs 1.15× slower

Parallelising made every op slower.

Why a bigger threshold is not the fix

The obvious repair is to raise PAR_MIN_LEN until the split only fires where
it still wins. It does not work, because the same sweep at
intra_op_num_threads = 1 says the opposite — there the split is a large win
from 1 Mi upwards:

op, 4 Mi f32 serial our rayon split
Erf 3551 µs 700 µs 5.07× faster
Gelu 5315 µs 1047 µs 5.08× faster

No constant satisfies both, because the variable that decides the answer is not
the slice length. It is how much of the machine the host is already using, and
only the host knows that.

Fix

Ask the host.

  • onnx-runtime-ep-api::host_parallel — a thread-local seam. It is the only
    crate both onnx-runtime-ep-cpu (which has the kernels) and
    onnx-runtime-ep-plugin (which has the OrtKernelContext) already depend on,
    so neither has to learn about the other. HostParallel::run(total, body) is
    blocking, which is what makes borrowing a &dyn Fn from the caller's frame
    sound; bodies run with in_host_task() set so a nested split can see it is
    already on a host thread.
  • onnx-runtime-ep-plugin::host_pool — the ORT-backed implementation over
    KernelContext_ParallelFor (num_batch = 0, so ORT's workers claim indices
    dynamically). install() returns an RAII guard; compute_execute holds one
    for the whole call.
  • simd_activations — run_chunked and run_chunked_rows prefer the host
    pool when one is installed, and only fall back to rayon when there is none.

Because we cannot ask ORT how wide its pool is, the host split cuts by size
rather than by thread count: MAX_HOST_CHUNKS = 64 pieces, floored at the
existing PAR_MIN_CHUNK. That also makes the chunk boundaries — and therefore
the bits — independent of the machine and of the session's thread count.

Nothing is handed to ORT's CPU EP

This is a threading change, not an assignment change. Every node our EP claims
is still computed by our kernels; the only thing that comes from ORT is the
threads they run on. No op is declined, no capability filter changes.

Correctness

cargo test -p onnx-runtime-ep-cpu --features mlas --lib → 1339 passed
(1332 on main + 7), -p onnx-runtime-ep-api → 51 (+9), -p onnx-runtime-ep-plugin
→ 238 (+7).

New coverage:

  • host_pool_split::{unary,bias}_kernels_match_the_unsplit_result — every unary
    and bias-fused kernel, split across a real four-thread stand-in host, is
    bit-identical to the unsplit result. A serial stand-in would prove the
    arithmetic but not the disjointness of the ranges, which is the part that
    would corrupt an output tensor.
  • the_rayon_pool_is_not_used_when_a_host_is_installed — the whole point of the
    change, asserted: exactly host_chunk_len(N).1 chunks dispatched, none to rayon.
  • a_nested_split_stays_serial — a kernel reached from inside a host task does
    not dispatch again.
  • serial_scope_still_suppresses_the_split — the f16/bf16 sandwich stays serial
    on the host path too.
  • host_chunk_policy_holds_across_lengths — whole vectors, never below
    SIMD_MIN_LEN, never a cut through a bias row, never an empty final range,
    never past the cap; swept to usize::MAX / 2.
  • host_pool::{a_refused_dispatch_still_runs_every_index, a_panicking_body_does_not_unwind_into_ort}
    — a failed KernelContext_ParallelFor still runs every index (a short write
    would leave the output tensor uninitialised), and a Rust panic is caught on
    the worker and re-raised on the calling thread rather than unwinding into C++.
  • host_parallel::{scope_restores_on_unwind, a_handle_is_not_visible_from_another_thread}
    — a leaked handle would be a dangling OrtKernelContext.

Benchmarks

The arm that decides this PR: the host-pool split vs staying serial. Every
number in "Root cause" compares our rayon split against serial, and serial won
every row — which says nothing about the change this PR actually makes, routing
the split onto ORT's own pool. That arm is measured here, and the
host-pool split beats serial by ~5.5–6.6×.

Measured on a 13th-gen i7-13800H (14 physical cores / 20 logical), release
build, 1 Mi f32, three arms in one process, interleaved per rep (p50 µs over
201 reps), repeated three times. intra_op is modelled at 14 rather than 16
because the stand-in's spinning workers are real threads: on a 14-core box,
15 spinners + the caller already oversubscribe and would starve the serial arm
before rayon even runs, so the width is matched to the cores. The stand-in
reproduces the one thing that matters — an intra-op pool that spins while
idle
, which is exactly why a coexisting rayon pool oversubscribes — and it
latches HOST_HELPED the honest way (a worker, not the dispatcher, runs a
chunk), so prefer_host reaches steady state through the real mechanism.

The rayon-split arm is a built-in control: its result is already known from
the table above (serial wins at intra_op = 16). It reproduces here — rayon is
slower than serial in every run — which validates the harness before the
host-pool number is trusted.

op run serial rayon split host-pool split rayon / serial host / serial
Sqrt 1 224 297 41 1.33× (loss) 0.18× (5.5× faster)
Sqrt 2 215 490 40 2.28× (loss) 0.18× (5.4× faster)
Sqrt 3 223 302 38 1.35× (loss) 0.17× (5.9× faster)
Gelu 1 1376 1514 209 1.10× (loss) 0.15× (6.6× faster)
Gelu 2 1395 1509 209 1.08× (loss) 0.15× (6.7× faster)
Gelu 3 1413 1510 213 1.07× (loss) 0.15× (6.6× faster)

The host-pool win (~5.5–6.6×) is an order of magnitude larger than each arm's
own p10..p90 spread (e.g. Sqrt host 34..62 µs, serial 187..256 µs), so it is a
real signal, not noise. Process CPU time (contention-immune) reproduced to
16.6 / 16.2 / 12.8 s across the three runs.

Why so much larger than the rayon split's best case? Because the host cut floors
at HOST_MIN_CHUNK (4 Ki), not the rayon path's 256 Ki: 1 Mi becomes 64 tasks
the host's threads claim dynamically, where the rayon path would make four. The
warm, already-spinning pool turns that into near-linear speedup with none of the
wake-up cost that made a second pool lose.

The no-host fall-through still wins at intra_op = 1

With no host installed (native executor, ORT < 1.17, or a null context) the
kernels keep the rayon path. On a free machine that path is still the right one
(p50 µs, same box):

op serial rayon split
Gelu, 4 Mi 4012 µs 1687 µs 2.38× faster
Sqrt, 4 Mi 2044 µs 914 µs 2.24× faster
Gelu, 1 Mi 1012 µs 631 µs 1.60× faster
Sqrt, 1 Mi 181 µs 172 µs 1.05× (memory-bound; unchanged from main)

(The absolute multiples are smaller than the EPYC's 5× because this box has
fewer, heterogeneous P+E cores; the direction — rayon split wins on a free
machine — is what the fall-through must preserve, and it does.)

Conclusion. Dispatching the elementwise split onto the host's own pool
beats staying serial by ~5.5–6.6× at intra_op = 16, while the rayon split
loses under the same busy pool. The simpler "stay serial under a host"
alternative would leave that entire ~6× on the table — so the host-pool split,
KernelContext_ParallelFor and all, is the correct fix, not merely the more
complex one. The PR merges as written.

The harness is committed as two #[ignore]d tests in
simd_activations::three_arm_bench, run with
EP_BENCH=1 cargo test --release -p onnx-runtime-ep-cpu --lib three_arm_bench -- --ignored --nocapture --test-threads=1.

Limitations

Our elementwise kernels split long slices across a rayon pool. That is
right for the native executor, which owns the machine, and wrong inside
an ORT session: ORT already has an intra-op pool and its workers spin,
so ours is a second pool on the same cores.

Measured at `intra_op_num_threads = 16`, 1 Mi f32, p50 (serial vs our
rayon split, same binary):

| op       | serial | split   |
|----------|--------|---------|
| Sqrt     |  252us |  777us  |
| Sigmoid  |  521us |  993us  |
| FastGelu | 1040us | 1479us  |

Every op lost by parallelising. Raising the length threshold does not
fix it, because at `intra_op = 1` the same split is a 2-5x *win* from
1 Mi upwards -- the variable that matters is how much of the machine
the host is already using, and only the host knows that.

So ask it. `OrtApi::KernelContext_ParallelFor` runs a callback on the
session's own intra-op pool. This adds a `host_parallel` seam in
`onnx-runtime-ep-api` (the only crate both the kernels and the plugin
already depend on), an ORT-backed implementation in the plugin, and
installs it for the dynamic extent of each `compute_execute`. One pool,
sized by whatever the user configured, and the oversubscription is gone
by construction rather than by tuning.

The native executor installs nothing and keeps its rayon path.

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

codecov Bot commented Aug 17, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 79.01829% with 218 lines in your changes missing coverage. Please review.
✅ Project coverage is 80.00%. Comparing base (11043a0) to head (17ec3d5).
⚠️ Report is 1 commits behind head on main.

Files with missing lines Patch % Lines
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 58.96% 204 Missing and 2 partials ⚠️
crates/onnx-runtime-ep-api/src/host_parallel.rs 96.33% 7 Missing and 1 partial ⚠️
crates/onnx-runtime-ep-plugin/src/host_pool.rs 98.73% 2 Missing and 2 partials ⚠️
Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1143      +/-   ##
==========================================
- Coverage   80.53%   80.00%   -0.54%     
==========================================
  Files         368      370       +2     
  Lines      161313   162350    +1037     
  Branches   161313   162350    +1037     
==========================================
- Hits       129921   129885      -36     
- Misses      26659    27730    +1071     
- Partials     4733     4735       +2     
Flag Coverage Δ
cli-ort-linux 83.79% <ø> (ø)
cli-ort-windows 83.40% <ø> (ø)
mlas 85.67% <ø> (-0.14%) ⬇️
offline 79.75% <79.01%> (-0.57%) ⬇️

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

Files with missing lines Coverage Δ
crates/onnx-runtime-ep-plugin/src/compute.rs 78.94% <100.00%> (+0.01%) ⬆️
crates/onnx-runtime-ep-plugin/src/lib.rs 100.00% <ø> (ø)
crates/onnx-runtime-ep-plugin/src/host_pool.rs 98.73% <98.73%> (ø)
crates/onnx-runtime-ep-api/src/host_parallel.rs 96.33% <96.33%> (ø)
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 88.58% <58.96%> (-8.54%) ⬇️

... and 7 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 17, 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
🔴 gather/medium_f32_threads=1-internal/32768 3.75 µs 6.39 µs +70.2%
🔴 block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 503.32 µs 744.26 µs +47.9%
🔴 gather/small_bf16_threads=1-internal/4096 442.8 ns 654.2 ns +47.8%
🔴 matmul/large_generic_bf16_threads=8/32x1024x1024 1.27 ms 1.83 ms +44.7%
🔴 add/large_bf16_threads=1-internal/4194304 1.53 ms 2.05 ms +33.9%
⚠️ add/large_f16_threads=1-internal/4194304 1.53 ms 1.98 ms +29.9%
⚠️ matmul/small_generic_f16_threads=8/1x256x256 27.92 µs 34.91 µs +25.1%
⚠️ reduce_mean/large_f32_threads=1-internal/262144 919.93 µs 1.13 ms +22.9%
⚠️ tokenization/decode_tokens_per_second 5.68 ms 6.96 ms +22.6%
⚠️ matmul/small_generic_bf16_threads=8/1x256x256 29.15 µs 35.44 µs +21.6%
⚠️ reduce_mean/medium_f32_threads=1-internal/65536 231.33 µs 280.20 µs +21.1%
⚠️ grammar_masking/llguidance_compute_mask/32 69.26 µs 83.21 µs +20.1%
⚠️ gather/small_f16_threads=1-internal/4096 456.7 ns 547.6 ns +19.9%
⚠️ tokenization/encode_tokens_per_second 344.02 µs 411.05 µs +19.5%
⚠️ qwen3_sampling_processors/top_p_fast_after_top_k 482.50 µs 574.58 µs +19.1%
⚠️ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.26 ms 3.77 ms +15.9%
⚠️ gather/large_f32_threads=1-internal/131072 26.19 µs 30.31 µs +15.7%
⚠️ sampling_latency/min_p_per_token 192.43 µs 221.35 µs +15.0%
✅ matmul/small_generic_f16_threads=1/1x256x256 29.04 µs 33.37 µs +14.9%
✅ reduce_mean/small_f32_threads=1-internal/4096 16.06 µs 18.42 µs +14.7%
✅ matmul/small_generic_bf16_threads=1/1x256x256 27.85 µs 31.91 µs +14.6%
✅ qwen3_sampling_processors/top_k_partial_selection 132.12 µs 150.81 µs +14.1%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 426.78 µs 484.29 µs +13.5%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.25 ms 5.92 ms +12.7%
✅ block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 40.70 µs 45.70 µs +12.3%
✅ gather/medium_bf16_threads=1-internal/32768 2.23 µs 2.49 µs +11.4%
✅ matmul/medium_generic_f32_threads=8/32x512x512 882.79 µs 982.94 µs +11.3%
✅ add/medium_bf16_threads=1-internal/262144 95.64 µs 106.05 µs +10.9%
✅ sampling_latency/greedy_per_token 2.99 µs 3.30 µs +10.2%
✅ add/large_f32_threads=1-internal/4194304 665.45 µs 732.08 µs +10.0%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 155.31 µs 170.67 µs +9.9%
✅ sampling_latency/top_p_per_token 354.60 µs 384.50 µs +8.4%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 1.99 ms 2.14 ms +7.9%
✅ qwen3_sampling_processors/top_k_top_p_fast 611.53 µs 658.39 µs +7.7%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 1.99 ms 2.11 ms +6.3%
✅ logit_processing/seven_processor_chain_per_step 299.29 µs 316.80 µs +5.9%
✅ matmul/medium_generic_f16_threads=8/32x512x512 27.39 µs 28.89 µs +5.5%
✅ gather/small_f32_threads=1-internal/4096 627.2 ns 661.7 ns +5.5%
✅ kv_cache/alloc_dealloc_pages 36.63 µs 38.54 µs +5.2%
✅ matmul/large_generic_f32_threads=8/32x1024x1024 3.65 ms 3.82 ms +4.7%
✅ block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 40.42 µs 42.12 µs +4.2%
✅ matmul/small_generic_f32_threads=8/1x256x256 48.46 µs 50.00 µs +3.2%
✅ matmul/small_generic_f32_threads=1/1x256x256 33.73 µs 34.70 µs +2.9%
✅ gather/large_f16_threads=1-internal/131072 13.60 µs 13.96 µs +2.6%
✅ add/small_f16_threads=1-internal/1024 428.1 ns 439.1 ns +2.6%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 78.74 µs 80.69 µs +2.5%
✅ sampling_latency/top_k_per_token 48.46 µs 49.27 µs +1.7%
✅ add/medium_f16_threads=1-internal/262144 99.69 µs 101.13 µs +1.4%
✅ matmul/medium_generic_bf16_threads=8/32x512x512 369.46 µs 374.28 µs +1.3%
✅ matmul/medium_generic_f16_threads=1/32x512x512 27.70 µs 27.97 µs +1.0%
✅ gather/medium_f16_threads=1-internal/32768 2.60 µs 2.62 µs +0.8%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.13 ms 2.14 ms +0.6%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 8.77 ms 8.80 ms +0.4%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 74.19 µs 74.46 µs +0.4%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 492.20 µs 486.83 µs -1.1%
✅ add/small_f32_threads=1-internal/1024 206.1 ns 194.5 ns -5.6%
✅ add/small_bf16_threads=1-internal/1024 470.9 ns 422.1 ns -10.4%
🟢 gather/large_bf16_threads=1-internal/131072 14.80 µs 12.48 µs -15.7%
🟢 add/medium_f32_threads=1-internal/262144 30.71 µs 23.45 µs -23.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.84 3.51 5.20 }
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 and others added 2 commits August 17, 2026 21:31
The previous commit always preferred `KernelContext_ParallelFor` over our
own rayon pool. That is right when ORT's intra-op pool is wide -- ours
would be a second pool fighting it for the same cores -- but wrong when
the session was built with `intra_op = 1`. There the host is not using
the machine at all, and borrowing its single thread ran 2-9x slower over
1-4 Mi than splitting across rayon, which is what main did.

We cannot ask ORT how wide its pool is, so observe it: the first dispatch
with more than one index records whether any index ran on a thread other
than the caller's, and stores the verdict in a cell owned by the fused
node (not a global -- one process may hold both a 1-thread and a
16-thread session, and the right answer is the opposite for each).
Until that verdict exists the host path is taken, which costs at most one
dispatch on a serial host.

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

Copy link
Copy Markdown
Owner Author

The seam is the right design and I want to say why before the objection: the plugin installs it itself.

compute.rs takes the guard from the OrtKernelContext ORT handed the call, not from anything our engine sets. That makes it the structural opposite of IN_SPMD_SCOPE (#1138), where only onnx-genai-engine set the thread-local and every third-party host inherited the costs with none of the benefit. Here any ORT host gets it with no cooperation. That is the contract shape #1138 is asking for, and #1142 only solved half of.

Gates pass on my box: -p onnx-runtime-ep-cpu --lib 1310 passed / 0 failed, -p onnx-runtime-ep-api 48 passed. The fall-through is also right -- with no host installed the rayon path is unchanged, so the native executor keeps the 5.07x it gets at intra_op = 1.

The blocking gap: the benchmark that decides this PR has not been run

The Benchmarks section says:

Filling in: interleaved same-machine A/B against ORT at intra_op 1 and 16, before and after, for decode and prefill sizes.

Everything measured so far compares our rayon split against staying serial, and serial wins every row -- Sqrt 252 vs 777 us, and eight more. No measurement in this PR shows that dispatching to ORT's pool beats staying serial. That is the change being made.

So the evidence supports a strictly simpler alternative that the PR never tests:

when a host pool is installed, stay serial

...which needs no KernelContext_ParallelFor, no 64-chunk policy, no disjointness argument, no panic-catching across the FFI boundary, and no ORT-1.17 fallback -- roughly 700 of the 1060 added lines. If the ORT-pool split does not beat serial at intra_op = 16, that one-line version is the correct fix and this PR is mostly cost.

I am not asserting it loses. ORT's pool is the right pool, its workers are already warm, and num_batch = 0 dynamic claiming is the right call. It is entirely plausible it beats both. But the arm that would tell us apart is the one arm not in the table, and this is the fourth time this month the measured mechanism has contradicted the predicted one -- including on #1142, where the intuitive repair (parallelise the shards) measured slower than doing nothing, for the same fork-join reason that motivates this PR.

Please add one row: Sqrt and Gelu, 1 Mi, intra_op = 16, three arms in one process, interleaved -- serial / rayon split / host-pool split. Two ops is enough; the nine-row sweep already told us the direction is uniform. If host-pool wins, the PR merges as written and that row is the whole justification. If it does not, we delete most of it and keep the seam only where it pays.

Smaller points, none blocking

  • MAX_HOST_CHUNKS = 64 cutting by size rather than by pool width is a good call for a reason worth keeping in the comment: it makes the bit pattern independent of the session's thread count, so results do not vary with intra_op. That is a correctness property, not a tuning one.
  • a_refused_dispatch_still_runs_every_index and a_panicking_body_does_not_unwind_into_ort are the two tests I would have asked for. A short write leaving an output tensor uninitialised is exactly the failure that would surface as a wrong logit three layers later.
  • the_rayon_pool_is_not_used_when_a_host_is_installed asserts the actual point of the change rather than a proxy for it. Good.
  • The four-thread stand-in host rather than a serial one is the right choice -- a serial stand-in proves the arithmetic but not the disjointness, and disjointness is the part that corrupts tensors.
  • Limitation 2 (other kernels still start rayon under an ORT session) is the real remaining scope. Once the seam is proven, MatMul and the int4 paths matter far more than elementwise. Please file that rather than grow this PR.

justinchuby pushed a commit that referenced this pull request Aug 17, 2026
The weight-cache guard failed the test it exists to pass. #1133 parked a
QLinearMatMul i32 accumulator in a thread_local! RefCell<Vec<i32>> bounded by a
per-buffer 32 MiB constant, while the buffer is retained on every worker thread
for the life of the process: 640 MiB on a 20-thread box, 1 GiB on a 32-vCPU one,
4 GiB on a 128-vCPU one. Run against that PR's 459 added lines the existing
regex matched zero, because it only knew OnceLock/OnceCell/LazyLock.

That is the guard committing, one level up, the defect it was written to catch:
a check that is green but structurally incapable of failing on the case that
motivated it.

The new job does not merely also-match thread_local. It asks the question the
first job does not -- what multiplies this buffer -- because the underlying
error has now cost four rounds: #1051 reported 247 MB against 592 MB measured,
#1100's ratio test drove a single instantiation so it could not observe the x2,
and #1133 bounded one copy of an N-per-thread buffer. Each comment was correct
about one copy and silent about N. Under-reporting is worse than reporting zero:
zero is obviously blind and gets caught at review, whereas a plausible 32 MiB
passes admission and then overruns.

The error text also requires that any test for such a buffer drive it from more
than one thread, since a single-instantiation test cannot observe an xN
multiplier -- which is exactly how #1100 shipped with the factor unmeasured.

Falsified against real history rather than assumed to work:

  #1133  (must fail)  matches=1  -> flagged
  #1143  (must pass)  matches=0  -> passes
  #1142  (must pass)  matches=0  -> passes

The pattern is rare in this tree (3 occurrences), so the false-positive cost is
low, and the 'per-thread-bound-reviewed' label records a deliberate judgement
rather than blocking.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
justinchuby and others added 3 commits August 17, 2026 23:00
The previous commit inferred the host pool's width from one dispatch: if
every body ran on the calling thread, ORT had no workers. Measurement
says that inference is unsound. On a 16-thread session 35% of dispatches
ran entirely on the calling thread -- in runs of up to 60 -- because ORT
hands its indices out dynamically and an unstalled caller drains them
before a worker wakes. Acting on that would start our pool alongside
ORT's sixteen, the 3-10x pathology this seam exists to remove.

So stop inferring and require positive evidence: a body seen running on a
thread that was not the one that dispatched it. Only a pool with workers
can produce that, so the verdict is permanent and cannot be faked by a
serial session. Until it arrives the kernels stay on their own pool --
exactly what they did before this seam existed -- except on probe
dispatches, which hold the caller's first index open for 100 us so a
worker that exists has time to claim another. Probes run back to back for
the first 32 dispatches of a session, then back off geometrically to one
in a thousand, so a session whose pool really is serial pays almost
nothing and still recovers if the opening burst was unlucky.

Measured at 16 threads (ORT/ours, 4 interleaved rounds, taskset 0-15):
every session now latches, and 1 Mi goes from 0.07-0.14 to 0.44-1.35.
At intra_op=1 the branch matches main within noise (0.45-2.79 vs
0.46-2.99), i.e. the rayon split is kept exactly where it was winning.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Writing the host decision inline in `run_chunked` cost `Relu` 34% at 1 Mi
on one thread -- 236 -> 315 us, reproducible to the microsecond across
runs -- with no change to the path it actually took: with `intra_op = 1`
nothing ever latches, so every one of those dispatches went down the same
serial call as before. Adding an `eprintln!` to diagnose it moved `Relu`
back to 237 us and pushed `Tanh` from 440 to 1110, which is the tell:
this is the codegen-unit repartitioning already documented for
`clip_chunked`, not a runtime effect.

The branch has no business being inline anyway. It decides with one
relaxed load and the split behind it only happens inside a session whose
pool has proved parallel, while `run_chunked`'s callers are the hottest
elementwise kernels in the crate. `#[inline(never)]` on `try_host` and
`try_host_rows`, with a test that keeps it there.

Measured at intra_op=1, rayon=1, 4 interleaved rounds (branch vs main, us
at 1 Mi): Relu 236.4/236.3, Clip 268.1/314.2, Gelu 1260.8/1303.9, Tanh
443.7/472.9, Erf 894.7/894.8 -- i.e. the regression is gone and Clip
gained. The 16-thread win is unchanged: Clip 1 Mi 697.5 -> 85.5 us,
Relu 693.7 -> 79.0, Erf 1398.0 -> 165.9, Gelu 1937.4 -> 240.6.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Add a gated (EP_BENCH=1, #[ignore]d) three-arm microbenchmark in
simd_activations::three_arm_bench that measures serial vs rayon split vs
host-pool split in one interleaved process, driving the real kernel paths
against a spinning stand-in pool that models ORT's always-hot intra-op pool.
The host-pool split beats serial ~5.5-6.6x at intra_op=16 on 1 Mi f32; the
rayon-split arm reproduces its known loss as a built-in control; the no-host
fall-through still wins at intra_op=1.

Also strengthen the MAX_HOST_CHUNKS comment: cutting by size rather than pool
width keeps the chunk boundaries (and thus the bit pattern) independent of the
session's thread count, which is a correctness property, not a tuning one.

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

Copy link
Copy Markdown
Owner Author

Three-arm measurement: the host-pool split beats serial ~5.5–6.6×

The one arm that decides this PR — host-pool split vs staying serial — is now measured. It wins. The "stay serial under a host" simplification would leave ~6× on the table, so the KernelContext_ParallelFor machinery is justified, not merely more complex. The PR merges as written.

i7-13800H (14 physical / 20 logical), release, 1 Mi f32, three arms in one process, interleaved per rep, p50 over 201 reps, ×3 runs. intra_op modelled at 14: the stand-in's spinning workers are real threads, so 15 spinners + caller would oversubscribe a 14-core box and starve the serial arm before rayon even runs — width matched to cores. The stand-in reproduces the one thing that matters, an intra-op pool that spins while idle, and latches HOST_HELPED honestly (a worker, not the dispatcher, runs a chunk).

The rayon-split arm is a built-in control: its loss at busy intra_op is already known from the PR's nine-row EPYC table. It reproduces here (rayon slower than serial every run), which validates the harness before the host-pool number is trusted.

op run serial µs rayon split µs host-pool split µs rayon / serial host / serial
Sqrt 1 224 297 41 1.33× (loss) 0.18× (5.5× faster)
Sqrt 2 215 490 40 2.28× (loss) 0.18× (5.4× faster)
Sqrt 3 223 302 38 1.35× (loss) 0.17× (5.9× faster)
Gelu 1 1376 1514 209 1.10× (loss) 0.15× (6.6× faster)
Gelu 2 1395 1509 209 1.08× (loss) 0.15× (6.7× faster)
Gelu 3 1413 1510 213 1.07× (loss) 0.15× (6.6× faster)

The win (~5.5–6.6×) is an order of magnitude larger than each arm's own p10..p90 spread (Sqrt host 34..62 µs, serial 187..256 µs) — real signal, not noise. Process CPU time (contention-immune, per the profiling SKILL) reproduced to 16.6 / 16.2 / 12.8 s across runs while another agent built on the same box.

Why so much larger than the rayon split's best case? The host cut floors at HOST_MIN_CHUNK (4 Ki), not the rayon path's 256 Ki: 1 Mi becomes 64 tasks the host's warm threads claim dynamically (num_batch = 0), where the rayon path would make four and pay wake-up cost on a second pool.

No-host fall-through still wins at intra_op = 1

With no host installed the kernels keep the rayon path, and on a free machine it is still right:

op serial µs rayon split µs
Gelu, 4 Mi 4012 1687 2.38× faster
Sqrt, 4 Mi 2044 914 2.24× faster
Gelu, 1 Mi 1012 631 1.60× faster
Sqrt, 1 Mi 181 172 1.05× (memory-bound; unchanged from main)

(Smaller absolute multiples than the EPYC's 5× because this box has fewer, heterogeneous P+E cores; the direction — rayon wins on a free machine — is what the fall-through must preserve, and does.)

Harness

Committed as two #[ignore]d, EP_BENCH-gated tests in simd_activations::three_arm_bench. It drives the real kernel entry points (sqrt_f32_slice, erf_gelu_f32_slice) and the real HostParallel seam (probe latch and all), not a reimplementation:

EP_BENCH=1 cargo test --release -p onnx-runtime-ep-cpu --lib three_arm_bench -- --ignored --nocapture --test-threads=1

The same commit also strengthens the MAX_HOST_CHUNKS comment to record why the cut is by size not pool width: it keeps the chunk boundaries — and thus the output bit pattern — independent of the session's thread count. That is a correctness property (a tensor produces the same bits at intra_op = 1 or 16), not a tuning one.

Gates (this branch, after rebase onto origin/deckard/host-parallel @ ee217724)

  • cargo test -p onnx-runtime-ep-cpu --lib → 1316 passed, 0 failed, 16 ignored
  • cargo test -p onnx-runtime-ep-api --lib → 53 passed, 0 failed
  • cargo test -p onnx-runtime-ep-plugin --lib → 246 passed, 0 failed
  • cargo clippy -p onnx-runtime-ep-cpu -p onnx-runtime-ep-api -p onnx-runtime-ep-plugin --all-targets -- -D warnings → clean

Follow-up

Limitation 2 (other kernels — MatMul, int4 — still start rayon work under an ORT session) is filed as #1145, referencing #1138 and this PR, rather than growing this PR's scope.

@justinchuby

Copy link
Copy Markdown
Owner Author

Reviewed the harness rather than the table, because the table is only as good as the arm that produced it. It holds up. Merging.

What convinces me

The control arm. The rayon-split arm reproduces its already-known loss (slower than serial in all six runs, matching the nine-row EPYC table). That is the strongest evidence shape available here: an arm the change provably cannot help, moving the way prior data says it should, before the headline number is read. Without it, "host-pool is 6x serial" would just be "14 threads beat 1 thread", which nobody disputed.

The stand-in is honest about what it models. I read three_arm_bench on 17ec3d5a. It is not a reimplementation of the kernel:

  • workers claim indices dynamically off a cursor, as ParallelFor with num_batch = 0 does;
  • the calling thread drains too (drain(total, is_worker)), because ORT runs tasks on the caller -- so the caller is one of the intra_op threads, not a spectator. Getting this wrong would have inflated the host arm by a thread;
  • idle workers spin_loop(), which is the single property that makes a coexisting rayon pool oversubscribe -- i.e. the mechanism under test is present, not assumed;
  • worker_helped is set by a worker, never the dispatcher, so the probe latches on real positive evidence.

Sizing workers + 1 == intra_op at 14 on a 14-physical-core box is the right call and the reasoning is stated: 15 spinners plus a caller would have starved the serial arm before rayon ever ran, which would have manufactured the result.

The margin dwarfs the spread. 5.5-6.6x against a p10..p90 of 34..62 us (host) and 187..256 us (serial). And CPU time was reported alongside wall clock, which is what makes a number taken while another agent was building on the same box readable at all.

What is still not measured, and why it cannot flip this

The pool is a stand-in, not ORT's KernelContext_ParallelFor. Real dispatch overhead, real batching behaviour under num_batch = 0, and real worker wake-up policy are all modelled rather than observed.

That is a genuine residual, and it should be said out loud rather than left implicit in the word "stand-in". But it cannot change the decision this PR turns on. The choice was between the ParallelFor machinery and the one-line "stay serial under a host" simplification. For serial to win, the real host pool would have to be more than 5.5x worse than the model -- not 20% worse, not 2x worse. Nothing about ORT's pool suggests that. The decision is robust to the modelling error by an order of magnitude, which is the only reason a modelled arm is admissible here at all.

Two things worth keeping for their own sake

MAX_HOST_CHUNKS cutting by size rather than pool width makes the chunk boundaries -- and therefore the output bit pattern -- independent of the session's intra_op. A tensor produces the same bits at 1 thread or 16. The comment now says so. That is a correctness property masquerading as a tuning constant, and it is exactly the kind of thing that gets "optimised" away by someone who reads it as tuning.

And scoping Limitation 2 out to #1145 instead of growing a 2150-line PR is the right instinct.

Verified read-only; I did not re-run the bench myself, because a second agent was building on this box and a contended re-run would be worse evidence than the interleaved one already taken, not better.

@justinchuby
justinchuby marked this pull request as ready for review August 18, 2026 00:27
@justinchuby
justinchuby merged commit 266a6fe into main Aug 18, 2026
12 of 19 checks passed
@justinchuby
justinchuby deleted the deckard/host-parallel branch August 18, 2026 00:28
justinchuby added a commit that referenced this pull request Aug 19, 2026
perf(cpu): run the MLAS prefill tiling on the CPU task runtime

The `m > 1` MLAS SQNBit prefill tiling (`run_mlas_shards`) issued a
single
`tiles.par_iter()` fan-out on global Rayon, sized by
`rayon::current_num_threads()`. With a co-resident ORT intra-op pool
spinning on the same cores, every parked-Rayon wake-up lands behind a
spinning thread, which is what made `gemm_nbits_*_t8` at t=32 the worst
cell in the benchmark ledger. This routes that fan-out through the CPU
task runtime from #1201 instead, with a work-size policy for the one
case
where the SMT-capped pool leaves hardware threads idle.

## MLAS remains opt-in and non-load-bearing

This does **not** enable MLAS anywhere. The changed code lives entirely
inside the pre-existing `m > 1 && active > 1 && !mlas_prefill_serial()`
path that already called `mlas_sys::sqnbit_gemm_into`. No Cargo feature,
`#[cfg(feature = "mlas")]` gate, or default-feature set is touched:
`mlas`
is still opt-in (`default = ["full"]`, and `full` does not include
`mlas`). The new policy helpers (`prefill_fan_out`,
`prefill_tile_grain`,
`PrefillFanOut`) are pure integer arithmetic marked
`#[cfg_attr(not(feature = "mlas"), allow(dead_code))]` so they compile
and
their unit tests run on the default (mlas-off) CI lane even though only
the MLAS path consults them. MLAS stays a labelled reference arm.

## What changed

- `prefill_fan_out(macs, lanes, wide)`: below `WIDE_PREFILL_MACS` (512
Mi
  MACs), or whenever global Rayon is not actually wider than the pool,
  fan out on the task runtime (cheap ~5 us dispatch, topology-aware,
  no fight with a co-resident ORT pool). Above it, use the wider global
  Rayon path -- a prefill tile is a multi-ms MLAS call whose dequantise
  step has enough load latency that SMT siblings pay off, so the SMT cap
  costs more than a 226 us park wake-up (0.25% of a 90 ms fan-out).
- `prefill_tile_grain`: a per-task tile floor so no task gets less than
  `MIN_PREFILL_TASK_MACS` (512 Ki) of arithmetic.

## Verification

Hardware: Intel Core i7-13800H, 14 physical / 20 logical (6 P + 8 E).
Baseline: `main` at the rebase point. Toolchain: cargo 1.97.1.

- **Bit-identity (the important one).** This is pure scheduling: the
`run_tile` closure and the `sqnbit_gemm_into` call are byte-for-byte the
  same, only the executor and grain differ, and every tile writes a
  disjoint `[row, row+rows) x [shard.start, shard.start+len)` window so
  order cannot matter. Confirmed empirically under `--features mlas`:
  `mlas_prefill_parallel_dispatch_matches_serial` and
  `mlas_prefill_dispatch_parity_subprocess` pass -- the routed parallel
  tiling matches the serial reference.
- **Policy tests (default features).** All six `prefill_*` unit tests
pass.
- **Falsified.** Flipping the threshold comparison in `prefill_fan_out`
  from `<` to `<=` turns `large_prefill_work_takes_the_wide_fan_out` RED
  (`left: TaskRuntime, right: Wide` at exactly `WIDE_PREFILL_MACS`) --
  the test is non-vacuous and guards the boundary. Restored to green.
- **`--features mlas` compiles clean;** clippy `--all-targets -D
warnings`
  clean on default features.

## Perf

The §35 (Phase 15) tables in the benchmark doc record up to 13.5x at
t=32
on the small int4 cells, dropping `gemm_nbits_*_t8` from 22-34x ORT to
2.6-2.8x. Those tables mix the author's EPYC 9V74 (16c/32t) and the
laptop measurements; per the repo's measurement rule they are peers,
named
by hardware. The mechanism (the work-size policy, the grain floor, the
disjoint-window safety) and correctness are verified, and the policy
decisions are reproduced in unit tests; no full ORT A/B sweep was re-run
on the laptop, so the headline speedup magnitudes are the author's, not
independently re-measured there.

## Rebase and convergence (2026-08-18)

#1201/#1202/#1207 and #1143 were **squash-merged**, so this branch's
ancestry
no longer reached `main`. Its three commits were cherry-picked onto
`main` at
`c55a3fab3` and applied cleanly. The diff is now self-contained: 2
files,
+355/-25 (`matmul_nbits.rs` and the benchmark doc). It is **not**
stacked on
#1232 any more.

Section numbering: #1232 lands first and takes §34/Phase 14, so this
PR's
section was renumbered to **§35/Phase 15**. (The `34.1×` figure in the
§35.4
matrix is a speed ratio, not a section reference, and is unchanged.)

Re-validated on the rebased head (rustc 1.97.1, the toolchain CI
resolves):

* `cargo test -p onnx-runtime-ep-cpu --lib` — **1433 passed, 0 failed,
17 ignored** (on top of #1232)
* `cargo clippy -p onnx-runtime-ep-cpu --all-targets` — clean
* `cargo fmt --all -- --check` — clean
* under `--features mlas`, the parity test that actually guards this
change,
  `mlas_prefill_parallel_dispatch_matches_serial`, **passes**

### The two `--features mlas` failures are pre-existing on `main`

Running with `--features mlas` fails two tests:
`feature_default_guard::mlas_is_not_a_default_feature` (it panics
*because*
`--features mlas` was passed explicitly) and

`kernels::simd_activations::mlas_ab::mlas_matches_rust_simd_on_special_values`
(a 1-ULP Erf disagreement between the MLAS and pure-Rust SIMD routes).
**Both
reproduce identically on `main` at `c55a3fab3` with this PR's changes
absent**,
so they are baseline-equivalent, not regressions, and they are out of
scope
here. Neither runs on any required lane: `mlas` is not a default
feature.

### The red-criterion comment is runner noise

The criterion report flags regressions including
`tokenization/decode_tokens_per_second` -47.6%. This diff cannot reach
the
tokenizer, and every changed line of `matmul_nbits.rs` is inside the
`#[cfg(feature = "mlas")]` `run_mlas_shards` path, which the benchmark
build
(default features) does not compile in. The default-feature build is
behaviourally identical to `main`; the deltas are shared-runner
variance.

### CI

The repository's Actions queue is saturated (every recent run is
`queued`,
nothing `in_progress`), so the two required checks — `Fast (Linux
x86_64)` and
`Rust quality` — cannot report. Both lanes were reproduced locally, step
for
step, from `.github/workflows/ci.yml`. `main` itself fails **both** of
them
today; #1346 is the fix for that and lands first.

Closes #1238. Working as sebastian (CPU perf).

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 19, 2026
…ler's work (#1154)

## 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 `.so`s 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.

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 19, 2026
…d MatMulNBits prefill (#959, #1091) (#1176)

## Summary

On the **default (`mlas` OFF) build** the `dequant-kn` prefill phase now
goes to **zero** for every `MatMulNBits` path that reaches the dense
fallback: the native CPU EP gets a transposed-B ("NT") SGEMM and
non-MLAS prefill routes through it, reusing the already-cached
contiguous `Nk` weight instead of materializing a second, transposed
`Kn` copy.

Refs #959, #1091. **Does not close #959** — the superlinear per-token
*decode* cost (~102 s/token at 14B) is a separate open question this
does not touch. This closes the *prefill* `Kn`-materialization term.

## The problem (#959)

Int4/int8 `MatMulNBits` prefill (`m > 1`) on the default build
dequantized the weight to f32 a **second** time in the transposed `Kn`
(`[k, n]`) layout — a strided-scatter transpose (each K step written at
stride N), **uncached**, then a dense NN GEMM. #959 measured it at
**~2.9x** the contiguous `Nk` pass and *degrading with N* (22 s/GB at
0.5B → 38 s/GB at 14B; **266 s of ~357 s** time-to-first-token on
qwen2.5-14b).

MLAS hosts already avoided this (`try_prefill_mlas_nt`) by feeding the
cached `Nk` weight to MLAS's cache-tiled `sgemm` with `trans_b`. The
default build had **no** transposed-B GEMM, so its `#[cfg(not(feature =
"mlas"))]` sibling returned `false` and fell back to the slow `Kn` path.

## What changed

MLAS's advantage here is not magic — it is a packer that reads the B
operand row-wise from `[n, k]`. Ported natively:

- **`x86_sgemm.rs`** — `sgemm_simd_nt(a, b_nk, c, m, k, n)` computing
`C[m,n] = A[m,k] · B_nk[n,k]^T`. It reuses the entire packed GEBP path
(`pack_a`, KC/strip blocking, `micro_6x16`) and adds **`pack_b_nt`**,
which gathers each output column from its *contiguous* `b_nk` row (unit
stride over `k`) into the L1-resident pack tile at stride `NR`. The
transpose becomes a **pack-time reshape of an in-cache tile**, not a
full-array stride-`N` scatter — exactly why MLAS wins.
- **`matmul.rs`** — `nt_gemm_supported(backend)` is the **single**
legality predicate (no two `cfg` arms to drift); `gemm_nt_with_backend`
dispatches `Mlas → trans_b`, `SimdX86 → sgemm_simd_nt`.
- **`matmul_nbits.rs`** — the two cfg-split `try_prefill_mlas_nt` arms
collapse into one **`try_prefill_nk_nt`** gated on `nt_gemm_supported`.
It dequantizes once into the pre-existing `weight_nk` `OnceLock` (the
same slot decode caches into, so a constant weight pays **one** dequant,
not two) and calls `gemm_nt_with_backend`. When the NT route runs, the
`dequant-kn` profile phase is skipped by construction.
- **Host pool (#1143)** — column strips dispatch onto the installed ORT
intra-op pool when present, rather than forking a rayon pool beside it;
rayon otherwise. Strip decomposition is numerically transparent.

## Correctness — bit-identical

**Byte-for-byte identical** to the existing `Kn` dense route, not
"within 1e-5". `pack_b_nt` produces the *same* packed panels `pack_b`
would from the `[k, n]` transpose (`b_kn[p·n+j] == b_nk[j·k+p]`); with
identical A-pack, K-panel order, and microkernel, every output element's
f32 accumulation sequence is unchanged. Strip count and pool choice do
**not** affect per-element reduction order, so bit-identity holds
regardless.

Verified (`assert_eq!` on `f32::to_bits`):
- Odd/tail shapes: `m ∈ {1, 2, 7, 33}`, `n ∈ {1, 3, 63, 64, 65}`, plus
tile-exact and multi-KC.
- 64 randomized shapes (differential NN-vs-NT).
- End-to-end by the existing 8-bit prefill oracle test
(`matmulnbits_8bit_prefill_batched_matches_dequant_f32_oracle`), which
reaches the dense fallback and thus the NT route.

This matches the standard the MLAS NT route already held itself to
(bit-identity to the no-transpose dense GEMM, recorded at
`matmul_nbits.rs` ~2048).

## Measurement

**`dequant-kn` → 0 is structural.** The
`mm_profile::time_prepack("dequant-kn", …)` call lives in the `if
!used_fast_nt {…}` branch; whenever the NT route returns `true` that
branch is skipped, so the phase is eliminated by construction.
`dequant-nk` is unchanged (still one pass, now the only one).

**Finding vs the #959 premise:** on current local builds, 4-bit `acc0`
prefill takes the borrowed-int4 in-place path (#979/#1117) and 4-bit
`acc4` uses the SDOT prepack — *neither reaches the dense fallback*, so
no local q4 model emits a `[mm_prepack] phase=dequant-kn` line to drive
to zero end-to-end (confirmed empirically with `ONNX_GENAI_PROFILE_MM=1`
on qwen05b q4 / q4-acc4 / symzp). The beneficiaries of this change are
therefore: **8-bit** weights (`m>1`), **grouped** quantization,
**weight_prepacked**, and **4-bit with `accuracy_level != 0`** that
falls to the dense fallback. No local 8-bit/grouped model was available
for an end-to-end token-to-first-token arm; per the profiling skill I do
not report a contended wall-clock figure I cannot defend.

**Per-phase microbench** (`nt_prefill_bench`, `#[ignore]`; best-of-7,
`--test-threads=1`, release; **contended box — other agents building
concurrently**). Bit-identity asserted in the same harness. The
`dequant-kn*` arm is a *plain f32 transpose* standing in for the strided
`Kn` materialization the NT route removes (the real int4 dequant is
~2.9x heavier per #959):

| shape (m=16) | dequant-kn* (transpose) | NN gemm | NT gemm | old
(kn*+NN) → new (NT) |
|---|---|---|---|---|
| k=5120, n=5120 (100 MiB) | 337 ms | 5.7 ms | 5.0 ms | 343 ms → **5.0
ms** |
| k=5120, n=13824 (270 MiB) | 641 ms | 15.0 ms | 11.8 ms | 656 ms →
**11.8 ms** |
| k=13824, n=5120 (270 MiB) | 1307 ms | 15.3 ms | 11.7 ms | 1323 ms →
**11.7 ms** |

The eliminated transpose term dominates and grows with size (as #959
predicted); the NT GEMM itself is even **slightly faster than NN** here
(contiguous per-column B reads pack better). Per #1132, a
native-faster-than-MLAS result on some shape is a graduation event for
`benches/native_vs_mlas.rs` — noting it, but **not** changing default
routing without that gate's measurement.

**RSS.** No new long-lived allocation (reuses `weight_nk`;
`apack`/`bpack` are per-call local scratch, freed at return). Removing
the second full f32 weight materialization cuts the transient f32
footprint of a prefill that hits the dense fallback by one full `[k, n]`
copy (e.g. **270 MiB** at k=13824/n=5120). Not measured end-to-end (no
local model reaches that path); the reduction is structural.

## Memory rules

No new field that outlives a call or scales with weight size — the NT
route reuses the pre-existing `weight_nk` `OnceLock`. `apack`/`bpack`
are per-call `vec![]` scratch. The added lines are not matched by
`weight-cache-guard.yml` (no `OnceLock<…Vec>` /
`(RefCell|Cell)<Vec|Box|Arc>` introduced); its regex and path filter are
untouched.

## Gates (exact counts)

- `cargo test -p onnx-runtime-ep-cpu --lib` → **1324 passed, 0 failed,
17 ignored**
- `cargo test -p onnx-runtime-ep-cpu --lib --features mlas` → **1354
passed, 0 failed, 28 ignored** (MLAS route still works and still wins
where enabled)
- `cargo clippy -p onnx-runtime-ep-cpu --all-targets -- -D warnings` →
**clean**
- `cargo fmt --check` → the three changed files are clean (verified with
`rustfmt --edition 2024 --check`). Pre-existing diffs remain in three
*unrelated* files (`governed_accumulator_budget.rs`,
`qlinear_matmul.rs`, `simd_activations.rs`) from a local rustfmt version
skew vs CI — left untouched to keep this PR surgical.

---
🤖 Generated with Squad. Flagged **needs review** — please have a squad
member review the kernel packing/tail handling before merge.


---

## Update (2026-08-18, Roy) — merged current `main`, plus a
production-path A/B

### Merge

Merged `main` (`c55a3fab3`), which had since made the #1091 M=1 GEMV the
unconditional `SimdX86` route and dropped the
`ONNX_GENAI_CPU_MM_SIMD_M1_GEMV` toggle this branch still carried.
Conflict resolved by keeping **both**: main's default-route test
(`the_default_entry_point_routes_m1_to_the_gemv`) and this branch's NT
kernel + bit-identity tests. Two follow-on fixes:

- **aarch64 cross-arch lane.** `gemm_nt_with_backend` compiled with
neither the `mlas` nor the x86 arm has no reader for any parameter, so
the `-D warnings` cross-arch pass rejected all six. Bound them in the
unsupported arm. (This is what the old `Rust quality` red was:
`Cross-target compile check`, nothing else.)
- **`mm_profile` gemv phase.** The MLAS route used to time this GEMM on
the `gemv` phase; after the two call sites collapsed into one, that
timer was lost. Restored — the default build gets a pass-through
(`tick()`, the reporter, is MLAS-only), so the shared call site stays
`cfg`-free and MLAS profiling is unchanged.

### Production-path A/B (new harness,
`benches/matmul_nbits_prefill_ab.rs`)

The original body was right that no local model reaches this route, and
honest about not reporting a number it could not defend. That gap is now
closed the way #1013 closed its own: drive the **real kernel through the
EP's own `get_kernel`/`execute`**, at the shapes and inputs that *do*
reach the dense fallback — 8-bit prefill, and 4-bit with `g_idx`. The
harness uses no symbol this branch introduces, so the identical file
runs on `main` and here; both arms were built and run **interleaved**, 3
repetitions each, on the same box.

Host: 32-core x86_64, AVX2, default build (**`mlas` off**), release.
**Contended** (other agents building; load ~12), so medians of per-arm
medians are reported and the ratios — not the absolute ms — are the
claim.

**Steady state** (weight already resident, per-call prefill cost):

| case | k | n | m | main (ms) | this PR (ms) | speedup |
|---|---:|---:|---:|---:|---:|---:|
| int8 dense fallback | 2048 | 2048 | 8 | 12.114 | **0.548** | **22.1x**
|
| int8 dense fallback | 2048 | 2048 | 64 | 12.276 | **1.375** | **8.9x**
|
| int8 dense fallback | 4096 | 11008 | 8 | 53.858 | **4.170** |
**12.9x** |
| int8 dense fallback | 4096 | 11008 | 64 | 56.560 | **7.595** |
**7.4x** |
| int4 + `g_idx` fallback | 2048 | 2048 | 8 | 31.583 | **0.563** |
**56.1x** |
| int4 + `g_idx` fallback | 2048 | 2048 | 64 | 32.405 | **1.419** |
**22.8x** |

**Cold** (fresh kernel per repetition, so the one-time weight dequant is
inside the measurement — the TTFT term #959 attacked):

| case | k | n | m | main (ms) | this PR (ms) | speedup |
|---|---:|---:|---:|---:|---:|---:|
| int8 dense fallback | 2048 | 2048 | 8 | 6.470 | 6.008 | 1.08x |
| int8 dense fallback | 2048 | 2048 | 64 | 12.519 | 7.974 | 1.57x |
| int8 dense fallback | 4096 | 11008 | 8 | 53.761 | 37.260 | 1.44x |
| int8 dense fallback | 4096 | 11008 | 64 | 57.725 | 38.503 | 1.50x |
| int4 + `g_idx` fallback | 2048 | 2048 | 8 | 31.636 | 23.200 | 1.36x |
| int4 + `g_idx` fallback | 2048 | 2048 | 64 | 32.083 | 25.883 | 1.24x |

**Why steady moves 7–56x and cold only ~1.1–1.6x — and why that is the
real finding.** The `Kn` route has **no cache**:
`dequantize_weight(WeightLayout::Kn)` is called *inside* `execute`, so
every prefill call re-materializes the whole transposed f32 weight. The
NT route dequantizes into the pre-existing `weight_nk` `OnceLock`, which
a constant weight fills **once**. So this change removes not one
transpose but *every repeat of it*. Cold (first call) improves by the
layout alone — a contiguous `Nk` write instead of the stride-`N`
scatter, 1.2–1.6x at these sizes; steady improves by the caching the
`Nk` layout makes possible.

Confirmed structurally with `ONNX_GENAI_PROFILE_MM=1` over the same
harness run: **main emits 72 `phase=dequant-kn` lines, this PR emits 0 —
and 24 `phase=dequant-nk`** (one per kernel instance, i.e. the cold arms
only; every steady call pays none).

**Bit-identity, across builds.** The harness prints an FNV-style digest
of the raw output bits. All six rows have the **identical digest on both
arms** (`db6ff07f991d431`, `b4c1df9fd8883789`, `7bd7418eacdcb870`,
`2eed012609649617`, `d5834afbecf42494`, `cd8fa77262302969`) —
bit-identity of the production `execute` result, not just of the kernel
driver, verified across two separately compiled builds.

### Scope, restated honestly

On the default build, 4-bit `accuracy_level=0` with contiguous
(borrowable) inputs takes the zero-copy borrowed int4 path
(#979/#1117/#1126) for both decode *and* prefill and returns before the
dense fallback. This PR therefore changes: **8-bit prefill**, **4-bit
with `g_idx`**, **`weight_prepacked`/non-borrowable** inputs, and
**4-bit `accuracy_level != 0`** that falls through. Those are exactly
the cases measured above. MLAS builds already had the NT route; their
behaviour is unchanged.

### Gates (re-run after the merge)

- `cargo test -p onnx-runtime-ep-cpu --lib` → **1420 passed, 0 failed,
18 ignored**
- `cargo clippy -p onnx-runtime-ep-cpu --all-targets -- -D warnings` →
clean
- `cargo clippy --locked --target aarch64-unknown-linux-gnu
--all-targets -p onnx-runtime-ep-cpu -- -D warnings` → clean (the lane
that was red)
- `cargo fmt --all -- --check` → the files this PR touches are clean;
one *inherited* diff remains in
`onnx-runtime-ep-cuda/standard_attention.rs` from `main`, fixed
separately in #1347.

---------

Co-authored-by: justinchuby <223556219+Copilot@users.noreply.github.com>
Co-authored-by: Roy <roy@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.

2 participants