Skip to content

perf(ep-cpu): route the activation transcendentals through MLAS - #1111

Merged
justinchuby merged 3 commits into
mainfrom
deckard/mlas-transcendentals
Aug 17, 2026
Merged

justinchuby merged 3 commits into
mainfrom
deckard/mlas-transcendentals

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 17, 2026 •

Copy link
Copy Markdown
Owner

What ships, and what does not

docs/performance/ABSORBING_MLAS.md (direction set 2026-08-16, after this PR
was opened) is explicit: we do not bundle MLAS by default, we absorb its
mechanisms into native kernels, and MLAS stays in the tree as a measurement
reference
. This PR is now scoped to that role, and the framing below is
corrected accordingly.

mlas is not in the default features of onnx-runtime-ep-cpu, the CLI or the
server. Every route added here sits behind #[cfg(feature = "mlas")], so
none of the speedups in this PR reach a default build. The tables below are
measurements of a reference implementation, not a shipping improvement, and
should be read as "here is how far our native kernel is from a good one".

What they are good for is picking the next absorption. The gap they expose is:

family our native kernel vs MLAS absorbed?
Tanh, Sigmoid 1.03–1.17x yes — #1121, in the default build
Erf, exact Gelu 1.22–1.57x not yet; exp_ps in the erf_ps tail is the suspect

#1121 came directly out of reading tanh.cpp and logistic.cpp for this PR:
our rational was already MLAS's, and the difference turned out to be four
redundant vector ops on our side. That absorption ships; this PR is the
instrument that found it.

Root cause

Our pure-Rust AVX2 activation kernels are compute-bound where ORT is
memory-bound
. Measured on this host (AMD EPYC 9V74, AVX2+FMA, AVX-512
masked off), single-threaded, 1 Mi f32:

ns/elem effective GB/s (4 B in + 4 B out)
our tanh_avx2 0.46 17.4
MlasComputeTanh 0.32 25.0

25 GB/s is single-core stream bandwidth on this part. MLAS is at the memory
wall; we are ~30% short of it and spending the difference in the polynomial.
That means no amount of loop tuning, unrolling or prefetching closes the gap —
the arithmetic itself is the cost.

MLAS is already vendored in this repo (crates/mlas-sys/vendor/mlas) and is
the library ONNX Runtime's own CPU activation kernels call. So this is not
"adopt a faster approximation"; it is stop maintaining a second, slower
approximation of the function ORT already ships
.

What changed

mlas-sys gains compute_tanh, compute_erf and compute_gelu_erf
next to the existing compute_logistic. simd_activations routes Tanh,
Sigmoid, Erf and exact Gelu to them at or above SIMD_MIN_LEN.

The pure-Rust path is not deleted. It remains the fallback: it is what
builds without the mlas feature, what runs off x86_64, and what still serves
tensors below SIMD_MIN_LEN, where the FFI call and MLAS's own dispatch cost
more than the arithmetic saved.

Correctness — two real divergences, both found by falsifiers

I did not assume MLAS matched us. Two tests were written to try to break the
claim, and both succeeded on the first run.

1. MLAS's tanh and logistic are not range-preserving.
monotonicity_within_documented_slack failed immediately with
tanh(-8.990999) = -1.0000001. A dense sweep of [-20, 20] at 2^20 points:

op max ulp vs our path points outside the mathematical range
Tanh 2 2934 / 1048576
Sigmoid 1 2 / 1048576
Erf 0 0
exact Gelu 0 0

ORT calls these directly and inherits the overshoot. Our kernels have always
guaranteed the range and downstream code is entitled to rely on it —
sqrt(1 - tanh²) turns a 1-ulp overshoot into a NaN. So we clamp. Clamping is
never less accurate: the true value lies strictly inside the open interval,
so pulling an out-of-range result to the endpoint moves it toward the exact
answer. NaN survives because every comparison against it is false.

2. MlasComputeGeluErf(-inf) returns NaN. It evaluates
x·0.5·(1 + erf(x/√2)) without special-casing the limit, so it computes
-inf · 0. The limit is 0, which is what our path already returned. Repaired
with a branch-free NaN scan over the block MLAS just wrote — a NaN there can
only come from a NaN input (on which both paths agree) or from this case, so
repairing only the -inf lanes is exact.

This makes us deliberately disagree with ORT on Gelu(-inf). Returning the
limit is the better answer and NaN would propagate through the rest of the
graph, so this is not something to "fix" by matching ORT.

After both repairs the MLAS route is bit-identical on this host (max ulp 0) to the
pure-Rust route across the dense sweep and across NaN (both signs), ±Inf,
±0, ±MIN_POSITIVE, the smallest subnormals, MAX/MIN and the exp
saturation thresholds — pinned by
mlas_ab::mlas_matches_rust_simd_on_special_values, which fails if a future
MLAS bump changes any of it.

Benchmark

Same binary, both routes compiled once, interleaved one iteration each per
round
so drift hits both sides equally, best-of-15. The
activation_bench.rs header documents a uniform cross-build offset on
byte-identical kernels, which an in-binary A/B cannot suffer.

cargo test -p onnx-runtime-ep-cpu --features mlas --release --lib mlas_ab -- --ignored --nocapture

Median of 3 full runs (min-max shown where it matters):

op 1 Ki 16 Ki 256 Ki 1 Mi 4 Mi
Erf 1.56x 1.57x 1.45x 1.46x 1.45x
exact Gelu 1.38x 1.38x 1.31x 1.30x 1.22x
Tanh 1.15x 1.17x 1.08x 1.08x 1.07x (1.03-1.07)
Sigmoid 1.12x 1.15x 1.04x 1.04x 1.03x (0.99-1.04)

Honest reading — this is not a uniform win.

  • Erf (1.45-1.57x) and exact Gelu (1.22-1.38x) are real, repeatable
    wins
    at every size.
  • Tanh and Sigmoid win 12-17% on small and medium tensors, but at 256 Ki
    and above they are bandwidth-bound ties: 1.03-1.08x, and Sigmoid at
    4 Mi measured as low as 0.99x. At those sizes both routes saturate
    memory, so the polynomial stops mattering and the fix-up pass is pure
    overhead. Claiming a win there would be claiming noise.

An earlier revision of this description said "no size loses". That was wrong,
and review caught it: it came from a single non-interleaved run. The
interleaved median above is the corrected measurement.

Two negative results, kept because they were measured:

  • A hand-written branch-free bit-select clamp was slower than
    f32::clamp (0.55 vs 0.43 ns/elem for Tanh at 4 Mi). LLVM already lowers
    clamp to vmaxps/vminps. The measured winner is the one in the diff.
  • Blocking the fix-up passes to keep them in cache was neutral, not a win
    (the cost was the extra pass, not a DRAM round trip). Kept anyway because it
    is free and bounds the working set.

Scope

This does not enable the mlas feature anywhere it was not already
enabled.
Turning it on for the plugin cdylib flips 280 #[cfg(feature = "mlas")] sites across matmul, gemm, sdpa, softmax, conv, pooling and the
quantized paths — a far larger behavioural change than activations, needing its
own conformance and benchmark evidence. That is a separate PR.

So the direct beneficiaries here are builds that already opt in
(onnx-genai-cli, onnx-genai-engine, onnx-runtime-session, onnx-genai-bench).

End-to-end: does this help inside a real session?

The table above is an in-binary A/B of two kernels. That is the right way to
measure a kernel, but it cannot answer whether the win survives ORT's node
dispatch. So the same build was run through a real ORT session both ways --
same single-node graph, same input, matched intra-op threads, p50 of 200
interleaved runs, with our EP's assignment asserted from ORT's own profiler --
with the mlas feature as the only difference:

op elements our us, mlas off our us, mlas on faster by
Tanh 65536 38.3 32.2 1.19x
Tanh 4194304 2285.8 1625.7 1.41x
Sigmoid 65536 42.9 35.2 1.22x
Sigmoid 4194304 2240.6 1735.5 1.29x
Erf 65536 91.2 63.2 1.44x
Erf 4194304 5141.6 3623.1 1.42x
exact Gelu 65536 111.0 86.4 1.28x
exact Gelu 4194304 6701.1 5433.2 1.23x

The win survives, and it is if anything larger end-to-end than in the
microbenchmark.

Where that leaves us against ORT. It is a real improvement but not yet a
win. In the same runs, expressed as ORT's time over ours (higher is better):

op 4096 65536 262144 4194304
Tanh 0.73 -> 0.76 0.71 -> 0.84 0.61 -> 0.62 0.96 -> 0.85
Sigmoid 0.74 -> 0.78 0.67 -> 0.81 0.84 -> 0.90 0.74 -> 0.86
Erf 0.71 -> 0.84 0.66 -> 0.95 0.80 -> 1.02 0.71 -> 1.38
exact Gelu 0.70 -> 0.81 0.67 -> 0.86 0.68 -> 0.88 0.75 -> 0.91

Erf crosses over. The rest close most of the gap but do not. Two things are
still in the way and neither is the polynomial:

  1. The fix-up pass. ORT calls MlasComputeTanh once; we call it and then
    clamp, which is a second pass over the output. Blocking it to 2048 elements
    and doing it with vmaxps/vminps recovered most but not all of that --
    about 12% remains at 65 Ki. Removing it entirely means giving up the range
    guarantee that monotonicity_within_documented_slack asserts, which is not
    a trade this PR makes.
  2. Per-node plugin overhead, ~1.7-2.7 us, which is most of the loss at
    4096 elements and none of it at 4 Mi.

The 4194304 column also moves around between runs because ORT's own time
there varied 1388-2199 us across repeats; our absolute times were stable to
~2%, which is why the absolute table above is the one to trust.

Independent review

Reviewed by an independent Opus reviewer: GO WITH FINDINGS. It reproduced
every number in this description, verified the FFI against mlas.h (including
that MlasComputeTanh is a template and that the shim's <float>
instantiation links), confirmed the MlasComputeGeluErf no-overlap
precondition is discharged by output_direct_write_eligible even for
in-place-capable nodes, and ran its own adversarial sweep (|x| out to
3.4e38, the +/-88 exp-saturation band, subnormals, and the SIMD_MIN_LEN seam)
finding 0 mismatches. It also confirmed this PR does not widen the
pre-existing scalar-vs-vector seam.

All four findings are addressed on this branch: the ISA-specific
bit-reproducibility caveat, the corrected benchmark claim, the new mlas-sys
smoke tests, and the interleaved A/B harness.

Limitations

  • Numbers are from one host. The dispatch is unconditional above
    SIMD_MIN_LEN rather than ISA-tuned, which is safe because MLAS does its own
    runtime ISA dispatch internally.
  • Bit-reproducibility is host-specific. Because MLAS dispatches by ISA, an
    AVX-512 or non-AVX2 machine may pick a different kernel than the AVX2+FMA one
    the pure-Rust path mirrors, so the two routes need not agree bitwise there.
    Before this change, mlas-on and mlas-off builds agreed bitwise for these
    four ops on every host; that is what is traded for the speed. The
    special-value test pins whatever ISA it runs on.
  • f16/bf16 still widen to f32 first; this PR does not change that path.
  • The clamp and NaN-scan fix-ups cost ~0.1 ns/elem. They are the price of
    keeping our stronger numerical contract, and are why Tanh/Sigmoid show
    1.0-1.2x rather than the 1.3-1.55x the raw MLAS kernels reach.

Validation

  • cargo fmt --all -- --check
  • cargo clippy -p onnx-runtime-ep-cpu --features mlas --all-targets — clean
    (one pre-existing needless_return in matmul_nbits.rs, untouched here)
  • cargo test -p onnx-runtime-ep-cpu --lib — 1323 passed (feature off)
  • cargo test -p onnx-runtime-ep-cpu --features mlas --lib — 1341 passed
  • cargo test -p mlas-sys — 39 passed, plus 5 new
    tests/transcendentals.rs cases pinning the new FFI at the crate boundary
    against a correctly-rounded f64 reference (so they check the function,
    not MLAS's own polynomial), including empty-slice and length-mismatch cases

Our pure-Rust AVX2 activations are compute-bound where ORT is memory-bound.
Measured on this host (AMD EPYC 9V74, AVX2+FMA, AVX-512 masked), a
single-threaded `Tanh` over 1 Mi f32 runs at 0.46 ns/elem in our kernel and
0.32 ns/elem in MLAS — and 0.32 ns/elem is ~25 GB/s for a 4 B-in/4 B-out
stream, i.e. bandwidth. The gap is the polynomial, not the loop, so no amount
of loop tuning closes it.

MLAS is already vendored in-tree and is the library ORT's own CPU activation
kernels call, so this shares ORT's polynomial rather than carrying a second,
slower approximation of the same function.

`mlas-sys` gains `compute_tanh`, `compute_erf` and `compute_gelu_erf`
alongside the existing `compute_logistic`, and `simd_activations` routes
`Tanh`, `Sigmoid`, `Erf` and exact `Gelu` to them at or above `SIMD_MIN_LEN`.
The pure-Rust path stays as the fallback: it is what builds without the
feature, what runs off x86, and what still serves short tensors.

Two behavioural differences were found by falsifiers and repaired, so the
MLAS route is bit-identical to the pure-Rust route on a dense sweep of
[-20, 20] and on every special value:

* `MlasComputeTanh` and `MlasComputeLogistic` are not range-preserving —
  2934 and 2 of 1048576 sweep points land outside [-1, 1] / [0, 1] by up to
  2 ulp (`tanh(-8.990999) = -1.0000001`). ORT inherits that. We clamp, which
  is never less accurate because the true value is inside the open interval.
* `MlasComputeGeluErf(-inf)` returns NaN (it computes `-inf * 0`) where the
  limit is 0. Repaired via a branch-free NaN scan, blocked so the scan stays
  in cache.

Speedups, same binary, interleaved, worst..best across 1 Ki..4 Mi:
Erf 1.45-1.61x, exact Gelu 1.21-1.31x, Tanh 1.08-1.21x, Sigmoid 1.04-1.26x.
No size loses.

`f32::clamp` beat a hand-written branch-free bit-select (0.43 vs 0.55
ns/elem); the measured version is the one kept.

This does not enable the `mlas` feature anywhere it was not already enabled
— doing that for the plugin cdylib flips 280 sites across matmul, gemm,
sdpa, softmax, conv and the quantized paths, and needs its own evidence.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@codecov

codecov Bot commented Aug 17, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 80.53%. Comparing base (b7fa5e1) to head (cb6ecf9).
⚠️ Report is 3 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1111      +/-   ##
==========================================
+ Coverage   79.86%   80.53%   +0.66%     
==========================================
  Files         367      369       +2     
  Lines      157553   160375    +2822     
  Branches   157553   160375    +2822     
==========================================
+ Hits       125825   129152    +3327     
+ Misses      27006    26482     -524     
- Partials     4722     4741      +19     
Flag Coverage Δ
cli-ort-linux 83.79% <ø> (ø)
cli-ort-windows 83.40% <ø> (ø)
mlas 84.83% <100.00%> (?)
offline 80.33% <100.00%> (+0.61%) ⬆️

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

Files with missing lines Coverage Δ
crates/mlas-sys/src/lib.rs 84.13% <100.00%> (ø)
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 98.15% <100.00%> (ø)

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

✅ Benchmarks — No Regression

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
✅ matmul/large_generic_bf16_threads=8/32x1024x1024 1.42 ms 1.49 ms +4.8%
✅ block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 46.19 µs 46.60 µs +0.9%
✅ tokenization/encode_tokens_per_second 407.77 µs 409.79 µs +0.5%
✅ add/medium_f16_threads=1-internal/262144 119.73 µs 119.51 µs -0.2%
✅ gather/large_bf16_threads=1-internal/131072 18.13 µs 18.03 µs -0.6%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 197.13 µs 189.15 µs -4.0%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 2.09 ms 1.99 ms -4.6%
✅ block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 46.30 µs 43.56 µs -5.9%
✅ add/medium_f32_threads=1-internal/262144 29.57 µs 27.77 µs -6.1%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 499.49 µs 468.86 µs -6.1%
✅ logit_processing/seven_processor_chain_per_step 351.97 µs 326.82 µs -7.1%
✅ block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 624.46 µs 574.54 µs -8.0%
✅ kv_cache/alloc_dealloc_pages 41.58 µs 37.96 µs -8.7%
✅ sampling_latency/min_p_per_token 239.24 µs 215.91 µs -9.8%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 91.96 µs 81.47 µs -11.4%
✅ qwen3_sampling_processors/top_k_partial_selection 157.15 µs 138.67 µs -11.8%
✅ grammar_masking/llguidance_compute_mask/32 84.44 µs 74.27 µs -12.0%
✅ qwen3_sampling_processors/top_k_top_p_fast 755.53 µs 657.37 µs -13.0%
🟢 matmul/medium_generic_f32_threads=1/32x512x512 2.74 ms 2.31 ms -15.8%
🟢 qwen3_sampling_processors/top_p_fast_after_top_k 614.15 µs 516.79 µs -15.9%
🟢 qwen3_sampling_processors/top_k_full_sort_baseline 2.54 ms 2.13 ms -16.0%
🟢 matmul/large_generic_f16_threads=8/32x1024x1024 107.84 µs 89.43 µs -17.1%
🟢 matmul/large_generic_f32_threads=1/32x1024x1024 11.42 ms 9.40 ms -17.6%
🟢 gather/medium_f32_threads=1-internal/32768 4.96 µs 4.06 µs -18.3%
🟢 matmul/small_generic_bf16_threads=8/1x256x256 39.38 µs 31.82 µs -19.2%
🟢 matmul/medium_generic_bf16_threads=8/32x512x512 470.00 µs 378.43 µs -19.5%
🟢 qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 4.41 ms 3.50 ms -20.7%
🟢 sampling_latency/top_p_per_token 482.15 µs 381.43 µs -20.9%
🟢 add/small_bf16_threads=1-internal/1024 626.2 ns 494.3 ns -21.1%
🟢 matmul/small_generic_bf16_threads=1/1x256x256 41.40 µs 32.49 µs -21.5%
🟢 sampling_latency/greedy_per_token 4.10 µs 3.20 µs -22.1%
🟢 matmul/small_generic_f32_threads=1/1x256x256 48.00 µs 37.19 µs -22.5%
🟢 matmul/medium_generic_bf16_threads=1/32x512x512 676.31 µs 519.38 µs -23.2%
🟢 matmul/medium_generic_f32_threads=8/32x512x512 1.22 ms 932.20 µs -23.4%
🟢 matmul/medium_generic_f16_threads=1/32x512x512 41.11 µs 30.42 µs -26.0%
🟢 tokenization/decode_tokens_per_second 8.25 ms 6.07 ms -26.4%
🟢 reduce_mean/small_f32_threads=1-internal/4096 20.50 µs 14.98 µs -26.9%
🟢 reduce_mean/medium_f32_threads=1-internal/65536 336.30 µs 244.31 µs -27.4%
🟢 sampling_latency/top_k_per_token 75.32 µs 54.03 µs -28.3%
🟢 reduce_mean/large_f32_threads=1-internal/262144 1.42 ms 996.31 µs -29.8%
🟢 matmul/large_generic_f32_threads=8/32x1024x1024 5.59 ms 3.86 ms -30.9%
🟢 gather/small_f16_threads=1-internal/4096 706.3 ns 474.6 ns -32.8%
🟢 add/medium_bf16_threads=1-internal/262144 156.45 µs 104.69 µs -33.1%
🟢 add/large_f16_threads=1-internal/4194304 2.53 ms 1.67 ms -34.0%
🟢 add/large_bf16_threads=1-internal/4194304 2.49 ms 1.64 ms -34.3%
🟢 gather/small_bf16_threads=1-internal/4096 730.2 ns 474.6 ns -35.0%
🟢 gather/small_f32_threads=1-internal/4096 1.06 µs 669.4 ns -37.1%
🟢 add/large_f32_threads=1-internal/4194304 1.07 ms 668.64 µs -37.4%
🟢 matmul/medium_generic_f16_threads=8/32x512x512 47.70 µs 29.39 µs -38.4%
🟢 qwen3_sampling_processors/top_k_top_p_full_sort_baseline 9.27 ms 5.67 ms -38.9%
🟢 gather/medium_f16_threads=1-internal/32768 3.98 µs 2.43 µs -38.9%
🟢 add/small_f32_threads=1-internal/1024 356.0 ns 217.1 ns -39.0%
🟢 gather/medium_bf16_threads=1-internal/32768 4.16 µs 2.42 µs -41.9%
🟢 gather/large_f16_threads=1-internal/131072 23.33 µs 12.42 µs -46.8%
🟢 matmul/small_generic_f16_threads=1/1x256x256 58.12 µs 30.78 µs -47.0%
🟢 matmul/small_generic_f16_threads=8/1x256x256 59.85 µs 30.43 µs -49.2%
🟢 add/small_f16_threads=1-internal/1024 946.6 ns 473.0 ns -50.0%
🟢 gather/large_f32_threads=1-internal/131072 68.23 µs 33.65 µs -50.7%
🟢 matmul/small_generic_f32_threads=8/1x256x256 82.09 µs 40.06 µs -51.2%

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: { 3.85 3.51 7.16 }
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 2 commits August 17, 2026 08:23
Four findings from independent review, all measurement- or
contract-related rather than correctness bugs:

1. The "bit-identical to the pure-Rust path" claim was unconditional.
   MLAS dispatches by ISA at runtime, so an AVX-512 or pre-AVX2 host can
   pick a different kernel than the AVX2+FMA one our Rust path mirrors.
   The claim is now scoped to "on hosts where MLAS selects the AVX2
   kernel", which is what the tests actually pin.

2. The benchmark harness took best-of-N for one route and then the
   other, so any thermal or frequency drift between the two phases
   landed entirely on one side. `best_ns_pair` now interleaves a single
   iteration of each route per round. Re-measured, this moves Tanh and
   Sigmoid at 256 Ki and above from "small win" to "tie" -- the honest
   result, and the PR description now says so.

3. mlas-sys had no test coverage for the three new entry points. Added
   tests/transcendentals.rs, which checks them against a correctly
   rounded f64 reference (so it validates the binding, not MLAS's
   polynomial), plus empty-slice and length-mismatch behaviour. This
   adds libm as a dev-dependency of mlas-sys; it is already in the
   workspace lock.

4. Noted that `dispatch_mlas!` intentionally omits the
   `vector_path_available()` guard the pure-Rust `dispatch!` uses,
   because MLAS performs its own runtime ISA dispatch and has a scalar
   path of its own.

cargo test -p mlas-sys: 44 passed (39 + 5 new).
cargo test -p onnx-runtime-ep-cpu --lib: 1323 passed.
cargo test -p onnx-runtime-ep-cpu --features mlas --lib: 1341 passed.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby
justinchuby marked this pull request as ready for review August 17, 2026 11:03
@justinchuby
justinchuby merged commit 7c7897d into main Aug 17, 2026
13 of 18 checks passed
@justinchuby
justinchuby deleted the deckard/mlas-transcendentals branch August 17, 2026 11:03
justinchuby pushed a commit that referenced this pull request Aug 17, 2026
Resolves the additive-test-module conflict with #1111's mlas_ab reference
module: both modules are kept, and the saturation removal is re-applied on
top of main's dispatch_mlas routing.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby pushed a commit that referenced this pull request Aug 17, 2026
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
## What

Delete the MLAS route for `Tanh` and `Sigmoid`. Both ops now run the
pure-Rust AVX2 kernel in every build, `mlas` on or off.

## Root cause

#1111 added MLAS routes for `Tanh`, `Sigmoid`, `Erf` and exact `Gelu`
because MLAS was measurably faster. That was true at the time. Two
things have changed since:

1. **The polynomial was never the win.** Reading `mlas/lib/tanh.cpp` and
`logistic.cpp` showed both are Eigen rational approximations `P(v²)·v /
Q(v²)` with the range clamp fused into the same pass. Our
`avx2::tanh_ps`/`sigmoid_ps` already used the *identical* rational and
the *identical* constants. The route was buying MLAS's loop, not a
better approximation.
2. **#1121 removed the loop deficit.** Our kernels carried ~6 redundant
vector ops (a result clamp plus two saturation `cmpps`+`blendvps` pairs)
that were proven bit-identically redundant over 2 086 666 240 values.
With those gone the pure-Rust kernel is no longer the slower loop.

Meanwhile the MLAS route always had a tax the native path does not:
`MlasComputeTanh`/`MlasComputeLogistic` are **not range-preserving**.
Over a dense sweep of [-20, 20], tanh lands outside `[-1, 1]` for 2934
of 1048576 points (up to 2 ulp, `tanh(-8.990999) = -1.0000001`) and
logistic outside `[0, 1]` for 2. Our kernels guarantee the range,
`monotonicity_within_documented_slack` asserts it, and downstream code
relies on it — `sqrt(1 - tanh²)` turns a 1-ulp overshoot into a NaN. So
the route had to re-read every block MLAS had just written and clamp it.

Net: the fix-up pass now costs more than the polynomial saves.

## Benchmark

Same machine (AMD EPYC 9V74, AVX2+FMA, AVX-512 masked), same harness, 1
thread, `--iters 30 --reps 5`, medians of 3 interleaved rounds
alternating the two builds, EP assignment asserted from ORT's own
profiler (zero `NOT-ASSIGNED`). p50 µs, lower is better.

| op | n | `mlas` **off** | `mlas` **on** (before) | on/off |
|---|---:|---:|---:|---:|
| Tanh | 4 Ki | **7.58** | 7.68 | 1.01x slower |
| Tanh | 16 Ki | **11.93** | 12.91 | 1.08x slower |
| Tanh | 64 Ki | **32.54** | 36.20 | 1.11x slower |
| Tanh | 256 Ki | **105.64** | 146.73 | **1.39x slower** |
| Tanh | 1 Mi | **438.43** | 469.44 | 1.07x slower |
| Tanh | 4 Mi | **1896.32** | 2226.19 | 1.17x slower |
| Sigmoid | 4 Ki | **7.27** | 7.49 | 1.03x slower |
| Sigmoid | 16 Ki | **11.99** | 13.35 | 1.11x slower |
| Sigmoid | 64 Ki | **31.16** | 38.26 | **1.23x slower** |
| Sigmoid | 256 Ki | **113.69** | 129.34 | 1.14x slower |
| Sigmoid | 1 Mi | **408.89** | 525.83 | **1.29x slower** |
| Sigmoid | 4 Mi | **1887.40** | 2268.40 | 1.20x slower |

The route lost at **every** size measured, for both ops.

The route was reachable for any length `>= SIMD_MIN_LEN`, which is
**32**, so the table above left the band `[32, 4096)` unmeasured. Opus
review flagged that "no crossover to preserve" therefore rested on
extrapolation, so it was measured too:

| op | n | `mlas` **off** | `mlas` **on** (before) | on/off |
|---|---:|---:|---:|---:|
| Tanh | 64 | **5.61** | 5.84 | 1.04x slower |
| Tanh | 256 | **5.46** | 5.73 | 1.05x slower |
| Tanh | 1024 | **5.99** | 6.23 | 1.04x slower |
| Sigmoid | 64 | **5.37** | 5.64 | 1.05x slower |
| Sigmoid | 256 | **5.45** | 5.60 | 1.03x slower |
| Sigmoid | 1024 | **6.09** | 6.27 | 1.03x slower |

Still no crossover — MLAS loses across the entire reachable range, 64 to
4 Mi. (Below ~4 Ki both builds sit on a ~5.4 µs per-node floor, so these
differences are small in relative terms; the point is only that the sign
never flips.) This is therefore a deletion rather than a threshold.

### Why this matters now

Since #1115 the wheel ships with `mlas` enabled on x86_64
(`python/nxrt-ep-cpu/setup.py::_mlas_features`). This was not a
measurement-only artefact — it was the shipped configuration. That is
exactly the failure mode `docs/performance/ABSORBING_MLAS.md` names:
"the configuration we measured was not the configuration we shipped".

### Effect against ORT

Restated as ORT p50 ÷ ours (>1 = we win), 1 thread:

| op | n | shipped before | shipped after |
|---|---:|---:|---:|
| Tanh | 256 Ki | 0.71 | **0.99** |
| Tanh | 4 Mi | 0.99 | **1.16** |
| Sigmoid | 64 Ki | 0.75 | **0.93** |
| Sigmoid | 1 Mi | 0.75 | **0.96** |

Still short of parity at several sizes — that work continues — but this
recovers most of a self-inflicted gap.

## What is kept

`Erf` and exact `Gelu` keep their MLAS routes, because the arithmetic is
different:

- **`Erf`** needs no fix-up at all — MLAS's result is used verbatim —
and still wins **1.40x at 64 Ki** and **1.59x at 4 Mi** (5654 → 3551
µs).
- **exact `Gelu`**'s fix-up is a *non-writing* compare scan (repairing
MLAS's `-inf → NaN`, which is `-inf · 0`), not a read-modify-write pass,
and it still wins **1.17x at 4 Mi** (6485 → 5550 µs).

So the rule this establishes is narrower than "MLAS is slower": **an
MLAS route that needs a writing fix-up pass no longer pays for itself;
one that needs no fix-up, or only a scan, still does.**

## Correctness

Nothing about the numerics of the shipped path changes for
`Tanh`/`Sigmoid` — it becomes the pure-Rust path, which is what a
default (`mlas`-off) build has always run and what every dense numeric
sweep in the file already tests. This strictly *reduces* cross-build
divergence: `mlas`-on and `mlas`-off builds are now bit-identical for
these two ops on every host, which they were not while MLAS's
ISA-dispatched kernel was in play.

The clamped MLAS reference is preserved verbatim inside the `mlas_ab`
test module as `tanh_mlas_ref`/`sigmoid_mlas_ref`, so all three
comparisons keep running against it:

- `mlas_matches_rust_simd_on_special_values` — bit-equality on NaN (both
signs), ±Inf, ±0, ±`MIN_POSITIVE`, smallest subnormals, `MAX`/`MIN`,
±88.5, ±1e-30, padded past `SIMD_MIN_LEN` and repeated so every value
lands at every lane offset.
- `mlas_vs_rust_dense_ulp_sweep` — dense ULP sweep over [-20, 20].
- `mlas_vs_rust_simd` — the benchmark that produced the table above.

Had these been deleted along with the route, a future MLAS bump could
have diverged silently.

## Tests

- `cargo test -p onnx-runtime-ep-cpu --features mlas --lib` → **1309
passed, 0 failed**
- `cargo build -p onnx-runtime-ep-cpu` (default, no `mlas`) → clean
- `cargo fmt --all -- --check` → clean

## Limitations

- Measured on one host (AMD EPYC 9V74, Zen 4, AVX2+FMA, AVX-512 masked
off). On a machine where MLAS dispatches an AVX-512 kernel and we stay
on AVX2, the balance could differ — but the fix-up pass is pure memory
traffic and does not get cheaper with a wider ISA, and the clamp is
required regardless, so the direction should hold.
- Only 1-thread numbers are quoted. The 16-thread picture is dominated
by a much larger threading gap that this PR does not touch and that is
tracked separately.
- This does not make `Tanh`/`Sigmoid` beat ORT at every size. It removes
a regression; the remaining gap is real and open.

## Review

Independent Opus review: **GO WITH FINDINGS**, three findings, all
addressed in `4b47896`:

1. *(MINOR)* The rationale comment had merged into `MLAS_FIXUP_BLOCK`'s
doc comment, so rustdoc rendered "Why `Tanh` and `Sigmoid` do not have
an MLAS route" as the summary of an 8192-element block-size constant —
and, being `#[cfg(feature = "mlas")]`-gated, the explanation for the
route's *absence* disappeared from a default build. It is now a
free-standing comment.
2. *(NIT)* `dispatch_mlas!`'s doc still said "these four ops" when the
macro now has two callers. Reworded, with a pointer to the removal note.
3. *(NIT)* "lost at every size" was every size *measured*, leaving `[32,
4096)` to extrapolation. Measured — see the second table above.

The reviewer independently verified the route removal is complete
(`mlas_sys::compute_tanh`/`compute_logistic` are unreachable from
shipped code), that `_mlas_features` really does ship `mlas` on
linux-x86_64, that `erf_gelu_mlas`'s scan is genuinely non-writing on
the all-finite path, that `mlas_clamped_ref` is faithful to the deleted
code rather than a straw man, and that the 1.39x / 1.59x / 1.17x
arithmetic checks out.

## CI

`Fast (Linux x86_64)` (the required check), `EP conformance (Linux
x86_64)`, `Rust quality`, `Miri unsafe-crate soundness`, `Rust coverage
(Linux x86_64)`, `audit` and both codecov lanes are green.

The red lanes — `CLI ORT (Linux/Windows x86_64)`, `CUDA compile
(Linux/Windows x86_64)`, `Rust coverage (macOS arm64)`, `Rust (Windows
ARM64)` — are red on pristine `origin/main` too, at every one of the
last five commits. Within `CLI ORT` the activation-relevant step fails
on
`shape_inference_coverage::every_registered_op_has_a_shape_rule_or_is_a_known_gap`,
which reproduces verbatim on an untouched `origin/main` checkout and
names six `Nchwc*` ops (`NchwcConv`, `NchwcMaxPool`, ...) with no shape
rule. Nothing to do with activations, and not introduced here.

Worth flagging for whoever owns those ops: that test failing means
`GetCapability` **is** silently handing six registered ops to ORT's CPU
EP, which the project's architecture policy forbids. The two
activation-side falsifiers in the same file —
`no_activation_or_norm_op_is_left_to_ort` and
`activation_and_norm_ops_clear_every_capability_filter` — pass.

---------

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.

1 participant