Repository navigation
perf(plugin): stop routing host intermediates through ORT scratch, and recycle them - #1073
Merged
Merged
Conversation
`x * 0.5` and `x + b` ran at ~7.7 ns/element, about 60x slower than ORT's CPU kernels. That is not a corner case: when ORT has no kernel for a contrib op at a given dtype it inlines the ONNX function, and the inlined body is almost entirely scalar-broadcast Mul and Add. Two separate causes, both structural. First, `binary_contiguous` only fires when all three shapes are identical, so any broadcast -- including a scalar operand -- fell to `broadcast_apply`, which per element does a dot product over the tensor rank, a `next_index` carry chain and a closure call, and additionally allocates a whole-tensor accumulator that it then walks three times. `Add` had no non-mlas fast path at all, so even a same-shape f32 Add took that route. Add `binary_broadcast_contiguous`, which recognises an operand whose shape is a right-aligned suffix of the output shape -- a scalar, or the row-bias case `[B, S, C] op [C]` -- and walks it as a repeated contiguous block. An interior unit axis such as `[B, 1, C]` is not a suffix and is still declined to the general path. `Add` now shares this primitive rather than growing a second copy of the walk. Second, `BinOp::apply` matches on a runtime value, and left inside the element loop that match is re-evaluated per element and blocks auto-vectorisation outright. With only the first fix an f32 scalar multiply was still 1.1 ns/element. `dispatch_binop!` resolves the combiner once, outside the loop, for both the new broadcast walk and the existing same-shape one. Measured, 1 thread, interleaved, ratio = ORT ns / plugin ns: Mul by scalar n=1048576 0.016 -> 0.961 (8577 us -> 137 us) Add by scalar n=1048576 0.017 -> 0.967 (8058 us -> 137 us) Equivalence with the general walk is pinned bitwise: every case is run twice over the same logical values, once with a contiguous broadcast operand and once with a strided view of padded storage that `dense_operand` rejects, and the raw output bytes must match. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
The two new fast-path tests asserted that the process-global `ADD_SCALAR_TEST_HITS` counter did *not* advance across an `AddKernel::execute` call. Cargo runs tests in parallel and several other tests in the same binary legitimately increment that counter, so the equality assertion raced and failed intermittently (observed ~1 run in 3 in debug). Assert the dispatch predicate directly instead: call `add_dense_fast_path` and require it to accept (and produce the right values), which is deterministic. `AddKernel::execute` only reaches the counter after that predicate declines, and the two arms ahead of it (`mlas`, vDSP) both require identical operand shapes, so a broadcasting input the predicate accepts provably never reaches the fallback. Adds the decline half as its own test: an interior unit axis must be refused outright and must leave the output untouched. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
A 16-bit binary elementwise op computes `from_acc(fold(to_acc(a), to_acc(b)))` with `Acc = f32`, so every element paid two `half` software conversions in and one out. That was ~3.6 ns/element -- roughly 12x slower than ORT -- and it dominated every f16 graph, including the bodies ORT inlines for the f16 contrib activations (FastGelu, QuickGelu, Gelu), which are mostly Cast/Mul/Add. Widen both operands in `HALF_STAGE_CHUNK`-element passes with the existing F16C/AVX2 bulk converters, fold in f32, and narrow back. The scalar-operand shape -- what an inlined activation emits -- skips the second staging buffer entirely and folds against a register constant. f16 @1m elements, vs ORT, interleaved A/B, 1 thread: Mul dense 0.082 -> 1.278 Add dense 0.100 -> 1.308 Mul scalar 0.086 -> 1.170 Add scalar 0.083 -> 1.167 f16 FastGelu 0.030 -> 0.524, QuickGelu 0.030 -> 0.730. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
…ementwise # Conflicts: # crates/onnx-runtime-ep-cpu/src/kernels/elementwise.rs
The comment claimed three f32 buffers totalling 12 KiB, but the staged loop allocates two (one per operand); the narrow step writes back through the left buffer. Corrected to two buffers / 8 KiB and cross-referenced F16_STAGE_CHUNK in dense_elementwise, which stages the unary paths with the same chunk size. Found by independent review. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
…into deckard/f16-contrib-assignment
…d recycle them
A multi-node partition was paying ~10x for every intermediate it produced.
Intermediates were allocated with KernelContext_GetScratchBuffer using the
memory info of kernel-context input 0, which on a host EP is CPU memory.
ORT services that through an aligned allocation, and glibc always maps a
fresh region for an aligned request of this size, so each intermediate was
a new mapping: first-touch page faults while the producing kernel wrote it,
and an unmap at the end of the Run. Nothing was ever reused, and a chain of
N nodes touched N cold megabytes.
Two changes, both confined to host-resident partitions:
* Intermediates go through ORT scratch only when the resolved memory info
is a *device*. That is the case scratch exists for - a device kernel
handed a host pointer dereferences it as device memory - and it is
untouched. Host partitions take a plain Vec instead.
* Host intermediates are recycled by liveness. A buffer is retired as soon
as the last node that reads it has run, so the next allocation reuses
storage that is still in cache, and a thread-local pool carries the
storage across Runs.
Measured on an 8-node f32 Relu chain, one thread, 15 interleaved rounds,
as a fraction of ORT CPU EP session latency:
elements before after
1024 0.183 0.623
16384 0.081 0.782
262144 0.066 0.801
f16 FastGelu, which ORT inlines into a 15-node primitive body, goes from
0.43-0.52x to 0.72-0.78x.
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #1073 +/- ##
==========================================
+ Coverage 79.71% 79.73% +0.01%
==========================================
Files 369 369
Lines 160706 160844 +138
Branches 160706 160844 +138
==========================================
+ Hits 128113 128242 +129
- Misses 27860 27865 +5
- Partials 4733 4737 +4
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
|
… end Two guardrails for the non-zeroing relaxation, both from review. Debug builds fill a reused buffer with 0xFF before handing it back: NaN in every float width, -1 in every signed integer width, so a kernel that leaves part of its output unwritten cannot have the gap absorbed by a tolerance comparison. Release builds still skip the write, which is the cost the change exists to avoid. The new end-to-end test runs the three-node Add/Mul/Add fixture six times against real ORT with changing inputs. Only the first Run gets zeroed storage; every later one is served recycled buffers, so a partial write would surface as the previous iteration's answer. A single Run cannot catch that. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
…ssignment # Conflicts: # crates/onnx-runtime-ep-cpu/src/kernels/elementwise.rs
justinchuby
marked this pull request as ready for review
August 16, 2026 19:56
justinchuby
added a commit
that referenced
this pull request
Aug 16, 2026
## What `cargo fmt --all -- --check` currently **fails on `origin/main`** (`bce03cabb`) in five places across `crates/onnx-runtime-ep-cpu/src/kernels/gemm.rs` and `crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs`. This is the `cargo fmt --all` output and nothing else. ## Why it happened Nobody wrote badly-formatted code. #1073, #1079 and #1080 each touched these two files and each was fmt-clean against its own base. Squash-merging them produced a combined text that rustfmt formats differently: - `TRANSPOSE_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner())` is 61 characters, which exceeds rustfmt's default `chain_width` of 60 once it sits at test-body indentation (two sites). - The `onnx_runtime_ir` import list grew past `max_width` and now wants braces on their own lines. - One `Gemm`/`transB` attribute chain became short enough to fit on a single line after a neighbouring edit. - One stray double blank line at end of a `mod tests`. `main` is unprotected, so no required check re-ran fmt on the merge result and the breakage landed silently. ## How it was found It blocked #1086: that PR's `Fast (Linux x86_64)` and `Rust quality` jobs failed on the PR **merge ref** with diffs in files #1086 does not touch. Reproduced independently by checking out `origin/main` into a clean worktree and running `cargo fmt --all -- --check`. ## Verification - `cargo fmt --all -- --check` -> clean (was: 5 diffs). - Diff is whitespace/line-breaking only; `git diff -w` on the two files is empty apart from the import-brace move. No logic, no behaviour, no test changes. ## Risk None. Mechanical formatter output. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
Every multi-node subgraph this EP claims was paying roughly 10x for each
intermediate tensor it produced.
Routed subgraphs allocate their intermediates with
KernelContext_GetScratchBuffer, using the memory info returned bydevice_mem_info. On a device EP that is correct and necessary: a device kernelhanded a host pointer dereferences it as device memory. On a host EP,
device_mem_infofinds no device-resident input and falls through to the memoryinfo of kernel-context input 0 — CPU memory — so the host path was going through
ORT's scratch allocator too.
That allocator services the request through an aligned allocation, and glibc
maps a fresh region for an aligned request of this size regardless of the
M_MMAP_THRESHOLDtuning that makes ordinarymallocreuse the heap. So everyintermediate was a brand-new mapping: first-touch page faults charged to the
producing kernel's write, and an unmap when the
Runended. Nothing was everreused, and an N-node chain touched N cold megabytes per
Run.The cost is invisible in a single-node subgraph — which is why it survived this
long — and it is exactly the shape ORT hands us for the float16 contrib
activations, because ORT has no float16 CPU kernel for them and inlines their
function bodies into 15-node
Cast/Mul/Add/Tanhprimitive graphs.How it was found
An 8-node
Reluchain, f32, one thread. Same kernel, same size, but position inthe subgraph changed the time by more than 10x:
Phase timing inside
compute_executeput 98% of it insideexecute_with_workspace, not in the allocation call — consistent with pagefaults being charged to the kernel's first write rather than to the allocator.
Instrumenting
try_dense_elementwiseconfirmed the SIMD fast path was firingfor every node and the overlap guard never tripped, so it was not a dispatch
regression. Forcing the host-buffer arm dropped the same graph from 3246 us to
383 us.
Change
Both parts are confined to host-resident partitions. Device EPs keep the exact
behaviour they have today.
intermediate_scratch— intermediates go through ORT scratch only whenthe resolved memory info is a device (
mem_info_is_device). Otherwise theyare plain
Vec<u8>.scratch_mem_infoitself is unchanged and still feedsprepare_workspace's placement fallback.Liveness-based recycling —
last_reader_per_buffercomputes, from therouting table, the highest node index that reads each intermediate. A buffer
is retired the moment its last reader has run, so the next allocation gets
storage that is still in cache, and a bounded thread-local pool carries that
storage across
Runs. An unread buffer or an out-of-range index maps toNone, which means "nothing keeps this alive" — the conservative answer,since a buffer is only ever retired after its recorded last reader.
Reused storage is not re-zeroed in release builds. That is not a weakening of what kernels may
assume: an output routed to an ORT sink already arrives as whatever
KernelContext_GetOutputreturned, which ORT does not zero, and thesingle-kernel path has always worked that way. Making buffer sinks behave the
same way means a kernel that fails to write its whole output now fails
identically wherever it sits in a partition, instead of only when it happens to
be last. Re-zeroing costs a full
memsetper intermediate — measured at 0.586xof ORT against 0.801x without it on the 1 MiB chain.
Debug builds do write reused storage, with
0xFFpoison rather than zeros —NaN in every float width,
-1in every signed integer width, so it cannot hideinside a tolerance comparison. A kernel that leaves part of its output unwritten
therefore fails loudly in every test run instead of inheriting something
plausible.
Benchmarks
AMD EPYC 9V74, AVX2 (no AVX-512, masked by the hypervisor), 1 intra-op thread,
taskset -c 0-15, 15 interleaved A/B rounds per point, ORT 1.28.0. Ratio isORT CPU EP session latency / plugin session latency, so >1.0 means the plugin
is faster.
"Before" is this commit's parent,
b0fd8a040— notgit merge-base HEAD origin/main. The merge base is the first parent of the f16 elementwise mergethis branch carries, so measuring float16 there reports ~0.10x and attributes
another PR's win to this one.
Numbers are steady state. The first large tensor a process touches pays a
one-off penalty the per-session warmup does not fully absorb, so a single
low-repetition run can read ~2x low on whichever size happens to go first.
8-node f32
Reluchain (a routed subgraph with 7 intermediates)float16 activations that ORT inlines into multi-node bodies
21 interleaved rounds. p90 within 0.005 of p50 on every cell.
FastGeluFastGeluFastGeluQuickGeluQuickGeluQuickGeluGelu(tanh)Gelu(tanh)Gelu(tanh)QuickGelu's inlined body is shorter and produces fewer intermediates, so itmoves less.
Gelu(tanh)is unchanged, as expected — it was already a singleclaimed subgraph whose ratio is set by the kernels, not the plumbing.
Single-node subgraphs are untouched by construction (no intermediates): f32
Reluat 262144 is 0.944 before and after.Limitations
ORT's allocation planner reuses two arena slots for the whole chain and its
elementwise kernels thread across the intra-op pool while ours are
single-threaded. Those are separate problems and are not claimed to be fixed
here.
aligned large allocations bypassing allocator reuse — is glibc-specific in its
details, so the size of the win will differ elsewhere. The direction should
not: reusing warm storage cannot be worse than mapping cold storage.
than that falls back to fresh allocation for the excess rather than growing
without limit.
Tests
Eight new unit tests in
compute.rs:last_reader_marks_the_final_consumer_of_each_bufferlast_reader_takes_the_highest_index_when_a_buffer_is_read_twice— thefalsifier for the dangerous failure mode, a buffer freed before its second
reader
last_reader_is_none_for_unread_and_out_of_range_buffersrecycled_intermediate_storage_is_reused_without_reallocating— assertsaddress reuse, which is the entire point
a_recycled_buffer_serves_a_smaller_request_at_the_requested_length— thelength must be the requested one, since
byte_lenbounds everyfrom_raw_partsbuilt from the buffera_request_larger_than_every_pooled_buffer_allocates_fresh_zeroed_storagescratch_backed_buffers_are_not_pooledandthe_pool_is_boundedEach pool test drains the thread-local pool first, so they are deterministic
under both the parallel harness and
--test-threads=1.One new end-to-end test against real ORT, in
plugin_ort_e2e.rs:conformance_chain_add_mul_repeated_runs_do_not_leak_stale_intermediates—the three-node
Add/Mul/Addfixture, sixRuns with changing inputs. Thefirst
Rungets freshly zeroed storage; every later one is served recycled,dirty buffers, so an element a kernel failed to write would surface as the
previous iteration's answer. In debug builds it is served
0xFFpoisoninstead, i.e. NaN, which no tolerance admits.
cargo test -p onnx-runtime-ep-plugin --lib→ 224 passed, parallel and--test-threads=1.NXRT_REQUIRE_ORT_TESTS=1 cargo test -p onnx-runtime-ep-cpu-plugin→ 32passed against real ORT (31 before, plus the new one), with the debug poison
active.
cargo test -p onnx-runtime-session --lib→ 161 passed.cargo clippy -p onnx-runtime-ep-plugin --all-targets→ clean.cargo fmt --all --check→ clean.Independent review
Reviewed independently: GO WITH FINDINGS, no blockers. The reviewer built
both sides in a separate worktree and reproduced every cell:
Reluchain=8, 1024Reluchain=8, 16384Reluchain=8, 262144FastGeluf16, 3072FastGeluf16, 262144FastGeluf16, 1048576QuickGeluf16, 1048576Reluf32, 262144They also:
compute.rs:2065against every new.take()and confirmed theTensorView<'static>aliases are dead before anyretirement runs — moving a
Vecinto the pool preserves its heap address, anda drop can only happen after the last use.
reachable as an interior node; confirmed
absent_scratchstill allocateszeroed (
compute.rs:2117).Reluchain=8 and chain=15,Sigmoidx12,Tanhx10, and f16FastGelu/QuickGelu, 50 iterations each sothe pool serves dirty storage throughout — all matched ORT (f32 max diff 0,
f16 within 2e-3).
--test-threads=1looking for pool-related flakiness; none.Findings applied: the debug poison fill and the repeated-run e2e test are theirs
(they asked for a guardrail on the non-zeroing relaxation); the benchmark
section now names the correct before-commit and warns about first-touch
dispersion; the test count is corrected.