Repository navigation
perf(mlas): let MlasGemmBatch use the registered parallel backend (4.4x dense f32 MatMul, reaches ORT parity) - #1045
Merged
Conversation
`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 Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ 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
Flags with carried forward coverage won't be shown. Click here to find out more.
🚀 New features to boost your workflow:
|
🔴 Benchmark Regression DetectedComparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).
Visual flags: Host infoWhat this cannot catch
|
… 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
marked this pull request as ready for review
August 15, 2026 23:53
This was referenced Aug 16, 2026
MLAS-routed speedups do not reach a default build, and the strategy is to absorb them natively
#1091
Open
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
mlas_sgemmandmlas_sgemm_packedpassed/*ThreadPool=*/nullptrtoMlasGemmBatch. In the standalone (BUILD_MLAS_NO_ONNXRUNTIME) build thatargument is not an inert handle — it is the parallelism enable flag:
So every f32 GEMM routed through the
mlasbackend ran single-threaded.Because
CpuBackend::auto_detect()prefersMlaswhen the feature is enabled,building with
--features mlasmade dense f32 MatMul slower than thebuilt-in
SimdX86backend 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
MlasGemmBatchsites.The bug, measured
Backend comparison, MatMul f32
K=3584 N=3584 M=128,native_minmedians:genericsimdmlas(default with the feature on)mlashad the best single-thread kernel and the only broken scaling — and at8 threads it was 3.6x slower than the pure-Rust
SimdX86backend itdisplaces.
Result
Interleaved A/B,
native_minmedians, 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:Per-shape at 8 threads, 3 interleaved reps each:
Parity
PASSon 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
8; we do not. That is a separate partitioning/pool issue and is not fixed
here. Reported, not hidden.
was never the problem; it was already at ORT parity.
Gemmis unaffected (measured 1.01x / 1.02x / 0.96x — noise).Gemmdoes 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. Includedhere as a control showing the change is scoped to what it claims. Those same
runs show
Gemmat 199–216x andGemm(transB)at 1342x of ORT, which isthe strongest argument yet for Gemm: route the f32 path through the shared GEMM backend (up to 178x faster) #1035.
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), returnsMlasStandaloneMaxThreads().MlasTrySimpleParallel—MLAS_UNREFERENCED_PARAMETER(ThreadPool), onlyThreadPool != nullptr.ArmKleidiAI::MlasGemmBatch(arm64 SME/SME2 only,
USE_KLEIDIAI), passes it on to exactly those twohelpers 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 atiming check. Asserts
sgemm_nnincrements the backend'sparallel_for_callscounter. Verified non-vacuous: with the sentinel reverted it fails with
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_multithreadprobeprinted the serial numbers next to a comment recording ORT's ~4.4x scaling,
but asserted nothing and is
#[ignore]d.Follow-ups (deliberately not in this PR)
mlas_qnbit_gemm_pack_bandMlasReorderOutputNchwstill passnullptr.Those are one-time packing/reorder rather than the measured hot path;
they deserve their own measurement.
CpuBackend::auto_detect()preferringMlaswas actively harmful beforethis fix. Worth a guard so a backend can never be selected when it is
measurably slower than the built-in one.
Reproduce
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 findingresolved 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::MlasGemmBatchis never compiled at all (sgemm_kleidiai.cppis not in
build.rs), so the onlyMlasSGemmBatchOverridethat could forwardthe sentinel does not exist in this build. Findings fixed in
44fa4f1d7:crate's unit tests, which cargo runs concurrently in one process. A
concurrent
sqnbit_gemm— which passes its own sentinel and does drivethe backend — could have bumped the counter between samples and let a broken
build pass. Moved to
tests/sgemm_threading.rs, its own binary/process.small-magnitude outputs.
Round 2 found the replacement tolerance was now too tight. Fixed in
00caeb923:8 * EPSILONis only valid fork <= 8; fork = 193the standard boundis
gamma_k ~ k * EPSILON, ~12x looser. An over-tight bound would be flakyon any ISA that reassociates the sum differently (AVX-512, SVE, different
blocking). Now scales with
k.degree < 2early-return would have made it a silent no-op on asingle-core runner. It now forces
ONNX_GENAI_MLAS_THREADPOOL_THREADS=4before first use and asserts, rather than skipping.
Both tests re-verified non-vacuous after every change:
0 -> 0 calls(64,100),|error|0.789 vs tolerance 0.0075 (105x margin), confirming the looser bound still
catches tile corruption
CI
mainis red before this branch for pre-existing, unrelated reasons. Thisbranch is based on
869c24b83, which includes #1043 (cargo fmt --alloverthe workspace), so the
Check formattingfailures that affect my older PRs donot apply here.
Locally green:
CI baseline
Rebuilt on top of
main@400fbe246(a plaingit merge origin/main, no history rewrite).Unmodified
main@400fbe246fails exactly these 6 jobs(run 31914831964):
mainCLI ORT (Linux x86_64)CLI ORT (Windows x86_64)CUDA compile (Linux x86_64)CUDA compile (Windows x86_64)Rust (Windows ARM64)Rust coverage (macOS arm64)None are touched by this PR.
Fast (Linux x86_64)andRust qualitypreviously failed onmaintoo (a repo-widecargo fmtdrift, fixed onmainby #1043); after merging currentmaininto 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,auditandcodecov-- are green.Ratio convention (added post-merge for clarity)
speedupis ours-before/ours-after.after/ORTis ours/ORT: >1 means we are slower, so1.00x/1.01x/1.03xare parity and1.76xat 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).