Skip to content

perf(ep-cpu): port MLAS's erf polynomial to AVX2, 3-28x on Erf and exact GELU - #1074

Merged
justinchuby merged 3 commits into
mainfrom
deckard/erf-avx2
Aug 16, 2026
Merged

justinchuby merged 3 commits into
mainfrom
deckard/erf-avx2

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 16, 2026 •

Copy link
Copy Markdown
Owner

Root cause

Erf and Gelu(approximate="none") were the last two activations in this crate still evaluating a libm transcendental per element, in f64. Every other activation family was vectorised in #1037. These two were explicitly left behind, with this comment in kernels/gelu.rs:

Exact GELU stays on libm::erf: the conformance suite's Erf reference is tighter than any f32 polynomial we would want to ship here.

That was wrong on both halves:

  1. conformance/run_onnx_tests.py compares at rtol=1e-4, atol=1e-5. A faithfully-rounded f32 erf is ~1e-7 — three orders of magnitude inside it.
  2. ORT's own CPU kernel is an f32 polynomial. Erf and Gelu(approximate="none") both call MlasComputeErf, whose FMA3 kernel evaluates a two-branch polynomial. So "matching ORT" meant adopting that polynomial, not avoiding it.

The cost was the largest single gap in the whole activation matrix: at 1048576 float32 elements a standalone Erf node took 39.8 ms against ORT's 0.89 ms.

Change

Ports MlasErfKernel from onnxruntime/core/mlas/lib/erf.cpp — already vendored at crates/mlas-sys/vendor/mlas/ — to core::arch::x86_64 in kernels/simd_activations.rs, alongside the tanh/logistic rationals ported in #1037. The mlas cargo feature is off by default, so linking MlasComputeErf directly was not an option for the default build; this is the same reason #1037 ported rather than linked.

  • erf_ps — the two-branch kernel: |x| <= 0.921875 uses x·(1 + P(x²)); above it, 1 - exp(-R(|x|)). Both branches are evaluated for every lane and merged with or, which works because the inactive branch is forced to +0.0 (the small branch by andnot, the big branch because zeroing its input collapses 1 - exp(-0) to 0). Branch-free is what buys the speedup — far more than the 8× the SIMD width alone gives.
  • exp_ps — the small range-reduced exp the big branch needs, using MlasErfConstants' own exp parameters, plus power_of_2_ps (MlasPowerOf2Float32x4).
  • erf_gelu_ps / erf_gelu_bias_f32_slice — exact GELU with the x/√2 scale fused so the intermediate never reaches memory, and the BiasGelu bias folded in-register exactly as FastGelu's already is.

Wired into four call sites that previously ran the scalar f64 form: Erf (elementwise.rs), ai.onnx::Gelu(approximate="none") and com.microsoft::Gelu (gelu.rs), and BiasGelu (contrib_fused.rs). kernels::gelu::exact_gelu had no remaining callers and is deleted; its f64 form survives as simd_activations::erf_gelu_scalar, the non-AVX2 fallback.

libm::erf remains the fallback on every non-AVX2 target, so no existing platform changes accuracy or speed.

Benchmarks

Session-level A/B (real ORT session latency, not kernel time, so ORT's per-node overhead and our plugin's are both included). Same host, same process, interleaved arms, taskset -c 0-15, 1 intra-op thread, reps=21 (the 1 048 576 row was re-measured at reps=31 after review found it dispersion-limited). Ratio = ORT ns / ours; >1 means we are faster. Before = ab6cb0168, after = this branch. Both arms built identically; the float32 rows use a scratch NXRT_EXP_CLAIM_ALL override (never committed) to force assignment, because the policy defers float32 today.

float32 Erf (measurement-only claim)

elements ORT (ns) before (ns) after (ns) before x after x kernel speedup
512 4 466 22 036 6 039 0.203 0.738 3.6x
3 072 6 665 102 342 9 455 0.065 0.706 10.8x
16 384 17 732 563 117 26 961 0.031 0.660 20.9x
65 536 60 937 2 402 053 89 150 0.025 0.680 26.9x
262 144 228 758 9 711 873 336 857 0.024 0.679 28.8x
1 048 576 889 823 40 278 112 1 295 386 0.022 0.687 31.1x

float32 Gelu(approximate="none") (measurement-only claim)

elements before x after x kernel speedup
512 0.251 0.758 3.0x
3 072 0.090 0.700 7.8x
16 384 0.052 0.671 12.9x
65 536 0.040 0.684 17.1x
262 144 0.039 0.687 17.6x
1 048 576 0.039 0.699 17.9x

float16 Gelu(approximate="none") — the range this EP actually claims

The assignment policy claims float16 Gelu unconditionally, because ORT has no float16 Gelu CPU kernel: declining makes ORT inline the function body and this EP then picks up the ungoverned constituents, measured at 0.024–0.049x. So this is a range where the EP really was serving users a 17x-slower kernel.

elements ORT (ns) before (ns) after (ns) before x after x
512 7 639 18 680 6 455 0.414 1.184
3 072 12 150 83 816 11 417 0.145 1.064
16 384 33 755 418 354 37 179 0.082 0.908
65 536 120 750 1 914 397 131 620 0.063 0.917
262 144 468 070 7 765 043 504 909 0.061 0.927
1 048 576 1 811 361 30 773 937 2 016 784 0.059 0.898

Control: float16 Erf, which the policy defers

1.035 / 1.018 / 1.001 / 1.009 / 0.996 / 1.004 — flat 1.0, i.e. the deferral still hands the node to ORT and this PR does not accidentally claim it.

What this does not claim

  • float32 Erf and Gelu(none) still lose to ORT (0.66–0.76x) and this PR does not change their assignment — they stay deferred. What changed is that they now sit on the same ~0.70x plateau as Tanh (0.70–0.75x) and Sigmoid (0.70–0.73x) instead of being 30x worse than their neighbours. That residual plateau is not the transcendental: it is this crate's elementwise kernels being single-threaded while ORT spreads the same work over its intra-op pool. assignment_policy.rs's rationale is updated to say so, with the new numbers replacing the stale 0.023-0.77x.
  • Measured on one machine (AMD EPYC 9V74, AVX2+FMA, AVX-512 masked by the hypervisor). No dispatch threshold is derived from these numbers, so there is nothing here to overfit — the change is unconditional on AVX2+FMA and strictly faster at every size measured.
  • The vector path is less accurate than the scalar one it replaces (faithful vs correctly rounded). That is deliberate and is what ORT does; see below.

Correctness

libm::erf is correctly rounded; MLAS's polynomial is faithfully rounded. Measured worst error over 400 003-point sweeps, scaled by max(1, |x|):

function sweep worst error
erf [-6, 6] + both branch boundaries + clamp + subnormals 5.96e-8
erf [-1.5, 1.5] (dense, where erf is steep) 5.96e-8
exact GELU [-25, 25] 1.19e-7

5.96e-8 is exactly one ulp below 1.0 — i.e. faithful, as advertised. The conformance suite's tolerance is rtol=1e-4.

New tests (11 in simd_activations, all also exercised through the kernels):

  • erf_dense_sweep_matches_f64_reference, erf_dense_sweep_near_origin, erf_gelu_dense_sweep_matches_f64_reference — the sweeps above, against libm::erf in f64, asserting a 3e-7 / 4e-7 bound.
  • erf_special_values, erf_gelu_special_values — -Inf -> -1, +Inf -> 1, ±0 keeps its sign, NaN propagates, ±MAX and ±1e30 saturate. Exact GELU pins -Inf -> 0 (the mathematical limit) as the tanh form already does.
  • erf_saturates_to_exactly_one_past_the_clamp — the 3.925 clamp must not be observable: every input past it returns exactly ±1.0, not 1 - eps.
  • erf_is_exactly_odd — the sign is applied by an or after the polynomial rather than carried through it, so exact antisymmetry over 4096 points is a falsifier for the sign mask leaking into the arithmetic.
  • erf_tail_lanes_match_the_vector_body — bit equality between every length in [32, 40) and the full run, so results cannot depend on tensor length mod 8.

Existing coverage that had to keep passing: erf_known_values, erf_odd_symmetry_and_limits, erf_bf16_reaches_dtype_without_touching_formula, gelu_*, bias_gelu_*, and the EP conformance lane.

Validation

  • cargo test -p onnx-runtime-ep-cpu --lib — 1224 passed (debug) and 1224 passed (release).
  • cargo clippy -p onnx-runtime-ep-cpu --all-targets -- -D warnings — clean.
  • cargo fmt --all --check — clean.

Limitations

  • AVX2+FMA only. Other ISAs keep libm::erf, unchanged in both speed and accuracy. Results therefore differ by ~1 ulp between an AVX2 host and a non-AVX2 host — the same ISA dependence this module already documents for tanh/sigmoid, and the same one ORT has.
  • Float64 inputs are untouched; they keep the exact f64 path.
  • Does not address the ~0.70x single-thread plateau shared by every activation here.

Independent review

Reviewed by an independent claude-opus-4.8 Rubber Duck against an adversarial brief (prove the branch merge, prove the constants, prove the special values, prove the benchmarks, find the regression). Verdict GO WITH FINDINGS, no blockers. The reviewer rebuilt both arms from scratch and re-derived rather than re-read:

  • Constants — bit-compared all 26 against erf.cpp. 25 bit-identical. The 26th, EXP_LOWER_RANGE, was -88.376_264 (0xc2b0c0a6) against MLAS's 0xc2b0c0a5; the reviewer also proved it unreachable for erf (the big branch's R maxes at 17.375652 at the |x| = 3.925 clamp, so the argument never approaches -88). Transcribed exactly anyway in 866b95c, since the module's contract is "verbatim MLAS".
  • Branch merge — confirmed analytically and empirically that the inactive big-branch lane is exactly +0.0, so the or merge cannot contaminate the small branch.
  • Numerics vs ORT — max ULP difference = 0 over 4M+ points for Erf, Gelu(none) and BiasGelu at widths 1/8/37/40, against ORT 1.28.0 on the same host. NaN sign, NaN payload and sNaN quieting are bit-identical too.
  • Benchmarks — reproduced the tables. Found the float32 Erf 1 048 576 cell (0.576) pessimistic against its own re-run; re-measured at reps=31 and corrected the row to 0.687 (p90 0.604) with a matching ORT baseline, which also raises the kernel speedup to 31.1x.
  • Regression found — an earlier edit in this branch had swallowed the #[inline] attribute on quick_gelu_ps into the preceding doc comment (neither fmt nor clippy catches this). Restored in 866b95c; unrelated to erf but a real defect this PR introduced.
  • Length seam — flagged that Erf now inherits the SIMD_MIN_LEN = 32 scalar/vector seam: erf(0.901059926) is 0x3f4c2503 in a 31-element tensor and 0x3f4c2504 in a 40-element one, and 286 of 2000 random values move by <=1 ulp purely with tensor length. Tanh and Sigmoid have had this since CPU EP: vectorise the approximate activation family (Tanh/Sigmoid/FastGelu/QuickGelu) on AVX2+FMA #1037 and ORT has it too; documented in the module rather than paying 20x on short tensors to remove it.

Out of scope, recorded here so it is not lost: ad-hoc benchmark scripts that let an InferenceSession outlive the plugin registration segfault at interpreter teardown. .work/sess_ab.py disposes correctly and exits 0; the reviewer's scratch scripts did not. Nothing in the shipped code is implicated.

…act GELU

`Erf` and `Gelu(approximate="none")` were the two remaining activations that
evaluated a `libm` transcendental per element, and they did it in `f64`. At
1048576 float32 elements the standalone `Erf` node took 39.8 ms against ORT's
0.89 ms — 0.022x. Every other activation in this crate had already been
vectorised; these two had not, because an earlier comment asserted that no
`f32` polynomial could meet the conformance suite's tolerance.

That assertion was wrong twice over. The suite compares at `rtol=1e-4`, and
ORT's own CPU `Erf` and `Gelu(none)` kernels call `MlasComputeErf`, which is a
faithfully-rounded `f32` polynomial — so matching ORT means *using* that
polynomial, not avoiding it.

Ports `MlasErfKernel` (`onnxruntime/core/mlas/lib/erf.cpp`, already vendored in
`crates/mlas-sys`) to `core::arch::x86_64`, adds the small `exp` it needs, and
routes `Erf`, `Gelu(none)`, `com.microsoft::Gelu` and `BiasGelu` through it.
The scalar `libm::erf` stays as the non-AVX2 fallback, exactly as `tanh` and
`sigmoid` keep theirs.

Session-level A/B against ORT 1.28.0, same host, interleaved, 1 thread,
reps=21, base ab6cb01 (ratio = ORT/ours, >1 means we win):

  float32 Erf         0.203 -> 0.738 (512) ... 0.022 -> 0.576 (1048576)
  float32 Gelu(none)  0.251 -> 0.758 (512) ... 0.039 -> 0.699 (1048576)
  float16 Gelu(none)  0.414 -> 1.184 (512) ... 0.059 -> 0.898 (1048576)

float16 `Gelu` is the range that matters for assignment honesty: the policy
claims it unconditionally because ORT has no float16 `Gelu` kernel to defer to,
so that 0.059x was a range this EP really was serving 17x slower than ORT.

Accuracy, measured over 400 003-point sweeps including both branch boundaries
and the saturation clamp: worst error 5.96e-8 (1 ulp below 1.0) for `erf`,
1.19e-7 scaled for exact GELU. Sign, signed zero, +/-Inf saturation to exactly
+/-1, NaN propagation and exact oddness are pinned by tests.

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

codecov Bot commented Aug 16, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 96.15385% with 10 lines in your changes missing coverage. Please review.
✅ Project coverage is 79.54%. Comparing base (f85373c) to head (6d6e528).

Files with missing lines Patch % Lines
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 96.66% 7 Missing and 1 partial ⚠️
...s/onnx-runtime-ep-cpu/src/kernels/contrib_fused.rs 88.88% 1 Missing ⚠️
...tes/onnx-runtime-ep-cpu/src/kernels/elementwise.rs 50.00% 0 Missing and 1 partial ⚠️
Additional details and impacted files

Impacted file tree graph

@@           Coverage Diff           @@
##             main    #1074   +/-   ##
=======================================
  Coverage   79.53%   79.54%           
=======================================
  Files         366      366           
  Lines      156827   156889   +62     
  Branches   156827   156889   +62     
=======================================
+ Hits       124737   124793   +56     
- Misses      27387    27396    +9     
+ Partials     4703     4700    -3     
Flag Coverage Δ
cli-ort-linux 83.70% <ø> (ø)
cli-ort-windows 83.21% <ø> (ø)
offline 79.39% <96.15%> (+<0.01%) ⬆️

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

Files with missing lines Coverage Δ
...rates/onnx-runtime-ep-cpu/src/assignment_policy.rs 98.56% <ø> (ø)
crates/onnx-runtime-ep-cpu/src/kernels/gelu.rs 88.23% <100.00%> (-0.16%) ⬇️
...s/onnx-runtime-ep-cpu/src/kernels/contrib_fused.rs 86.16% <88.88%> (-0.09%) ⬇️
...tes/onnx-runtime-ep-cpu/src/kernels/elementwise.rs 90.33% <50.00%> (-0.98%) ⬇️
...nnx-runtime-ep-cpu/src/kernels/simd_activations.rs 98.68% <96.66%> (-0.65%) ⬇️

... and 6 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 16, 2026 •

Copy link
Copy Markdown

⚠️ Benchmark Change 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
⚠️ grammar_masking/llguidance_compute_mask/32 75.50 µs 97.94 µs +29.7%
⚠️ block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 788.57 µs 949.15 µs +20.4%
⚠️ logit_processing/seven_processor_chain_per_step 320.52 µs 383.99 µs +19.8%
⚠️ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.67 ms 6.61 ms +16.5%
✅ kv_cache/alloc_dealloc_pages 38.95 µs 44.43 µs +14.1%
✅ qwen3_sampling_processors/top_k_top_p_fast 664.28 µs 755.02 µs +13.7%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 522.70 µs 582.36 µs +11.4%
✅ add/small_f16_threads=1-internal/1024 460.2 ns 492.0 ns +6.9%
✅ qwen3_sampling_processors/top_k_partial_selection 139.53 µs 146.89 µs +5.3%
✅ sampling_latency/top_p_per_token 383.93 µs 397.23 µs +3.5%
✅ tokenization/encode_tokens_per_second 384.86 µs 397.71 µs +3.3%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 177.12 µs 182.39 µs +3.0%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 2.06 ms 2.11 ms +2.5%
✅ matmul/large_generic_bf16_threads=8/32x1024x1024 2.09 ms 2.13 ms +2.1%
✅ tokenization/decode_tokens_per_second 6.44 ms 6.55 ms +1.7%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 535.43 µs 544.26 µs +1.6%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.34 ms 2.38 ms +1.6%
✅ matmul/small_generic_bf16_threads=8/1x256x256 37.56 µs 38.15 µs +1.6%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 594.50 µs 602.98 µs +1.4%
✅ block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 85.85 µs 86.87 µs +1.2%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.51 ms 3.55 ms +1.1%
✅ sampling_latency/min_p_per_token 210.63 µs 212.92 µs +1.1%
✅ matmul/small_generic_bf16_threads=1/1x256x256 35.83 µs 36.15 µs +0.9%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 89.85 µs 90.66 µs +0.9%
✅ sampling_latency/greedy_per_token 3.20 µs 3.22 µs +0.8%
✅ matmul/small_generic_f16_threads=8/1x256x256 37.86 µs 38.05 µs +0.5%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 2.16 ms 2.17 ms +0.4%
✅ sampling_latency/top_k_per_token 51.87 µs 51.99 µs +0.2%
✅ matmul/large_generic_f32_threads=8/32x1024x1024 6.16 ms 6.17 ms +0.1%
✅ add/large_bf16_threads=1-internal/4194304 1.66 ms 1.66 ms -0.0%
✅ matmul/medium_generic_f16_threads=8/32x512x512 42.48 µs 42.16 µs -0.8%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 99.22 µs 98.46 µs -0.8%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.54 ms 9.44 ms -1.1%
✅ reduce_mean/large_f32_threads=1-internal/262144 1.01 ms 996.41 µs -1.3%
✅ matmul/medium_generic_f16_threads=1/32x512x512 35.44 µs 34.96 µs -1.4%
✅ matmul/small_generic_f16_threads=1/1x256x256 36.39 µs 35.83 µs -1.5%
✅ reduce_mean/medium_f32_threads=1-internal/65536 251.93 µs 247.96 µs -1.6%
✅ reduce_mean/small_f32_threads=1-internal/4096 15.13 µs 14.88 µs -1.7%
✅ gather/medium_bf16_threads=1-internal/32768 2.40 µs 2.35 µs -2.0%
✅ gather/medium_f16_threads=1-internal/32768 2.40 µs 2.34 µs -2.6%
✅ gather/small_bf16_threads=1-internal/4096 494.3 ns 475.5 ns -3.8%
✅ gather/large_f16_threads=1-internal/131072 13.89 µs 13.33 µs -4.0%
✅ gather/medium_f32_threads=1-internal/32768 4.10 µs 3.90 µs -4.7%
✅ add/large_f16_threads=1-internal/4194304 1.79 ms 1.70 ms -4.8%
✅ matmul/small_generic_f32_threads=1/1x256x256 43.90 µs 40.87 µs -6.9%
✅ matmul/medium_generic_f32_threads=8/32x512x512 1.65 ms 1.53 ms -7.6%
✅ matmul/small_generic_f32_threads=8/1x256x256 50.42 µs 46.45 µs -7.9%
✅ matmul/medium_generic_bf16_threads=8/32x512x512 635.35 µs 579.47 µs -8.8%
✅ gather/large_bf16_threads=1-internal/131072 14.07 µs 12.80 µs -9.1%
✅ add/medium_bf16_threads=1-internal/262144 113.25 µs 102.47 µs -9.5%
✅ add/medium_f32_threads=1-internal/262144 27.52 µs 24.50 µs -10.9%
✅ add/medium_f16_threads=1-internal/262144 116.55 µs 103.27 µs -11.4%
✅ gather/large_f32_threads=1-internal/131072 39.88 µs 34.84 µs -12.6%
✅ gather/small_f32_threads=1-internal/4096 792.4 ns 692.2 ns -12.7%
✅ add/large_f32_threads=1-internal/4194304 949.24 µs 828.48 µs -12.7%
✅ gather/small_f16_threads=1-internal/4096 572.1 ns 489.3 ns -14.5%
🟢 add/small_bf16_threads=1-internal/1024 533.0 ns 446.5 ns -16.2%
🟢 block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 116.12 µs 95.67 µs -17.6%
🟢 add/small_f32_threads=1-internal/1024 275.1 ns 193.3 ns -29.7%

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: { 6.43 4.14 5.63 }
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)

Independent review of #1074 returned four findings; this applies all four.

* Restore the `#[inline]` attribute on `quick_gelu_ps`. An earlier edit in
  this branch joined it onto the end of the preceding doc-comment line, so it
  silently became comment text. Neither `fmt` nor `clippy` catches that.
* Transcribe `EXP_LOWER_RANGE` exactly as MLAS spells it
  (`-88.3762626647949f`, bits `0xc2b0c0a5`). The previous `-88.376_264`
  rounded to `0xc2b0c0a6`, the only constant of the 26 that was not
  bit-identical. It is provably unreachable for `erf` -- the big branch's
  `R` maxes at 17.375652 at the `|x| = 3.925` clamp -- so no output changes,
  but the module's contract is verbatim MLAS.
* Document the `SIMD_MIN_LEN` seam that `Erf` now inherits: below 32
  elements the correctly-rounded scalar fallback runs, so the same value can
  differ by <=1 ulp with tensor length. `Tanh` and `Sigmoid` have had this
  since #1037 and ORT has it too.
* Correct the float32 `Erf` 1048576 ratio in the policy rationale from a
  dispersion-limited 0.58 to a reps=31 re-measurement of 0.687.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby
justinchuby marked this pull request as ready for review August 16, 2026 15:40
@justinchuby
justinchuby merged commit 8635e4d into main Aug 16, 2026
11 of 17 checks passed
@justinchuby
justinchuby deleted the deckard/erf-avx2 branch August 16, 2026 16:07
justinchuby added a commit that referenced this pull request Aug 18, 2026
…ap_bias_ps (#1227)

## What this is

`map_ps` is the loop driver behind every unary activation on AVX2. A
merge
(the `#1037` / `#1074` layering) inserted `map_bias_ps` *between*
`map_ps`'s
doc comment and `map_ps` itself, so the comment and **both of its
attributes**
ended up on the wrong function:

```rust
/// Apply an 8-lane kernel across a slice. ...   <- map_ps's doc
#[inline]
#[target_feature(enable = "avx2,fma")]           <- map_ps's attributes
/// Like [`map_ps`], but adds a bias row ...
pub(super) unsafe fn map_bias_ps(...)            <- ...on map_bias_ps

pub(super) unsafe fn map_ps(...)                 <- bare
```

`map_bias_ps` needs those attributes too, so nothing was broken outright
—
but `map_ps` was left with no `#[inline]` and no `avx2,fma`.

## Why the missing `target_feature` is a hazard

LLVM will not inline a callee that requires a feature its caller lacks
(a
*strict* superset; equal feature sets inline fine). Every kernel closure
`map_ps` takes (`erf_ps`, `tanh_ps`,
`sigmoid_ps`, ...) is `avx2,fma`. With `map_ps` compiled at baseline
features,
those closures are not inlinable into its loop.

It is harmless **today** only by luck of ordering: every caller
(`tanh_avx2`, `erf_avx2`, ...) is itself `avx2,fma`, so `map_ps` gets
inlined
*upwards* into the caller first, and once there the closure folds in.
That is
an inlining decision, not a guarantee. Grow `map_ps` past the inline
threshold
— which is exactly what an unroll experiment does — and the closure
becomes a
real call per 8 elements, and a kernel like `erf_ps` has to
re-materialise all
of its constants on every one of those calls.

I hit this while unrolling `map_ps`, which is how it surfaced.

## Evidence that this changes nothing today

Disassembled `tanh_avx2` out of the built rlib before and after:

```
diff <(objdump -d ... base) <(objdump -d ... fixed)   ->  identical
```

Byte-identical, and `erf_avx2` / `sigmoid_avx2` keep their instruction
counts
(216 / 117). No `call` to any kernel remains in the lib. So this is a
latent-hazard and documentation fix with **zero** codegen delta — not a
perf
change, and it needs no A/B.

## Also recorded here: a negative result on the divide

While looking at these loops I measured one algorithmic change and it
lost, so
it is written down rather than left for someone to retry.

`tanh_ps` and `sigmoid_ps` each end in `_mm256_div_ps`. `vdivps ymm` is
the
only instruction in either loop that is not fully pipelined (Zen 4
retires one
about every 4.5 cycles), so replacing it with `vrcpps` + one
Newton-Raphson
step — `y1 = y0 + y0*(1 - q*y0)`, four pipelined uops for one blocking
one —
looked like the obvious win.

It is not:

| case | ORT p50 (control) | divide | rcp + NR |
|---|---|---|---|
| `bench_tanh_f32_4k` | 0.0031 ms, all 6 runs | 1.518 / 1.525 / 1.526 |
1.550 / 1.570 / 1.570 |

Interleaved rebuild-and-alternate, 3 rounds, 1 thread, `taskset -c
8-15`.
The 4k case is the only trustworthy one here — it is L1-resident and
ORT's own
p50 was identical to four decimal places across all six runs, whereas on
the
1M cases ORT's p50 moved 4.6x mid-session and neither arm is quotable.

Consistent ~3% **regression**. The `rcp` + 2 FMA dependency chain is
about as
long as the divide it replaces, so a latency-bound kernel gains nothing
and
just pays two extra uops.

It also breaks four exactness proofs that #1121 landed: `tanh(±Inf)`
comes back
as `0.99999994` instead of exactly `±1.0`, because the input clamp makes
`p`
and `q` bit-equal at `|v| = 9` and only an exact divide turns that into
exactly
`1.0`. Buying that back would mean restoring the saturation blend #1121
proved
redundant over 1.05 billion inputs — to fund a change that is already
slower.
Not pursued. No fallback involved; the kernel keeps its own divide.

## Validation

- `cargo test --release -p onnx-runtime-ep-cpu --lib` — 1337 passed, 0
failed
- `cargo clippy --release -p onnx-runtime-ep-cpu --all-targets` — clean
- `cargo fmt --all --check` — clean
- `tanh_avx2` disassembly identical to `main`

## Review

Opus review: no blockers, no should-fixes.

It independently confirmed the three things the PR rests on: every one
of the
16 `map_ps` and 2 `map_bias_ps` call sites is already inside a
`#[target_feature(enable = "avx2,fma")]` function (so nothing becomes
unsound
or fails to compile), `map_bias_ps` keeps exactly the attributes it had,
and
there are **no bare/function-pointer references** to either.

It also did the scan I most wanted: all ~130 `fn` definitions in the
file were
checked for the same orphaned-doc-comment-after-attribute signature that
caused
this bug. **Zero other anomalies** — every `*_ps` helper and every
`*_avx2`
dispatcher carries the attributes it should. So this is the only
instance.

- **NIT (applied)** — the doc comment said LLVM won't inline a callee
whose
features are "a superset" of the caller's. Read non-strictly that would
forbid inlining at *equal* feature sets, which is the opposite of what
the
fix relies on. Reworded to "requires a feature its caller lacks — a
*strict*
  superset; equal feature sets inline fine".

---------

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