Repository navigation
Native transposed-B SGEMM: eliminate the Kn dequant from default-build MatMulNBits prefill (#959, #1091) - #1176
Conversation
…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 Report❌ Patch coverage is Additional details and impacted files@@ 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
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
|
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>
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>
Second reconciliation: #1176 and #1403 both claim 8-bit prefillContinuing the ten-PR merge rehearsal, #1176 and #1403 collide the same way
Two of #1403's own tests fail on the merged stack while passing on its branch. Resolution: order by cache admission, not by merge orderThe two routes are not simply better or worse than each other; they amortize
So the dispatch now decides on 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 stackAll ten PRs merged together — #1176, #1013, #1356, #1431, #1365, #1381, #1403,
Other conflicts found in the rehearsal, and how they resolve
None of these require a change on any individual branch; they are merge-commit |
# Conflicts: # crates/onnx-runtime-ep-cpu/Cargo.toml # crates/onnx-runtime-ep-cpu/src/kernels/matmul.rs
Final validation on latest main — mergingHead Retraction: my earlier "6 of 12 cells regress" report was wrongI 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 Two independent methods now confirm this:
What this PR actually does — measured on 8-bit cellsThe
Native-alone, min-of-N, ORT threadpool not running (
Numerics: Cross-architecture
The Local gate matrix — all green
Miri ( No-MLAS default: 0 Note |
#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>
…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>
Summary
On the default (
mlasOFF) build thedequant-knprefill phase now goes to zero for everyMatMulNBitspath 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 contiguousNkweight instead of materializing a second, transposedKncopy.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
MatMulNBitsprefill (m > 1) on the default build dequantized the weight to f32 a second time in the transposedKn([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 contiguousNkpass 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 cachedNkweight to MLAS's cache-tiledsgemmwithtrans_b. The default build had no transposed-B GEMM, so its#[cfg(not(feature = "mlas"))]sibling returnedfalseand fell back to the slowKnpath.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)computingC[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 addspack_b_nt, which gathers each output column from its contiguousb_nkrow (unit stride overk) into the L1-resident pack tile at strideNR. The transpose becomes a pack-time reshape of an in-cache tile, not a full-array stride-Nscatter — exactly why MLAS wins.matmul.rs—nt_gemm_supported(backend)is the single legality predicate (no twocfgarms to drift);gemm_nt_with_backenddispatchesMlas → trans_b,SimdX86 → sgemm_simd_nt.matmul_nbits.rs— the two cfg-splittry_prefill_mlas_ntarms collapse into onetry_prefill_nk_ntgated onnt_gemm_supported. It dequantizes once into the pre-existingweight_nkOnceLock(the same slot decode caches into, so a constant weight pays one dequant, not two) and callsgemm_nt_with_backend. When the NT route runs, thedequant-knprofile phase is skipped by construction.Correctness — bit-identical
Byte-for-byte identical to the existing
Kndense route, not "within 1e-5".pack_b_ntproduces the same packed panelspack_bwould 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!onf32::to_bits):m ∈ {1, 2, 7, 33},n ∈ {1, 3, 63, 64, 65}, plus tile-exact and multi-KC.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. Themm_profile::time_prepack("dequant-kn", …)call lives in theif !used_fast_nt {…}branch; whenever the NT route returnstruethat branch is skipped, so the phase is eliminated by construction.dequant-nkis unchanged (still one pass, now the only one).Finding vs the #959 premise: on current local builds, 4-bit
acc0prefill takes the borrowed-int4 in-place path (#979/#1117) and 4-bitacc4uses the SDOT prepack — neither reaches the dense fallback, so no local q4 model emits a[mm_prepack] phase=dequant-knline to drive to zero end-to-end (confirmed empirically withONNX_GENAI_PROFILE_MM=1on qwen05b q4 / q4-acc4 / symzp). The beneficiaries of this change are therefore: 8-bit weights (m>1), grouped quantization, weight_prepacked, and 4-bit withaccuracy_level != 0that 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. Thedequant-kn*arm is a plain f32 transpose standing in for the stridedKnmaterialization the NT route removes (the real int4 dequant is ~2.9x heavier per #959):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/bpackare 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_nkOnceLock.apack/bpackare per-callvec![]scratch. The added lines are not matched byweight-cache-guard.yml(noOnceLock<…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 ignoredcargo 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→ cleancargo fmt --check→ the three changed files are clean (verified withrustfmt --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/BMerge
Merged
main(c55a3fab3), which had since made the #1091 M=1 GEMV the unconditionalSimdX86route and dropped theONNX_GENAI_CPU_MM_SIMD_M1_GEMVtoggle 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:gemm_nt_with_backendcompiled with neither themlasnor the x86 arm has no reader for any parameter, so the-D warningscross-arch pass rejected all six. Bound them in the unsupported arm. (This is what the oldRust qualityred was:Cross-target compile check, nothing else.)mm_profilegemv phase. The MLAS route used to time this GEMM on thegemvphase; 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 stayscfg-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 withg_idx. The harness uses no symbol this branch introduces, so the identical file runs onmainand here; both arms were built and run interleaved, 3 repetitions each, on the same box.Host: 32-core x86_64, AVX2, default build (
mlasoff), 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):
g_idxfallbackg_idxfallbackCold (fresh kernel per repetition, so the one-time weight dequant is inside the measurement — the TTFT term #959 attacked):
g_idxfallbackg_idxfallbackWhy steady moves 7–56x and cold only ~1.1–1.6x — and why that is the real finding. The
Knroute has no cache:dequantize_weight(WeightLayout::Kn)is called insideexecute, so every prefill call re-materializes the whole transposed f32 weight. The NT route dequantizes into the pre-existingweight_nkOnceLock, 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 contiguousNkwrite instead of the stride-Nscatter, 1.2–1.6x at these sizes; steady improves by the caching theNklayout makes possible.Confirmed structurally with
ONNX_GENAI_PROFILE_MM=1over the same harness run: main emits 72phase=dequant-knlines, this PR emits 0 — and 24phase=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 productionexecuteresult, not just of the kernel driver, verified across two separately compiled builds.Scope, restated honestly
On the default build, 4-bit
accuracy_level=0with 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 withg_idx,weight_prepacked/non-borrowable inputs, and 4-bitaccuracy_level != 0that 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 ignoredcargo clippy -p onnx-runtime-ep-cpu --all-targets -- -D warnings→ cleancargo 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 inonnx-runtime-ep-cuda/standard_attention.rsfrommain, fixed separately in fix(ci): restore the green Rust quality lane on main #1347.