Repository navigation
Parallelise the f32 elementwise activation kernels - #1105
Merged
Merged
Conversation
Every activation in `simd_activations.rs` ran on one core regardless of tensor size. At prefill widths that left most of the machine idle: a 512x4096 activation is 2M independent elements and a handful of cycles each, so it is pure throughput work with nothing to serialise it. Route the `dispatch!` macro, `quick_gelu_f32_slice` and the two fused bias kernels through a chunked runner. Measured on this 32-vCPU EPYC, same binary at 1 vs 16 threads, median of three interleaved repeats: Erf/f32/prefill512x4096 1.5865 -> 0.2552 ns/elem 6.22x Sqrt/f32/prefill512x4096 0.5768 -> 0.0988 5.84x FastGelu/f32/prefill512x4096 1.1671 -> 0.2149 5.43x Gelu/f32/prefill512x4096 1.1655 -> 0.2192 5.32x QuickGelu/f32/prefill512x4096 0.9693 -> 0.1870 5.18x Sigmoid/f32/prefill512x4096 0.7823 -> 0.1673 4.68x Tanh/f32/prefill512x4096 0.7874 -> 0.1746 4.51x Decode and small shapes are unchanged (0.93-1.14x, inside this bench's noise band): the length check runs before rayon is touched at all, so short calls never reach the pool. Three things this had to get right. The length test must precede `rayon::current_num_threads()`. That call reaches the global registry and initialising it spawns the pool, ~1.4us, which is more than a whole 4096-element activation. An earlier revision measured a uniform 0.6-0.8x regression on every short case until the check moved above it. Chunking must not change results. The vector and scalar paths round differently, so the path is chosen once for the whole slice and every chunk inherits it; chunks are floored at `PAR_MIN_CHUNK` and rounded to whole vectors so none can fall out of the vector path. The bias kernels index `bias[i % width]`, so their chunks are whole multiples of `width` -- a mid-row cut would rotate the bias for everything after it. f16/bf16 are deliberately left serial. They widen into an f32 scratch, compute, then narrow back, and parallelising only the middle layer measured *slower*: Sqrt/f16 0.59x, Tanh/f16 0.79x, with bf16 prefill swinging 1.6-3.4 ns/elem across repeats where f32 held to +/-5%. Spreading 8MB of scratch across sixteen private caches for a serial narrow to pull back costs more than the arithmetic saves. The narrow arm of `write_mapped_reading` now runs under `serial_scope`, which pins those paths to their previous behaviour (measured 0.98-1.01x). Fusing widen/compute/narrow per chunk is the real fix and belongs in its own change. Tests assert bit-identical output between a one-thread pool and the multi-threaded global pool for every affected kernel, over row widths chosen to be coprime with the lane count. The chunk policy is a pure function so it is swept separately over thread counts and lengths this host cannot produce. All three guards were verified to fail when the corresponding invariant is broken: `+1` on the row chunk gives "chunk 8194 cuts row width 3 in half" and a real numeric divergence at element 8194, and dropping the vector floor gives "chunk 4096 could drop below the vector threshold". Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #1105 +/- ##
==========================================
+ Coverage 79.91% 80.53% +0.61%
==========================================
Files 368 366 -2
Lines 159662 157044 -2618
Branches 159662 157044 -2618
==========================================
- Hits 127594 126474 -1120
+ Misses 27351 25865 -1486
+ Partials 4717 4705 -12
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
|
Review noted the guard was restored by a plain assignment after the closure returned, so an unwind would leave it set. That is not a correctness problem -- serial and parallel results are bit-identical -- which is exactly what makes it worth fixing: the thread would just run every later activation serially, silently, with nothing wrong in the output to point at it. Restore from a `Drop` guard instead. The test that documented the old leak now asserts the opposite, and silences the panic hook so the expected panic does not print a scary backtrace in a passing run. A second test covers nesting, which the assignment form also got right but which nothing pinned. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
The thresholds in this PR were chosen from a benchmark loop, and a benchmark loop is the one workload where they are safe. It calls the kernel back-to-back, so rayon's workers never park and the fork/join looks nearly free. A real ORT session does the opposite: one short burst per node, with the rest of the graph in between, so almost every call wakes the pool from scratch. Measured through a real ORT session -- same graph, same input, our EP against ORT's CPU EP, one intra-op thread each, p50 of 200 interleaved runs -- the old 16 Ki threshold was not a small loss but a 5x one: | elements | PAR_MIN_LEN 16 Ki | PAR_MIN_LEN 1 Mi | |---|---|---| | 16384 | 0.19x | 1.20x | | 32768 | 0.17x | 1.43x | | 65536 | 0.16x | 1.60x | | 131072 | 0.21x | 1.68x | | 262144 | 0.98x | 1.80x | Sqrt/f32, which is the cleanest case because it has no fix-up pass. The wake-up costs ~50 us; the work below a megabyte is worth less than that. So the split now starts at 1 Mi with a 256 Ki floor per chunk, which keeps every size where it wins and hands the rest back to the serial path that was already faster than ORT. This does not change any result: the split is still exactly chunk-independent, and thread_invariance still asserts bit-identity. cargo test -p onnx-runtime-ep-cpu --lib: 1319 passed. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Re-review pointed out that PAR_MIN_LEN was derived from Sqrt -- the cheapest, most bandwidth-bound kernel here -- and then applied to every kernel, and that break-even scales with per-element cost, so the transcendentals are probably under-threaded. Measured, and they are. At 16 intra-op threads, lowering the threshold to 256 Ki makes Gelu 1.26-2.35x faster across 256 Ki - 2 Mi, Erf 1.23-1.78x and FastGelu 1.28-2.03x -- while making Sqrt at 256 Ki 2.3x slower, which is exactly why the constant is where it is. The fix is a per-kernel cost class, which needs plumbing through four generic entry points; this commit records the measurement and the reasoning next to the constant rather than guessing at a compromise value. Too high costs throughput; too low cost 5x, so the conservative value stays for now. Also corrects two documentation inaccuracies found in the same review: the thread configuration behind the headline table is now stated exactly (both sides at intra_op_num_threads=1, ours additionally at RAYON_NUM_THREADS=32, so it is our parallel kernel against single-threaded ORT), and the thread_invariance N comment no longer claims coprimality with every row width when N is divisible by 3. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Both branches appended a test module at end of file, which git spliced together. Resolved by keeping #1111's mlas_ab module intact and re-appending thread_invariance whole, so neither the MLAS dispatch code nor the chunking invariance tests are lost. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby
added a commit
that referenced
this pull request
Aug 17, 2026
…1121) ## What this is `tanh_ps` and `sigmoid_ps` are the two AVX2 primitives most of the CPU activation family is built on: `Tanh`, `Sigmoid`, SiLU / Swish, `QuickGelu` and `FastGelu` all end up in one of them. Both already carry the Eigen rational that MLAS uses — the polynomial was never the problem. What they also carried was a redundant saturation step, and it is not free. The shape of both kernels is: clamp the *input* to the rational's valid range (`±9` for `tanh`, `±18` for `logistic`), evaluate `P(v²)·v / Q(v²)`, clamp the *output* to the function's mathematical range. On top of that, both then ran a second saturation: ```rust let above = _mm256_cmp_ps(x, _mm256_set1_ps(tanh_c::UPPER), _CMP_GT_OQ); let below = _mm256_cmp_ps(x, _mm256_set1_ps(tanh_c::LOWER), _CMP_LT_OQ); let r = _mm256_blendv_ps(poly, _mm256_set1_ps(1.0), above); _mm256_blendv_ps(r, _mm256_set1_ps(-1.0), below) ``` Four vector ops, two of them `vblendvps`, which is two uops on Zen. **It never changes a bit.** The input clamp happens *first*, so the largest argument the rational ever sees is the clamp point itself — and there it has already reached the output bound: | | at the clamp point | output clamp gives | |---|---|---| | `tanh`, `v = 9` | `p` and `q` are **bit-equal** (`0x3fcd33e9`), so `p/q` is exactly `1.0` | `1.0` | | `tanh`, `v = -9` | exactly `-1.0` | `-1.0` | | `logistic`, `v = 18` | `p/q + 0.5` is exactly `1.0` | `1.0` | | `logistic`, `v = -18` | `p/q + 0.5` is `-5.96e-8` | `0.0` | Three of the four are *equality*, not slack, and the inclusive `minps`/`maxps` clamp passes equality through unchanged. (The rational's overshoot to `1.0000001` happens strictly *inside* the range, near `|v| = 8.9999971` — that is what the output clamp is there for, and it is unrelated to saturation.) Every input past the range, `±Inf` included, therefore already exits with the saturated value. The blends were re-deriving a result the clamp had produced two instructions earlier. ## Proof, not argument `saturation_blend_is_redundant_exhaustively` keeps the deleted sequence as a reference implementation and checks the shipped kernels against it over **every finite `f32` beyond the clamp, both signs** — 1 047 527 424 values for `tanh` and 1 039 138 816 for `sigmoid`, 2 086 666 240 in total — plus `±Inf`, both signed `NaN`s, `f32::MAX`/`MIN`, and both clamp boundaries at ±2 ULP. Bit-identical everywhere. Not within tolerance — identical. The sweep runs under the default rounding mode. The property was also checked by hand under round-down, round-up and round-toward-zero, and holds in all three; only round-to-nearest is exercised in CI. It runs in 19 s in release and is `#[ignore]`d because an unoptimised test build takes minutes. Two non-ignored tests cover the decade above each boundary (32.5 M and 21.7 M values) and the special values, and run in 12 s in the default `cargo test` profile, so the property stays guarded on every CI run. ## Also The three `vblendvps`-against-zero that pin `-Inf → 0` in `tanh_gelu_ps`, `quick_gelu_ps` and `erf_gelu_ps` become `vandnps`. Blending against zero *is* an `andnot` when the mask is all-ones or all-zeros, which is exactly what `vcmpps` produces. One uop instead of two, same bits. ## Measured Same-machine, alternating build A/B through a **real ORT session** — one node, one input, `intra_op_num_threads=1`, `RAYON_NUM_THREADS=1`, p50 of 30 runs, median of 3 alternating rounds per build, with our EP's assignment asserted from ORT's own profiler on every row. AMD EPYC 9V74, AVX2+FMA, ORT 1.28.0. Measured on top of #1097 and #1105, because on `main` today the assignment policy declines these ops to ORT's CPU EP, so our kernel never runs and the measurement is vacuous. (The first run of this A/B *was* vacuous for exactly that reason — both columns were ORT to within 0.1%. The harness's assignment check is what caught it.) | op | elements | before, us | after, us | **speedup** | |---|---|---|---|---| | `Tanh` | 16384 | 13.49 | 11.71 | **1.15x** | | `Tanh` | 65536 | 38.33 | 31.65 | **1.21x** | | `Tanh` | 262144 | 131.87 | 105.48 | **1.25x** | | `Tanh` | 1048576 | 533.18 | 438.21 | **1.22x** | | `Tanh` | 4194304 | 2272.02 | 1889.22 | **1.20x** | | `Sigmoid` | 16384 | 13.91 | 11.94 | **1.17x** | | `Sigmoid` | 65536 | 39.96 | 35.27 | **1.13x** | | `Sigmoid` | 262144 | 143.08 | 123.31 | **1.16x** | | `Sigmoid` | 1048576 | 542.67 | 408.28 | **1.33x** | | `Sigmoid` | 4194304 | 2052.81 | 1883.35 | **1.09x** | | `FastGelu` | 65536 | 66.13 | 55.53 | **1.19x** | | `FastGelu` | 262144 | 244.24 | 201.70 | **1.21x** | | `FastGelu` | 1048576 | 969.63 | 782.34 | **1.24x** | | `FastGelu` | 4194304 | 3994.51 | 3419.07 | **1.17x** | | `QuickGelu` | 65536 | 49.36 | 42.27 | **1.17x** | | `QuickGelu` | 262144 | 176.89 | 148.68 | **1.19x** | | `QuickGelu` | 4194304 | 3016.00 | 2438.27 | **1.24x** | | exact `Gelu` | 65536 | 111.12 | 108.26 | 1.03x | | exact `Gelu` | 4194304 | 6898.34 | 6753.38 | 1.02x | | `Erf` *(control)* | 65536 | 86.78 | 86.88 | 1.00x | | `Erf` *(control)* | 1048576 | 1286.59 | 1286.52 | 1.00x | | `Sqrt` *(control)* | 65536 | 21.63 | 21.38 | 1.01x | | `Sqrt` *(control)* | 1048576 | 244.85 | 244.64 | 1.00x | `Erf` and `Sqrt` touch neither primitive and are flat, which is the check that the rest is real and not drift. Exact `Gelu` goes through `erf_ps`, so it only collects the `andnot`, and moves ~2-3%. ## Against ORT Same runs, ORT's CPU EP as the control, ORT time over ours — above `1.00` we win: | op | 4096 | 16384 | 65536 | 262144 | 1048576 | 4194304 | |---|---|---|---|---|---|---| | `Tanh` | 0.73 -> 0.77 | 0.72 -> 0.83 | 0.72 -> 0.87 | 0.76 -> 0.95 | 0.88 -> **1.08** | 0.96 -> **1.16** | | `Sigmoid` | 0.75 -> 0.81 | 0.73 -> 0.85 | 0.74 -> 0.84 | 0.75 -> 0.87 | 0.70 -> 0.92 | 0.73 -> 0.79 | | `QuickGelu` | 0.83 -> 0.88 | 0.89 -> **1.03** | 0.97 -> **1.14** | 1.03 -> **1.23** | 1.05 -> **1.18** | 1.02 -> **1.26** | | `FastGelu` | 0.73 -> 0.79 | 0.68 -> 0.76 | 0.67 -> 0.80 | 0.67 -> 0.81 | 0.66 -> 0.81 | 0.72 -> 0.84 | | exact `Gelu` | 0.71 -> 0.72 | 0.67 -> 0.70 | 0.67 -> 0.69 | 0.69 -> 0.71 | 0.69 -> 0.71 | 0.71 -> 0.73 | **This does not claim we now beat ORT.** It moves every one of these families toward it, takes `Tanh` past it at >=1 Mi and `QuickGelu` past it from 16 Ki up, and leaves `FastGelu`, exact `Gelu` and `Erf` still behind. The remaining gap is not in these two primitives, and the next steps are elsewhere: `erf_ps`'s `exp_ps` tail, and the ~1.7-2.7 us fixed per-node plugin overhead that dominates at 4096 elements. ## Direction This is an absorption in the sense of `docs/performance/ABSORBING_MLAS.md`: the win lands in the native kernel, in the **default build**, with no `mlas` feature and no dependency. Reading MLAS's `tanh.cpp` and `logistic.cpp` is what made the redundancy visible — MLAS does not have this step, because it does not promise the output range we promise, and comparing the two instruction sequences is what showed that our extra promise costs nothing to keep and four ops to re-state. ## Correctness - `saturation_blend_is_redundant_exhaustively` — 2 086 666 240 values, bit-identical. - Two boundary sweeps + special values, in the default test profile. - Full `onnx-runtime-ep-cpu` lib suite: 1326 passed. - No tolerance was relaxed and no reference output changed; every pre-existing `tanh`/`sigmoid`/`gelu` accuracy test passes unmodified. ## Limitations - x86-64 AVX2+FMA only. The scalar and NEON paths are untouched. - The A/B was run on one host. The instruction-count argument is machine-independent; the exact percentages are not. - The `1048576` and `4194304` rows sit where ORT's own timing is least stable (its control column moved up to 1.6x between rounds on `Tanh`); medians of three alternating rounds are reported, and the `Erf`/`Sqrt` controls are the evidence that the reported deltas are larger than that drift. ## Independent review Reviewed by an independent Opus reviewer, verdict **GO WITH FINDINGS**. The reviewer re-derived the exhaustive proof independently (2 095 054 848 `tanh` and 2 078 277 632 `sigmoid` values, zero disagreements), confirmed `andnot ≡ blendv` across quiet, signalling and non-canonical `NaN` payloads, `±0`, `±Inf` and subnormals, and confirmed the redundancy holds under all four MXCSR rounding modes. All findings are applied: 1. **The margin claim was wrong, and it was load-bearing.** The comment and this body said `p/q` at `v = 9` is `1.0000001`. It is exactly `1.0` — `p` and `q` are the same bit pattern — so the safety argument rests on *equality with an inclusive clamp*, not on slack, and the old wording also contradicted a correct comment fifteen lines above it. Both the comment and the body now state the exact values, and the sigmoid comment no longer claims the result lands "outside `[0, 1]` on both ends" when at `+18` it is exactly on the boundary. Corrected in `d32e52b`. 2. Signalling and non-canonical-payload `NaN`s added to the special-value test. 3. Rounding-mode assumption noted in the kernel comment and above. 4. The reviewer's remaining point is that this PR is necessary but not sufficient: on `main` the assignment policy declines these ops, so the kernels are not reached until #1097 lands. That is why the benchmark is stacked, and it is stated in the section above. --------- Co-authored-by: Deckard <deckard@users.noreply.github.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Same end-of-file module collision as the previous merge, now also against #1121's saturation_absorption module. Resolved the same way: main's modules kept intact, thread_invariance re-appended whole. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby
marked this pull request as ready for review
August 17, 2026 11:56
justinchuby
added a commit
that referenced
this pull request
Aug 17, 2026
## What
Make the MLAS activation routes go through `run_chunked`, the shared
chunking/parallel seam every pure-Rust route already uses.
One line:
```rust
- f(input, output);
+ run_chunked(input, output, f);
```
## Root cause
`dispatch_mlas!` called its kernel on the whole tensor and returned:
```rust
if input.len() >= SIMD_MIN_LEN {
let f: fn(&[f32], &mut [f32]) = $mlas;
f(input, output); // <-- straight past run_chunked
return;
}
```
`run_chunked` was 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_chunked` was a plain loop
then. #1105 taught it to parallelise above `PAR_MIN_LEN`, and from that
moment the two builds diverged: an `mlas`-off build scaled with the
pool, an `mlas`-on build did not. Since #1115 the wheel ships `mlas` on
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_invariance`
asserts 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 (zero
`NOT-ASSIGNED`). Both arms built `--features mlas --release` from the
same base. p50 µs.
### 16 threads — the ops that take the MLAS route
| op | n | before | after | speedup |
|---|---:|---:|---:|---:|
| Erf | 1 Mi | 897.3 | **393.9** | **2.28x** |
| Erf | 4 Mi | 3558.7 | **561.6** | **6.34x** |
| Gelu (exact) | 1 Mi | 1304.0 | **552.8** | **2.36x** |
| Gelu (exact) | 4 Mi | 5458.7 | **1134.3** | **4.81x** |
| Tanh | 1 Mi | 480.8 | **300.4** | **1.60x** |
| Tanh | 4 Mi | 2219.6 | **425.9** | **5.21x** |
| Sigmoid | 1 Mi | 503.5 | **298.6** | **1.69x** |
| Sigmoid | 4 Mi | 1991.6 | **430.0** | **4.63x** |
(`Tanh`/`Sigmoid` measured 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:
| op | n | before | after |
|---|---:|---:|---:|
| FastGelu | 1 Mi | 413.9 | 395.5 |
| FastGelu | 4 Mi | 600.2 | 605.4 |
| QuickGelu | 1 Mi | 329.7 | 347.1 |
| QuickGelu | 4 Mi | 529.1 | 545.9 |
| Sqrt | 1 Mi | 255.3 | 232.3 |
| Sqrt | 4 Mi | 352.8 | 375.1 |
### 1 thread — no regression
`run_chunked` returns before touching rayon below `PAR_MIN_LEN`, and
`par_chunk_len` declines to split when there is nothing to gain, so the
single-threaded path is untouched:
| op | n | before | after |
|---|---:|---:|---:|
| Erf | 4 Mi | 3628.0 | 3643.8 |
| Gelu | 4 Mi | 5466.7 | 5479.0 |
| Erf | 65 Ki | 62.2 | 63.5 |
| Gelu | 65 Ki | 86.7 | 86.4 |
Every cell is inside run-to-run noise (<0.5%).
## Correctness
`run_chunked` splits 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: element `i` depends only on input
`i`. The existing
`thread_invariance::unary_kernels_are_thread_count_invariant` already
covers `erf` and `erf_gelu` and 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_reachability` asserts the mechanism rather than the output: a
tensor over `PAR_MIN_LEN`, submitted from outside the pool, must
increment `run_chunked`'s parallel-branch counter.
Proven to falsify — reverting just the one line makes it fail with:
```
erf: a 1052675-element call did not reach run_chunked's parallel branch, so it
runs single-threaded no matter how large the pool is. A kernel that calls its
backend directly and returns will fail here.
```
while `pure_rust_kernels_go_through_run_chunked` keeps 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 failed**
- `cargo test -p onnx-runtime-ep-cpu --lib` (no `mlas`) → **1294 passed,
0 failed**
- `cargo clippy -p onnx-runtime-ep-cpu --features mlas --lib --tests` →
clean
- `cargo fmt --all -- --check` → clean
### A pre-existing intermittent crash, ruled out as mine
While validating I hit an intermittent `SIGSEGV` in the full
`onnx-runtime-ep-cpu` lib test binary. It is **not** from this change:
- Activation tests alone, 40 consecutive runs on this branch: **0
failures**.
- Full suite, 14 consecutive runs on this branch: **0 segfaults**.
- Full suite, 14 consecutive runs on pristine `origin/main`: **1
segfault**.
- `origin/main` also fails
`kernels::sdpa::tests::sdpa_dispatch_matches_scalar_oracle_across_shapes`
intermittently (~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
- The thresholds are still `PAR_MIN_LEN` = 1 Mi and `PAR_MIN_CHUNK` =
256 Ki. #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.
- This closes a self-inflicted gap; it does not by itself make these ops
beat ORT at 16 threads. At 4 Mi, `Erf` goes from 0.11x to 0.66x of ORT
and `Gelu` from 0.09x to 0.42x. Real wins, still short of parity, and
the remaining multi-thread gap stays open.
- Measured on one 32-vCPU host.
## Review
Independent Opus review: **GO WITH FINDINGS**.
The reviewer independently confirmed the parts that could have been
wrong: that `run_chunked` splits into disjoint sub-slices at identical
input/output offsets (so `erf_gelu_mlas`'s internal 8192-element
blocking and its `(x, y)`-paired NaN repair still line up inside a
chunk, trailing partial block included); that
`MlasComputeErf`/`MlasComputeGeluErf` are elementwise with no threadpool
argument and no shared mutable state, so concurrent calls from rayon
workers are sound; that the `SIMD_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; that `note_parallel_dispatch`
compiles to nothing in a non-test build; and that the speedup arithmetic
and the "serial against serial" claim are accurate.
Two findings:
1. *(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 call
`mlas_sys::compute_*` on the whole tensor with no chunking seam — and
`mlas-sys` documents those entry points as "Single threaded; callers
shard across threads themselves". `silu_f32_slice` is 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_rows` exist only in
`simd_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.
2. *(NIT)* `parallel_reachability` hard-fails rather than skips on a
single-threaded pool. This is deliberate and mirrors the pre-existing
`thread_invariance::assert_same` guard, 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.
---------
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
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
Every activation in
simd_activations.rsran on a single core no matter how large the tensor was. At prefill widths that leaves most of the machine idle: a 512x4096 activation is 2M independent elements at a few cycles each — pure throughput work with nothing forcing it to be serial.This routes the
dispatch!macro,quick_gelu_f32_slice, and the two fused bias kernels (tanh_gelu_bias_f32_slice,erf_gelu_bias_f32_slice) through a chunked runner backed by the rayon pool that the GEMM kernels already use.Measured
Same binary,
RAYON_NUM_THREADS=1vs=16, median of three interleaved repeats, AMD EPYC 9V74 (32 vCPU, avx2/f16c/fma). ns/element:No regressions anywhere else: decode4096 0.98–1.02x, decode3072 0.99–1.03x, small64 0.93–1.14x, tiny16 0.95–1.01x — all inside this bench's noise band. Short calls test their length before rayon is touched at all, so they never reach the pool.
Methodology — why this is a same-binary comparison
activation_bench.rs's own header documents a uniform 0.70–0.82x offset on byte-identical kernels between separately-built binaries. Interleaving does not remove it. An earlier revision of this change measured a "uniform regression" that was entirely this artifact — proven because untouchedQuickGeluregressed by the same factor.So the A/B here varies only the thread count within one binary. In the first run
QuickGeluwas left serial deliberately and came out at exactly 1.00x, confirming the harness attributes nothing to an unchanged kernel; it is parallelised in this PR and now scales at 5.18x.These are self-relative figures. They are not a claim against ORT — ORT parallelises too, and a direct matched-thread A/B against ORT is separate follow-up work.
The thresholds were wrong, and only a real session showed it
The table above is a self-relative measurement, and self-relative is
exactly the measurement that cannot see this bug. A benchmark loop calls the
kernel back-to-back, so rayon's workers never park and the fork/join looks
nearly free. A real ORT session does the opposite: one short burst per node,
with the rest of the graph in between, so almost every call wakes the pool from
scratch — measured at roughly 50 us.
Run through a real ORT session (same single-node graph, same input, our EP
against ORT's CPU EP, matched intra-op threads, p50 of 200 interleaved runs,
EP assignment asserted from ORT's own profiler), the originally proposed
PAR_MIN_LENof 16 Ki was not a small loss but a 5x one:PAR_MIN_LEN= 16 KiPAR_MIN_LEN= 1 Mi(
Sqrt/f32 — the cleanest case, because it has no fix-up pass. Higher isbetter; these are ORT time over our time.) Every one of those sizes was
already faster than ORT on the serial path. The split took a 1.2-1.8x win and
turned it into a 5x loss, on exactly the sizes a decode step uses.
So the thresholds are now
PAR_MIN_LEN = 1 MiandPAR_MIN_CHUNK = 256 Ki:at ~50 us of wake-up and ~0.3 ns/element, break-even is around 160 Ki elements
per chunk. Below 1 Mi the kernel stays on the serial path, which the table
above shows is where it belongs.
What this PR does and does not claim against ORT
With the corrected thresholds this change is a strict improvement on our own
serial path wherever it engages, and it no longer makes any size slower than it
was. It does not yet make us faster than ORT at high thread counts: at
sixteen matched threads ORT still scales better than a rayon pool can from
inside a plugin EP, because ORT's intra-op pool is already hot when the node
starts and ours is not. Measured at 4 Mi,
Tanh/f32, sixteen threads: ORT202 us, this PR 396 us, our serial path 1977 us. So the split is worth having —
it is 5x better than not splitting — and it is still not enough.
Closing that gap needs the host's pool rather than one of our own, via
OrtApi::KernelContext_ParallelFor. A prototype of that is measured and is alarge improvement over rayon at prefill sizes (
Tanh4 Mi: 816 us rayon ->307 us host pool), but it carries a ~40-105 us fixed cost per call of its own
that has to be understood before it can ship. That is separate follow-up work
and is not in this PR.
Three things this had to get right
Do not touch rayon before checking the length.
rayon::current_num_threads()reaches the global registry, and initialising it spawns the pool — about 1.4 µs, more than an entire 4096-element activation. An earlier revision measured a uniform 0.6–0.8x regression on every short case until that check moved above it.Chunking must not change numerics. The vector and scalar paths round differently, so the path is chosen once for the whole slice and every chunk inherits it. Chunks are floored at
PAR_MIN_CHUNKand rounded to whole vectors so none can fall out of the vector path. The bias kernels indexbias[i % width], so their chunks are whole multiples ofwidth— a mid-row cut would rotate the bias for everything after it.f16/bf16 are deliberately left serial. They widen into an f32 scratch, compute, then narrow back, and parallelising only the middle layer measured slower: Sqrt/f16 0.59x, Tanh/f16 0.79x, with bf16 prefill swinging 1.6–3.4 ns/element across repeats where f32 held to ±5%. Spreading 8 MB of scratch across sixteen private caches for a serial narrow to pull back costs more locality than the arithmetic saves. Parallelising the bulk conversions too was tried and was worse still. The narrow-output arm of
write_mapped_readingnow runs underserial_scope, pinning those paths to their previous behaviour — measured 0.98–1.01x.Fusing widen/compute/narrow into one pass per chunk is the real fix for f16/bf16 and belongs in its own change, measured on its own.
Correctness
Parallel output must be bit-identical to serial, not merely close — these kernels are exactly chunk-independent, so anything else is a bug.
unary_kernels_are_thread_count_invariant,quick_gelu_is_thread_count_invariant,bias_kernels_are_thread_count_invariantcompare a one-thread pool against the multi-threaded global pool over row widths (1, 3, 7, 11, 64, 4096, 4099) chosen to be coprime with the lane count and not to divide the chunk size.chunk_policy_holds_across_thread_counts_and_lengthssweeps it over thread counts and lengths this host cannot produce (up to 4096 threads, 4M elements).serial_scope_suppresses_the_splitandserial_scope_is_restored_after_a_panicpin the f16/bf16 guard, including documenting that it is not unwind-safe.A note on how the parallel side is obtained:
ThreadPool::installruns its closure on a pool worker, which trips the nesting guard and silently serialises. An earlier version of these tests usedinstallfor both sides and passed with the row alignment deliberately broken. The parallel run is now a plain direct call.All three guards were verified to fail when their invariant is broken:
+1on the row chunk →rows: chunk 8194 cuts row width 3 in half, and a real numeric divergence at element 8194.chunk 4096 could drop below the vector threshold.cargo fmt --check, scoped clippy, and 1316onnx-runtime-ep-cpulib tests are green.Limitations
large improvement on our own serial path, not a win against ORT.
PAR_MIN_LEN(1 Mi) andPAR_MIN_CHUNK(256 Ki) aredeliberately conservative, and the sub-threshold path is provably identical
to before — which is now most of the range.
Any pool we own rather than borrow will have it.
Co-authored-by: Copilot 223556219+Copilot@users.noreply.github.com
Independent re-review
The thresholds changed after the first review's GO, so it was re-reviewed.
Verdict GO WITH FINDINGS. The reviewer independently confirmed the parallel
path is sound (rayon's safe
par_chunks_mut/par_chunkszip, no manualSend/Sync, no raw pointers, path chosen once for the whole slice so no chunkcan drop to scalar), confirmed the bias kernels' chunks always start on a row
boundary, and confirmed the invariance test's chunk boundaries really are
adversarial (262 146 and 262 336 — not multiples of 8 — so a mid-row cut would
be caught). 1331 tests pass. No interaction with #1097's rule: these constants
are referenced only inside
simd_activations.rsand change how our kernelruns, never whether the node is ours.
Findings, and what was done:
One global threshold, derived from the cheapest kernel, applied to all of
them. The reviewer's point was that break-even scales with per-element
cost, so
Erfand exactGelu— 2-3x more work per element thanSqrt—are likely under-threaded between 256 Ki and 1 Mi. Measured, and correct.
At 16 intra-op threads, dropping
PAR_MIN_LENto 256 Ki gives:GeluFastGeluErfSqrtSo the reviewer is right for the transcendentals and the current value is
right for
Sqrt, which is the kernel that would be destroyed by a lower one.A single compromise value cannot serve both; the fix is a per-kernel cost
class, which needs plumbing through four generic entry points and is left to
a follow-up. This PR records the measurement and the reasoning beside the
constant (
967191438) instead of guessing, and keeps the conservative value:too high costs throughput, too low cost 5x.
The 256 Ki chunk floor caps a 1 Mi tensor at four workers. True, and it
is the same trade-off as finding 1 — it follows the cost class, so it is
tracked with it.
Thread configuration of the headline table was ambiguous. Fixed: the
doc-comment now states that both sides ran at
intra_op_num_threads = 1withour pool at
RAYON_NUM_THREADS = 32, so the "1.8x serial win" is againstsingle-threaded ORT, not multi-threaded ORT.
Ninthread_invarianceis divisible by 3, contradicting a commentclaiming coprimality with every row width. Comment corrected to say where the
adversarial coverage actually comes from.
What this PR does not fix
At 16 intra-op threads our elementwise kernels remain far behind ORT —
0.06-0.50x across these families, on both threshold settings. ORT scales these
ops ~14x from 1 to 16 threads; we manage ~6x. That gap is not a threshold
problem and is not addressed here; it is a property of owning a pool that has
to be woken per node, and the next step is to use the host's pool instead of
our own. Recorded so the merged state is not mistaken for a solved one.