Repository navigation
perf(cpu-ep): run the elementwise split on ORT's pool, not a second one - #1143
Conversation
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 Report❌ Patch coverage is Additional details and impacted files@@ 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
Flags with carried forward coverage won't be shown. Click here to find out more.
🚀 New features to boost your workflow:
|
🔴 Benchmark Regression DetectedComparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).
Visual flags: Host infoWhat this cannot catch
|
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>
|
The seam is the right design and I want to say why before the objection: the plugin installs it itself.
Gates pass on my box: The blocking gap: the benchmark that decides this PR has not been runThe Benchmarks section says:
Everything measured so far compares our rayon split against staying serial, and serial wins every row -- So the evidence supports a strictly simpler alternative that the PR never tests:
...which needs no I am not asserting it loses. ORT's pool is the right pool, its workers are already warm, and Please add one row: Smaller points, none blocking
|
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
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>
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 i7-13800H (14 physical / 20 logical), release, 1 Mi f32, three arms in one process, interleaved per rep, p50 over 201 reps, ×3 runs. The rayon-split arm is a built-in control: its loss at busy
The win (~5.5–6.6×) is an order of magnitude larger than each arm's own p10..p90 spread ( Why so much larger than the rayon split's best case? The host cut floors at No-host fall-through still wins at
|
| 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 ignoredcargo test -p onnx-runtime-ep-api --lib→ 53 passed, 0 failedcargo test -p onnx-runtime-ep-plugin --lib→ 246 passed, 0 failedcargo 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.
|
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 meThe 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
Sizing 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 thisThe pool is a stand-in, not ORT's 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 Two things worth keeping for their own sake
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. |
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>
…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>
…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>
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_ParallelForinstead.Root cause
simd_activations::run_chunkedsplits any f32 slice of ≥ 1 Mi across arayonpool. 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_THREADS1 vs 16):
SqrtReluClipSigmoidTanhQuickGeluErfFastGeluGeluParallelising made every op slower.
Why a bigger threshold is not the fix
The obvious repair is to raise
PAR_MIN_LENuntil the split only fires whereit still wins. It does not work, because the same sweep at
intra_op_num_threads = 1says the opposite — there the split is a large winfrom 1 Mi upwards:
ErfGeluNo 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 onlycrate both
onnx-runtime-ep-cpu(which has the kernels) andonnx-runtime-ep-plugin(which has theOrtKernelContext) already depend on,so neither has to learn about the other.
HostParallel::run(total, body)isblocking, which is what makes borrowing a
&dyn Fnfrom the caller's framesound; bodies run with
in_host_task()set so a nested split can see it isalready on a host thread.
onnx-runtime-ep-plugin::host_pool— the ORT-backed implementation overKernelContext_ParallelFor(num_batch = 0, so ORT's workers claim indicesdynamically).
install()returns an RAII guard;compute_executeholds onefor the whole call.
simd_activations—run_chunkedandrun_chunked_rowsprefer the hostpool 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 = 64pieces, floored at theexisting
PAR_MIN_CHUNK. That also makes the chunk boundaries — and thereforethe 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 unaryand 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 thechange, asserted: exactly
host_chunk_len(N).1chunks dispatched, none to rayon.a_nested_split_stays_serial— a kernel reached from inside a host task doesnot dispatch again.
serial_scope_still_suppresses_the_split— the f16/bf16 sandwich stays serialon the host path too.
host_chunk_policy_holds_across_lengths— whole vectors, never belowSIMD_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_ParallelForstill runs every index (a short writewould 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_opis modelled at 14 rather than 16because 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_HELPEDthe honest way (a worker, not the dispatcher, runs achunk), so
prefer_hostreaches 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 isslower than serial in every run — which validates the harness before the
host-pool number is trusted.
SqrtSqrtSqrtGeluGeluGeluThe host-pool win (~5.5–6.6×) is an order of magnitude larger than each arm's
own p10..p90 spread (e.g.
Sqrthost 34..62 µs, serial 187..256 µs), so it is areal 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 tasksthe 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 = 1With 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):
Gelu, 4 MiSqrt, 4 MiGelu, 1 MiSqrt, 1 Mi(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 splitloses 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_ParallelForand all, is the correct fix, not merely the morecomplex one. The PR merges as written.
The harness is committed as two
#[ignore]d tests insimd_activations::three_arm_bench, run withEP_BENCH=1 cargo test --release -p onnx-runtime-ep-cpu --lib three_arm_bench -- --ignored --nocapture --test-threads=1.Limitations
KernelContext_ParallelForandkeep the rayon path.
that start their own rayon work under an ORT session have the same
oversubscription and are not addressed here — filed as perf(cpu-ep): route the remaining kernels' parallelism through the host pool (MatMul, int4) — follow-up to #1143 #1145
(follow-up to The CPU EP's fast paths depend on a host-set scope: standalone and plugin hosts get the slow path #1138), because MatMul and the int4 paths matter far more than
elementwise and deserve their own measured change.