Repository navigation
perf(cpu-ep): probe the host pool with an empty dispatch, not the caller's work - #1154
Conversation
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #1154 +/- ##
==========================================
- Coverage 80.88% 80.14% -0.75%
==========================================
Files 364 364
Lines 160729 160978 +249
Branches 160729 160978 +249
==========================================
- Hits 130005 129012 -993
- Misses 26069 27310 +1241
- Partials 4655 4656 +1
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
|
361ec89 to
0a71f78
Compare
Rebased onto
|
| suite | result |
|---|---|
cargo test -p onnx-runtime-ep-api --lib |
57 passed / 0 failed |
cargo test -p onnx-runtime-ep-plugin --lib |
246 passed / 0 failed |
cargo test -p onnx-runtime-ep-cpu --lib |
1417 passed / 0 failed / 17 ignored |
cargo clippy --all-targets (all three) |
clean |
rustfmt --check (touched files) |
clean |
The Fast/Rust quality failures on the previous run were rustfmt drift in mlas-sys, governed_accumulator_budget.rs and qlinear_matmul.rs — files this PR does not touch, inherited from the old base and since repaired on main. They are gone on the rebase.
I did not re-run the ORT A/B sweep; the µs tables in the body are the author's, on the EPYC 9V74.
Review caught that the unknown state was being paid for with the caller's own work. `prefer_host` returned true on a probe dispatch, so the slice in front of it went to the host pool -- and on an `intra_op = 1` session, which never latches, that is ORT running 1 Mi on one thread where our own pool is 5x faster. It recovered after the opening burst, but the first 32 dispatches of every fused node paid for the question. Measured against a build with probing compiled out, at `intra_op = 1` with a 16-thread rayon pool: 5-17% at 1 Mi across seven ops. Ask with an empty eight-index dispatch instead. It is the scheduling behaviour we are measuring, not the arithmetic, so the probe does not need to carry anything, and its cost no longer scales with the caller's slice. That also makes the question far easier to answer, because nothing drains the indices ahead of the workers: instrumented over 15 sixteen-thread sessions, all latched, taking 1 probe (nine), 2 (five) or 3 (one), against ~2.5 for the old work-carrying probe. So the burst drops 32 -> 8 and the recovery period widens 16 -> 64. A burst of 8 x 100 us turned out to be too little on a loaded machine -- several sessions never latched at load 5-10 and kept their work on the wrong pool -- so the stall is now a *deadline*: hold the first index until another thread is seen, up to 400 us. A pool with workers pays its wake-up and no more; only a pool with nobody to wake pays it in full. At intra_op=1/rayon=16, 1 Mi, this branch vs main vs probing-compiled-out (us): Tanh 324/323/383, Erf 345/335/401, Swish 305/299/346, Gelu 502/465/499 -- i.e. within the control's own spread. The 16-thread win is unchanged: every op at 65 Ki and above still improves 2-10x over main. Also from review: the bit-identity note on `run_on_host` claimed every chunk is a multiple of eight lanes and at least SIMD_MIN_LEN, which the *final* chunk is not (65537 ends with a chunk of one). Restated for the real reason -- chunk starts are 8-aligned and every chunk runs the same masked-tail kernel. `a_nested_split_stays_serial` now uses a slice long enough for the rayon path to fire and asserts on the rayon counter, so it can actually fail; and `try_host` no longer bumps the rayon dispatch counter on a host dispatch. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> (cherry picked from commit 0af4789)
Two review defects in the probe tests, both of which let a broken schedule pass. `a_probe_does_not_carry_the_callers_work` built an `AtomicUsize` it never read and stored zero into it at the end of the test — dead state left over from an earlier draft. `the_probe_and_the_latch_agree` asserted `settled >= PROBE_MIN` one line after `assert_eq!(settled, PROBE_MAX)`, which is `1024 >= 64` between two constants: it holds no matter what the back-off does, and it replaced the count-based assertion the previous revision of this test had. The property it was meant to state — a serial-looking session keeps re-asking, so an unlucky opening burst is recoverable — is now measured where it is observable, on the pool: `counted_inline_parallel_for` counts the dispatches ORT's stand-in is actually handed, and the test bounds that count on both sides. 15 probes in 4096 dispatches (8 burst + 7 back-off). Falsified: raising the lower bound to 100000 turns it RED with `probing stopped after the opening burst (15 probes)`. `cargo test -p onnx-runtime-ep-api --lib` 57 passed, `-p onnx-runtime-ep-plugin --lib` 246 passed, `-p onnx-runtime-ep-cpu --lib` 1417 passed / 0 failed. Clippy and rustfmt clean on the touched files. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
0a71f78 to
bf30828
Compare
Convergence report — rebased onto current
|
| check | result |
|---|---|
cargo test -p onnx-runtime-ep-api |
57 passed, 0 failed |
cargo test -p onnx-runtime-ep-plugin |
246 passed, 0 failed |
cargo test -p onnx-runtime-ep-cpu |
1433 passed, 0 failed, 17 ignored |
cargo clippy … --all-targets -- -D warnings (all three) |
clean |
cargo fmt --all -- --check |
clean |
CI
The Actions queue is saturated (every recent run queued, nothing
in_progress), so Fast (Linux x86_64) and Rust quality cannot report. Both
were reproduced locally, step for step, out of .github/workflows/ci.yml. The
fmt-drift failures this PR was showing earlier were in files it does not touch;
they were main's, and #1346/#1352 fixed them.
Working as sebastian (CPU perf).
#1244) ## What this is The per-`Run` cost of dispatching a **single-node fused subgraph** through the plugin path, with the kernel removed from the picture. On a 4 Ki f32 elementwise node our arm spends ~3.5 µs per `Run` where plain ORT spends ~2.6 µs, and the kernels themselves account for well under a microsecond of that — the same kernel on a 1 Mi tensor runs at 0.51x–1.03x of ORT. The difference is fixed dispatch overhead, and it is why every cheap elementwise op sits at 1.1x–1.5x of ORT at 4 Ki while the identical kernel wins at 1 Mi. This PR removes four pieces of that overhead. No kernel is touched, no numerics change, and nothing about assignment or execution ownership changes: the CPU EP still executes every node it claims, locally, with no ORT CPU fallback anywhere. ## The four cuts **1. `device_mem_info` is no longer resolved eagerly (6 ORT FFI calls per `Run`).** `compute_execute` resolved `scratch_mem_info` at the top of *every* call: ```rust let scratch_mem_info = unsafe { device_mem_info(api_ref, kernel_context, exported.device_staging.as_ref()) }; ``` `device_mem_info` calls `KernelContext_GetInputCount`, then for each input `KernelContext_GetInput` + `GetTensorMemoryInfo` + `GetMemoryInfoDeviceType`, then falls through to two more calls for input 0. For a one-input node that is **six FFI calls**. Who reads it? Two consumers. `intermediate_scratch` — routed multi-node path only. And `PlacementSources::subgraph_fallback`, which `operand_mem_info` reads **only when the node binds no ORT operands at all**, and which `prepare_workspace` only reaches past its zero-byte and lifetime gates. An elementwise node has operands and needs no workspace, so it reached neither. Now the routed path resolves it once per `Run` (unchanged) and passes `SubgraphFallback::Resolved`; the single-node path passes `SubgraphFallback::Deferred(staging)` and `operand_mem_info` calls `device_mem_info` itself, with the same arguments, if it ever actually needs the answer. Deferring is safe: the function only *reads* input memory info, so running it after outputs are allocated cannot see a different answer. **2. `allocate_output` takes `want_mem_info` (1 ORT FFI call per output per `Run`).** `OwnedOutput::mem_info` has exactly one consumer in the tree — `stage_host_boundary_inputs`, called under `if let Some(staging) = exported.device_staging.as_ref()`. A host EP has no staging context, so every `Run` made a `GetTensorMemoryInfo` call per output and dropped the result. Both call sites now pass `exported.device_staging.is_some()`, so a device EP is bit-for-bit unchanged and a host EP stops making the call. **3. `staging_log(&format!(..))` → `staging_log!(..)` (one `String` per `Run`).** `staging_log` checks `ONNX_GENAI_PLUGIN_TRANSFER_TRACE` *inside* the function, so the `format!` was evaluated whether or not the trace was on — and one of these sites is at the top of `compute_execute`, formatting three fields into a heap `String` on every dispatch. The macro checks the gate first. All 12 sites converted, so the footgun is gone rather than papered over at one site. **4. Absent-output bookkeeping is built only when there are absent outputs (5 allocations per `Run`).** ```rust let absent_shapes: Vec<Vec<usize>> = output_shapes.clone(); let absent_strides_storage: Vec<Vec<i64>> = absent_shapes.iter().map(|s| contiguous_strides(s)).collect(); let mut ort_views: Vec<TensorMut<'_>> = owned_outputs.iter_mut().map(|o| o.view_mut()).collect(); let mut ort_view_iter = ort_views.drain(..); ``` That storage exists solely to back the `TensorMut`s of *absent* output slots. A node with no absent outputs — every elementwise op — cloned every output shape and built a stride vector per output to lend them to nobody. It is now conditional on `!absent_bufs.is_empty()`. The `ort_views` `Vec` was collected and immediately drained; the iterator is taken directly from `owned_outputs.iter_mut()`. Related: placement operands are now described by an `OrtOperands` enum, so the single-node path lends `entry.input_slots` (`&[Option<usize>]`, flattened lazily) instead of collecting a fresh `Vec<usize>` per call for a consumer that almost never runs. The routed path passes its already-resolved slice. Per `Run` on a 1-in/1-out elementwise node this is **7 fewer ORT FFI calls** (16 → 9) and **7 fewer heap allocations**. ## A/B, one thread `taskset -c 8-15`, this commit vs its parent, **interleaved** (A,B,A,B,A,B with a rebuild before each arm, so drift hits both arms equally), 400 iterations after 50 warmup, three rounds. Ratio is **ours/ORT, lower is better**; the p50 column is the median of the three rounds' p50 ratios, p90 likewise. | case | before p50 | after p50 | before p90 | after p90 | |---|---|---|---|---| | `sqrt_f32_4k` | 1.132 | **1.026** | 1.153 | **1.033** | | `sigmoid_f32_4k` | 1.484 | **1.346** | 1.517 | **1.362** | | `tanh_f32_4k` | 1.534 | **1.406** | 1.595 | **1.434** | | `erf_f32_4k` | 1.760 | **1.510** | 1.779 | **1.512** | Best-of-three (the contention-robust statistic on this shared box) agrees: sqrt 1.130 → 1.023, sigmoid 1.481 → 1.344, tanh 1.503 → 1.387, erf 1.570 → 1.480. 12 of 12 arm-pairs favour the change; there is no round in which any case regressed. In absolute terms `tanh_f32_4k` goes 0.0046 ms → 0.0042 ms, i.e. **~0.4 µs off a ~0.9 µs gap**. ## A/B, threaded Same protocol, `taskset -c 0-15`, two rounds, `NXRT_MM_BENCH_THREADS` = `ONNX_GENAI_MLAS_THREADPOOL_THREADS` = `RAYON_NUM_THREADS`. | case | 4t before | 4t after | 16t before | 16t after | |---|---|---|---|---| | `sqrt_f32_4k` | 1.257 | **1.129** | 1.262 | **1.156** | | `sigmoid_f32_4k` | 1.520 | **1.355** | 1.499 | **1.340** | | `tanh_f32_4k` | 1.551 | **1.418** | 1.548 | **1.455** | | `erf_f32_4k` | 1.776 | **1.671** | 1.777 | **1.624** | The overhead is per call, not per element or per worker, so the gain is the same absolute number of microseconds at every thread count. ## Drift control: the 1 Mi grid is unchanged Same protocol, `_f32_1m`, three interleaved rounds, one thread. A large tensor amortises the per-call cost away, so these must *not* move — and they don't: | case | before | after | |---|---|---| | `relu_f32_1m` | 1.032 | 1.034 | | `exp_f32_1m` | 1.021 | 1.022 | | `sigmoid_f32_1m` | 1.066 | 1.063 | | `tanh_f32_1m` | 1.117 | 1.115 | | `gelu_tanh_f32_1m` | 1.243 | 1.241 | | `gelu_exact_f32_1m` | 1.424 | 1.420 | | `fastgelu_f32_1m` | 1.239 | 1.224 | | `erf_f32_1m` | 1.461 | 1.439 | | `quickgelu_f32_1m` | 0.809 | 0.811 | | `sqrt_f32_1m` | 0.514 | 0.511 | Ten of ten within ±0.5 %, which is this box's noise floor. That is the shape of a per-call fix. ## Correctness **New test, with a verified falsifier.** `output_memory_info_is_queried_only_when_the_caller_asked_for_it` drives `allocate_output` against a hand-built `OrtApi` whose `GetTensorMemoryInfo` counts its calls, and asserts 0 calls for `want_mem_info: false` and 1 for `true`. Falsifier: making `allocate_output` ignore the flag fails it with `left: 1, right: 0` — checked by breaking the code, not by inspection. **Behaviour preserved, argued per cut.** (1) `device_mem_info` is called with identical arguments, only later and only when read; it reads inputs, which output allocation cannot change. (2) The gate on the memory-info query is *the same condition* as the gate on its only consumer. (3) The macro's only difference is when the `format!` runs. (4) The absent storage is only ever indexed for absent slots. **Suites.** `-p onnx-runtime-ep-plugin`: 247 unit tests pass (246 before, +1 new). `-p onnx-runtime-ep-cpu-plugin` with `NXRT_REQUIRE_ORT_TESTS=1`: the full e2e suite passes, including all 54 `plugin_ort_e2e` cases — the routed multi-node fixtures (`conformance_chain_add_mul`, `..._repeated_runs_do_not_leak_stale_intermediates`, `conformance_mixed_partition`) exercise the `SubgraphFallback::Resolved` arm, and `every_assigned_node_is_also_executed_by_this_ep` passes, so every node this EP claims is still executed here with ORT CPU fallback disabled. `cargo fmt --all` and `cargo clippy --release --all-targets -p onnx-runtime-ep-cpu -p onnx-runtime-ep-cpu-plugin -p onnx-runtime-ep-plugin` are clean. **Build identity.** Pure native CPU EP: no MLAS at runtime, no ORT CPU EP fallback, no new dependency. AVX2/FMA host (`avx2 fma f16c`, no AVX-512), so ORT/MLAS and we are on the same instruction footing. ## What is left The remaining ~0.5 µs is, as far as I can attribute it without CPU counters (`perf_event_paranoid` is 4 on this box, so `perf record` is not available and everything here is A/B attribution): * **~9 ORT FFI calls that are genuinely needed.** `read_inputs` costs 7 per input — `KernelContext_GetInput`, `GetTensorTypeAndShape` (which allocates an ORT object we then release), `GetTensorElementType`, `GetDimensionsCount`, `GetDimensions`, `ReleaseTensorTypeAndShapeInfo`, `GetTensorData` — and there is no cheaper spelling in the stable C API. ORT's own CPU kernels reach the same data through `OpKernelContext` with no FFI at all, which is a structural part of what a plugin EP pays. * **~7 remaining allocations**: `OwnedInput`'s shape and strides per input, `kernel_inputs`, `infer_shapes`'s `Vec<Vec<usize>>`, `slot_map`, `output_views`, `prepare_workspace`'s metadata, `allocate_output`'s dims, and the `Box<HostPool>` in `host_pool::install`. Each is worth ~25–35 ns. Removing them needs either an inline-capacity vector type or per-session caching of the parts that cannot change between `Run`s; both are worth doing and neither belongs in this PR. I deliberately did **not** touch `host_pool::install` — the per-call `Box` is one allocation, and that file is @sebastian's 16-thread scheduling work; a change there should come from him or after his PRs land. --- ## Refreshed against `main` (2026-08-18) The branch was behind `main` and its whole red CI wall came from that, not from this change: `crates/onnx-runtime-session/src/executor/mod.rs:175` failed `-D dead-code` on current stable, which `ca32b3adf` ("fix(ci): unbreak the Rust quality lane on current stable", #1239) fixed on `main` after this branch forked. `origin/main` (`c55a3fab3`) is merged in — no rebase, no force-push — and the diff this PR owns is unchanged at 2 files, +285/-79. Revalidated on the merge commit, AVX2/FMA host, no AVX-512: * `cargo test --release -p onnx-runtime-ep-plugin` — **247 passed, 0 failed**. * `NXRT_REQUIRE_ORT_TESTS=1 cargo test --release -p onnx-runtime-ep-cpu-plugin` — every suite green, including all **55** `plugin_ort_e2e` cases. The ones that matter to this change all pass on the merge: `every_assigned_node_is_also_executed_by_this_ep`, `no_supported_node_is_ever_left_to_the_ort_cpu_ep`, `no_matmul_family_node_escapes_to_the_ort_cpu_ep`, and `every_fixture_loads_with_cpu_fallback_disabled` — so assigned still equals executed with ORT CPU fallback off. * `cargo fmt` clean for the crates this PR touches. The one `cargo fmt --all` hunk on this tree is in `onnx-runtime-ep-cuda/src/kernels/standard_attention.rs`, which arrived from `main` untouched by this PR and is a local rustfmt-version difference, not a branch defect. ## Independent review Reviewed by **Claude Opus 4.8**, read-only, with the four cuts and the absent-slot history stated as the priority list. Verdict **APPROVE**, no blockers. It independently confirmed: * `SlotKind::Absent(idx)` is pushed into `slot_map` only in the same branch that pushes into `absent_bufs`, so `has_absent == false` implies `slot_map` holds no `Absent` and the skipped storage is never indexed; and the surviving indices are the same full-slot indices as before. * `ort_view_iter` from `owned_outputs.iter_mut()` has the identical borrow structure as the old `collect()` + `drain(..)`, so nothing borrows a temporary. * `OrtOperands::Slots(..).indices()` yields the same elements in the same order as the removed `.iter().flatten().copied().collect()`, `None` skipped. * `operand_mem_info` is the only reader of the deferred value, only reachable when the node binds **zero** ORT inputs, which makes the timing of the deferred `device_mem_info` moot rather than merely argued. * Both `allocate_output` call sites gate on the same predicate as the only reader of `OwnedOutput::mem_info`, so a device EP is unchanged. Its one correction is applied above: the prose said 14 `staging_log` sites; the real count in `origin/main` is **12**, and all 12 are converted. --- ## Refreshed against `main` @ `6a855d5e0`, and measured as a stack `origin/main` moved a long way while this sat in the CI queue (#1346, #1352 and #1361 on the quality lane; #1154, #1232, #1238 on the CPU side). Merged in normally — no rebase — and re-measured from scratch against the new baseline. Production pure-native A/B, plain ORT as the control arm. No MLAS, no ORT CPU fallback, no deferral. `taskset -c 8-15`, one thread, 400 iterations, five interleaved rounds out of two worktrees, started only once cores 8-15 were >=93% idle. **Ratio is ours/ORT, lower is better.** `before` is `main` at `6a855d5e0`; `after` is #1244 + #1246 together, since #1246 is stacked on #1244 and the pair is what a user gets. | case | ratio p50 main | ratio p50 stack | Δ | ratio p90 main | ratio p90 stack | ours us | ORT drift | rounds won | |---|---|---|---|---|---|---|---|---| | `thresholdedrelu_f32_4k` | 1.512 | **1.216** | -19.6% | 1.519 | **1.216** | 3.6 → **2.8** | -4.2% | 5/5 | | `tanh_f32_4k` | 1.511 | **1.279** | -15.4% | 1.513 | **1.284** | 4.6 → **3.9** | +0.0% | 5/5 | | `sigmoid_f32_4k` | 1.475 | **1.248** | -15.4% | 1.484 | **1.256** | 4.7 → **4.0** | +0.0% | 5/5 | | `erf_f32_4k` | 1.470 | **1.363** | -7.3% | 1.489 | **1.364** | 7.5 → **7.0** | +0.0% | 5/5 | | `hardsigmoid_f32_4k` | 1.416 | **1.127** | -20.4% | 1.426 | **1.133** | 3.5 → **2.8** | +0.0% | 5/5 | | `leakyrelu_f32_4k` | 1.361 | **1.094** | -19.6% | 1.371 | **1.104** | 3.5 → **2.8** | +0.0% | 5/5 | | `sqrt_f32_4k` | 1.141 | **0.946** | -17.1% | 1.150 | **0.955** | 4.0 → **3.4** | -2.8% | 5/5 | | `log_f32_4k` | 0.767 | **0.698** | -9.0% | 0.776 | **0.703** | 7.9 → **7.2** | +0.0% | 5/5 | | `selu_f32_4k` | 0.505 | **0.440** | -12.9% | 0.512 | **0.444** | 5.6 → **4.9** | +0.0% | 5/5 | | `elu_f32_4k` | 0.480 | **0.416** | -13.3% | 0.488 | **0.421** | 5.3 → **4.6** | +0.0% | 5/5 | | `celu_f32_4k` | 0.470 | **0.411** | -12.6% | 0.478 | **0.418** | 5.7 → **5.0** | +0.0% | 5/5 | | `mish_f32_4k` | 0.276 | **0.264** | -4.3% | 0.278 | **0.268** | 17.3 → **16.6** | +0.0% | 5/5 | **Every case, every round.** The two rows with a moving control (`sqrt` -2.8%, `thresholdedrelu` -4.2%) are reported rather than dropped; both won 5/5 anyway and their absolute time fell by the same ~0.7 us as everything else. That constant ~0.7 us is the point. It is not proportional to tensor size — the same absolute amount comes off `hardsigmoid` (3.5 -> 2.8 us) as off `mish` (17.3 -> 16.6 us) — which is what a fixed per-`Run` cost looks like when you remove some of it. It moves the cheap ops the most because they had the least to hide it behind, and `sqrt` crosses from 1.141 to **0.946**, from a loss to a win. ### Where the remaining time goes Measured directly, by instrumenting `compute_execute` segment by segment on top of this stack (temporary probe, not committed; `perf` is unavailable on this host — `perf_event_paranoid=4`). Per `Run`, one-in/one-out elementwise node, 4096 `f32`, microseconds: | segment | us | note | |---|---|---| | `KernelContext_GetOutput` | 0.35 | ORT's own API — ours to call, not to optimise | | `read_inputs` | 0.15 | 4 ORT FFI calls, already one shape call after #1246 | | rest of `allocate_output` | 0.13 | `GetTensorMutableData` + strides | | `prepare_workspace` | 0.09 | metadata vector + plan-cache lookup, for a kernel needing 0 bytes | | `host_pool::install` | 0.05 | @sebastian's, not touched | | `infer_shapes` | 0.05 | | | `kernel_inputs` | 0.04 | | | `output_views` | 0.04 | | Non-kernel node cost is **~1.25 us and near-constant across all twelve operators** (0.28 to 15.2 us of kernel time), which is the direct confirmation that small-node ratios on this EP are dispatch-bound rather than kernel-bound. There is no single large item left — the biggest, `KernelContext_GetOutput`, is ORT's. The rest is a long tail of 0.04-0.15 us items, which is what #1358 (`InlineVec`) starts on. ### And nothing breaks at 1 Mi Same harness, 1048576 elements, 120 iterations, 3 rounds. A fixed per-`Run` cost should be invisible here, and it is: | case | ratio p50 main | ratio p50 stack | ours us | ORT drift | |---|---|---|---|---| | `celu_f32_1m` | 0.141 | 0.140 | 382.1 → 377.9 | -0.0% | | `elu_f32_1m` | 0.138 | 0.136 | 347.7 → 343.6 | +0.1% | | `erf_f32_1m` | 0.671 | 0.670 | 595.2 → 594.6 | +0.1% | | `exp_f32_1m` | 0.606 | 0.590 | 247.0 → 240.5 | +0.1% | | `fastgelu_f32_1m` | 0.643 | 0.648 | 415.3 → 413.2 | -1.7% | | `gelu_exact_f32_1m` | 0.572 | 0.592 | 706.1 → 710.6 | -2.8% ⚠ | | `gelu_tanh_f32_1m` | 0.657 | 0.651 | 414.8 → 410.5 | +0.3% | | `hardsigmoid_f32_1m` | 0.376 | 0.372 | 88.0 → 69.6 | -0.9% | | `leakyrelu_f32_1m` | 0.421 | 0.430 | 90.1 → 87.3 | -2.2% ⚠ | | `log_f32_1m` | 0.270 | 0.274 | 616.4 → 612.4 | +1.7% | | `mish_f32_1m` | 0.105 | 0.105 | 1626.0 → 1626.0 | -0.2% | | `quickgelu_f32_1m` | 0.455 | 0.453 | 321.5 → 320.0 | -0.5% | | `relu_f32_1m` | 1.034 | 1.022 | 131.7 → 130.2 | +0.0% | | `selu_f32_1m` | 0.147 | 0.148 | 374.1 → 371.8 | -1.5% | | `sigmoid_f32_1m` | 0.471 | 0.606 | 231.6 → 229.0 | -23.4% ⚠ | | `sqrt_f32_1m` | 0.314 | 0.302 | 148.1 → 143.9 | +1.1% | | `tanh_f32_1m` | 0.644 | 0.631 | 226.2 → 221.9 | -0.2% | | `thresholdedrelu_f32_1m` | 0.500 | 0.486 | 70.9 → 68.6 | -0.3% | Flat, as predicted — 0.7 us against 70-1626 us of work. Absolute time is equal or better in 16 of 18 cases. The two ⚠ rows had the control move more than the effect: `sigmoid` is unusable (ORT itself moved -23.4%; our own absolute went 231.6 -> 229.0 us), and `gelu_exact`'s +0.6% absolute sits inside its -2.8% control. Reported rather than dropped. This is the coverage claim for the change: it buys ~0.7 us at every size, which is 20% of a small node and nothing at all of a large one, and it costs nothing anywhere. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
) ## What this is `read_inputs` runs **once per input per `Run`**, and it asked ORT for the element type and the dimensions through the classic five-call sequence: ``` GetTensorTypeAndShape → GetTensorElementType → GetDimensionsCount → GetDimensions → ReleaseTensorTypeAndShapeInfo ``` The first of those makes ORT **heap-allocate** an `OrtTensorTypeAndShapeInfo` and copy the shape into it; the last frees it. Five FFI crossings and one allocation on ORT's side of the boundary, per input, per `Run`, to read data the `OrtValue` already owns. `GetTensorElementTypeAndShapeDataReference` (ORT C API, since 1.24) returns the element type **and a reference to the value's own shape array** in one call, allocating nothing. This PR routes `read_inputs` through it. Stacked on #1244 (since merged; this branch now targets `main` directly). No kernel is touched, no numerics change, and node assignment and execution ownership are unchanged: the CPU EP still executes every node it claims, locally, with ORT CPU fallback off. ## The three cuts **1. Five ORT calls and one ORT-side allocation per input → one call, no allocation.** The plugin already fails closed below API 27, so the hook is always present in practice. The five-call sequence is kept as a fallback for a host that leaves it null, and a host offering **neither** now fails closed with a message naming both — it does not silently read garbage shapes. The borrowed pointer is used only to build the owned `shape`/`strides` of the `OwnedInput`; nothing retains it. ORT spells a scalar as a **null** pointer with count 0, and `slice::from_raw_parts` is UB on null, so that case is guarded explicitly — see *Correctness* below for how that guard is actually enforced. **2. `validate_dims` takes a lazily-formatted label, not `&str`.** Every call site passed `&format!("input {i}")`, i.e. a heap `String` built on the success path purely to label an error that almost never happens. Callers now pass `format_args!(..)` and the string is only materialised inside the error branches. One `String` per input per `Run` gone. All call sites converted, including `transfer.rs` and the unit tests; no message text changed. (The merge with `main` took main's `impl std::fmt::Display` signature here, which is a superset of this branch's original `Arguments<'_>` and keeps both callers working unchanged.) **3. `allocate_output` converts the output shape to `i64` on the stack.** `let dims: Vec<i64> = shape.iter().map(|&d| d as i64).collect()` was a heap allocation per output per `Run` whose entire lifetime was the `KernelContext_GetOutput` call below it. Rank ≤ 8 now goes to a stack array; a taller tensor still works, through the `Vec`. For a 1-in/1-out elementwise node that is **5 fewer ORT FFI calls**, one fewer ORT-side allocation, and **2 fewer heap allocations** per `Run`, on top of #1244. ## A/B, one thread `taskset -c 8-15`, this branch vs its merge parent (`origin/main` + #1244), **interleaved** (A,B,A,B,… so drift hits both arms equally), 400 iterations after 50 warmup, **five rounds**. Ratio is **ours/ORT, lower is better**; each column is the median of the five rounds. "ORT drift" is the control: ORT's own p50 measured in the two arms, which must not move for the comparison to mean anything. | case | before p50 | after p50 | before p90 | after p90 | ORT drift | rounds won | |---|---|---|---|---|---|---| | `sqrt_f32_4k` | 1.023 | **0.952** | 1.031 | **0.964** | +2.9 % | 4/5 | | `leakyrelu_f32_4k` | 1.202 | **1.097** | 1.210 | **1.104** | 0.0 % | 5/5 | | `hardsigmoid_f32_4k` | 1.213 | **1.144** | 1.226 | **1.149** | 0.0 % | 5/5 | | `thresholdedrelu_f32_4k` | 1.316 | **1.209** | 1.314 | **1.220** | 0.0 % | 5/5 | | `sigmoid_f32_4k` | 1.328 | **1.252** | 1.333 | **1.258** | 0.0 % | 5/5 | | `tanh_f32_4k` | 1.368 | **1.287** | 1.369 | **1.291** | 0.0 % | 5/5 | | `erf_f32_4k` | 1.383 | **1.327** | 1.387 | **1.329** | 0.0 % | 5/5 | | `celu_f32_4k` | 0.430 | **0.411** | 0.435 | **0.418** | 0.0 % | 4/5 | | `elu_f32_4k` | 0.437 | **0.416** | 0.443 | **0.422** | −1.8 % | 4/5 | | `selu_f32_4k` | 0.467 | **0.441** | 0.471 | **0.446** | 0.0 % | 5/5 | | `log_f32_4k` | 0.727 | **0.700** | 0.728 | **0.706** | 0.0 % | 5/5 | | `mish_f32_4k` | 0.268 | **0.264** | 0.271 | **0.266** | 0.0 % | 4/5 | 12 of 12 cases improve at p50 and p90; 56 of 60 individual arm-pairs favour the change, and no case regressed in its median. In absolute terms it is **0.2–0.3 µs off every `Run`**: `tanh_f32_4k` 4.20 µs → 3.90 µs against ORT's 3.10 µs, `thresholdedrelu_f32_4k` 3.10 µs → 2.80 µs against ORT's 2.30 µs. That is the shape of a fixed per-call cost being removed, which is what it is. `sqrt_f32_4k` crosses below 1.00 — the plugin path now dispatches that op faster than plain ORT does. ## A/B, four threads Same protocol, `NXRT_MM_BENCH_THREADS` = `ONNX_GENAI_MLAS_THREADPOOL_THREADS` = `RAYON_NUM_THREADS` = 4, three rounds. | case | before p50 | after p50 | before p90 | after p90 | |---|---|---|---|---| | `leakyrelu_f32_4k` | 1.069 | **0.966** | 1.075 | **0.976** | | `hardsigmoid_f32_4k` | 1.152 | **1.050** | 1.164 | **1.060** | | `thresholdedrelu_f32_4k` | 1.212 | **1.091** | 1.228 | **1.106** | | `sqrt_f32_4k` | 1.121 | **1.036** | 1.130 | **1.045** | | `sigmoid_f32_4k` | 1.346 | **1.264** | 1.359 | **1.263** | | `tanh_f32_4k` | 1.434 | **1.319** | 1.453 | **1.325** | | `erf_f32_4k` | 1.581 | **1.502** | 1.588 | **1.510** | | `celu_f32_4k` | 0.460 | **0.437** | 0.476 | **0.445** | | `elu_f32_4k` | 0.481 | **0.448** | 0.486 | **0.468** | | `selu_f32_4k` | 0.500 | **0.458** | 0.482 | **0.464** | | `log_f32_4k` | 0.766 | **0.731** | 0.781 | **0.734** | | `mish_f32_4k` | 0.332 | **0.322** | 0.345 | **0.339** | 36 of 36 arm-pairs favour the change. The cost is per call, not per element or per worker, so the gain is the same absolute microseconds at every thread count. ## Drift control: the 1 Mi grid does not move Same protocol, `_f32_1m`, 200 iterations after 30 warmup, five interleaved rounds, one thread. A 1 Mi tensor amortises a fixed per-call cost away, so **these must not move** — and they don't: | case | before | after | | case | before | after | |---|---|---|---|---|---|---| | `relu_f32_1m` | 1.024 | 1.028 | | `sqrt_f32_1m` | 0.313 | 0.309 | | `exp_f32_1m` | 0.580 | 0.576 | | `log_f32_1m` | 0.264 | 0.258 | | `sigmoid_f32_1m` | 0.588 | 0.584 | | `mish_f32_1m` | 0.100 | 0.099 | | `tanh_f32_1m` | 0.618 | 0.620 | | `celu_f32_1m` | 0.127 | 0.129 | | `erf_f32_1m` | 0.651 | 0.668 | | `elu_f32_1m` | 0.129 | 0.129 | | `gelu_exact_f32_1m` | 0.585 | 0.584 | | `selu_f32_1m` | 0.136 | 0.138 | | `gelu_tanh_f32_1m` | 0.602 | 0.646 | | `hardsigmoid_f32_1m` | 0.455 | 0.453 | | `fastgelu_f32_1m` | 0.608 | 0.639 | | `leakyrelu_f32_1m` | 0.409 | 0.417 | | `quickgelu_f32_1m` | 0.429 | 0.438 | | `thresholdedrelu_f32_1m` | 0.567 | 0.573 | The per-case round-win counts here are 0/5–4/5 with a median of 2/5, i.e. a coin flip — the signature of noise, not of an effect. The two largest movers, `gelu_tanh` (+7 %) and `fastgelu` (+5 %), are the two cases whose **ORT** side also drifted most in the same rounds (−4.4 % and −2.8 %), so the ratio moved because the denominator did. This host's floor is roughly ±5 % on the 1 Mi grid and I am not claiming anything below it. ## Correctness **Five new tests, three with a verified falsifier**, driving `read_inputs` and `allocate_output` against a hand-built `OrtApi` whose hooks count their calls. | test | what it pins | falsifier | |---|---|---| | `input_shapes_come_from_one_call_when_ort_offers_the_reference_hook` | 1 reference call, **0** legacy calls, and the shape/strides/dtype that come out | forcing the legacy route fails it with `left: 0, right: 1` (run) | | `the_five_call_fallback_produces_the_same_input_as_the_reference_hook` | the fallback is exercised and its `OwnedInput` matches the reference path field for field | — (it *is* the parity check) | | `a_borrowed_scalar_shape_never_dereferences_null` | ORT's documented scalar spelling — null pointer, count 0 — takes the borrowed route and yields rank 0 | deleting the null guard makes **Miri** report `out-of-bounds pointer use: null pointer is a dangling pointer` at `kernel_ctx.rs:280` (run) | | `read_inputs_fails_closed_when_no_shape_route_exists` | the error names **both** routes; no silent garbage shapes | — | | `output_dims_are_identical_on_the_inline_and_heap_ranks` | ranks 0/1/8/9/12 all reach ORT with every dimension intact | dropping the heap arm delivers rank 9 truncated to 8 dims (run) | **The null guard is now actually enforced, not just asserted.** Natively, `from_raw_parts(null, 0)` returns an empty slice, so a test asserting the resulting shape passes with or without the guard — the guarantee only exists if Miri sees it. `onnx-runtime-ep-plugin` was not in the Miri matrix, so this PR adds `kernel_ctx::` as a lane in `.github/workflows/miri.yml` (the crate's other modules dlopen ORT and are not Miri-tractable, which is why the lane is scoped to the module rather than the crate). 25 tests pass under Miri in 4.15 s; with the guard deleted the lane fails. That pairing is what makes the test non-vacuous. **Suites.** `-p onnx-runtime-ep-plugin`: **252** unit tests pass (247 before, +5 new). `-p onnx-runtime-ep-cpu-plugin` with `NXRT_REQUIRE_ORT_TESTS=1`: every suite green, including all **55** `plugin_ort_e2e` cases — `every_assigned_node_is_also_executed_by_this_ep`, `no_supported_node_is_ever_left_to_the_ort_cpu_ep`, `no_matmul_family_node_escapes_to_the_ort_cpu_ep` and `every_fixture_loads_with_cpu_fallback_disabled` all pass, so **assigned still equals executed** with ORT CPU fallback disabled. `cargo clippy --release --all-targets -p onnx-runtime-ep-plugin` and `cargo fmt` are clean. **Build identity.** Pure native CPU EP: no MLAS at runtime, no ORT CPU EP fallback, no new dependency. AVX2/FMA host (`avx2 fma f16c`, no AVX-512), so ORT and we are on the same instruction footing. ## Independent review Reviewed by **Claude Opus 4.8**, read-only, briefed with the exact ORT contract for the borrowed pointer and asked specifically to hunt UB and vacuous tests. Verdict **APPROVE**, no blockers. It confirmed the null/scalar guard is unreachable-by-construction for `from_raw_parts`, that the borrow cannot outlive the `OrtValue` (its only consumer copies), that `ReleaseTensorTypeAndShapeInfo` still runs on every legacy error path — and noted that moving `DataType::from_onnx` after the match incidentally closes a **pre-existing** leak of `type_shape` on the unsupported-dtype path. It raised two test-quality defects, both real and both **fixed** in `6027f6e16`: 1. The three counting tests shared two process-wide `AtomicUsize`es while asserting exact counts, and libtest runs them in parallel — one test's reset could land inside another's assertion window. They now take a shared `SHAPE_COUNTER_LOCK` that resets both counters under the guard. 2. `a_borrowed_scalar_shape_never_dereferences_null` was **vacuous** with respect to its name, for exactly the reason above, and the crate was not under Miri. Hence the Miri lane, and the test now also asserts the scalar went through the borrowed route so it cannot pass on the fallback. It also flagged that `OrtStatus` is not released on error paths in this file. That is pre-existing and repo-wide in `kernel_ctx.rs` — the old five-call code leaked identically — so it is not touched here; it belongs in its own change. ## What is left After #1244 and this PR, a 1-in/1-out elementwise `Run` is down to roughly `KernelContext_GetInputCount` + `GetInput` + the one shape reference + `GetTensorData` + `GetOutput` + `GetTensorMutableData`. What remains: * **`OwnedInput`'s `shape` and `strides` `Vec`s**, `kernel_inputs`, `infer_shapes`'s `Vec<Vec<usize>>`, `slot_map`, `output_views`, and the `Box<HostPool>` in `host_pool::install` — each worth ~25–35 ns. Removing them needs an inline-capacity vector type or per-session caching of the parts that cannot change between `Run`s. Both are worth doing; neither belongs here. * **The `HostPool` box** specifically is @sebastian's 16-thread scheduling file and should come from him or after his PRs land. --- ## Refreshed against `main` @ `6a855d5e0`, and measured as a stack `origin/main` moved a long way while this sat in the CI queue (#1346, #1352 and #1361 on the quality lane; #1154, #1232, #1238 on the CPU side). Merged in normally — no rebase — and re-measured from scratch against the new baseline. Production pure-native A/B, plain ORT as the control arm. No MLAS, no ORT CPU fallback, no deferral. `taskset -c 8-15`, one thread, 400 iterations, five interleaved rounds out of two worktrees, started only once cores 8-15 were >=93% idle. **Ratio is ours/ORT, lower is better.** `before` is `main` at `6a855d5e0`; `after` is #1244 + #1246 together, since #1246 is stacked on #1244 and the pair is what a user gets. | case | ratio p50 main | ratio p50 stack | Δ | ratio p90 main | ratio p90 stack | ours us | ORT drift | rounds won | |---|---|---|---|---|---|---|---|---| | `thresholdedrelu_f32_4k` | 1.512 | **1.216** | -19.6% | 1.519 | **1.216** | 3.6 → **2.8** | -4.2% | 5/5 | | `tanh_f32_4k` | 1.511 | **1.279** | -15.4% | 1.513 | **1.284** | 4.6 → **3.9** | +0.0% | 5/5 | | `sigmoid_f32_4k` | 1.475 | **1.248** | -15.4% | 1.484 | **1.256** | 4.7 → **4.0** | +0.0% | 5/5 | | `erf_f32_4k` | 1.470 | **1.363** | -7.3% | 1.489 | **1.364** | 7.5 → **7.0** | +0.0% | 5/5 | | `hardsigmoid_f32_4k` | 1.416 | **1.127** | -20.4% | 1.426 | **1.133** | 3.5 → **2.8** | +0.0% | 5/5 | | `leakyrelu_f32_4k` | 1.361 | **1.094** | -19.6% | 1.371 | **1.104** | 3.5 → **2.8** | +0.0% | 5/5 | | `sqrt_f32_4k` | 1.141 | **0.946** | -17.1% | 1.150 | **0.955** | 4.0 → **3.4** | -2.8% | 5/5 | | `log_f32_4k` | 0.767 | **0.698** | -9.0% | 0.776 | **0.703** | 7.9 → **7.2** | +0.0% | 5/5 | | `selu_f32_4k` | 0.505 | **0.440** | -12.9% | 0.512 | **0.444** | 5.6 → **4.9** | +0.0% | 5/5 | | `elu_f32_4k` | 0.480 | **0.416** | -13.3% | 0.488 | **0.421** | 5.3 → **4.6** | +0.0% | 5/5 | | `celu_f32_4k` | 0.470 | **0.411** | -12.6% | 0.478 | **0.418** | 5.7 → **5.0** | +0.0% | 5/5 | | `mish_f32_4k` | 0.276 | **0.264** | -4.3% | 0.278 | **0.268** | 17.3 → **16.6** | +0.0% | 5/5 | **Every case, every round.** The two rows with a moving control (`sqrt` -2.8%, `thresholdedrelu` -4.2%) are reported rather than dropped; both won 5/5 anyway and their absolute time fell by the same ~0.7 us as everything else. That constant ~0.7 us is the point. It is not proportional to tensor size — the same absolute amount comes off `hardsigmoid` (3.5 -> 2.8 us) as off `mish` (17.3 -> 16.6 us) — which is what a fixed per-`Run` cost looks like when you remove some of it. It moves the cheap ops the most because they had the least to hide it behind, and `sqrt` crosses from 1.141 to **0.946**, from a loss to a win. ### Where the remaining time goes Measured directly, by instrumenting `compute_execute` segment by segment on top of this stack (temporary probe, not committed; `perf` is unavailable on this host — `perf_event_paranoid=4`). Per `Run`, one-in/one-out elementwise node, 4096 `f32`, microseconds: | segment | us | note | |---|---|---| | `KernelContext_GetOutput` | 0.35 | ORT's own API — ours to call, not to optimise | | `read_inputs` | 0.15 | 4 ORT FFI calls, already one shape call after #1246 | | rest of `allocate_output` | 0.13 | `GetTensorMutableData` + strides | | `prepare_workspace` | 0.09 | metadata vector + plan-cache lookup, for a kernel needing 0 bytes | | `host_pool::install` | 0.05 | @sebastian's, not touched | | `infer_shapes` | 0.05 | | | `kernel_inputs` | 0.04 | | | `output_views` | 0.04 | | Non-kernel node cost is **~1.25 us and near-constant across all twelve operators** (0.28 to 15.2 us of kernel time), which is the direct confirmation that small-node ratios on this EP are dispatch-bound rather than kernel-bound. There is no single large item left — the biggest, `KernelContext_GetOutput`, is ORT's. The rest is a long tail of 0.04-0.15 us items, which is what #1358 (`InlineVec`) starts on. ### And nothing breaks at 1 Mi Same harness, 1048576 elements, 120 iterations, 3 rounds. A fixed per-`Run` cost should be invisible here, and it is: | case | ratio p50 main | ratio p50 stack | ours us | ORT drift | |---|---|---|---|---| | `celu_f32_1m` | 0.141 | 0.140 | 382.1 → 377.9 | -0.0% | | `elu_f32_1m` | 0.138 | 0.136 | 347.7 → 343.6 | +0.1% | | `erf_f32_1m` | 0.671 | 0.670 | 595.2 → 594.6 | +0.1% | | `exp_f32_1m` | 0.606 | 0.590 | 247.0 → 240.5 | +0.1% | | `fastgelu_f32_1m` | 0.643 | 0.648 | 415.3 → 413.2 | -1.7% | | `gelu_exact_f32_1m` | 0.572 | 0.592 | 706.1 → 710.6 | -2.8% ⚠ | | `gelu_tanh_f32_1m` | 0.657 | 0.651 | 414.8 → 410.5 | +0.3% | | `hardsigmoid_f32_1m` | 0.376 | 0.372 | 88.0 → 69.6 | -0.9% | | `leakyrelu_f32_1m` | 0.421 | 0.430 | 90.1 → 87.3 | -2.2% ⚠ | | `log_f32_1m` | 0.270 | 0.274 | 616.4 → 612.4 | +1.7% | | `mish_f32_1m` | 0.105 | 0.105 | 1626.0 → 1626.0 | -0.2% | | `quickgelu_f32_1m` | 0.455 | 0.453 | 321.5 → 320.0 | -0.5% | | `relu_f32_1m` | 1.034 | 1.022 | 131.7 → 130.2 | +0.0% | | `selu_f32_1m` | 0.147 | 0.148 | 374.1 → 371.8 | -1.5% | | `sigmoid_f32_1m` | 0.471 | 0.606 | 231.6 → 229.0 | -23.4% ⚠ | | `sqrt_f32_1m` | 0.314 | 0.302 | 148.1 → 143.9 | +1.1% | | `tanh_f32_1m` | 0.644 | 0.631 | 226.2 → 221.9 | -0.2% | | `thresholdedrelu_f32_1m` | 0.500 | 0.486 | 70.9 → 68.6 | -0.3% | Flat, as predicted — 0.7 us against 70-1626 us of work. Absolute time is equal or better in 16 of 18 cases. The two ⚠ rows had the control move more than the effect: `sigmoid` is unusable (ORT itself moved -23.4%; our own absolute went 231.6 -> 229.0 us), and `gelu_exact`'s +0.6% absolute sits inside its -2.8% control. Reported rather than dropped. This is the coverage claim for the change: it buys ~0.7 us at every size, which is 20% of a small node and nothing at all of a large one, and it costs nothing anywhere. --- ## Status after merging latest `main` (2026-08-19) Validated on latest `main` after the six-PR #1077 stack landed (#1387, #1409, #1412, #1430, #1433, #1472). Full evidence in the PR comment below; the measured summary: **Deterministic counters, `relu_1_tiny`** — `OrtFfiCall` **10 → 6 per `Run`**, `DispatchAlloc` **24 → 20**. Four fewer round trips for one input, i.e. the predicted 7 → 3 per input. **Timing A/B**, production build, 3 clean reps (a 4th discarded for contention — it read 0.973, which would have flattered us): | case | main | this PR | |---|---|---| | `relu_1_tiny` | 1.397 | **1.271** | | `relu_10_tiny` | 1.250 | **1.173** | | `relu_100_tiny` | 1.163 | **1.110** | The gain shrinks as depth grows — the signature of a fixed-cost fix, since this removes FFI calls per `Run`, not per node. Fixed per-`Run` overhead **1.38 → 1.23** against ORT; per-node slope unchanged. **Test coverage gap this merge exposed and closed:** the pinned costs (`1 + 7` calls, `1 + 3` allocations) run against a zeroed `fake_api()`, where the reference hook is null — so they describe the *legacy fallback*, not the fast path this PR exists to add. Added `the_reference_hook_path_costs_exactly_three_ort_calls_per_input` (pins `1 + 3`) and `the_reference_hook_path_allocates_less_than_the_legacy_path`. `dispatch_probe`'s FFI-coverage table updated to 12 members / 12 `ort_call()` sites, which is the guard that flagged the new API member in the first place. --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: Resch <resch@squad.local>
## The build we ship had no integer GEMM at all `QLinearMatMul` in the default build did this per call: 1. `read_quantized` widened operand `A` to a `Vec<i32>`, then did it again for operand `B`. For a 2048x2048 `B` that is a 16 MiB allocation and fill on every single call, thrown away at the end of it. 2. A scalar rank-1 update walked `A` row by row, and each row re-streamed the whole of `B`. At `m = 128` that is 512 MiB of traffic for 1 GFLOP of work. The result was 11.8x ORT at `m = 1` and 12.1x at `m = 128` — the largest single loss on the x86-64 CPU EP. The performance doc's `QLinearMatMul` rows never described this build: they were taken with `--features mlas`, which is a research build we do not ship. That is now called out in the doc. This adds `kernels/qgemm_native.rs`, a native byte-operand integer GEMM, and points `qlinear_matmul.rs` at it. Nothing defers, and nothing falls back. ## Two kernels, chosen by `m` | shape | kernel | why | | --- | --- | --- | | `m <= 4` (decode) | pack-free fused | one pass over `A` means a packed panel of `B` is never reused, so packing is pure cost. Accumulators stay in registers across a 256-row `k` block. | | `m > 4` (prefill) | packed | `KC/2` pairs of `NC` columns (`KC = 512`, `NC = 256`) is 256 KiB of `B`, which stays in L2 while every row of `A` sweeps it. | The inner tile is `vpmaddwd` over `NR = 16` columns and `MR = 4` rows, `k` consumed two rows at a time. ## Why `vpmaddwd` and not `vpmaddubsw` MLAS gets 32 MACs from two instructions using `vpmaddubsw`, which **saturates**: it needs a sign-domain translation of `B` and its intermediate is only nominally exact. `vpmaddwd` needs four instructions for the same 32 MACs, but with centred `a` in `[-255, 255]` and raw `b` in `[-128, 255]` a product is at most 65025 and a pair sum at most 130050, so it cannot saturate and cannot overflow. No sign-domain flip, no reasoning about clamped intermediates. That instruction-count difference is the whole of the residual gap at `m = 1`. Closing it means giving up exact integer arithmetic, which is not a trade I am willing to make for a quantized kernel whose entire value is that it is exact. ## Determinism is structural, not tested-in The kernel computes `sum_k (a - za)(b - zb)` as `sum_k (a - za) * b - zb * sum_k (a - za)`, with every accumulation a **wrapping** `i32` add. Wrapping addition is arithmetic mod 2^32, which is associative and commutative, so *any* blocking, tiling, column split, row split or thread count gives bit-identical output — including on overflow, where the wrap itself is reproducible. `wrapping_overflow_is_reordering_invariant` and `the_thread_count_cannot_change_the_result` assert exactly that, and the SIMD path is checked bit-for-bit against a portable scalar oracle (`the_simd_kernel_is_bit_identical_to_the_portable_loop`, and separately for the fused path). ## Numbers Session A/B against plain ORT, `K = N = 2048`, u8 x u8, ratio is `ours / ORT`, **lower is better**, p50 of 61 iterations. ORT's own timings moved under 1.5% between the two arms at 1 and 4 threads, which is the control that makes the comparison mean anything. | M | threads | before | after | ours before | ours after | | ---: | ---: | ---: | ---: | ---: | ---: | | 1 | 1 | 12.20x | **2.17x** | 1.402 ms | 0.226 ms | | 128 | 1 | 11.90x | **1.20x** | 99.84 ms | 9.99 ms | | 1 | 4 | 37.11x | **4.03x** | 1.379 ms | 0.170 ms | | 128 | 4 | 14.44x | **1.47x** | 31.03 ms | 3.14 ms | | 1 | 16 | 83.12x | 35.63x | 2.366 ms | 1.336 ms | | 128 | 16 | 42.28x | 15.05x | 51.05 ms | 16.55 ms | `i8_m1` goes 0.206 ms to **0.049 ms** at one thread. Kernel-level scaling (`bench_qgemm_ab`, `taskset -c 0-15`), with the portable scalar arm as the control: | shape | 1t | 2t | 4t | 8t | 16t | portable 1t | | --- | ---: | ---: | ---: | ---: | ---: | ---: | | 1x2048x2048 | 0.229 ms | 0.136 | 0.090 | 0.098 | 0.166 | 4.92 ms (21x) | | 4x2048x2048 | 0.565 ms | 0.311 | 0.199 | 0.237 | 0.345 | 4.58 ms (8.1x) | | 128x2048x2048 | 8.911 ms | 4.773 | 2.755 | 2.780 | 1.991 | — | | 128x5120x5120 | 53.56 ms | 27.06 | 14.35 | 8.98 | 11.51 | — | The task grid splits rows as well as columns. Columns alone gave only `n / NC` tasks — eight for `n = 2048` — so a sixteen-worker pool left half of itself spinning; `128x2048x2048` was 2.69 ms at sixteen threads against 1.62 ms at eight. Splitting columns further would shrink the panel and re-walk `B`; splitting rows duplicates only the pack, about a percent of the GEMM it feeds. ## Things I measured and rejected - **Software prefetch** of the next `B` rows (`PREFETCH_ROWS = 8`): a consistent **8% regression** with a stable `m = 128` control. The hardware prefetcher already has the sequential stream. - **Permuting inside the fused inner loop**: replaced by accumulators held in the permuted order with a single `vperm2i128` fixup per `k`-block flush. Saves eight instructions per 32 MACs. ## Left open, deliberately - **Constant-`B` packed cache.** The pack is repeated per call. Caching it would remove it from prefill entirely, but any new weight-derived cache has to go through `kernels/governed_weight_cache.rs` to satisfy the "New weight-derived caches must be governed" gate. That is a separate PR with its own eviction story, not a rider on this one. - **The session-level threading gap.** At four threads the session takes 0.170 ms while the kernel alone does 0.090 ms, and past eight threads both arms get worse. That is the pre-existing oversubscription item — it is present before and after this change, so it is not a regression here, and it is the next thing I am working on. ## Validation - `cargo test --release -p onnx-runtime-ep-cpu --lib` — 1340 passed, 0 failed. - Every `onnx-runtime-ep-cpu-plugin` suite with `NXRT_REQUIRE_ORT_TESTS=1`, including the 53-test `plugin_ort_e2e` ORT conformance suite with CPU fallback disabled. - `cargo clippy -p onnx-runtime-ep-cpu --all-targets` clean, `cargo fmt --all --check` clean. - `cargo check -p onnx-runtime-ep-cpu --lib --features mlas` — the research build still compiles. - Reviewed by Claude Opus 4.8 against the memory-safety, lane-semantics, determinism and edge-extent claims above; no blockers, two documentation fixes applied. --- ## Refreshed against `main` (2026-08-18) The branch was behind `main` and its red CI wall came from that, not from this change: `crates/onnx-runtime-session/src/executor/mod.rs:175` failed `-D dead-code` on current stable, fixed on `main` by `ca32b3adf` (#1239) after this branch forked. `origin/main` (`c55a3fab3`) is merged in — no rebase, no force-push. One conflict, in `docs/performance/CPU_MATMUL_ASSIGNMENT.md`, resolved as a **union**: this branch's `#### 3b` (the native integer GEMM) and `main`'s `### 4` (the f32 `M = 1` GEMV becoming the default, #1091) were both new sections appended after 3a. Both are kept, in that order. Taking either side would have silently deleted the other's record. Revalidated on the merge commit, AVX2/FMA host, no AVX-512: * `cargo test --release -p onnx-runtime-ep-cpu --lib` — **1424 passed, 0 failed**, 18 ignored, including `qgemm_i32_matches_the_integer_oracle_for_every_signedness`, `the_simd_kernel_is_bit_identical_to_the_portable_loop`, `wrapping_overflow_is_reordering_invariant` and `the_thread_count_cannot_change_the_result`. The measurements in this PR were taken before the merge; nothing in the merged range touches `qgemm_native.rs`, `qlinear_matmul.rs`, or the CPU threadpool, so they stand as recorded. The `main` change that did land in this range (#1091's f32 `M = 1` GEMV default) is on a different kernel family and is documented in the section-4 text kept above. --- ## Refreshed again against `main` @ `6a855d5e0`, and a real branch bug found `main` moved again while this was queued (#1346/#1352/#1361 on the quality lane, #1154/#1232/#1238 on the CPU side). Merged in normally — no rebase — and revalidated. The revalidation caught something the earlier ones had not. Running `-p onnx-runtime-ep-cpu --lib` in a **debug** profile rather than `--release` fails: ``` kernels::qgemm_native::tests::degenerate_extents_do_nothing assertion `left == right` failed left: 0 right: 4 ``` `degenerate_extents_do_nothing` called `qgemm` with an empty `b_zero_points` and `n == 4`. `qgemm` opens with `debug_assert_eq!(b_zero_points.len(), n)`, so that call is not one the function accepts — the test was exercising the `m == 0` early return through an argument list the contract forbids. It passed every previous run here only because `debug_assert` compiles out under `--release`, which is how I had been validating this branch locally. A debug test profile fails it, and this is branch-caused: `qgemm_native.rs` is new in this PR. Fixed in `9ca99e538` by sizing the test's zero points to `m` and `n`, not by weakening the assertion — the assertion states the contract the kernel's indexing depends on, and a caller whose `m` is zero still has `n` columns and still knows their zero points. **`-p onnx-runtime-ep-cpu --lib`, debug profile: 1440 passed, 0 failed** (was 1439 passed, 1 failed). This is the second time on this stack that the profile a test runs under decided whether it caught anything. Worth remembering: `--release` silently disables every `debug_assert` in the crate under test, so a local `cargo test --release` is not a substitute for what CI runs. --- ## Re-validated on latest `main` (`e0aedd0fa`), 2026-08-19 Latest `main` merged in normally (no rebase). Full re-measurement, 1 thread pinned, `K = N = 2048`, 61 iters / 10 warmup, 2 reps, `ours_p50 / ort_p50`: | case | `main` ours | `main` ratio | this PR ours | this PR ratio | speedup | |---|---|---|---|---|---| | `bench_qlinear_u8_m1` | 1.418 / 1.435 ms | 11.83x / 12.51x | **0.121 / 0.123 ms** | **1.16x / 1.18x** | **11.7x** | | `bench_qlinear_u8_m128` | 29.55 / 29.56 ms | 3.57x / 3.57x | **3.055 / 3.065 ms** | **0.372x / 0.373x** | **9.6x** | | `bench_qlinear_i8_m1` | 1.516 / 1.497 ms | 0.215x / 0.212x | **0.209 / 0.212 ms** | **0.030x / 0.030x** | **7.2x** | ORT-side drift between the two arms was 0.7% at `m = 128` and 0.0% on `i8`, which is the control that makes the comparison mean anything. **At `m = 128` we are now 2.7x faster than ORT outright**, and `m = 1` closes from 11.8x to 1.16x. These are better than the numbers originally posted above because the dispatch work in #1077 landed in between. ## Review fixes (`987aa0c5c`) An independent review found no blockers but two things worth fixing: 1. **aarch64 built with 5 warnings** — `NR`/`MR`/`NC`/`KC`/`FUSED_KC` are read only by the x86 kernels, so every non-x86 target warned on all five. CI builds with `-D warnings`, so this was a branch-caused CI failure waiting to happen; the local x86 clippy run could never have caught it. Now `#[cfg]`-gated alongside the code that uses them: **0 warnings on both x86-64 and aarch64**. 2. **The fused-parallel path had no end-to-end coverage.** Every `m <= 4` shape in `qlinear_matmul_reordered_accumulation_is_bit_identical` sat below `PARALLEL_MIN_WORK`, so the pack-free kernel's column split was only ever checked at the kernel level, never through `requantize_rows`. Added `(4, 1029, 1100)`, which forks both. The review independently re-derived the register-shuffle math in numpy (`cvtep*_epi16`, `permute4x64_epi64(0xD8)`, `unpacklo/hi_epi16`, `madd_epi16`, `permute2x128`) against a plain per-column dot product over 2000 tiles with extreme values — 0 mismatches — and confirmed the `vpmaddwd` non-saturation bound for all four operand combos, the wrapping-add determinism claim, and the absence of out-of-bounds access in every tail path. ## Validation on the merged base - `cargo fmt` clean; `cargo clippy --all-targets -D warnings` clean - **1553 `onnx-runtime-ep-cpu` tests**, debug profile (so `debug_assert`s are live) - **55 plugin conformance tests** (`NXRT_REQUIRE_ORT_TESTS=1`, release) - `every_assigned_node_is_also_executed_by_this_ep` and `every_fixture_loads_with_cpu_fallback_disabled` green — nothing defers, nothing falls back to the ORT CPU EP - **aarch64-unknown-linux-gnu** cross-check clean, 0 warnings --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: Resch <resch@squad.local>
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()returningtrueon a probe dispatch, so theslice in front of it went to the host pool. On a session that never latches —
intra_op_num_threads = 1, where ORT runsKernelContext_ParallelForinline —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 interleavedrounds) against a control build with probing compiled out:
QuickGeluSwishErfTanhSigmoidClip5-17% at 1 Mi across seven ops, and the
intra_op = 1benchmark in #1143'sbody 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:
against ~2.5 for the old work-carrying probe. So:
PROBE_BURST32 → 8 — still ~3x the worst case observed.PROBE_MIN16 → 64 — the burst is what answers the question; this isonly the recovery path for a session whose pool was busy through all eight.
PROBE_STALLis now a deadline, not a duration: hold the first indexuntil 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 (
ErfandGeluat 65 Ki,Clipat 256 Ki showed no improvementover
mainat all in that run). Early exit is what makes 400 µs affordable.Benchmarks
Same method as #1143:
.work/mt2.pyalternates the two.sos round by roundin 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.
controlis this branch with probing compiled out, whichbounds what the measurement noise on this shared box looks like:
TanhErfSwishSigmoidQuickGeluGeluClipReluSqrtFastGeluThe 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 166 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:
ClipErfFastGeluGeluQuickGeluReluSigmoidSqrtSwishTanhThat 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 be3 rounds, µs at 1 Mi:
Relu235.7 vs 312.7 on main,Sigmoid779.5 vs 855.4,Erf894.2 vs 892.3,Gelu1302.7 vs 1301.2,Sqrt666.5 vs 666.0,Tanh809.9 vs 814.6,Clip313.6 vs 313.7.Also from the review of #1143
run_on_hostclaimed every chunk is a multiple ofeight 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 stillbit-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_serialused a slice shorter thanPAR_MIN_LEN, so itpassed whether or not the
in_host_taskguard existed, and it only assertedon the host counter. It now uses
PAR_MIN_LEN + 4099and asserts the rayoncounter too — the guard's actual job.
try_hostbumpedPARALLEL_DISPATCHES(documented as rayon dispatches) ona host dispatch, while
try_host_rowsdid 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 fmtclean for the files this PR touches; clippy clean.
the_probe_and_the_latch_agreeis rewritten for the new semantics: it drivesprefer_hostthrough the realort_parallel_forover 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_MAXcap.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
open:
Tanh0.37x,Sigmoid0.41x,Relu0.51x,Clip0.52x,FastGelu0.54x at 1 Mi. Those are per-op kernel and memory-bandwidth problems.
HOST_MIN_LENand dominated by per-node pluginoverhead.
PROBE_BURST = 8and 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.