Skip to content

Parallelise the f32 elementwise activation kernels - #1105

Merged
justinchuby merged 8 commits into
mainfrom
deckard/parallel-elementwise
Aug 17, 2026
Merged

justinchuby merged 8 commits into
mainfrom
deckard/parallel-elementwise

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 17, 2026 •

Copy link
Copy Markdown
Owner

What

Every activation in simd_activations.rs ran 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=1 vs =16, median of three interleaved repeats, AMD EPYC 9V74 (32 vCPU, avx2/f16c/fma). ns/element:

case (2,097,152 elems) 1 thread 16 threads speedup
Erf/f32/prefill512x4096 1.5865 0.2552 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

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 untouched QuickGelu regressed by the same factor.

So the A/B here varies only the thread count within one binary. In the first run QuickGelu was 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_LEN of 16 Ki 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 — the cleanest case, because it has no fix-up pass. Higher is
better; 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 Mi and PAR_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: ORT
202 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 a
large improvement over rayon at prefill sizes (Tanh 4 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_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/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_reading now runs under serial_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_invariant compare 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.
  • The chunk policy is a pure function, so chunk_policy_holds_across_thread_counts_and_lengths sweeps it over thread counts and lengths this host cannot produce (up to 4096 threads, 4M elements).
  • serial_scope_suppresses_the_split and serial_scope_is_restored_after_a_panic pin the f16/bf16 guard, including documenting that it is not unwind-safe.

A note on how the parallel side is obtained: ThreadPool::install runs its closure on a pool worker, which trips the nesting guard and silently serialises. An earlier version of these tests used install for 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:

  • +1 on the row chunk → rows: chunk 8194 cuts row width 3 in half, and a real numeric divergence at element 8194.
  • dropping the vector floor → chunk 4096 could drop below the vector threshold.

cargo fmt --check, scoped clippy, and 1316 onnx-runtime-ep-cpu lib tests are green.

Limitations

  • f16/bf16 unchanged by design (see above).
  • Does not reach ORT's scaling at high thread counts (see above). It is a
    large improvement on our own serial path, not a win against ORT.
  • Tuned on one host. PAR_MIN_LEN (1 Mi) and PAR_MIN_CHUNK (256 Ki) are
    deliberately conservative, and the sub-threshold path is provably identical
    to before — which is now most of the range.
  • The wake-up cost is a property of a parked pool, not of rayon specifically.
    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_chunks zip, no manual
Send/Sync, no raw pointers, path chosen once for the whole slice so no chunk
can 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.rs and change how our kernel
runs, never whether the node is ours.

Findings, and what was done:

  1. 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 Erf and exact Gelu — 2-3x more work per element than Sqrt —
    are likely under-threaded between 256 Ki and 1 Mi. Measured, and correct.
    At 16 intra-op threads, dropping PAR_MIN_LEN to 256 Ki gives:

    op 262144 524288 1048576 2097152 4194304
    exact Gelu 1.94x 2.35x 1.58x 1.26x 0.96x
    FastGelu 1.32x 2.03x 1.33x 1.28x 0.92x
    Erf 1.28x 1.78x 1.52x 1.23x 0.79x
    Sqrt 0.43x 0.72x 1.00x 1.16x 0.97x

    So 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.

  2. 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.

  3. Thread configuration of the headline table was ambiguous. Fixed: the
    doc-comment now states that both sides ran at intra_op_num_threads = 1 with
    our pool at RAYON_NUM_THREADS = 32, so the "1.8x serial win" is against
    single-threaded ORT, not multi-threaded ORT.

  4. N in thread_invariance is divisible by 3, contradicting a comment
    claiming 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.

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

codecov Bot commented Aug 17, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 96.80365% with 7 lines in your changes missing coverage. Please review.
✅ Project coverage is 80.53%. Comparing base (4273c4e) to head (ded24cc).
⚠️ Report is 2 commits behind head on main.

Files with missing lines Patch % Lines
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 96.80% 6 Missing and 1 partial ⚠️
Additional details and impacted files

Impacted file tree graph

@@            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     
Flag Coverage Δ
cli-ort-linux 83.79% <ø> (ø)
cli-ort-windows 83.31% <ø> (ø)
mlas ?
offline 80.41% <96.80%> (+0.73%) ⬆️

Flags with carried forward coverage won't be shown. Click here to find out more.

Files with missing lines Coverage Δ
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 97.41% <96.80%> (+0.28%) ⬆️

... and 9 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@github-actions

github-actions Bot commented Aug 17, 2026 •

Copy link
Copy Markdown

🔴 Benchmark Regression Detected

Comparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).

ℹ️ Absolute times are informational only — they vary with runner load. The % change column is the reliable signal because both sides ran under identical conditions.

Status Scenario Base PR Change
🔴 block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 691.10 µs 1.71 ms +147.2%
🔴 matmul/small_generic_bf16_threads=8/1x256x256 46.05 µs 103.62 µs +125.0%
🔴 matmul/small_generic_f16_threads=8/1x256x256 40.61 µs 86.57 µs +113.2%
🔴 matmul/medium_generic_f16_threads=8/32x512x512 32.75 µs 67.48 µs +106.1%
🔴 matmul/small_generic_f32_threads=8/1x256x256 43.35 µs 84.06 µs +93.9%
🔴 block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 58.15 µs 103.24 µs +77.5%
🔴 matmul/medium_generic_bf16_threads=8/32x512x512 440.17 µs 703.58 µs +59.8%
🔴 matmul/large_generic_bf16_threads=8/32x1024x1024 1.44 ms 2.27 ms +57.4%
⚠️ sampling_latency/top_k_per_token 52.83 µs 65.57 µs +24.1%
⚠️ matmul/large_generic_f16_threads=8/32x1024x1024 85.89 µs 100.10 µs +16.5%
⚠️ sampling_latency/top_p_per_token 371.91 µs 432.02 µs +16.2%
✅ sampling_latency/min_p_per_token 201.57 µs 227.44 µs +12.8%
✅ tokenization/encode_tokens_per_second 383.28 µs 430.55 µs +12.3%
✅ sampling_latency/greedy_per_token 3.16 µs 3.52 µs +11.5%
✅ tokenization/decode_tokens_per_second 6.23 ms 6.92 ms +11.2%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 505.59 µs 556.50 µs +10.1%
✅ matmul/small_generic_f16_threads=1/1x256x256 32.79 µs 36.09 µs +10.1%
✅ grammar_masking/llguidance_compute_mask/32 74.00 µs 81.18 µs +9.7%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 580.48 µs 623.92 µs +7.5%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.45 ms 3.71 ms +7.3%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.62 ms 2.80 ms +6.9%
✅ matmul/large_generic_f32_threads=8/32x1024x1024 3.94 ms 4.21 ms +6.8%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.57 ms 5.95 ms +6.7%
✅ logit_processing/seven_processor_chain_per_step 320.73 µs 341.40 µs +6.4%
✅ qwen3_sampling_processors/top_k_top_p_fast 638.14 µs 678.20 µs +6.3%
✅ block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 59.69 µs 63.06 µs +5.7%
✅ kv_cache/alloc_dealloc_pages 38.41 µs 40.58 µs +5.6%
✅ qwen3_sampling_processors/top_k_partial_selection 141.21 µs 146.99 µs +4.1%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 2.09 ms 2.18 ms +4.0%
✅ add/large_f32_threads=1-internal/4194304 724.57 µs 747.83 µs +3.2%
✅ add/small_bf16_threads=1-internal/1024 434.0 ns 447.5 ns +3.1%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 544.36 µs 556.07 µs +2.2%
✅ add/small_f16_threads=1-internal/1024 459.9 ns 458.8 ns -0.2%
✅ add/medium_bf16_threads=1-internal/262144 102.85 µs 102.21 µs -0.6%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 80.12 µs 79.49 µs -0.8%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.51 ms 9.29 ms -2.3%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 179.11 µs 175.02 µs -2.3%
✅ add/medium_f16_threads=1-internal/262144 107.92 µs 104.30 µs -3.4%
✅ reduce_mean/medium_f32_threads=1-internal/65536 248.73 µs 238.29 µs -4.2%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 2.00 ms 1.91 ms -4.4%
✅ reduce_mean/small_f32_threads=1-internal/4096 15.39 µs 14.62 µs -5.0%
✅ matmul/medium_generic_f16_threads=1/32x512x512 32.48 µs 30.73 µs -5.4%
✅ gather/medium_f16_threads=1-internal/32768 2.50 µs 2.33 µs -6.7%
✅ reduce_mean/large_f32_threads=1-internal/262144 1.03 ms 960.91 µs -6.7%
✅ add/medium_f32_threads=1-internal/262144 25.32 µs 23.49 µs -7.2%
✅ gather/medium_bf16_threads=1-internal/32768 2.55 µs 2.35 µs -7.8%
✅ matmul/small_generic_f32_threads=1/1x256x256 42.50 µs 38.23 µs -10.0%
✅ gather/small_f32_threads=1-internal/4096 741.3 ns 647.6 ns -12.6%
✅ add/large_f16_threads=1-internal/4194304 1.82 ms 1.58 ms -13.1%
✅ add/small_f32_threads=1-internal/1024 253.7 ns 220.4 ns -13.1%
✅ add/large_bf16_threads=1-internal/4194304 1.84 ms 1.59 ms -13.9%
✅ gather/small_bf16_threads=1-internal/4096 534.5 ns 458.6 ns -14.2%
✅ gather/medium_f32_threads=1-internal/32768 4.93 µs 4.23 µs -14.2%
✅ gather/large_f16_threads=1-internal/131072 20.84 µs 17.79 µs -14.6%
🟢 gather/large_f32_threads=1-internal/131072 49.71 µs 42.20 µs -15.1%
🟢 gather/small_f16_threads=1-internal/4096 553.7 ns 460.1 ns -16.9%
🟢 matmul/small_generic_bf16_threads=1/1x256x256 40.03 µs 31.20 µs -22.1%
🟢 matmul/medium_generic_f32_threads=8/32x512x512 1.78 ms 1.32 ms -25.9%
🟢 gather/large_bf16_threads=1-internal/131072 23.14 µs 16.77 µs -27.5%

Visual flags: ⚠️ ≥ 15% slower, 🔴 ≥ 30% slower — calibrated against measured runner noise (~27% worst-case on multi-threaded matmul)

Host info
CPU: Apple M1 (Virtual)
Cores: 3
OS: Darwin 25.5.0 arm64
Rust: rustc 1.97.1 (8bab26f4f 2026-07-14)
Load avg: { 4.14 3.93 5.59 }
What this cannot catch
  • Regressions in code paths not covered by these benchmarks (e.g., end-to-end decode with a real model)
  • Sub-threshold regressions that compound over multiple PRs
  • Performance changes that only manifest under GPU execution
  • Latency changes in the ORT integration path (these benchmarks exercise the native Rust kernels)

justinchuby and others added 3 commits August 17, 2026 04:44
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>
justinchuby and others added 3 commits August 17, 2026 10:51
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
justinchuby marked this pull request as ready for review August 17, 2026 11:56
@justinchuby
justinchuby merged commit f3d68df into main Aug 17, 2026
12 of 18 checks passed
@justinchuby
justinchuby deleted the deckard/parallel-elementwise branch 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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants