Skip to content

MatMulNBits: gate accuracy-4 M=1 decode on the host actually having an int8 dot product (up to 15x) - #1028

Merged
justinchuby merged 3 commits into
mainfrom
squad/roy-int4-acc4-decode-isa
Aug 16, 2026
Merged

justinchuby merged 3 commits into
mainfrom
squad/roy-int4-acc4-decode-isa

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 15, 2026 •

Copy link
Copy Markdown
Owner

Summary

try_mlas_sqnbit short-circuits m < sqnbit_decode_min() back to the hand int4/int8 decode kernels, on the premise that they beat MLAS SQNBit CompInt8 at small m while also avoiding MLAS's one-time packing. That premise is a claim about the host, but it was encoded as a constant:

let hand_decode_is_fast =
    self.bits == 4 && self.accuracy_level == 4 && !prefer_arm64_mlas_decode;

The hand kernels are only fast where the int8 accumulation is a single instruction. On x86_64 that means AVX-VNNI or AVX-512-VNNI (vpdpbusd) — DotKernel::supports_int4_direct already requires VNNI for the int4-direct route, and without it the AVX2 fallback emulates the dot with vpmaddubsw + vpmaddwd + widening adds. On such a host the short-circuit sent every M=1 accuracy_level = 4 decode to a kernel an order of magnitude slower than the MLAS SQNBit kernel sitting directly behind the gate.

This is the decode hot path: M=1 is one generated token.

Measurement

Single-node com.microsoft::MatMulNBits models run through both our native CPU EP and the real ORT 1.27 CPU EP in one process, interleaved, via bench_generic.

  • Host: AMD EPYC 9V74 (Azure), 32 vCPU / 16 physical cores, AVX2 + FMA + F16C — no AVX-512, no AVX-VNNI, no AMX.
  • Thread-matched: --ort-intra-threads 8 vs ONNX_GENAI_CPU_DECODE_THREADS=8.
  • p50 of 15 runs after 5 warmups; before/after differ only by this commit.
  • ⚠️ Shared, contended runner (load average ~17). Absolute times are inflated; ratios are interleaved so they hold. Treat <1.2x as noise.
model (int4, symmetric) before after ORT before/ORT after/ORT gain
K=896 N=4864 M=1 block32 1.416 ms 0.094 ms 0.051 ms 21.6x 1.85x 15.1x
K=1024 N=3072 M=1 block16 1.074 ms 0.088 ms 0.045 ms 18.7x 1.96x 12.2x
K=1024 N=3072 M=1 block32 0.762 ms 0.078 ms 0.041 ms 16.1x 1.89x 9.8x
K=1024 N=3072 M=1 block128 0.570 ms 0.073 ms 0.038 ms 11.8x 1.91x 7.8x
K=3072 N=1024 M=1 block32 0.734 ms 0.101 ms 0.042 ms 14.6x 2.42x 7.3x
K=3584 N=3584 M=1 block32 1.386 ms 0.191 ms 0.112 ms 10.7x 1.71x 7.3x
K=1024 N=3072 M=1 block64 0.499 ms 0.074 ms 0.038 ms 10.8x 1.94x 6.7x

Against ORT these nodes move from 10.7x–21.6x slower to 1.7x–2.4x slower, parity PASS on every one.

m >= sqnbit_decode_min() (M=16, M=128) is unaffected — it was already above the crossover and measures flat (1.0x–1.4x, within noise on this host).

Design

hand_int8_decode_has_native_dot() reads the selected DotKernel rather than re-probing CPUID, so ONNX_GENAI_CPU_DOT_KERNEL-style overrides and the test harness stay consistent with what actually executes.

  • x86_64: requires uses_vnni_int4_direct() (AVX-VNNI / AVX-512-VNNI).
  • aarch64: unconditionally true. NEON is baseline on ARM64, so Neon/NeonDot always have a real dot product — ARM behaviour is unchanged.
  • other: false. The scalar DotKernel has no dot product, so MLAS (when present) should take the node.

On a VNNI x86 host the gate evaluates exactly as before, so this is a strict improvement for AVX2-only hardware and a no-op everywhere else.

Deliberately not changed — measured losses that remain

region after why declined
int4 asymmetric M=1 ~12.4x slower than ORT MLAS's SQ4BitGemmM1Kernel_CompInt8_avx2 asymmetric kernel is numerically broken (~46% disagreement, already refused in-tree by host_supports_mlas_sqnbit_m1_asym_int8). Correctness wins. Serving these with CompFp32 instead would need a second packed weight cache, since SQNBitPackedB carries its compute type and the session cache holds exactly one — that memory cost needs its own justification, so it is not bundled here.
bits = 8 M=1 13.8x–16.2x slower MLAS SQNBit has no 8-bit x86 kernel; there is nothing faster to defer to. Needs its own kernel work.

Both are reported rather than papered over: this PR does not claim to fix them.

Test changes

Two existing tests encoded the old unconditional premise:

  • matmulnbits_try_mlas_gates_decode_by_m_threshold — now expects Some(()) below the crossover where the host has no native dot, and None where it does. The decode/prefill split is still regression-locked; the decode half is now host-conditional, which is the actual invariant.
  • matmulnbits_accuracy4_prepack_reuses_selected_weight_format — the invariant under test ("one weight format, chosen once, reused across calls, never the f32 expansion") is unchanged; only which cache legitimately owns it differs by host, so the assertion now branches on the same predicate.

New: hand_int8_decode_native_dot_matches_selected_kernel pins the predicate to selected_dot_kernel() per architecture.

Validation

$ cargo test -p onnx-runtime-ep-cpu --features mlas --release
test result: ok. 1076 passed; 0 failed; 15 ignored
$ cargo test -p onnx-runtime-ep-cpu --release            # no mlas
test result: ok. 1079 passed; 0 failed; 10 ignored
$ cargo fmt --all                                        # clean
$ cargo clippy -p onnx-runtime-ep-cpu --features mlas --all-targets --release -- -D warnings
    0 errors
$ python3 scripts/check_feature_gate_coverage.py
✓ All feature-gated fast paths have fallback coverage.

Relationship to #1027

Independent. #1027 fixes accuracy_level = 0 (borrowed path pre-empting MLAS CompFp32); this fixes accuracy_level = 4 M=1 (decode crossover assuming an ISA the host may not have). They touch different gates and can land in either order.

Review

Independent Rubber Duck review (Opus, read-only): APPROVE, with two items I acted on.

1. Dynamic weights (raised by the reviewer, dismissed there, kept and fixed here). The reviewer flagged and then dismissed the possibility that routing decode to MLAS regresses non-constant weights. I did not dismiss it, because it is measurable: with can_prepack == false there is no session-lifetime buffer to hold MLAS's packed weight, so MLAS repacks the entire weight on every call — measured at 55.8 ms for a 6.4 MB int4 weight (ONNX_GENAI_PROFILE_MM=1) against a sub-millisecond hand decode. That is a large regression on exactly the case the old constant false gate happened to protect.

Fixed in bef29131b: the gate is now !can_prepack || hand_int8_decode_has_native_dot(). Dynamic weights keep the hand path on every ISA — a slow kernel beats repacking megabytes per token. matmulnbits_accuracy4_dynamic_weight_decode_keeps_hand_path locks that in, and matmulnbits_try_mlas_gates_decode_by_m_threshold now declares its weights constant, which is the only case where the ISA question is live.

2. Tautological test (MINOR). hand_int8_decode_native_dot_matches_selected_kernel cross-checked the predicate against selected_dot_kernel().uses_vnni_int4_direct() — the implementation restated. Fixed in the same commit: it now checks CPUID directly (avx512vnni || avxvnni), and skips when ONNX_GENAI_CPU_DOT_KERNEL overrides the selection, since there the predicate must follow what actually executes rather than what the hardware advertises.

CI

main is currently red for reasons that predate and are untouched by this PR — verified by checking out unmodified origin/main in this worktree, where cargo fmt --all -- --check already flags crates/onnx-genai-ort/src/lib.rs and crates/onnx-runtime-ep-cuda/src/kernels/matmul_nbits.rs. The last 8 CI runs on main all conclude failure (independent CLI ORT build, CUDA compile inventory/clippy and Rust (Windows ARM64) test failures). This PR changes exactly one file, crates/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs, and the fmt-flagged set is byte-identical before and after.

Locally green on this branch:

cargo test -p onnx-runtime-ep-cpu --features mlas --release --lib   # 1077 passed, 0 failed
cargo test -p onnx-runtime-ep-cpu --release --lib                   # 1079 passed, 0 failed
cargo clippy -p onnx-runtime-ep-cpu --features mlas --all-targets --release -- -D warnings   # clean
python3 scripts/check_platform_naming.py                            # PASS
python3 scripts/check_dispatch_reachability.py                      # PASS
python3 scripts/check_dispatch_manifest.py                          # PASS
python3 scripts/check_feature_gate_coverage.py                      # PASS
python3 .github/scripts/verify_documented_env_vars.py               # PASS

Verified pre-existing CI failing set

Re-checked against the unmodified baseline commit 0b872ed2f
(CI run 31902589284).
Exactly these eight jobs fail on main with no changes applied:

CLI ORT (Linux x86_64), CLI ORT (Windows x86_64),
CUDA compile (Linux x86_64), CUDA compile (Windows x86_64),
Fast (Linux x86_64), Rust (Windows ARM64),
Rust coverage (macOS arm64), Rust quality.

This is the same set this PR shows — no job fails here that does not already
fail on main. Fast (Linux x86_64) and Rust quality both fail in their
Check formatting step: cargo fmt --all -- --check reports the identical 7
diffs on this branch and on unmodified origin/main, all in onnx-genai-ort
and onnx-runtime-ep-cuda, none in a file this PR touches.

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)

Columns before/ORT and after/ORT are ours/ORT: >1 means we are slower. gain is ours-before/ours-after (this PR's own gain), not a comparison with ORT. p50 of 15 runs after 5 warmups, interleaved A/B, 8 threads on both sides, steady state. Shared, contended host: treat <1.2x as noise.

…he host has a native int8 dot

`try_mlas_sqnbit` short-circuits `m < sqnbit_decode_min()` back to the hand
int4/int8 decode kernels, on the premise that they beat MLAS SQNBit
CompInt8 at small `m` while also avoiding MLAS's one-time packing. That
premise is a claim about the *host*, but it was encoded as a constant:

    let hand_decode_is_fast =
        self.bits == 4 && self.accuracy_level == 4 && !prefer_arm64_mlas_decode;

The hand kernels are only fast where the int8 accumulation is a single
instruction. On x86_64 that means AVX-VNNI or AVX-512-VNNI (`vpdpbusd`);
`DotKernel::supports_int4_direct` already requires VNNI for the int4-direct
route, and without it the AVX2 fallback emulates the dot with
`vpmaddubsw` + `vpmaddwd` + widening adds. On such a host the short-circuit
sent every M=1 accuracy-4 decode to a kernel an order of magnitude slower
than the MLAS SQNBit kernel sitting right behind the gate.

Measured on an AMD EPYC 9V74 (AVX2/FMA/F16C, no AVX-512, no AVX-VNNI)
against ORT 1.27's CPU EP, both pinned to 8 intra-op threads, interleaved
A/B, p50 of 15 runs after 5 warmups:

  K=896  N=4864 M=1 block32   1.416 ms -> 0.094 ms  (15.1x)  vs ORT 0.051 ms
  K=1024 N=3072 M=1 block16   1.074 ms -> 0.088 ms  (12.2x)  vs ORT 0.045 ms
  K=1024 N=3072 M=1 block32   0.762 ms -> 0.078 ms   (9.8x)  vs ORT 0.041 ms
  K=1024 N=3072 M=1 block128  0.570 ms -> 0.073 ms   (7.8x)  vs ORT 0.038 ms
  K=3584 N=3584 M=1 block32   1.386 ms -> 0.191 ms   (7.3x)  vs ORT 0.112 ms
  K=3072 N=1024 M=1 block32   0.734 ms -> 0.101 ms   (7.3x)  vs ORT 0.042 ms

versus ORT those nodes move from 10.7x-21.6x slower to 1.7x-2.4x slower,
with parity PASS on every one.

`hand_int8_decode_has_native_dot` reads the *selected* `DotKernel` rather
than probing CPUID again, so `ONNX_GENAI_CPU_DOT_KERNEL`-style overrides and
the test harness stay consistent with what actually executes. aarch64 is
unconditionally true (NEON is baseline, so `Neon`/`NeonDot` always have a
real dot product) and non-x86/non-ARM is false (the scalar kernel has no dot
product at all), so ARM and VNNI x86 behaviour is unchanged.

Deliberately *not* changed, and still measured as losses on this host:

* asymmetric int4 M=1 (~12.4x slower than ORT). MLAS's
  `SQ4BitGemmM1Kernel_CompInt8_avx2` asymmetric kernel is numerically broken
  and is already refused by `host_supports_mlas_sqnbit_m1_asym_int8`, so
  these correctly stay on the hand path. Serving them with CompFp32 instead
  would need a second packed weight cache, since `SQNBitPackedB` carries its
  compute type and the session cache holds exactly one.
* `bits = 8` (13.8x-16.2x slower). MLAS SQNBit has no 8-bit x86 kernel, so
  there is nothing faster to defer to here.

Two existing tests encoded the old unconditional premise and now assert the
host-aware rule instead: `matmulnbits_try_mlas_gates_decode_by_m_threshold`
expects `Some(())` below the crossover where the host has no native dot, and
`matmulnbits_accuracy4_prepack_reuses_selected_weight_format` checks the
same "one weight format, chosen once, reused, never the f32 expansion"
invariant against whichever cache legitimately owns it.

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 78.78%. Comparing base (400fbe2) to head (8348af5).

Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1028      +/-   ##
==========================================
- Coverage   78.78%   78.78%   -0.01%     
==========================================
  Files         365      365              
  Lines      149050   149050              
  Branches   149050   149050              
==========================================
- Hits       117435   117434       -1     
- Misses      26978    26979       +1     
  Partials     4637     4637              
Flag Coverage Δ
cli-ort-linux 83.52% <ø> (ø)
cli-ort-windows 83.11% <ø> (ø)
mlas 80.87% <ø> (-0.05%) ⬇️
offline 78.57% <ø> (ø)

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 73.94% <ø> (ø)

... and 1 file 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
🔴 matmul/small_generic_f32_threads=8/1x256x256 32.18 µs 54.31 µs +68.8%
🔴 gather/large_f16_threads=1-internal/131072 9.28 µs 15.55 µs +67.5%
🔴 matmul/small_generic_f16_threads=8/1x256x256 28.85 µs 45.01 µs +56.0%
🔴 matmul/small_generic_f16_threads=1/1x256x256 28.30 µs 40.08 µs +41.6%
🔴 matmul/medium_generic_f16_threads=8/32x512x512 28.02 µs 39.54 µs +41.1%
🔴 matmul/medium_generic_f32_threads=8/32x512x512 1.06 ms 1.47 ms +38.7%
🔴 block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 381.05 µs 517.08 µs +35.7%
🔴 gather/large_bf16_threads=1-internal/131072 12.19 µs 16.47 µs +35.0%
⚠️ gather/large_f32_threads=1-internal/131072 30.26 µs 38.63 µs +27.7%
⚠️ matmul/small_generic_bf16_threads=1/1x256x256 29.04 µs 35.99 µs +23.9%
⚠️ kv_cache/alloc_dealloc_pages 40.16 µs 48.72 µs +21.3%
⚠️ reduce_mean/large_f32_threads=1-internal/262144 980.36 µs 1.17 ms +19.6%
⚠️ gather/medium_f16_threads=1-internal/32768 2.42 µs 2.83 µs +16.9%
✅ matmul/small_generic_bf16_threads=8/1x256x256 29.72 µs 33.81 µs +13.8%
✅ matmul/small_generic_f32_threads=1/1x256x256 33.83 µs 38.36 µs +13.4%
✅ gather/small_f32_threads=1-internal/4096 680.5 ns 759.8 ns +11.7%
✅ grammar_masking/llguidance_compute_mask/32 76.48 µs 85.34 µs +11.6%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.14 ms 2.38 ms +11.4%
✅ qwen3_sampling_processors/top_k_top_p_fast 670.68 µs 746.45 µs +11.3%
✅ gather/medium_f32_threads=1-internal/32768 3.81 µs 4.24 µs +11.2%
✅ matmul/medium_generic_f16_threads=1/32x512x512 27.65 µs 30.58 µs +10.6%
✅ gather/small_bf16_threads=1-internal/4096 488.1 ns 535.7 ns +9.7%
✅ reduce_mean/medium_f32_threads=1-internal/65536 240.95 µs 260.22 µs +8.0%
✅ tokenization/encode_tokens_per_second 406.06 µs 437.71 µs +7.8%
✅ sampling_latency/top_k_per_token 51.51 µs 54.62 µs +6.0%
✅ gather/small_f16_threads=1-internal/4096 530.8 ns 555.0 ns +4.5%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 78.41 µs 81.22 µs +3.6%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 6.17 ms 6.36 ms +3.2%
✅ block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 45.02 µs 45.52 µs +1.1%
✅ tokenization/decode_tokens_per_second 6.90 ms 6.97 ms +0.9%
✅ sampling_latency/greedy_per_token 3.27 µs 3.28 µs +0.3%
✅ logit_processing/seven_processor_chain_per_step 337.60 µs 337.28 µs -0.1%
✅ gather/medium_bf16_threads=1-internal/32768 3.16 µs 3.12 µs -1.4%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.67 ms 9.48 ms -2.0%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.81 ms 3.68 ms -3.3%
✅ add/small_f32_threads=1-internal/1024 203.3 ns 195.4 ns -3.9%
✅ add/small_bf16_threads=1-internal/1024 12.53 µs 12.00 µs -4.2%
✅ add/small_f16_threads=1-internal/1024 12.76 µs 12.05 µs -5.6%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 551.70 µs 519.06 µs -5.9%
✅ add/medium_f32_threads=1-internal/262144 2.51 ms 2.36 ms -6.3%
✅ add/large_f16_threads=1-internal/4194304 41.27 ms 38.63 ms -6.4%
✅ reduce_mean/small_f32_threads=1-internal/4096 15.80 µs 14.75 µs -6.7%
✅ add/large_bf16_threads=1-internal/4194304 45.47 ms 42.21 ms -7.2%
✅ add/medium_bf16_threads=1-internal/262144 2.69 ms 2.49 ms -7.6%
✅ add/medium_f16_threads=1-internal/262144 2.60 ms 2.39 ms -8.0%
✅ add/large_f32_threads=1-internal/4194304 41.21 ms 37.42 ms -9.2%
✅ sampling_latency/top_p_per_token 436.72 µs 386.59 µs -11.5%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 2.20 ms 1.89 ms -13.8%
🟢 qwen3_sampling_processors/top_k_partial_selection 185.51 µs 154.52 µs -16.7%
🟢 matmul/medium_generic_bf16_threads=8/32x512x512 606.06 µs 503.66 µs -16.9%
🟢 qwen3_sampling_processors/top_p_fast_after_top_k 626.88 µs 498.16 µs -20.5%
🟢 matmul/large_generic_bf16_threads=8/32x1024x1024 1.55 ms 1.23 ms -20.6%
🟢 sampling_latency/min_p_per_token 281.63 µs 211.44 µs -24.9%
🟢 matmul/large_generic_f32_threads=8/32x1024x1024 6.56 ms 4.83 ms -26.4%
🟢 matmul/large_generic_f16_threads=1/32x1024x1024 117.99 µs 76.26 µs -35.4%
🟢 qwen3_sampling_processors/top_k_full_sort_baseline 3.33 ms 2.09 ms -37.2%
🟢 matmul/large_generic_f16_threads=8/32x1024x1024 162.57 µs 91.07 µs -44.0%
🟢 block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 67.19 µs 36.90 µs -45.1%
🟢 block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 958.86 µs 499.10 µs -47.9%

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: { 3.73 3.45 4.97 }
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)

Review follow-up on the host-aware decode gate.

Routing decode to MLAS when the host has no `vpdpbusd` is right only when the
packed weight can be cached. With `can_prepack == false` there is no
session-lifetime buffer to hold it, so MLAS repacks the entire weight on every
call: measured 55.8 ms for a 6.4 MB int4 weight (`ONNX_GENAI_PROFILE_MM=1`),
against a sub-millisecond hand decode. That is a large regression for
dynamic-weight nodes -- exactly the case the previous constant `false` gate
happened to protect.

The gate now reads `!can_prepack || hand_int8_decode_has_native_dot()`: dynamic
weights keep the hand path on every ISA, because a slow kernel beats repacking
megabytes per token. `matmulnbits_accuracy4_dynamic_weight_decode_keeps_hand_path`
locks that in, and the M-threshold test now declares its weights constant, which
is the only case where the ISA question is live.

Also de-circularizes the capability test: it cross-checked
`hand_int8_decode_has_native_dot()` against
`selected_dot_kernel().uses_vnni_int4_direct()`, which is the implementation
restated. It now checks CPUID directly, skipping when
`ONNX_GENAI_CPU_DOT_KERNEL` overrides the selection -- there the predicate must
follow what actually executes, not what the hardware advertises.

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

Copy link
Copy Markdown
Owner Author

Related measurement from #1027, which routes the other accuracy level to the same MLAS SQNBit machinery -- please read it before this lands.

On my host (RTX 4060 laptop, 20 logical CPUs, 68.5 GB RAM, AVX2 + FMA + F16C + AVX-VNNI, no AVX-512) I measured MLAS's packed B buffer as a second private copy of the weights held beside the still-resident mapped file:

model borrowed path MLAS route delta
qwen05b-symzp (367 MB weights) 380 MB 993 MB +613 MB
qwen14b-symzp (8.55 GB weights) 8.17 GB 25.5 GB +17.4 GB

It is not per-thread (flat at 2/4/8/16 threads) and the memory ledger does not see it: resident_f32_cache_bytes=0 in both modes. #1027 is blocked on making those bytes visible and admitting them through the memory strategy; whatever mechanism lands there should cover this PR's route too.

Two things specific to this PR:

Its gate is a no-op on hosts that have VNNI, which includes mine. supports_int4_direct requires VNNI, and this host has AVX-VNNI, so the hand decode is retained and I cannot reproduce your speed-up locally -- I can only confirm no regression. That is worth stating in the PR body: the win is specific to AVX2-without-VNNI hosts, which is the right call, but it means the change is unmeasurable on a large class of machines and needs a host-capability line in the results table.

The premise is now correctly a host question rather than a constant, which is the right shape. No objection to the diagnosis. Please rebase on #1027 once its accounting lands so the two routes share one admission policy rather than each growing its own.

@justinchuby
justinchuby marked this pull request as ready for review August 15, 2026 23:52
@justinchuby
justinchuby merged commit 553eed1 into main Aug 16, 2026
11 of 17 checks passed
@justinchuby
justinchuby deleted the squad/roy-int4-acc4-decode-isa branch August 16, 2026 00:37
justinchuby added a commit that referenced this pull request Aug 16, 2026
…Fp32 (up to 33x, 78x->2x vs ORT) (#1027)

## Summary

`Kernel::execute` evaluates the `bits == 4 && accuracy_level == 0`
**borrowed zero-copy** branch *before* `try_mlas_sqnbit`. That branch
order made the `accuracy_level = 0` → MLAS SQNBit **CompFp32** route
unreachable dead code — the very route `try_mlas_sqnbit`'s own doc
comment describes as intended:

> For every other accuracy level the hand fallback is a slow
full-f32-dequant GEMV, so prefer MLAS SQNBit (CompFp32) whenever MLAS
actually has a kernel — this matches ORT/onnxruntime-genai, which treat
`accuracy_level` 0/1 as CompFp32.

On x86_64 the consequence is severe: `borrowed_affine_int4_matmul`'s
vectorized block dot is `#[cfg(target_arch = "aarch64")]`, so **on x86
every `bits=4, accuracy_level=0` node ran a scalar nibble-unpack GEMV**,
and for `m > 1` it re-streamed the whole weight once per activation row
with no GEMM blocking. `accuracy_level = 0` is exactly what Foundry
`cuda-gpu` int4 exports emit, so this is the common case, not a corner.

## Reproduction and measurement method

Single-node `com.microsoft::MatMulNBits` ONNX models (constant
B/scales/zero-points) executed through **both** our native CPU EP and
the **real ORT 1.27 CPU EP** in one process, interleaved, via
`bench_generic`.

- Host: AMD EPYC 9V74 (Azure), 32 vCPU / 16 physical cores, **AVX2 + FMA
+ F16C, no AVX-512, no AVX-VNNI, no AMX**.
- **Thread-matched**: `--ort-intra-threads 8` against
`ONNX_GENAI_CPU_DECODE_THREADS=8`.
- p50 of 9 runs after 3 warmups, interleaved A/B.
- Before/after are the **same binary**, toggled with
`ONNX_GENAI_CPU_MM_MLAS_QNBIT=0/1`, so compiler, ISA, allocator and host
load are identical across the pair.
- ⚠️ The runner is **shared and contended** (load average ~17 from other
jobs). Absolute times are inflated; the ratios are interleaved so they
hold, but treat sub-1.2x differences as noise.

## Results

| model (block32 unless noted) | before | after | ORT | before/ORT |
after/ORT | speed-up |
|---|---|---|---|---|---|---|
| K=1024 N=3072 M=1 | 1.587 ms | **0.169 ms** | 0.097 ms | 15.6x | 1.75x
| 9.4x |
| K=3072 N=1024 M=1 | 1.867 ms | **0.205 ms** | 0.087 ms | 16.3x | 2.36x
| 9.1x |
| K=3584 N=3584 M=1 | 4.625 ms | **0.525 ms** | 0.341 ms | 15.7x | 1.54x
| 8.8x |
| K=896 N=4864 M=1 | 2.147 ms | **0.246 ms** | 0.124 ms | 17.1x | 1.98x
| 8.7x |
| K=1024 N=3072 M=1 asym | 1.671 ms | **0.225 ms** | 0.121 ms | 15.8x |
1.86x | 7.4x |
| K=1024 N=3072 M=1 block128 | 1.563 ms | **0.211 ms** | 0.089 ms |
16.0x | 2.38x | 7.4x |
| K=1024 N=3072 M=1 block16 | 1.344 ms | **0.269 ms** | 0.108 ms | 12.3x
| 2.50x | 5.0x |
| K=1024 N=3072 M=16 | 18.118 ms | **0.959 ms** | 0.344 ms | 51.1x |
2.79x | 18.9x |
| K=3584 N=3584 M=16 | 67.283 ms | **3.340 ms** | 1.278 ms | 48.1x |
2.61x | 20.1x |
| K=896 N=4864 M=16 | 25.452 ms | **2.011 ms** | 0.444 ms | 51.6x |
4.53x | 12.7x |
| K=1024 N=3072 M=128 | 116.348 ms | **4.619 ms** | 1.746 ms | 70.5x |
2.65x | 25.2x |
| K=1024 N=3072 M=128 block128 | 106.963 ms | **3.236 ms** | 1.512 ms |
63.8x | 2.14x | 33.1x |
| K=3072 N=1024 M=128 | 123.315 ms | **4.039 ms** | 1.515 ms | 67.2x |
2.67x | 30.5x |
| K=3584 N=3584 M=128 | 385.308 ms | **13.139 ms** | 7.121 ms | 52.0x |
1.84x | 29.3x |
| K=896 N=4864 M=128 | 136.841 ms | **4.887 ms** | 2.241 ms | 62.3x |
2.18x | 28.0x |

`bits = 8` rows are untouched by the gate (`bits == 4`) and measured
flat, as expected.

**Against ORT the same node moves from 12x–78x slower to 1.5x–4.5x
slower.** This PR does not claim parity with ORT — see "What this PR
does *not* fix".

### Correctness

Parity vs ORT becomes **bit-identical**, because both runtimes now
execute the same MLAS kernels:

```
parity_output[0]: max_abs=0.000000e0 max_rel=0.000000e0 PASS
top1: native=2096 ort=2096 AGREE
```

### Setup cost is reported, not hidden

The one-time MLAS shard pack for the K=3584 N=3584 block-32 weight, from
`ONNX_GENAI_PROFILE_MM=1`:

```
[mm_prepack] phase=mlas-shards calls=1 prepack_total=55.8ms cum_bytes=6422528 this=55.80ms this_bytes=6422528
result: native=0.540 ms  native_p90=0.542 ms  native_min=0.526 ms  native_spread=1.00
```

55.8 ms once for 6.4 MB, against ~4.1 ms saved per call → repaid in ~14
calls, then amortized for the session. For decode workloads that is a
fraction of one token.

## Design: an explicit ownership predicate, not a branch reorder

`mlas_sqnbit_owns_fp32_compute(can_prepack, has_zero_points)` states the
rule where it can be read and tested:

- **MLAS takes the node only when `can_prepack`** — B/scales/zero-points
are graph constants — so the packed buffer is built once per session.
With **dynamic weights** MLAS would repack on *every* call with no cache
to amortize against, so those explicitly stay on the borrowed zero-copy
path. This is the deliberate decline half of the policy, covered by
`matmulnbits_int4_acc0_dynamic_weight_keeps_borrowed_path`.
- **`sqnbit_packed_b_size` is MLAS's own "do I have a kernel for this
shape on this host" probe.** When MLAS says no (e.g. `block_size = 8`),
the borrowed path remains the fallback — nothing is left stranded.
- **Without the `mlas` feature the predicate is `false`**, so non-MLAS
builds are byte-for-byte unchanged.

### #979 is preserved

#979 removed the resident f32 `weight_nk` expansion (~8x the file size
in RAM). MLAS's packed buffer is int4-sized, so that expansion is still
never built. The #979 regression tests now assert the memory invariant
*directly* (`weight_nk` stays empty) alongside the route actually taken,
where the expected route is derived from the **same predicate production
code uses** — so on hosts/shapes where MLAS declines they stay exact
assertions rather than degrading into tautologies. Under `--features
mlas` the route proof reads **per-kernel** state (`mlas_shards` /
`mlas_packed`), not the process-global test counters, which are shared
with tests running in parallel and cannot support a negative assertion.

## What this PR does *not* fix (measured, reported, not hidden)

| region | after this PR | why |
|---|---|---|
| int4 `accuracy_level = 0`, all shapes | still **1.5x–4.5x** slower
than ORT | our static N-shard split vs MLAS's dynamic tile partitioning;
separate investigation |
| int4 `accuracy_level = 4`, M=1 | **10.7x–18.7x** slower |
`sqnbit_decode_min()` short-circuits M=1 onto the hand int8 kernel,
which needs AVX-VNNI to be fast. Follow-up PR. |
| int4 `accuracy_level = 4`, M=1, **asymmetric**, AVX2-only | ~11.7x
slower | MLAS's `SQ4BitGemmM1Kernel_CompInt8_avx2` asymmetric kernel is
numerically broken (already guarded in-tree). Correct decline;
documented gap. |
| `bits = 8` `accuracy_level = 4` | 13.8x–14.6x slower | MLAS SQNBit has
no 8-bit x86 kernel; separate work. |
| default (non-`mlas`) builds | unchanged, i.e. still 12x–78x | **the
`mlas` feature is not enabled for published wheels/CLI** — see below. |

### The `mlas` feature gap is the biggest remaining exposure

This fix only helps builds compiled with `--features mlas`, which is
**opt-in and not enabled for shipped artifacts**. On a default x86 build
the borrowed scalar GEMV is still what runs. Two follow-ups are needed
and are being handled separately: (1) an AVX2 int4 CompFp32 GEMV/GEMM so
default builds are not scalar, and (2) a decision on enabling `mlas` by
default for `x86_64-unknown-linux-gnu` artifacts.

## Validation

```
$ cargo test -p onnx-runtime-ep-cpu --features mlas --release
test result: ok. 1077 passed; 0 failed; 15 ignored
$ cargo test -p onnx-runtime-ep-cpu --release            # no mlas
test result: ok. 1080 passed; 0 failed; 10 ignored
$ cargo fmt --all                                        # clean
$ cargo clippy -p onnx-runtime-ep-cpu --features mlas --all-targets --release -- -D warnings
    (no warnings)
$ python3 scripts/check_feature_gate_coverage.py
✓ All feature-gated fast paths have fallback coverage.
```

Both `cfg(feature = "mlas")` and `cfg(not(feature = "mlas"))` variants
of the predicate exist, so the feature-gate coverage rule is satisfied.

## Review

Independent Rubber Duck review (Opus, read-only): **APPROVE**.

One MINOR finding: `Int4Acc0RouteProbe::assert_fast_route` derived its
expectation from `mlas_sqnbit_owns_fp32_compute` — the very predicate
under test — so it caught "the code disagrees with its own policy" but
not "the policy regressed to never choosing MLAS".

**Fixed** in `4c599eafd`:
`int4_acc0_constant_weight_reaches_mlas_on_x86_64` pins the concrete
case this PR exists for with a hardcoded expectation — on x86_64 with
the vendored MLAS, a constant-weight symmetric int4 `accuracy_level = 0`
node at block size 32 **must** reach MLAS SQNBit. Reintroducing the
branch order that made that route dead code now fails a test even if the
predicate is edited to agree with it.

## CI

`main` is currently red for reasons that predate and are untouched by
this PR. Verified by checking out unmodified `origin/main` in this
worktree:

* `cargo fmt --all -- --check` flags `crates/onnx-genai-ort/src/lib.rs`
and `crates/onnx-runtime-ep-cuda/src/kernels/matmul_nbits.rs` on
`origin/main` itself. This PR changes exactly one file,
`crates/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs`, and the
flagged set is byte-identical before and after.
* Independent pre-existing failures on the same runs: `CLI ORT
(Linux/Windows)` → *Build onnx-genai-cli*, `CUDA compile (Linux)` →
*Verify CUDA test inventory*, `CUDA compile (Windows)` → *Clippy CUDA
EP*, `Rust (Windows ARM64)` → *Test cross-platform offline crates*. The
last 8 `CI` runs on `main` all conclude `failure`.

Locally green on this branch:

```
cargo test -p onnx-runtime-ep-cpu --features mlas --release --lib   # 1078 passed, 0 failed
cargo test -p onnx-runtime-ep-cpu --release --lib                   # 1080 passed, 0 failed
cargo clippy -p onnx-runtime-ep-cpu --features mlas --all-targets --release -- -D warnings   # clean
python3 scripts/check_platform_naming.py                            # PASS
python3 scripts/check_dispatch_reachability.py                      # PASS
python3 scripts/check_dispatch_manifest.py                          # PASS
python3 scripts/check_feature_gate_coverage.py                      # PASS
python3 .github/scripts/verify_documented_env_vars.py               # PASS
```

---

### Verified pre-existing CI failing set

Re-checked against the unmodified baseline commit `0b872ed2f`
([CI run
31902589284](https://github.com/justinchuby/onnx-genai/actions/runs/31902589284)).
Exactly these eight jobs fail on `main` with no changes applied:

`CLI ORT (Linux x86_64)`, `CLI ORT (Windows x86_64)`,
`CUDA compile (Linux x86_64)`, `CUDA compile (Windows x86_64)`,
`Fast (Linux x86_64)`, `Rust (Windows ARM64)`,
`Rust coverage (macOS arm64)`, `Rust quality`.

This is the same set this PR shows — no job fails here that does not
already
fail on `main`. `Fast (Linux x86_64)` and `Rust quality` both fail in
their
`Check formatting` step: `cargo fmt --all -- --check` reports the
identical 7
diffs on this branch and on unmodified `origin/main`, all in
`onnx-genai-ort`
and `onnx-runtime-ep-cuda`, none in a file this PR touches.

## Re-validated against merged `main` (`eac0ee24e`)

The other seven PRs in this series are now merged. I rebuilt both arms
from source on the
same host and confirmed the binaries actually differ (`md5 a61b3130` for
`main`,
`8e842c67` with this PR) before trusting the numbers, then measured
`bench_generic --runs 40 --warmups 10 --native-threads 8
--ort-intra-threads 8`:

| Shape (`accuracy_level=0`, b4, bs32, sym) | `main` vs ORT | this PR vs
ORT | gain |
|---|---|---|---|
| K=1024 N=3072 **M=1** (decode) | **15.21x slower** | **1.86x slower**
| **8.2x** |
| K=1024 N=3072 **M=128** (prefill) | **70.49x slower** | **2.30x
slower** | **30.6x** |

`accuracy_level=0` is the *only* remaining int4 regression on `main` --
`acc4` is already
1.92x (M=1) / 2.38x (M=128) after #1028. Numeric parity is `PASS`
(bit-identical) in every
run. Absolute timings on this shared box drift 1.5-2x between runs, so
the ratios matter,
not the milliseconds.


## 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](https://github.com/justinchuby/onnx-genai/actions/runs/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.

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant