Skip to content

perf(mlas): let MlasGemmBatch use the registered parallel backend (4.4x dense f32 MatMul, reaches ORT parity) - #1045

Merged
justinchuby merged 4 commits into
mainfrom
squad/roy-mlas-sgemm-threadpool
Aug 16, 2026
Merged

justinchuby merged 4 commits into
mainfrom
squad/roy-mlas-sgemm-threadpool

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 15, 2026 •

Copy link
Copy Markdown
Owner

Summary

mlas_sgemm and mlas_sgemm_packed passed /*ThreadPool=*/nullptr to
MlasGemmBatch. In the standalone (BUILD_MLAS_NO_ONNXRUNTIME) build that
argument is not an inert handle — it is the parallelism enable flag:

// vendor/mlas/.../lib/threading.cpp — MlasTrySimpleParallel
MlasStandaloneParallelFor(Iterations, &Work, ThreadPool != nullptr);
//                                           ^^^^^^^^^^^^^^^^^^^^^ enable_backend
// vendor/shim.cpp — MlasStandaloneParallelFor
if (enable_backend && g_parallel_for != nullptr && iterations > 1) {
    g_parallel_for(...);        // registered Rust work-stealing pool
} else {
    for (tid = 0; tid < iterations; ++tid) fn(tid);   // SERIAL
}

So every f32 GEMM routed through the mlas backend ran single-threaded.
Because CpuBackend::auto_detect() prefers Mlas when the feature is enabled,
building with --features mlas made dense f32 MatMul slower than the
built-in SimdX86 backend at every thread count above one.

The QNBit shim immediately below already documents this exact hazard and passes
a non-null sentinel; SGEMM never got the same treatment. This lifts that
sentinel into a named constant and applies it to both MlasGemmBatch sites.

The bug, measured

Backend comparison, MatMul f32 K=3584 N=3584 M=128, native_min medians:

backend 1 thread 8 threads scaling
generic 741.9 ms 167.8 ms 4.4x
simd 52.6 ms 11.8 ms 4.5x
mlas (default with the feature on) 38.2 ms 42.5 ms 0.90x

mlas had the best single-thread kernel and the only broken scaling — and at
8 threads it was 3.6x slower than the pure-Rust SimdX86 backend it
displaces.

Result

Interleaved A/B, native_min medians, 25 runs + 5 warmups, thread-matched,
real ORT CPU EP measured in the same process.

Thread scaling, MatMul f32 K=3584 N=3584 M=128:

threads before after speedup ORT after/ORT
1 38.25 ms 38.27 ms 1.00x 38.27 ms 1.00x
2 42.17 37.14 1.14x 21.15 1.76x
4 38.28 18.79 2.04x 18.60 1.01x
8 42.29 9.60 4.40x 9.33 1.03x
16 40.30 11.52 3.50x 4.84 2.38x

Per-shape at 8 threads, 3 interleaved reps each:

shape before after speedup ORT before/ORT after/ORT
MatMul f32 K=3584 N=3584 M=128 42.51 ms 9.60 ms 4.43x 9.35 4.5x 1.0x
MatMul f32 K=3584 N=3584 M=128 42.39 9.54 4.44x 9.35 4.5x 1.0x
MatMul f32 K=1024 N=3072 M=128 11.01 2.46 4.47x 2.51 4.4x 1.0x
MatMul f32 K=1024 N=3072 M=128 11.16 3.18 3.51x 1.58 7.1x 2.0x

Parity PASS on every run.

Dense f32 MatMul now matches the ORT CPU EP at 4 and 8 threads, where it
was 4.4–7.1x behind.

Honest limits

  • 16 threads still lags (11.5 vs ORT 4.84 = 2.4x). ORT keeps scaling past
    8; we do not. That is a separate partitioning/pool issue and is not fixed
    here. Reported, not hidden.
  • 1 thread is unchanged (1.00x) — as it must be. The single-thread kernel
    was never the problem; it was already at ORT parity.
  • Gemm is unaffected (measured 1.01x / 1.02x / 0.96x — noise). Gemm
    does not route through this path on main; that is Gemm: route the f32 path through the shared GEMM backend (up to 178x faster) #1035's job. Included
    here as a control showing the change is scoped to what it claims. Those same
    runs show Gemm at 199–216x and Gemm(transB) at 1342x of ORT, which is
    the strongest argument yet for Gemm: route the f32 path through the shared GEMM backend (up to 178x faster) #1035.
  • No precision change: identical kernel, identical accumulation type. Only the
    partitioning across threads differs, and MLAS assigns each thread a disjoint
    output tile.

Safety of the sentinel

The sentinel is never dereferenced in the standalone build:

  • MlasGetMaximumThreadCount — MLAS_UNREFERENCED_PARAMETER(ThreadPool), returns
    MlasStandaloneMaxThreads().
  • MlasTrySimpleParallel — MLAS_UNREFERENCED_PARAMETER(ThreadPool), only
    ThreadPool != nullptr.
  • The single override that forwards it, ArmKleidiAI::MlasGemmBatch
    (arm64 SME/SME2 only, USE_KLEIDIAI), passes it on to exactly those two
    helpers and nothing else.

This is the same sentinel value and the same reasoning already used by
mlas_qnbit_gemm, mlas_conv, and the NCHWc entry points in this shim.

Tests

sgemm_nn_drives_the_registered_parallel_backend — deterministic, not a
timing check. Asserts sgemm_nn increments the backend's parallel_for_calls
counter. Verified non-vacuous: with the sentinel reverted it fails with

sgemm_nn did not enter the parallel-for backend (0 -> 0 calls);
MlasGemmBatch was likely handed a null MLAS_THREADPOOL, which forces
MLAS's serial fallback

Skips cleanly when mlas_threading_degree() < 2.

sgemm_nn_is_correct_when_parallelized — odd, non-tile-multiple dims
(129x257x193) against a scalar oracle, guarding against torn or
doubly-written output tiles from the now-active partitioning.

Why this was not caught: the pre-existing perf_sgemm_multithread probe
printed the serial numbers next to a comment recording ORT's ~4.4x scaling,
but asserted nothing and is #[ignore]d.

cargo test -p mlas-sys --release --lib     # 28 passed, 0 failed, 3 ignored

Follow-ups (deliberately not in this PR)

  • mlas_qnbit_gemm_pack_b and MlasReorderOutputNchw still pass nullptr.
    Those are one-time packing/reorder rather than the measured hot path;
    they deserve their own measurement.
  • The >8-thread scaling gap above.
  • CpuBackend::auto_detect() preferring Mlas was actively harmful before
    this fix. Worth a guard so a backend can never be selected when it is
    measurably slower than the built-in one.

Reproduce

ORT_ROOT=<ort-prebuilt> cargo build --release -p onnx-genai-bench \
  --features mlas --bin bench_generic
LD_LIBRARY_PATH=<ort-prebuilt>/lib ./bench_generic \
  --model matmul_f32_k3584_n3584_m128.onnx \
  --runs 25 --warmups 5 --native-threads 8 --ort-intra-threads 8

Host: AMD EPYC 9V74, AVX2+FMA+F16C, no AVX-512/VNNI/AMX. Shared and contended,
so ratios are trustworthy and absolutes are not; every number above is an
interleaved A/B median.

Requires #1025 (harness) for the thread-matching flags used above.

Review

Two independent review rounds, both APPROVE WITH FINDINGS; every finding
resolved and re-reviewed.

Round 1 verified the root cause, the disjointness of MLAS's output tiles,
and the sentinel's safety — and established something stronger than I had:
ArmKleidiAI::MlasGemmBatch is never compiled at all (sgemm_kleidiai.cpp
is not in build.rs), so the only MlasSGemmBatchOverride that could forward
the sentinel does not exist in this build. Findings fixed in 44fa4f1d7:

  • The threading guard read a process-global counter while living in the
    crate's unit tests, which cargo runs concurrently in one process. A
    concurrent sqnbit_gemm — which passes its own sentinel and does drive
    the backend — could have bumped the counter between samples and let a broken
    build pass. Moved to tests/sgemm_threading.rs, its own binary/process.
  • The correctness oracle's tolerance could mask a zeroed tile at
    small-magnitude outputs.

Round 2 found the replacement tolerance was now too tight. Fixed in
00caeb923:

  • 8 * EPSILON is only valid for k <= 8; for k = 193 the standard bound
    is gamma_k ~ k * EPSILON, ~12x looser. An over-tight bound would be flaky
    on any ISA that reassociates the sum differently (AVX-512, SVE, different
    blocking). Now scales with k.
  • The guard's degree < 2 early-return would have made it a silent no-op on a
    single-core runner. It now forces ONNX_GENAI_MLAS_THREADPOOL_THREADS=4
    before first use and asserts, rather than skipping.

Both tests re-verified non-vacuous after every change:

  • revert the sentinel → threading guard fails 0 -> 0 calls
  • inject a zeroed 16x16 output tile → oracle fails at (64,100), |error|
    0.789 vs tolerance 0.0075 (105x margin), confirming the looser bound still
    catches tile corruption

CI

main is red before this branch for pre-existing, unrelated reasons. This
branch is based on 869c24b83, which includes #1043 (cargo fmt --all over
the workspace), so the Check formatting failures that affect my older PRs do
not apply here.

Locally green:

cargo test -p mlas-sys --release          # 27 + 1 passed, 0 failed, 3 ignored
cargo test -p mlas-sys                    # 27 + 1 passed, 0 failed (debug)
cargo fmt -p mlas-sys -- --check          # clean
cargo clippy --release -p mlas-sys --all-targets   # zero warnings

CI baseline

Rebuilt on top of main @ 400fbe246 (a plain git merge origin/main, no history rewrite).

Unmodified main @ 400fbe246 fails exactly these 6 jobs
(run 31914831964):

Job Fails on unmodified main
CLI ORT (Linux x86_64) yes
CLI ORT (Windows x86_64) yes
CUDA compile (Linux x86_64) yes
CUDA compile (Windows x86_64) yes
Rust (Windows ARM64) yes
Rust coverage (macOS arm64) yes

None are touched by this PR. Fast (Linux x86_64) and Rust quality previously failed on
main too (a repo-wide cargo fmt drift, fixed on main by #1043); after merging current
main into this branch both are green here, which confirms those earlier reds were never mine.

The jobs this PR is actually accountable for -- Fast (Linux x86_64), Rust quality,
EP conformance (Linux x86_64), Rust coverage (Linux x86_64), Miri unsafe-crate soundness,
audit and codecov -- are green.


Ratio convention (added post-merge for clarity)

speedup is ours-before/ours-after. after/ORT is ours/ORT: >1 means we are slower, so 1.00x/1.01x/1.03x are parity and 1.76x at 2 threads is a loss. "4.4x" in the title is before/after within this EP, not against ORT. p50, interleaved, per-row thread counts as tabulated, steady state (one-time packing measured separately and not folded in).

`mlas_sgemm` and `mlas_sgemm_packed` passed `/*ThreadPool=*/nullptr` to
`MlasGemmBatch`. In the standalone (`BUILD_MLAS_NO_ONNXRUNTIME`) build that
argument is not merely an unused handle: `MlasTrySimpleParallel` forwards
`ThreadPool != nullptr` to `MlasStandaloneParallelFor` as the *enable* flag,
so a null pointer selects MLAS's serial fallback loop and the registered
work-stealing backend is never entered.

The result was that every f32 GEMM routed through the `mlas` backend ran
single-threaded. Because `CpuBackend::auto_detect()` prefers `Mlas` when the
feature is enabled, building with `--features mlas` made dense f32 MatMul
*slower* than the built-in `SimdX86` backend at every thread count above one.

Measured on AVX2+FMA (AMD EPYC 9V74, 8 threads), MatMul f32 K=3584 N=3584
M=128, `native_min` medians:

    backend    1 thread    8 threads   scaling
    generic     741.9 ms     167.8 ms    4.4x
    simd         52.6 ms      11.8 ms    4.5x
    mlas         38.2 ms      42.5 ms    0.90x   <- serial

The QNBit shim directly below already documents this exact hazard and passes
a non-null sentinel; SGEMM never got the same treatment. This lifts that
sentinel into a named constant and applies it to both `MlasGemmBatch` sites.

The sentinel is never dereferenced in the standalone build:
`MlasGetMaximumThreadCount` ignores it and reports `MlasStandaloneMaxThreads()`,
and `MlasTrySimpleParallel` only compares it against null. The one override
that forwards it (`ArmKleidiAI::MlasGemmBatch`, arm64 SME only) likewise only
passes it to those same two helpers.

Adds a deterministic regression test that asserts `sgemm_nn` increments the
backend's `parallel_for_calls` counter, rather than a timing check. Verified
non-vacuous: with the sentinel reverted it fails with `0 -> 0 calls`. The
pre-existing `perf_sgemm_multithread` probe printed the serial numbers but
asserted nothing, which is why this went unnoticed.

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

codecov Bot commented Aug 15, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 79.40%. Comparing base (400fbe2) to head (123a23a).
⚠️ Report is 1 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1045      +/-   ##
==========================================
+ Coverage   78.78%   79.40%   +0.61%     
==========================================
  Files         365      365              
  Lines      149050   149070      +20     
  Branches   149050   149070      +20     
==========================================
+ Hits       117435   118371     +936     
+ Misses      26978    26060     -918     
- Partials     4637     4639       +2     
Flag Coverage Δ
cli-ort-linux 83.52% <ø> (ø)
cli-ort-windows 83.11% <ø> (ø)
mlas 82.14% <100.00%> (+1.22%) ⬆️
offline 79.20% <ø> (+0.63%) ⬆️

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 80.75% <100.00%> (+1.37%) ⬆️

... and 8 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 15, 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
🔴 gather/large_f16_threads=1-internal/131072 12.56 µs 26.23 µs +108.8%
🔴 gather/large_f32_threads=1-internal/131072 32.73 µs 52.31 µs +59.8%
🔴 matmul/medium_generic_f32_threads=8/32x512x512 919.54 µs 1.34 ms +46.0%
🔴 matmul/small_generic_f32_threads=8/1x256x256 34.32 µs 48.92 µs +42.6%
🔴 matmul/small_generic_bf16_threads=8/1x256x256 31.24 µs 42.15 µs +34.9%
🔴 reduce_mean/medium_f32_threads=1-internal/65536 247.26 µs 325.11 µs +31.5%
⚠️ gather/medium_bf16_threads=1-internal/32768 2.43 µs 3.11 µs +27.9%
⚠️ matmul/small_generic_f16_threads=8/1x256x256 30.05 µs 37.99 µs +26.4%
⚠️ matmul/small_generic_f16_threads=1/1x256x256 31.01 µs 38.27 µs +23.4%
⚠️ reduce_mean/small_f32_threads=1-internal/4096 15.38 µs 18.81 µs +22.3%
⚠️ add/large_bf16_threads=1-internal/4194304 43.08 ms 52.35 ms +21.5%
⚠️ gather/large_bf16_threads=1-internal/131072 12.70 µs 15.42 µs +21.4%
⚠️ gather/medium_f32_threads=1-internal/32768 3.82 µs 4.61 µs +20.7%
⚠️ gather/medium_f16_threads=1-internal/32768 2.44 µs 2.85 µs +16.5%
⚠️ matmul/small_generic_f32_threads=1/1x256x256 36.49 µs 42.26 µs +15.8%
✅ sampling_latency/top_p_per_token 473.10 µs 532.39 µs +12.5%
✅ gather/small_f16_threads=1-internal/4096 491.2 ns 533.7 ns +8.7%
✅ sampling_latency/top_k_per_token 54.18 µs 57.39 µs +5.9%
✅ matmul/small_generic_bf16_threads=1/1x256x256 32.09 µs 33.96 µs +5.8%
✅ add/large_f32_threads=1-internal/4194304 42.14 ms 44.49 ms +5.6%
✅ sampling_latency/greedy_per_token 3.69 µs 3.88 µs +5.1%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.40 ms 2.49 ms +3.9%
✅ gather/small_bf16_threads=1-internal/4096 488.2 ns 506.8 ns +3.8%
✅ add/large_f16_threads=1-internal/4194304 41.46 ms 42.69 ms +3.0%
✅ reduce_mean/large_f32_threads=1-internal/262144 1.07 ms 1.10 ms +2.0%
✅ kv_cache/alloc_dealloc_pages 40.60 µs 39.90 µs -1.7%
✅ sampling_latency/min_p_per_token 250.94 µs 245.36 µs -2.2%
✅ add/small_bf16_threads=1-internal/1024 14.94 µs 14.24 µs -4.7%
✅ tokenization/decode_tokens_per_second 7.20 ms 6.84 ms -5.1%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.95 ms 3.73 ms -5.4%
✅ matmul/medium_generic_f16_threads=1/32x512x512 40.28 µs 37.44 µs -7.0%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 600.43 µs 557.78 µs -7.1%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 583.98 µs 538.84 µs -7.7%
✅ matmul/medium_generic_bf16_threads=8/32x512x512 419.23 µs 384.12 µs -8.4%
✅ add/medium_bf16_threads=1-internal/262144 2.88 ms 2.64 ms -8.4%
✅ qwen3_sampling_processors/top_k_partial_selection 174.10 µs 158.51 µs -9.0%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 93.63 µs 84.83 µs -9.4%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 6.29 ms 5.70 ms -9.4%
✅ add/small_f16_threads=1-internal/1024 15.58 µs 13.99 µs -10.2%
✅ qwen3_sampling_processors/top_k_top_p_fast 734.44 µs 657.95 µs -10.4%
✅ grammar_masking/llguidance_compute_mask/32 83.95 µs 75.20 µs -10.4%
✅ add/medium_f16_threads=1-internal/262144 2.96 ms 2.65 ms -10.5%
✅ tokenization/encode_tokens_per_second 473.97 µs 421.79 µs -11.0%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 93.30 µs 82.78 µs -11.3%
✅ logit_processing/seven_processor_chain_per_step 364.85 µs 321.91 µs -11.8%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 103.99 µs 89.56 µs -13.9%
✅ matmul/medium_generic_f16_threads=8/32x512x512 45.82 µs 39.44 µs -13.9%
✅ gather/small_f32_threads=1-internal/4096 845.0 ns 725.1 ns -14.2%
🟢 block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 117.28 µs 98.48 µs -16.0%
🟢 matmul/large_generic_bf16_threads=1/32x1024x1024 2.35 ms 1.96 ms -16.6%
🟢 matmul/large_generic_f32_threads=8/32x1024x1024 4.53 ms 3.73 ms -17.8%
🟢 add/medium_f32_threads=1-internal/262144 3.07 ms 2.50 ms -18.4%
🟢 add/small_f32_threads=1-internal/1024 288.9 ns 208.2 ns -27.9%
🟢 block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 553.96 µs 393.53 µs -29.0%
🟢 matmul/large_generic_f32_threads=1/32x1024x1024 13.96 ms 9.35 ms -33.0%
🟢 qwen3_sampling_processors/top_k_full_sort_baseline 3.70 ms 2.36 ms -36.3%
🟢 matmul/large_generic_bf16_threads=8/32x1024x1024 2.32 ms 1.35 ms -41.7%
🟢 block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 1.26 ms 569.37 µs -54.8%
🟢 block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 135.64 µs 57.48 µs -57.6%

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.10 4.42 5.96 }
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 15, 2026 23:06
… oracle

Review findings from the independent review of #1045.

1. The `parallel_for_calls` counter is process-global and `cargo test` runs a
   crate's unit tests concurrently in one process, so a concurrent
   `sqnbit_gemm` -- which passes its own non-null sentinel and does drive the
   backend -- could bump the counter between the two samples and let a broken
   build pass. Moved the guard into `tests/sgemm_threading.rs`, which cargo
   compiles into its own binary, so the sampled calls are the only MLAS work
   in the process. Re-verified non-vacuous after the move: reverting the
   sentinel still fails it with `0 -> 0 calls`.

2. The correctness oracle used `1e-3 * max(|want|, 1.0)`. That bound grows
   looser as outputs grow and, more importantly, collapses to a fixed 1e-3
   for small outputs, so a zeroed tile covering small-magnitude entries could
   pass. Replaced with a per-element bound derived from the quantity f32
   rounding actually accumulates, `8 * EPSILON * sum|a*b|`.

3. Documented the inherited nested-parallelism constraint: the work-stealing
   pool takes a dispatch lock, so invoking these entry points from inside a
   work item already running on that pool would deadlock. No call site does
   this today; noted for future callers.

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

Second review round on #1045.

1. The oracle's `8 * EPSILON * sum|a*b|` bound was not justified for k=193.
   The standard forward error bound for a length-k f32 dot product is
   `gamma_k = k*u/(1 - k*u)` with `u = EPSILON/2`, i.e. roughly
   `k * EPSILON` -- about 12x looser than what was written. MLAS is free to
   reassociate the sum (blocked accumulation, FMA contraction, wider vectors
   on AVX-512 or SVE), so a bound derived from one host's accumulation order
   would be flaky on another. Now scales with k.

   Re-verified the looser bound still does its job by injecting a zeroed
   16x16 output tile: caught at (64,100) with |error| 0.789 against a
   tolerance of 0.0075, a 105x margin.

2. The guard skipped itself when `mlas_threading_degree() < 2`, which would
   turn it into a no-op on a single-core runner -- precisely where a
   threading regression would hide. Since the test now owns its process and
   the pool is a `OnceLock` built on first use, it sets
   `ONNX_GENAI_MLAS_THREADPOOL_THREADS=4` up front and asserts the degree
   instead of skipping.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby
justinchuby marked this pull request as ready for review August 15, 2026 23:53
@justinchuby
justinchuby merged commit 15037e6 into main Aug 16, 2026
10 of 16 checks passed
@justinchuby
justinchuby deleted the squad/roy-mlas-sgemm-threadpool branch August 16, 2026 00:48
justinchuby added a commit that referenced this pull request Aug 17, 2026
…t copy) (#1116)

## Summary

#1045 won **4.4x on dense f32 MatMul with `--features mlas`**; #1091
asks to make that real for a **default (no-mlas) build**, by absorbing
the mechanism into our own `SimdX86` kernel rather than shipping behind
MLAS. This PR measures the gap on this host, finds where it comes from,
proves it is reachable **without** MLAS's session-lifetime packed
buffer, and ports it. Same shape as #1104.

**The 4.4x prefill number does not reproduce on this host.** In one
binary containing both paths (same-binary A/B via the existing
`NXRT_CPU_GEMM_BACKEND=mlas|simd` toggle), at **M=128 prefill our
built-in `SimdX86` 6×16 packed microkernel is already at parity with
MLAS** (0.87–1.15x, net slightly favoring `SimdX86`). #1045's 4.4x was
an AMD EPYC without AVX-512; here it's gone.

**The entire reproducible gap is at M=1 decode GEMV: 2.2–4.6x.**

| shape (M×K×N) | M | simd/mlas (before) |
| --- | --- | --- |
| 1×5120×5120 (o_proj) | 1 | **4.59x** |
| 1×5120×7168 (qkv) | 1 | 3.58x |
| 1×5120×13824 (gate/up) | 1 | 3.44x |
| 1×13824×5120 (down) | 1 | 2.49x |
| 1×5120×152064 (lm_head) | 1 | 2.23x |
| 128×5120×5120 | 128 | 0.88x (simd faster) |
| 128×5120×13824 | 128 | 0.87x (simd faster) |
| 128×13824×5120 | 128 | 1.15x |

## Mechanism — source-cited, both sides

MLAS `sgemm.cpp` (vendored, `MlasSgemmOperation`):

```
// Handle the special case of a small M. The data from matrix B is not
// referenced multiple times, so using a local packed buffer is a wasted
// memory copy.
if (M == 1 && TransA == CblasNoTrans && alpha == 1.0f && ...) {
    SgemmKernelM1Routine(A, B, C, K, N, ldb, beta);   // reads B in place at stride ldb
    return;
}
```

MLAS routes M==1 to `SgemmKernelM1Avx.asm`, which streams B **in place**
— K unrolled ×4 (`ProcessRowLoop4`), N swept contiguously
(`ProcessColumnLoop`) — **no pack, no resident buffer**.

Ours (`x86_sgemm.rs::sgemm_simd`) calls `pack_b` into a `bpack` scratch
**unconditionally**. At M=1 there is a single A-panel, so each packed B
panel is reused **zero** times — the pack is a wasted full read+write
copy of B (K·N f32), ≈3× the memory traffic of a straight GEMV. **It is
memory traffic, not arithmetic and not layout**, and the fix needs no
resident buffer.

## How much is reachable without a resident copy — all of it

`sgemm_simd_m1`: for M==1, stream B exactly once (K unrolled ×4,
sequential N sweep, C accumulated in cache, Rayon over disjoint column
strips). **No `pack_b`, no scratch, no `OnceLock`, no
`GovernedWeightCache`** — it actually *removes* the `bpack` allocation
at M=1. Exactly #1104's "no resident copy" property; nothing to
admit/decline under #1056.

The first attempt (column-major, C in registers) regressed lm_head to
3.72x because it strided B by N (608 KB stride → TLB thrash). Matching
MLAS's **K-outer / N-inner sequential** traversal fixed it — the layout
that matters for wide outputs.

## A/B result — process CPU time and peak RSS

Same binary, `SimdX86` M=1 route toggled by
**`ONNX_GENAI_CPU_MM_SIMD_M1_GEMV`** (default **off**, like #1104's
`ONNX_GENAI_CPU_MM_INT4_NBLK`). One arm per process; peak RSS polled by
PID every 150 ms; **process CPU time** (`TotalProcessorTime`); 5 decode
shapes, min-of-30.

| arm | process CPU time | peak RSS |
| --- | --- | --- |
| MLAS (`SgemmKernelM1`) | 52.5 s | 2978 MB |
| ours, packed (toggle off) | 169.4 s | 2982 MB |
| **ours, GEMV (toggle on)** | **57.3 s** | **2977 MB** |

**2.96× faster than the packed path** (169.4 → 57.3 s), **within 1.09×
of MLAS**, at **identical** peak RSS. Recovered fraction of the MLAS
gap: `(169.4 − 57.3)/(169.4 − 52.5)` = **95.9%**, with **zero added
footprint**. Per-shape `simd/mlas` after: 5120×5120 1.39x, 5120×7168
1.12x, 5120×13824 1.22x, 13824×5120 1.10x, lm_head 1.11x (all down from
2.2–4.6x).

## Numerical output

**Not byte-identical to the packed path** — the GEMV reassociates the
f32 sum (K-unrolled-by-4 running accumulation vs the packed KC-panel
order). It matches the naive f64 / Generic reference within the **same
tolerance the existing `SimdX86`-vs-reference tests use**
(`1e-3·(1+|e|)`), and a new test asserts GEMV-vs-packed agreement within
that bound (they differ only by summation order, never in which products
are summed). This is reported as a numerical change, not shipped
silently: the toggle defaults **off**.

## What could not be ported / caveats

- **No f32 model exercises this path on this host.** `qwen2.5-14b-f32`,
`qwen2.5-14b-onnx`, and every qwen05b variant route their weights
through `MatMulNBits` (int4), which does not touch the dense f32 GEMM.
So there is no end-to-end token-identity check here; the A/B is a
**synthetic in-binary driver** (`bench_f32_gemm_ab`, `#[ignore]`),
reported honestly as such rather than as a model number that never took
the path.
- **Prefill (M>1) is unchanged** — it is already at parity, so this PR
deliberately scopes to M==1, exactly as MLAS special-cases only M==1.
- Default-off toggle means a default build is not yet faster; recommend
flipping it on for `SimdX86` in a follow-up once the reassociation is
signed off, which is what makes the #1091 win reach users.

## Gates (on this host, not CI)

- `cargo test -p onnx-runtime-ep-cpu --lib` ×5: **1321 / 1321 / 1321 /
1321 / 1321 passed, 0 failed, 12 ignored** each.
- `cargo clippy -p onnx-runtime-ep-cpu --lib -- -D warnings`: clean.
Also `--tests --features mlas`: clean.
- New unit tests: `m1_gemv_shapes`,
`m1_route_matches_packed_within_tolerance`.

Refs #1091 #1045. Precedent #1104.

---------

Co-authored-by: justinchuby <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
justinchuby pushed a commit that referenced this pull request Aug 17, 2026
… the deliverable

Updates the ledger: dense f32 M=1 GEMV is absorbed (#1116), and #1045's headline
4.4x is recorded as **not reproducing** on this host -- `simd/mlas` was already
0.57-1.05x at M=128, so the entire reproducible gap was M=1 decode. Inheriting
that figure would have sent someone optimising prefill, which was not the problem.
The fix was to stop packing B at M==1, matching MLAS's own reasoning that packing
a matrix referenced once is a wasted copy: a win from doing less work.

Also records a pattern that has now held three times in a row. Each brief
predicted a mechanism and the measurement found a different one -- #1104 expected
layout and found register blocking, #1116 expected a 4.4x prefill gap and found
the gap was entirely at M=1, #1126 expected missing GEMM blocking and found
per-row dispatch and allocation overhead. In all three the correction was worth
more than the patch.

The point is not that briefs are unreliable: each hypothesis was specific enough
to direct a measurement that could refute it, which is what a hypothesis is for.
The point is to ask for the mechanism *before* the kernel, because a plan is
cheap to change then and expensive afterwards -- #1104's transient-tile design was
abandoned as unnecessary rather than built and then found unnecessary.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: d60eb808-7cc6-4abc-b48d-2a6dd3841624
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