Repository navigation
perf(cpu): run the MLAS activation routes through run_chunked - #1127
Merged
Merged
Conversation
The MLAS route called its kernel directly on the whole tensor and returned, so it never went through `run_chunked` -- the shared chunking/parallel seam every pure-Rust route uses. That was invisible until #1105 taught `run_chunked` to parallelise: from then on an `mlas`-on build ran these ops single-threaded no matter how large the pool, while an `mlas`-off build scaled. Since #1115 the wheel ships `mlas` on x86_64, so this was shipping. Measured at 16 threads, p50 us, medians of 3 interleaved rounds with EP assignment asserted: 1 Mi 4 Mi Erf 897 -> 394 (2.3x) 3559 -> 562 (6.3x) Gelu 1304 -> 553 (2.4x) 5459 -> 1134 (4.8x) Tanh 481 -> 300 (1.6x) 2220 -> 426 (5.2x) Sigmoid 504 -> 299 (1.7x) 1992 -> 430 (4.6x) Ops with no MLAS route are the control and are unchanged within noise (FastGelu 414 -> 395 / 600 -> 605, QuickGelu 330 -> 347 / 529 -> 546, Sqrt 255 -> 232 / 353 -> 375). At 1 thread nothing moves: Erf 3628 -> 3644, Gelu 5467 -> 5479, every other cell inside run-to-run noise. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
The bypass produced correct results, so nothing caught it -- the thread-invariance tests pass trivially for a route that never splits. This asserts the mechanism instead: a tensor over PAR_MIN_LEN, submitted from outside the pool, must increment run_chunked's parallel-branch counter. Verified to fail on the pre-fix code (`erf: a 1052675-element call did not reach run_chunked's parallel branch`) while the pure-Rust control kept passing, so it is not a blanket assertion. The counter is thread-local and `#[cfg(test)]`, so concurrent tests cannot bump each other's count and non-test builds get an empty inlined function. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #1127 +/- ##
==========================================
+ Coverage 80.71% 80.73% +0.01%
==========================================
Files 368 368
Lines 159904 160093 +189
Branches 159904 160093 +189
==========================================
+ Hits 129066 129244 +178
- Misses 26120 26131 +11
Partials 4718 4718
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
|
justinchuby
added a commit
that referenced
this pull request
Aug 17, 2026
#1130) ## Root cause `run_chunked` is the seam that #1105 taught to split activation work across the rayon pool. #1127 fixed `dispatch_mlas!`, which called its kernel directly and returned, bypassing that seam entirely — so every MLAS-routed op ran single threaded no matter how many threads were configured. The independent review on #1127 reported as a **MAJOR** finding that three more callers have the identical pathology, in other files, so they were out of scope there: | caller | file | |---|---| | `silu_f32_slice` | `kernels/activations.rs` | | `relu_contiguous_f32_mlas` | `kernels/relu.rs:144` | | `Clip` (selection) | `kernels/selection.rs:182` | | `Clip` (conv epilogue) | `kernels/conv.rs:732` | `mlas-sys` documents `compute_silu`, `compute_relu` and `compute_clip` as *"Single threaded; callers shard across threads themselves."* Nobody sharded. This PR wraps all four call sites in `run_chunked`. SiLU needed one extra step. Its MLAS route is followed by a correction scan over the whole tensor (MLAS's `compute_silu` is inaccurate outside `±SILU_MLAS_SAFE_BOUND`). Run whole-tensor, that scan streams the buffer a second time from DRAM. It is now blocked at `SILU_CORRECTION_BLOCK = 8192` so each block stays in L2, and the scan is a branch-free OR-reduction **over the input only** — the predicate `!x.is_finite() || x.abs() > SILU_MLAS_SAFE_BOUND` depends solely on the input, so the common all-in-band case skips the write loop entirely. ## Benchmarks Session level through the plugin `.so`, base = `origin/main` @ `b5309f799`, 16 threads, 3 interleaved rounds, randomised order, `# NOT-ASSIGNED: 0` on every run (no node was left to ORT's CPU EP). µs, p50. | op | n | base | this PR | speedup | ORT | ORT-rel before | ORT-rel after | |---|---:|---:|---:|---:|---:|---:|---:| | Clip | 1 Mi | 351.36 | 256.98 | **1.37×** | 34.73 | 0.099 | 0.135 | | Clip | 4 Mi | 1284.90 | 516.92 | **2.49×** | 88.91 | 0.069 | 0.172 | | Relu | 1 Mi | 334.00 | 252.63 | **1.32×** | 40.42 | 0.121 | 0.160 | | Relu | 4 Mi | 1095.22 | 494.47 | **2.22×** | 78.41 | 0.072 | 0.159 | | Swish | 1 Mi | 1415.31 | 404.02 | **3.50×** | 226.21 | 0.160 | 0.560 | | Swish | 4 Mi | 5629.53 | 549.99 | **10.24×** | 474.99 | 0.084 | **0.864** | `Swish` (default domain, opset 24) is the ORT-visible spelling of SiLU and is supported by ORT 1.28, so SiLU does have a real single-node session-level A/B after all — the earlier note that it did not was wrong, and it is the op that gains the most here. Kernel-level SiLU, `serial_scope` vs parallel in-process, 32 threads, so the MLAS route is compared against itself with only the split changed: | n | serial | parallel | speedup | |---:|---:|---:|---:| | 1 Mi | 2051.4 | 641.6 | 3.20× | | 4 Mi | 8231.0 | 1266.0 | 6.50× | | 16 Mi | 33589.2 | 3380.0 | 9.94× | ## Correctness - `blocked_correction_matches_the_whole_tensor_loop_bit_for_bit` — the blocked, OR-reduced scan is compared bit for bit against the original whole-tensor loop over in-band values, out-of-band values, `±Inf`, NaN, `±0`, denormals and values sitting exactly on `SILU_MLAS_SAFE_BOUND`, at lengths that straddle the block boundary. - `silu_reaches_run_chunked_parallel_branch` — asserts the **mechanism**, not the output, using the `PARALLEL_DISPATCHES` counter added in #1127. Verified to falsify: reverting the `run_chunked` wrapper makes it fail. - `silu_is_thread_count_invariant` — identical results across pool sizes. - No tolerance was relaxed anywhere. No numerical behaviour changes: this PR only changes *who* runs the arithmetic, plus a blocking/reduction rewrite that is proven bit identical. `cargo test -p onnx-runtime-ep-cpu --features mlas --lib` → **1322 passed, 0 failed**. Both feature configurations build. `cargo fmt` clean. ## Limitations - **Clip, Relu and SiLU still lose to ORT** at these sizes (0.135–0.864×). This PR is a 2.2–10.2× step toward the architectural requirement that our CPU EP beat ORT on every op it accepts; it does not finish the job, and no fallback was added. The remaining gap is the general 16-thread scaling gap tracked in `docs/performance/CPU_ACTIVATION_GAPS.md` — ORT scales these ops ~14× from 1→16 threads, we manage ~6×, because we split over our own rayon pool rather than ORT's intra-op pool. The `host_parallel` seam over `KernelContext_ParallelFor` is the next step. - 1 Mi is exactly `PAR_MIN_LEN`, so gains there are smaller and noisier than at 4 Mi. - 16-thread medians on this shared machine are noisy; untouched control ops swung up to 36% across 3 rounds. The 4 Mi wins are far outside that band. The 1 Mi Clip/Relu numbers are closer to it and should be read as directional. - `run_chunked`, `PAR_MIN_LEN` and `parallel_dispatches` are widened to `pub(crate)` because the three other callers live in sibling modules. --------- Co-authored-by: Deckard <deckard@users.noreply.github.com> 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.
What
Make the MLAS activation routes go through
run_chunked, the shared chunking/parallel seam every pure-Rust route already uses.One line:
Root cause
dispatch_mlas!called its kernel on the whole tensor and returned:run_chunkedwas the only thing that split work across the pool. So every MLAS-routed op ran single-threaded, regardless of how many threads the session had.This was harmless when it was written —
run_chunkedwas a plain loop then. #1105 taught it to parallelise abovePAR_MIN_LEN, and from that moment the two builds diverged: anmlas-off build scaled with the pool, anmlas-on build did not. Since #1115 the wheel shipsmlason x86_64, so the shipped configuration was the non-scaling one.Nothing caught it because the results were still correct, and correctness is all the existing tests checked.
thread_invarianceasserts a kernel gives the same answer serial and parallel — which a route that never goes parallel satisfies trivially.Benchmark
Same machine (AMD EPYC 9V74, 32 vCPU, AVX2+FMA),
ab_ort.py,--iters 30 --reps 5, medians of 3 interleaved rounds alternating the two builds, EP assignment asserted from ORT's own profiler (zeroNOT-ASSIGNED). Both arms built--features mlas --releasefrom the same base. p50 µs.16 threads — the ops that take the MLAS route
(
Tanh/Sigmoidmeasured on a base that predates #1124; they no longer take the MLAS route, but they exercised the same mechanism and are kept here as evidence of it.)16 threads — ops with no MLAS route, as the control
Unchanged within noise, which is what says the win is the routing change and not drift:
1 thread — no regression
run_chunkedreturns before touching rayon belowPAR_MIN_LEN, andpar_chunk_lendeclines to split when there is nothing to gain, so the single-threaded path is untouched:Every cell is inside run-to-run noise (<0.5%).
Correctness
run_chunkedsplits a slice into disjoint sub-slices and calls the same function on each. The MLAS activation entry points (MlasComputeErf,MlasComputeGeluErf) are elementwise and take no threadpool argument, so splitting cannot change a result: elementidepends only on inputi. The existingthread_invariance::unary_kernels_are_thread_count_invariantalready coverserfanderf_geluand asserts serial and parallel agree bit for bit — it now actually exercises the parallel branch for them, where before it compared serial against serial.Regression test
parallel_reachabilityasserts the mechanism rather than the output: a tensor overPAR_MIN_LEN, submitted from outside the pool, must incrementrun_chunked's parallel-branch counter.Proven to falsify — reverting just the one line makes it fail with:
while
pure_rust_kernels_go_through_run_chunkedkeeps passing, so it is not a blanket assertion that would pass on anything. It also guards the two preconditions that could make it vacuous (pool must have ≥2 threads; must not already be inside the pool).The counter is
#[cfg(test)]and thread-local, not a global atomic, so concurrently running tests cannot bump each other's count. Non-test builds get an empty#[inline(always)]function.Tests
cargo test -p onnx-runtime-ep-cpu --features mlas --lib→ 1312 passed, 0 failedcargo test -p onnx-runtime-ep-cpu --lib(nomlas) → 1294 passed, 0 failedcargo clippy -p onnx-runtime-ep-cpu --features mlas --lib --tests→ cleancargo fmt --all -- --check→ cleanA pre-existing intermittent crash, ruled out as mine
While validating I hit an intermittent
SIGSEGVin the fullonnx-runtime-ep-cpulib test binary. It is not from this change:origin/main: 1 segfault.origin/mainalso failskernels::sdpa::tests::sdpa_dispatch_matches_scalar_oracle_across_shapesintermittently (~1 in 8 full-suite runs), unrelated to activations.Flagging for whoever owns
sdpa/qmoe; not chased here as it is outside this PR's scope and reproduces without it.Limitations
PAR_MIN_LEN= 1 Mi andPAR_MIN_CHUNK= 256 Ki. Parallelise the f32 elementwise activation kernels #1105 recorded that one global constant cannot serve kernels with different per-element costs, and this PR does not change that — it only makes the MLAS routes obey whatever the constants say. Both are now on the same footing, which is the precondition for tuning them per kernel.Erfgoes from 0.11x to 0.66x of ORT andGelufrom 0.09x to 0.42x. Real wins, still short of parity, and the remaining multi-thread gap stays open.Review
Independent Opus review: GO WITH FINDINGS.
The reviewer independently confirmed the parts that could have been wrong: that
run_chunkedsplits into disjoint sub-slices at identical input/output offsets (soerf_gelu_mlas's internal 8192-element blocking and its(x, y)-paired NaN repair still line up inside a chunk, trailing partial block included); thatMlasComputeErf/MlasComputeGeluErfare elementwise with no threadpool argument and no shared mutable state, so concurrent calls from rayon workers are sound; that theSIMD_MIN_LEN(32) /PAR_MIN_LEN(1 Mi) composition has no gap and short tensors still return before touching rayon; that the counter increments on the calling thread and cannot be cross-contaminated by concurrent tests; thatnote_parallel_dispatchcompiles to nothing in a non-test build; and that the speedup arithmetic and the "serial against serial" claim are accurate.Two findings:
(MAJOR, pre-existing, outside this file) The same pathology exists in three sibling elementwise activations.
silu_f32_slice(kernels/activations.rs:332),relu_contiguous_f32_mlas(kernels/relu.rs:144) and the two Clip paths (kernels/selection.rs:182,kernels/conv.rs:732) all callmlas_sys::compute_*on the whole tensor with no chunking seam — andmlas-sysdocuments those entry points as "Single threaded; callers shard across threads themselves".silu_f32_sliceis worse than the ops fixed here: it follows the MLAS call with a second full-tensor serial correction loop. SiLU is the SwiGLU activation, so this is a hot path.Confirmed by grep:
run_chunked/run_chunked_rowsexist only insimd_activations.rs, and every entry point inside that file does route through the seam — so this PR is complete for its stated scope. Being a different subsystem with its own correctness wrinkle (the correction pass), it gets its own PR rather than being folded in here.(NIT)
parallel_reachabilityhard-fails rather than skips on a single-threaded pool. This is deliberate and mirrors the pre-existingthread_invariance::assert_sameguard, so it adds no environmental assumption the suite did not already make — a test that silently skipped would be worse, since the bug it guards is invisible in the output.