Skip to content

Native transposed-B SGEMM: eliminate the Kn dequant from default-build MatMulNBits prefill (#959, #1091) - #1176

Merged
justinchuby merged 11 commits into
mainfrom
squad/native-nt-sgemm
Aug 19, 2026
Merged

justinchuby merged 11 commits into
mainfrom
squad/native-nt-sgemm

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 18, 2026 •

Copy link
Copy Markdown
Owner

Summary

On the default (mlas OFF) build the dequant-kn prefill phase now goes to zero for every MatMulNBits path that reaches the dense fallback: the native CPU EP gets a transposed-B ("NT") SGEMM and non-MLAS prefill routes through it, reusing the already-cached contiguous Nk weight instead of materializing a second, transposed Kn copy.

Refs #959, #1091. Does not close #959 — the superlinear per-token decode cost (~102 s/token at 14B) is a separate open question this does not touch. This closes the prefill Kn-materialization term.

The problem (#959)

Int4/int8 MatMulNBits prefill (m > 1) on the default build dequantized the weight to f32 a second time in the transposed Kn ([k, n]) layout — a strided-scatter transpose (each K step written at stride N), uncached, then a dense NN GEMM. #959 measured it at ~2.9x the contiguous Nk pass and degrading with N (22 s/GB at 0.5B → 38 s/GB at 14B; 266 s of ~357 s time-to-first-token on qwen2.5-14b).

MLAS hosts already avoided this (try_prefill_mlas_nt) by feeding the cached Nk weight to MLAS's cache-tiled sgemm with trans_b. The default build had no transposed-B GEMM, so its #[cfg(not(feature = "mlas"))] sibling returned false and fell back to the slow Kn path.

What changed

MLAS's advantage here is not magic — it is a packer that reads the B operand row-wise from [n, k]. Ported natively:

  • x86_sgemm.rs — sgemm_simd_nt(a, b_nk, c, m, k, n) computing C[m,n] = A[m,k] · B_nk[n,k]^T. It reuses the entire packed GEBP path (pack_a, KC/strip blocking, micro_6x16) and adds pack_b_nt, which gathers each output column from its contiguous b_nk row (unit stride over k) into the L1-resident pack tile at stride NR. The transpose becomes a pack-time reshape of an in-cache tile, not a full-array stride-N scatter — exactly why MLAS wins.
  • matmul.rs — nt_gemm_supported(backend) is the single legality predicate (no two cfg arms to drift); gemm_nt_with_backend dispatches Mlas → trans_b, SimdX86 → sgemm_simd_nt.
  • matmul_nbits.rs — the two cfg-split try_prefill_mlas_nt arms collapse into one try_prefill_nk_nt gated on nt_gemm_supported. It dequantizes once into the pre-existing weight_nk OnceLock (the same slot decode caches into, so a constant weight pays one dequant, not two) and calls gemm_nt_with_backend. When the NT route runs, the dequant-kn profile phase is skipped by construction.
  • Host pool (perf(cpu-ep): run the elementwise split on ORT's pool, not a second one #1143) — column strips dispatch onto the installed ORT intra-op pool when present, rather than forking a rayon pool beside it; rayon otherwise. Strip decomposition is numerically transparent.

Correctness — bit-identical

Byte-for-byte identical to the existing Kn dense route, not "within 1e-5". pack_b_nt produces the same packed panels pack_b would from the [k, n] transpose (b_kn[p·n+j] == b_nk[j·k+p]); with identical A-pack, K-panel order, and microkernel, every output element's f32 accumulation sequence is unchanged. Strip count and pool choice do not affect per-element reduction order, so bit-identity holds regardless.

Verified (assert_eq! on f32::to_bits):

  • Odd/tail shapes: m ∈ {1, 2, 7, 33}, n ∈ {1, 3, 63, 64, 65}, plus tile-exact and multi-KC.
  • 64 randomized shapes (differential NN-vs-NT).
  • End-to-end by the existing 8-bit prefill oracle test (matmulnbits_8bit_prefill_batched_matches_dequant_f32_oracle), which reaches the dense fallback and thus the NT route.

This matches the standard the MLAS NT route already held itself to (bit-identity to the no-transpose dense GEMM, recorded at matmul_nbits.rs ~2048).

Measurement

dequant-kn → 0 is structural. The mm_profile::time_prepack("dequant-kn", …) call lives in the if !used_fast_nt {…} branch; whenever the NT route returns true that branch is skipped, so the phase is eliminated by construction. dequant-nk is unchanged (still one pass, now the only one).

Finding vs the #959 premise: on current local builds, 4-bit acc0 prefill takes the borrowed-int4 in-place path (#979/#1117) and 4-bit acc4 uses the SDOT prepack — neither reaches the dense fallback, so no local q4 model emits a [mm_prepack] phase=dequant-kn line to drive to zero end-to-end (confirmed empirically with ONNX_GENAI_PROFILE_MM=1 on qwen05b q4 / q4-acc4 / symzp). The beneficiaries of this change are therefore: 8-bit weights (m>1), grouped quantization, weight_prepacked, and 4-bit with accuracy_level != 0 that falls to the dense fallback. No local 8-bit/grouped model was available for an end-to-end token-to-first-token arm; per the profiling skill I do not report a contended wall-clock figure I cannot defend.

Per-phase microbench (nt_prefill_bench, #[ignore]; best-of-7, --test-threads=1, release; contended box — other agents building concurrently). Bit-identity asserted in the same harness. The dequant-kn* arm is a plain f32 transpose standing in for the strided Kn materialization the NT route removes (the real int4 dequant is ~2.9x heavier per #959):

shape (m=16) dequant-kn* (transpose) NN gemm NT gemm old (kn*+NN) → new (NT)
k=5120, n=5120 (100 MiB) 337 ms 5.7 ms 5.0 ms 343 ms → 5.0 ms
k=5120, n=13824 (270 MiB) 641 ms 15.0 ms 11.8 ms 656 ms → 11.8 ms
k=13824, n=5120 (270 MiB) 1307 ms 15.3 ms 11.7 ms 1323 ms → 11.7 ms

The eliminated transpose term dominates and grows with size (as #959 predicted); the NT GEMM itself is even slightly faster than NN here (contiguous per-column B reads pack better). Per #1132, a native-faster-than-MLAS result on some shape is a graduation event for benches/native_vs_mlas.rs — noting it, but not changing default routing without that gate's measurement.

RSS. No new long-lived allocation (reuses weight_nk; apack/bpack are per-call local scratch, freed at return). Removing the second full f32 weight materialization cuts the transient f32 footprint of a prefill that hits the dense fallback by one full [k, n] copy (e.g. 270 MiB at k=13824/n=5120). Not measured end-to-end (no local model reaches that path); the reduction is structural.

Memory rules

No new field that outlives a call or scales with weight size — the NT route reuses the pre-existing weight_nk OnceLock. apack/bpack are per-call vec![] scratch. The added lines are not matched by weight-cache-guard.yml (no OnceLock<…Vec> / (RefCell|Cell)<Vec|Box|Arc> introduced); its regex and path filter are untouched.

Gates (exact counts)

  • cargo test -p onnx-runtime-ep-cpu --lib → 1324 passed, 0 failed, 17 ignored
  • cargo test -p onnx-runtime-ep-cpu --lib --features mlas → 1354 passed, 0 failed, 28 ignored (MLAS route still works and still wins where enabled)
  • cargo clippy -p onnx-runtime-ep-cpu --all-targets -- -D warnings → clean
  • cargo fmt --check → the three changed files are clean (verified with rustfmt --edition 2024 --check). Pre-existing diffs remain in three unrelated files (governed_accumulator_budget.rs, qlinear_matmul.rs, simd_activations.rs) from a local rustfmt version skew vs CI — left untouched to keep this PR surgical.

🤖 Generated with Squad. Flagged needs review — please have a squad member review the kernel packing/tail handling before merge.


Update (2026-08-18, Roy) — merged current main, plus a production-path A/B

Merge

Merged main (c55a3fab3), which had since made the #1091 M=1 GEMV the unconditional SimdX86 route and dropped the ONNX_GENAI_CPU_MM_SIMD_M1_GEMV toggle this branch still carried. Conflict resolved by keeping both: main's default-route test (the_default_entry_point_routes_m1_to_the_gemv) and this branch's NT kernel + bit-identity tests. Two follow-on fixes:

  • aarch64 cross-arch lane. gemm_nt_with_backend compiled with neither the mlas nor the x86 arm has no reader for any parameter, so the -D warnings cross-arch pass rejected all six. Bound them in the unsupported arm. (This is what the old Rust quality red was: Cross-target compile check, nothing else.)
  • mm_profile gemv phase. The MLAS route used to time this GEMM on the gemv phase; after the two call sites collapsed into one, that timer was lost. Restored — the default build gets a pass-through (tick(), the reporter, is MLAS-only), so the shared call site stays cfg-free and MLAS profiling is unchanged.

Production-path A/B (new harness, benches/matmul_nbits_prefill_ab.rs)

The original body was right that no local model reaches this route, and honest about not reporting a number it could not defend. That gap is now closed the way #1013 closed its own: drive the real kernel through the EP's own get_kernel/execute, at the shapes and inputs that do reach the dense fallback — 8-bit prefill, and 4-bit with g_idx. The harness uses no symbol this branch introduces, so the identical file runs on main and here; both arms were built and run interleaved, 3 repetitions each, on the same box.

Host: 32-core x86_64, AVX2, default build (mlas off), release. Contended (other agents building; load ~12), so medians of per-arm medians are reported and the ratios — not the absolute ms — are the claim.

Steady state (weight already resident, per-call prefill cost):

case k n m main (ms) this PR (ms) speedup
int8 dense fallback 2048 2048 8 12.114 0.548 22.1x
int8 dense fallback 2048 2048 64 12.276 1.375 8.9x
int8 dense fallback 4096 11008 8 53.858 4.170 12.9x
int8 dense fallback 4096 11008 64 56.560 7.595 7.4x
int4 + g_idx fallback 2048 2048 8 31.583 0.563 56.1x
int4 + g_idx fallback 2048 2048 64 32.405 1.419 22.8x

Cold (fresh kernel per repetition, so the one-time weight dequant is inside the measurement — the TTFT term #959 attacked):

case k n m main (ms) this PR (ms) speedup
int8 dense fallback 2048 2048 8 6.470 6.008 1.08x
int8 dense fallback 2048 2048 64 12.519 7.974 1.57x
int8 dense fallback 4096 11008 8 53.761 37.260 1.44x
int8 dense fallback 4096 11008 64 57.725 38.503 1.50x
int4 + g_idx fallback 2048 2048 8 31.636 23.200 1.36x
int4 + g_idx fallback 2048 2048 64 32.083 25.883 1.24x

Why steady moves 7–56x and cold only ~1.1–1.6x — and why that is the real finding. The Kn route has no cache: dequantize_weight(WeightLayout::Kn) is called inside execute, so every prefill call re-materializes the whole transposed f32 weight. The NT route dequantizes into the pre-existing weight_nk OnceLock, which a constant weight fills once. So this change removes not one transpose but every repeat of it. Cold (first call) improves by the layout alone — a contiguous Nk write instead of the stride-N scatter, 1.2–1.6x at these sizes; steady improves by the caching the Nk layout makes possible.

Confirmed structurally with ONNX_GENAI_PROFILE_MM=1 over the same harness run: main emits 72 phase=dequant-kn lines, this PR emits 0 — and 24 phase=dequant-nk (one per kernel instance, i.e. the cold arms only; every steady call pays none).

Bit-identity, across builds. The harness prints an FNV-style digest of the raw output bits. All six rows have the identical digest on both arms (db6ff07f991d431, b4c1df9fd8883789, 7bd7418eacdcb870, 2eed012609649617, d5834afbecf42494, cd8fa77262302969) — bit-identity of the production execute result, not just of the kernel driver, verified across two separately compiled builds.

Scope, restated honestly

On the default build, 4-bit accuracy_level=0 with contiguous (borrowable) inputs takes the zero-copy borrowed int4 path (#979/#1117/#1126) for both decode and prefill and returns before the dense fallback. This PR therefore changes: 8-bit prefill, 4-bit with g_idx, weight_prepacked/non-borrowable inputs, and 4-bit accuracy_level != 0 that falls through. Those are exactly the cases measured above. MLAS builds already had the NT route; their behaviour is unchanged.

Gates (re-run after the merge)

  • cargo test -p onnx-runtime-ep-cpu --lib → 1420 passed, 0 failed, 18 ignored
  • cargo clippy -p onnx-runtime-ep-cpu --all-targets -- -D warnings → clean
  • cargo clippy --locked --target aarch64-unknown-linux-gnu --all-targets -p onnx-runtime-ep-cpu -- -D warnings → clean (the lane that was red)
  • cargo fmt --all -- --check → the files this PR touches are clean; one inherited diff remains in onnx-runtime-ep-cuda/standard_attention.rs from main, fixed separately in fix(ci): restore the green Rust quality lane on main #1347.

…ough it

MatMulNBits int4/int8 prefill on the default (mlas OFF) build materialized
the weight to f32 twice per generate: once contiguous (`Nk`, cached) for
decode, and once transposed (`Kn`, uncached) for the m>1 dense GEMM. The
`Kn` pass is a strided-scatter transpose (each K step written at stride N)
that #959 measured at ~2.9x the contiguous `Nk` pass and degrading with N
(22 s/GB at 0.5B, 38 s/GB at 14B; 266 s of ~357 s TTFT on qwen2.5-14b).

MLAS hosts already avoided it (`try_prefill_mlas_nt`) by feeding the cached
`Nk` weight to MLAS's cache-tiled `sgemm` with `trans_b`. The default build
had no transposed-B GEMM, so it fell back to the slow direct-`Kn` path.

This ports MLAS's idea natively:

- x86_sgemm.rs: add `sgemm_simd_nt` (c[m,n] = a[m,k] @ b_nk[n,k]^T). It
  reuses the whole packed GEBP path (`pack_a`, KC/strip blocking,
  `micro_6x16`) and adds `pack_b_nt`, which gathers each output column from
  its contiguous `b_nk` row (unit stride over k) into the L1-resident pack
  tile at stride NR. That is MLAS's advantage: the transpose becomes a
  pack-time reshape of an in-cache tile, not a full-array stride-N scatter.
  Bit-identical to the NN path by construction (identical packed panels,
  A-pack, K-panel order, microkernel => identical per-element accumulation).
- matmul.rs: add `nt_gemm_supported(backend)` (the single legality
  predicate, no cfg drift) and `gemm_nt_with_backend` (Mlas -> trans_b,
  SimdX86 -> sgemm_simd_nt).
- matmul_nbits.rs: replace the two cfg-split `try_prefill_mlas_nt` arms with
  one `try_prefill_nk_nt` gated on `nt_gemm_supported`. It dequantizes once
  into the existing `weight_nk` OnceLock (the same slot decode caches into,
  so a model pays ONE dequant, not two) and calls `gemm_nt_with_backend`.
  When the NT route runs, the `dequant-kn` profile phase is skipped
  entirely.
- Column strips dispatch onto the host pool (#1143) when an ORT intra-op
  pool is installed, instead of forking a rayon pool beside it; rayon
  otherwise. Strip decomposition is numerically transparent.

Correctness: NT output is byte-for-byte identical to the existing `Kn`
dense route (asserted on `f32::to_bits`) across odd shapes (m in {1,2,7,33},
n in {1,3,63,64,65}, tile-exact and multi-KC) and 64 randomized shapes, and
end-to-end by the existing 8-bit prefill oracle test.

Memory: no new long-lived buffer; reuses `weight_nk`. `apack`/`bpack` are
per-call local scratch. Does not touch the weight-cache guard's regex/path.

Refs #959, #1091. Does not close #959 (the superlinear per-token decode
cost is separate and untouched).

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

codecov Bot commented Aug 18, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 73.26007% with 73 lines in your changes missing coverage. Please review.
✅ Project coverage is 80.11%. Comparing base (4a9f4ec) to head (7672f7b).
⚠️ Report is 71 commits behind head on main.

Files with missing lines Patch % Lines
...rates/onnx-runtime-ep-cpu/src/kernels/x86_sgemm.rs 72.01% 61 Missing ⚠️
crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs 70.83% 7 Missing ⚠️
...es/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs 83.87% 2 Missing and 3 partials ⚠️
Additional details and impacted files

Impacted file tree graph

@@             Coverage Diff             @@
##             main    #1176       +/-   ##
===========================================
- Coverage   82.10%   80.11%    -2.00%     
===========================================
  Files          12      377      +365     
  Lines        5471   164937   +159466     
  Branches     5471   164937   +159466     
===========================================
+ Hits         4492   132134   +127642     
- Misses        780    27970    +27190     
- Partials      199     4833     +4634     
Flag Coverage Δ
cli-ort-linux 82.60% <ø> (?)
cli-ort-windows 82.10% <ø> (ø)
offline 80.02% <73.26%> (?)

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

Files with missing lines Coverage Δ
...es/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs 76.68% <83.87%> (ø)
crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs 79.30% <70.83%> (ø)
...rates/onnx-runtime-ep-cpu/src/kernels/x86_sgemm.rs 89.09% <72.01%> (ø)

... and 363 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.
  • 📦 JS Bundle Analysis: Save yourself from yourself by tracking and limiting bundle sizes in JS merges.

@github-actions

github-actions Bot commented Aug 18, 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
🔴 matmul/medium_generic_f32_threads=8/32x512x512 913.79 µs 1.86 ms +104.0%
🔴 tokenization/decode_tokens_per_second 6.11 ms 9.94 ms +62.7%
🔴 matmul/large_generic_f32_threads=8/32x1024x1024 3.65 ms 5.24 ms +43.6%
🔴 tokenization/encode_tokens_per_second 352.82 µs 495.41 µs +40.4%
⚠️ matmul/small_generic_bf16_threads=8/1x256x256 29.58 µs 38.06 µs +28.7%
⚠️ matmul/medium_generic_bf16_threads=8/32x512x512 416.52 µs 532.34 µs +27.8%
⚠️ matmul/medium_generic_f16_threads=1/32x512x512 27.61 µs 34.74 µs +25.8%
⚠️ matmul/medium_generic_f16_threads=8/32x512x512 35.07 µs 42.03 µs +19.9%
⚠️ matmul/medium_generic_f32_threads=1/32x512x512 2.22 ms 2.65 ms +19.4%
✅ sampling_latency/greedy_per_token 3.16 µs 3.63 µs +14.8%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.37 ms 3.86 ms +14.6%
✅ sampling_latency/top_k_per_token 48.72 µs 55.22 µs +13.3%
✅ add/small_f16_threads=1-internal/1024 449.7 ns 502.5 ns +11.8%
✅ matmul/small_generic_bf16_threads=1/1x256x256 31.45 µs 35.14 µs +11.7%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 485.51 µs 533.41 µs +9.9%
✅ qwen3_sampling_processors/top_k_partial_selection 134.38 µs 147.61 µs +9.8%
✅ kv_cache/alloc_dealloc_pages 36.70 µs 40.19 µs +9.5%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 2.09 ms 2.28 ms +9.0%
✅ qwen3_sampling_processors/top_k_top_p_fast 618.55 µs 672.08 µs +8.7%
✅ sampling_latency/min_p_per_token 196.57 µs 209.02 µs +6.3%
✅ gather/large_f32_threads=1-internal/131072 24.74 µs 26.17 µs +5.8%
✅ logit_processing/seven_processor_chain_per_step 313.31 µs 326.80 µs +4.3%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.53 ms 5.76 ms +4.3%
✅ gather/large_bf16_threads=1-internal/131072 10.03 µs 10.44 µs +4.1%
✅ sampling_latency/top_p_per_token 415.17 µs 429.86 µs +3.5%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.76 ms 9.98 ms +2.2%
✅ gather/medium_f32_threads=1-internal/32768 3.66 µs 3.71 µs +1.2%
✅ grammar_masking/llguidance_compute_mask/32 68.99 µs 69.30 µs +0.4%
✅ matmul/small_generic_f16_threads=8/1x256x256 30.43 µs 30.24 µ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 88.39 µs 87.17 µs -1.4%
✅ gather/medium_f16_threads=1-internal/32768 2.32 µs 2.28 µs -1.7%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 589.16 µs 577.94 µs -1.9%
✅ add/small_bf16_threads=1-internal/1024 449.4 ns 440.2 ns -2.1%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 83.25 µs 81.09 µs -2.6%
✅ matmul/small_generic_f16_threads=1/1x256x256 28.59 µs 27.76 µs -2.9%
✅ matmul/small_generic_f32_threads=8/1x256x256 34.58 µs 33.08 µs -4.3%
✅ gather/small_bf16_threads=1-internal/4096 482.6 ns 458.2 ns -5.1%
✅ add/medium_bf16_threads=1-internal/262144 102.07 µs 94.79 µs -7.1%
✅ matmul/small_generic_f32_threads=1/1x256x256 39.05 µs 36.20 µs -7.3%
✅ reduce_mean/small_f32_threads=1-internal/4096 14.95 µs 13.86 µs -7.3%
✅ add/medium_f32_threads=1-internal/262144 24.68 µs 22.80 µs -7.6%
✅ reduce_mean/medium_f32_threads=1-internal/65536 253.41 µs 233.48 µs -7.9%
✅ add/large_bf16_threads=1-internal/4194304 1.71 ms 1.54 ms -9.7%
✅ add/medium_f16_threads=1-internal/262144 107.52 µs 96.87 µs -9.9%
✅ gather/large_f16_threads=1-internal/131072 11.06 µs 9.71 µs -12.2%
✅ reduce_mean/large_f32_threads=1-internal/262144 1.05 ms 914.44 µs -13.3%
✅ add/large_f16_threads=1-internal/4194304 1.80 ms 1.55 ms -13.8%
✅ add/large_f32_threads=1-internal/4194304 629.43 µs 539.75 µs -14.2%
🟢 add/small_f32_threads=1-internal/1024 213.4 ns 179.8 ns -15.7%
🟢 gather/medium_bf16_threads=1-internal/32768 2.70 µs 2.25 µs -16.7%
🟢 gather/small_f32_threads=1-internal/4096 758.6 ns 618.8 ns -18.4%
🟢 matmul/large_generic_f16_threads=8/32x1024x1024 115.40 µs 93.10 µs -19.3%
🟢 block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 64.84 µs 51.75 µs -20.2%
🟢 block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 769.37 µs 597.75 µs -22.3%
🟢 matmul/large_generic_bf16_threads=1/32x1024x1024 2.72 ms 2.06 ms -24.2%
🟢 gather/small_f16_threads=1-internal/4096 587.0 ns 440.2 ns -25.0%
🟢 block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 489.36 µs 362.90 µs -25.8%
🟢 block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 64.23 µs 45.46 µs -29.2%
🟢 matmul/large_generic_bf16_threads=8/32x1024x1024 3.32 ms 1.38 ms -58.4%

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.49 3.49 4.82 }
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)

Resolves x86_sgemm.rs: main made the #1091 M=1 GEMV the unconditional
route (dropping the ONNX_GENAI_CPU_MM_SIMD_M1_GEMV toggle this branch
still carried) while this branch added the NT packed path beside it.
Keep both: the NT kernel/tests and main's default-route test.

Time the NT GEMM on the same mm_profile gemv phase the MLAS route used
before the two call sites merged; the default build gets a pass-through
so the shared call site stays cfg-free.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby and others added 3 commits August 18, 2026 23:44
The cross-arch lane (aarch64, -D warnings) rejected gemm_nt_with_backend:
with neither the mlas nor the x86 arm compiled in, every parameter is
unused. Bind them in the unsupported arm.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Times the real CPU EP kernel through get_kernel/execute (not a GEMM
driver), splitting cold (fresh kernel, pays the one-time dequant) from
steady (warm weight cache), and prints a bit-digest of the output so two
builds can be compared for bit-identity as well as speed. Uses no symbol
introduced by this branch, so the same file runs on main and here and the
two runs are a true A/B.

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

Copy link
Copy Markdown
Owner Author

Second reconciliation: #1176 and #1403 both claim 8-bit prefill

Continuing the ten-PR merge rehearsal, #1176 and #1403 collide the same way
#1356 and #1431 do — invisibly to either PR's own CI, because the interaction
only exists in the merged tree.

try_prefill_nk_nt is introduced by #1176 only (neither main nor #1403
has it). In the merged tree it runs before #1403's try_int8_prefill_gebp
and claims 8-bit prefill unconditionally, so the fused GEBP never executes.
This is not a subtle degradation — it fails outright:

matmulnbits_int8_gebp_is_bit_identical_to_the_dequant_route     FAILED
matmulnbits_int8_prefill_gebp_matches_reference_and_ran         FAILED
  k=96 n=20 m=2 asymmetric=false: the fused 8-bit GEBP must be the route that ran
  left: 0  right: 1

Two of #1403's own tests fail on the merged stack while passing on its branch.

Resolution: order by cache admission, not by merge order

The two routes are not simply better or worse than each other; they amortize
differently, and which one wins is decided by a condition already present in
the code:

So the dispatch now decides on nt_keeps_weight_resident rather than on
whichever PR merged last, which keeps both wins instead of silently discarding
one:

let nt_keeps_weight_resident = can_prepack && resident_dequant_f32_cache_enabled();
let used_fused_int8 = !nt_keeps_weight_resident && self.try_int8_prefill_gebp(...);
let used_fast_nt = !used_fused_int8 && self.try_prefill_nk_nt(...)?;

Verified on the fully integrated stack

All ten PRs merged together — #1176, #1013, #1356, #1431, #1365, #1381, #1403,
#1417, #1436, #1446 — plus both reconciliations:

  • 1487 tests pass, 0 failed (was 1485/2 before this fix).
  • check_dispatch_reachability, check_dispatch_manifest,
    check_feature_gate_coverage, check_profile_table, check_platform_naming,
    check_publish_order, verify_documented_env_vars — all PASS.
  • The three cargo fmt diffs on the integrated tree are in files untouched by
    any of these PRs and identical to origin/main; that is a local rustfmt
    version difference, not a merge artifact.

Other conflicts found in the rehearsal, and how they resolve

pair file nature resolution
#1356 x #1365 x86_sgemm.rs both append an independent GEBP at the same anchor; pure insertions (723/0 and 529/0) apply both insertion sets; 699 + 723 + 529 = 1951 lines exactly
#1356 x #1365 x #1403 Cargo.toml each adds a [[bench]] keep all
#1381 x #1417 half_gemv.rs, gemm.rs #1381 generalized simd_available/gemv_f16_kn/f16v to take a HalfFormat; #1417 and #1436 still call the old signatures keep #1417's trans_b branch, route the non-transposed arm through gemv_half_kn(HalfFormat::F16, ...), update 12 call sites
#1403 matmul_nbits.rs #1403's doc calls INT4_PREFILL_GEBP_MIN_ROWS "a crossover against the row-serial kernel" corrected — after #1431 it is a crossover against the column-blocked kernel at 12
#1431/#1436/#1446 CPU_MATMUL_ASSIGNMENT.md three PRs each add a "section 5" renumbered 5/6/7/8

None of these require a change on any individual branch; they are merge-commit
work, and I will carry them as each PR lands.

@justinchuby
justinchuby enabled auto-merge (squash) August 19, 2026 11:51
@justinchuby

Copy link
Copy Markdown
Owner Author

Final validation on latest main — merging

Head 7672f7b03 = this branch with origin/main (7de4bb1dc) merged. True footprint vs merge-base is unchanged: 5 files, all in crates/onnx-runtime-ep-cpu/.

Retraction: my earlier "6 of 12 cells regress" report was wrong

I previously reported that this PR regressed 6 of 12 MatMulNBits cells (0.77x–0.94x). That result was an artifact and I retract it. Root cause: every cell in that matrix was a 4-bit MatMulNBits model, and the code this PR changes (try_prefill_nk_nt, reached from the m > 1 dense-dequant fallback) is never reached for 4-bit — those nodes are owned by the borrowed int4 path. I proved this directly with ONNX_GENAI_PROFILE_MM=1: the [mm_prepack] line is absent in both arms on every 4-bit cell. Both arms were executing identical code, so the whole matrix was measuring host noise (intra-arm spread on this shared box reaches 2–29x).

Two independent methods now confirm this:

  • A same-binary kill-switch A/B (dispatch as the only variable) reproduced none of the regressions: mlp_t128 0.78x→1.02x, mlp_t1 0.77x→1.00x, mlp_t8 0.82x→1.06x, qkv_t128 0.80x→1.17x.
  • 4-bit behaviour is unchanged by construction, not merely within noise.

What this PR actually does — measured on 8-bit cells

The m > 1 branch this PR rewrites serves "cases not owned by the 4-bit path — notably 8-bit weights". I generated 8-bit MatMulNBits cells (bits=8, block_size=32) and measured the real effect. ONNX_GENAI_PROFILE_MM=1 shows the mechanism:

  • before: phase=dequant-kn calls=2 — the strided-scatter Kn dequant runs on every prefill call
  • after: phase=dequant-nk calls=1 — contiguous Nk dequant, cached once in the weight_nk slot

Native-alone, min-of-N, ORT threadpool not running (--native-only):

cell (8-bit) steady main steady #1176 speedup cold main cold #1176 speedup
qwen3_0p6b_qkv_t1 0.177 0.185 0.96x 0.256 0.259 0.99x
qwen3_0p6b_qkv_t8 5.903 0.197 29.96x 2.067 0.465 4.45x
qwen3_0p6b_qkv_t32 5.570 0.555 10.04x 2.323 0.599 3.88x
qwen3_0p6b_qkv_t128 7.014 1.264 5.55x 3.717 1.123 3.31x
qwen3_0p6b_mlp_t8 6.476 0.401 16.15x 4.478 0.596 7.51x
qwen3_0p6b_mlp_t32 9.045 0.609 14.85x 3.103 0.779 3.98x
qwen3_0p6b_mlp_t128 9.714 1.432 6.78x 4.284 1.625 2.64x
llama3_8b_qkv_t8 26.057 2.041 12.77x 20.418 2.482 8.23x
llama3_8b_qkv_t32 25.219 3.239 7.79x 21.099 3.063 6.89x
llama3_8b_qkv_t128 26.978 7.862 3.43x 25.900 7.251 3.57x
llama3_8b_mlp_t8 66.190 4.142 15.98x 51.492 4.458 11.55x
llama3_8b_mlp_t32 66.354 6.001 11.06x 52.785 6.295 8.39x
llama3_8b_mlp_t128 76.817 15.977 4.81x 61.471 14.829 4.15x
(all four _t1 cells) 0.96–1.01x 0.98–0.99x

m == 1 is neutral as expected — decode uses the separate else if m == 1 branch, which this PR does not touch. The kill-switch matrix independently gave the same shape (3.2x–14.8x steady).

Numerics: parity=PASS on 16/16 8-bit cells. The NT route is bit-identical to the Kn dense route by construction — B_nk[n, k] is exactly the weight the Kn path stores transposed.

Cross-architecture

nt_gemm_supported and gemm_nt_with_backend carry identical cfg arms (feature = "mlas" / any(target_arch = "x86", target_arch = "x86_64")). On aarch64 without mlas, nt_gemm_supported returns false and callers keep the previous direct-Kn path, so non-x86 behaviour is unchanged. cargo clippy --target aarch64-unknown-linux-gnu -D warnings passes.

The aarch64-pc-windows-msvc check cannot complete on this Linux host: onnx-genai-ort-sys's bindgen needs the Windows SDK (fatal error: 'stdlib.h' file not found). I verified this fails identically on plain origin/main in a clean worktree, so it is a host limitation, not a regression from this PR. Recording the exact scope rather than silently claiming coverage.

Local gate matrix — all green

fmt --all · offline-crate build · -p onnx-runtime-ep-cpu tests (1515 passed, 0 failed) · clippy offline -D warnings · --features mlas tests · -p onnx-genai-engine --features native-backend clippy · aarch64-linux clippy -D warnings · --no-default-features · --all-features · default_artifacts_are_mlas_free · check_cross_compile.sh · 7 guard scripts · workspace_test_packages verify (49 tested, 5 denied).

Miri (strided 8, provider 16, dtype 12, task_runtime 29) = 65 tests, 0 failures.

No-MLAS default: 0 mlas refs in the default dep tree; 0 MLAS symbols in libonnx_runtime_ep_cpu_plugin.so; both bench binaries 0 MLAS symbols.

Note cargo fmt --all --check now passes on this branch — the pre-existing main-wide fmt breakage was fixed by #1492.

@justinchuby
justinchuby merged commit 9a452d9 into main Aug 19, 2026
6 checks passed
@justinchuby
justinchuby deleted the squad/native-nt-sgemm branch August 19, 2026 18:52
justinchuby added a commit that referenced this pull request Aug 19, 2026
#1176 landed the native transposed-B SGEMM, so `try_prefill_nk_nt` now
succeeds on the default x86-64 build for 8-bit prefill. That made this
PR's fused GEBP unreachable: it was consulted only after the NT route
declined, and the NT route no longer declines.

The two are not redundant, though. NT wins whenever it keeps its
dequantized `Nk` weight resident -- one dequant per session, then a
cache-tiled NT GEMM. When the #971 memory governor declines that cache,
NT has nothing to amortize into and rebuilds the whole `k * n` f32 panel
on every call (measured: `dequant-nk calls=5` for five prefills), which
is exactly the cost this fused pack removes, in exactly the
configuration where memory was already too tight to hold it.

So order on `nt_keeps_weight_resident` rather than on NT declining.

Measured, 8-bit cells, native-alone, min of 5:
  governor declined:  2.97x - 21.53x faster than main
  cache admitted:     0.97x - 1.05x (neutral; NT keeps priority)

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 20, 2026
…locked half GEMM (22-65x loss to a win) (#1417)

## What

An f16 `Gemm` at `M = 1` with `transB = 1` measured **32-48 ms against
ONNX Runtime's 0.16-1.5 ms** at `K = N = 3584` — **22x to 65x slower** —
and did not improve with thread count at all. It now runs in **0.3-1.5
ms**: a **38x to 90x** absolute speedup, up to **148x** on per-run
minima, and a **win against ORT at 1 through 8 threads**.

`transB = 1` is not a corner case. It is the layout every `nn.Linear`
export produces, so this is what a QKV, an output projection and an MLP
gate look like whenever a model is exported through `Gemm` rather than
`MatMul`.

## Why it was slow

`GemmKernel::execute` disqualified both f16 fast paths if either
transpose flag was set:

```rust
let half_fast_path = if self.trans_a || self.trans_b { None } else { ... };
```

with the reasoning, in a comment directly above it, that both "read B in
its stored `[K, N]` order, and materialising a transpose first would
give back what they save".

The premise is right and the conclusion does not follow. You only need a
transpose if you insist on reusing the `[K, N]` kernel. A `[N, K]`
weight does not need one — it is the **better** GEMV layout:

| | `[K, N]` (`transB = 0`) | `[N, K]` (`transB = 1`) |
|---|---|---|
| one output element is | a strided gather over `k` | one **contiguous**
`k`-run |
| parallel granularity | a stripe of columns, min 32 (a cache line) |
**one row**, any width |
| accumulator working set | `W` f32 live across the whole `k` sweep |
one f32 |

So the fix is a second kernel, not a transpose. The path it replaced was
the portable blocked half GEMM, which packs both operands into `MR x NR`
panels — at `M = 1` there is no reuse to amortise that against. The same
file already recorded exactly this failure for the *untransposed* case
("the worst dense region measured anywhere in this EP ... 10.07 ms at 1
thread, 10.26 ms at 8"). That case was fixed; the transposed one kept
the note and the behaviour.

## Measurements

`K = N = 3584` (Qwen3-8B hidden), ORT and native alternating **in one
process** so every ratio is paired. Host AMD EPYC 9V74, 32 vCPU / 16
cores, AVX2+FMA+F16C only, ORT 1.27.0, native build (no `mlas`).

Tabulated run is the quietest and largest of three sweeps: 9 trials,
load 9-11.

| threads | before ms | after ms | speedup | before `ours/ORT` | after
p50 | after p90 |
|---:|---:|---:|---:|---:|---:|---:|
| 1 | 36.606 | **0.968** | 37.8x | 22.46 | **0.655 win** | 0.674 |
| 2 | 36.383 | **0.633** | 57.5x | 29.75 | **0.634 win** | 0.733 |
| 4 | 32.587 | **0.468** | 69.6x | 36.86 | **0.840 win** | 1.021 |
| 8 | 36.234 | **0.532** | 68.1x | 49.50 | **0.757 win** | 0.929 |
| 16 | 48.096 | **0.534** | 90.1x | 65.26 | 1.706 | 1.224 |

On per-run minima: 32.203 → 0.814, 31.465 → 0.461, 31.628 → 0.294,
36.076 → 0.243, 35.394 → 0.302 ms (**40x to 148x**).

The before column does not move with thread count (36.6 at 1, 36.2 at 8,
48.1 at 16). The blocked path never scaled on a single-row problem; the
ratio degraded purely because ORT did scale.

### The ratio is now the noisy part, and the claim is narrowed
accordingly

Post-fix these cells run in 0.3-1.5 ms — short enough that this shared
host's other tenants move the ratio more than the kernel does. Three
independent sweeps of the same five cells:

| threads | 1 | 2 | 4 | 8 | 16 |
|---|---|---|---|---|---|
| sweep A (7 trials, load 6-13) | 0.588 | 0.592 | 1.224 | 1.833 | 1.455
|
| sweep B (7 trials, load 12-18) | 0.884 | 1.513 | 1.234 | 1.040 | 2.680
|
| **sweep C (9 trials, load 9-11, tabulated)** | **0.655** | **0.634** |
**0.840** | **0.757** | 1.706 |

`t = 16` swings 1.46-2.68 and `t = 8` swings 0.76-1.83, so neither is
worth two digits. The defensible statement is **a consistent win at 1-8
threads and a `t = 16` ratio unresolvable on this host**. A p90 taken at
load 12-18 read 8.8 at two threads and 19.0 at sixteen — evidence that
the tail statistic measures the other tenants at these durations, not
this kernel.

What is not in doubt is the absolute change, three to four orders of
magnitude above that noise: **32-48 ms → 0.3-1.5 ms in every sweep**.

**Controls** — every other `M = 1` cell, same binary, same session,
unchanged within noise:

| model | before (t=1/4/16) | after (t=1/4/16) |
|---|---|---|
| `MatMul` f16 | 1.79 / 1.49 / 2.24 | 1.42 / 1.59 / 2.79 |
| `MatMul` f32 | 1.01 / 3.13 / 2.56 | 1.04 / 2.38 / 2.71 |
| `MatMulNBits` 4-bit | 2.20 / 3.84 / 7.06 | 1.90 / 3.79 / 6.78 |
| `MatMulNBits` 8-bit | 0.19 / 0.22 / 0.19 | 0.16 / 0.16 / 0.17 |

End-to-end parity against ORT: `max_rel = 6.6e-4`, PASS.

## Numerics — the one deliberate departure

`gemv_f16_kn` is bit-identical to a naive sequential loop.
**`gemv_f16_nk` is not**, and I would rather say so than bury it: it
contracts *along* the lanes rather than across them, so the sum splits
into 32 partial sums. Slightly *more* accurate, but different.

The order is fully specified in the function's doc comment and pinned
two ways:

- `dot_row_scalar` is an executable statement of that order, and
`nk_simd_and_scalar_rows_agree_bit_for_bit` asserts the vector path
matches it **exactly** at every `k` around the 32- and 8-element loop
boundaries (`1, 7, 8, 9, 31, 32, 33, 39, 40, 64, 65, 96, 127, 128,
3584`).
- `nk_is_independent_of_the_thread_count` asserts the answer does not
move across pools of 1, 2, 3 and 8 workers — `ROW_STRIPE` partitions the
output, never the contraction.

Routing is asserted rather than inferred, via a **thread-local**
`NK_GEMV_CALLS` counter: values alone cannot tell which route ran,
because the blocked fallback computes the same thing, only ~40x slower.
Thread-local rather than global on purpose — the count is taken at
entry, before any rayon dispatch, so it is exact without every caller in
the crate having to agree on a lock.

## What this does not fix (recorded in the work list, not left implied)

- **`transB` f16 *prefill* is still badly broken.** At `M = 128` it
measures **4.04x at 1 thread and 16.95x at 8** (156 ms vs 39 ms). Those
cells run 100-160 ms and are reproducible, unlike the decode ratios. The
GEMV is a decode kernel and correctly declines `M > 1` —
`half_prefill_gemm_does_not_take_the_nk_gemv` asserts exactly that — so
prefill still takes the blocked path. Closing it needs a packed **NT**
half GEMM, the f16 analogue of #1176's transposed-B SGEMM. Separate
work, not a comment change.
- **The residual loss at high thread counts is a pre-existing ceiling,
not this kernel.** Every one of our `M = 1` GEMVs flattens at ~0.7-1.1
ms past 4-8 threads while ORT keeps scaling (4-bit 3.60 → 0.79 ms over
1..16 threads vs ORT 1.59 → 0.13; f32 1.73 → 0.94 vs 1.76 → 0.17). Worth
being precise about what that rules out: we do **not** get slower as
threads are added, so it is not fork/join overhead — we stop getting
*faster*. **It is also not task granularity — I built that fix and
measured it away.** `gemv_f16_kn`'s fixed `STRIPE = 512` does yield only
7 tasks at `n = 3584`, so it was the obvious suspect. Making the width
adaptive (`n / (2 * threads)`, rounded up to a multiple of 32) raises
that to 8/16/28 tasks at 4/8/16 threads and moves nothing. Two
interleaved A/B runs of the same binary pair, each with a null control,
disagree in sign:

| run | host load | t=4 | t=8 | t=16 |
|---|---|---|---|---|
| 1 (5 trials) | 5.0 | — | +49.1% *(noise 51.9%)* | **+41.0%** *(noise
0.5%)* |
| 2 (9 trials) | 9.9 | -15.1% *(noise 8.9%)* | +1.1% *(noise -1.3%)* |
+9.2% *(noise 33.5%)* |

The effect is smaller than this host's between-run variability, so **the
change is not in this PR**. The absolute numbers settle it rather than
the ratios: native `p50` stayed at 1.15–1.64 ms in *both* arms at
*every* thread count — 4x the tasks, same time. Commit `667e4c0a6`
records the rejection in the work list so the lead is not retried.

The suspect it leaves standing is the loop nest: `gemv_f16_kn` holds a
`w`-wide `f32` accumulator *in memory* with `p` outermost, so every FMA
is a load-modify-**store** against L1 — `w` overflows the 16 available
`ymm` registers at 512 lanes and at 128 alike, which is exactly why
narrowing the stripe changed nothing. Inverting the nest onto a
register-resident accumulator tile (what `gemv_f16_nk` already does for
`[N, K]`) is the next thing I intend to take. It is a kernel rewrite and
it is unproven.
- The 8-bit `MatMulNBits` "win" in the control table is mostly ORT being
slow (30.6 ms at 1 thread), not us being fast.
- AVX-512 / VNNI untested — this host has neither.
- `trans_a` still declines. At `M = 1` transposing a single-row A is
meaningless.

## Gates

`cargo fmt --all`, **1455** ep-cpu lib tests, clippy `--all-targets -D
warnings` on x86_64, `verify_documented_env_vars.py`,
`workspace_test_packages.py verify`.

aarch64 clippy fails with three "never used" errors — I reproduced them
on a **pristine detached worktree of `origin/main`** before assuming
they were mine, and filed #1415. This branch adds no new ones.

## Files

- `crates/onnx-runtime-ep-cpu/src/kernels/half_gemv.rs` — `gemv_f16_nk`,
`dot_row_simd`, `dot_row_scalar`, `ROW_STRIPE`, the test counter, 6 new
tests.
- `crates/onnx-runtime-ep-cpu/src/kernels/gemm.rs` —
`try_half_fast_path` takes `trans_b` and picks a kernel; `execute` no
longer disqualifies on it; 2 new routing tests.
- `docs/benchmarks/2026-08-19-f16-gemm-transb-decode.md` — full record.
- `docs/performance/CPU_MATMUL_ASSIGNMENT.md` — 7 new `transB` rows
(they were absent entirely) + root-cause section 5.
- `scripts/ort_ab/gen_decode.py`, `scripts/ort_ab/sweep_decode.py` — the
harness; `gen_gemm.py` had no f16 `MatMul` or `Gemm` at all.

---------

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.

Native CPU path costs ~43 s/GB of weights to load: 14B is 561 s to first tokens, while decode itself is fine

2 participants