Repository navigation
MatMulNBits: gate accuracy-4 M=1 decode on the host actually having an int8 dot product (up to 15x) - #1028
Conversation
…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 Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ 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
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
|
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>
|
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:
It is not per-thread (flat at 2/4/8/16 threads) and the memory ledger does not see it: Two things specific to this PR: Its gate is a no-op on hosts that have VNNI, which includes mine. 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. |
…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>
Summary
try_mlas_sqnbitshort-circuitsm < sqnbit_decode_min()back to the hand int4/int8 decode kernels, on the premise that they beat MLAS SQNBit CompInt8 at smallmwhile also avoiding MLAS's one-time packing. That premise is a claim about the host, but it was encoded as a constant: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_directalready requires VNNI for the int4-direct route, and without it the AVX2 fallback emulates the dot withvpmaddubsw+vpmaddwd+ widening adds. On such a host the short-circuit sent every M=1accuracy_level = 4decode 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::MatMulNBitsmodels run through both our native CPU EP and the real ORT 1.27 CPU EP in one process, interleaved, viabench_generic.--ort-intra-threads 8vsONNX_GENAI_CPU_DECODE_THREADS=8.Against ORT these nodes move from 10.7x–21.6x slower to 1.7x–2.4x slower, parity
PASSon 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 selectedDotKernelrather than re-probing CPUID, soONNX_GENAI_CPU_DOT_KERNEL-style overrides and the test harness stay consistent with what actually executes.uses_vnni_int4_direct()(AVX-VNNI / AVX-512-VNNI).true. NEON is baseline on ARM64, soNeon/NeonDotalways have a real dot product — ARM behaviour is unchanged.false. The scalarDotKernelhas 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
SQ4BitGemmM1Kernel_CompInt8_avx2asymmetric kernel is numerically broken (~46% disagreement, already refused in-tree byhost_supports_mlas_sqnbit_m1_asym_int8). Correctness wins. Serving these with CompFp32 instead would need a second packed weight cache, sinceSQNBitPackedBcarries its compute type and the session cache holds exactly one — that memory cost needs its own justification, so it is not bundled here.bits = 8M=1Both 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 expectsSome(())below the crossover where the host has no native dot, andNonewhere 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_kernelpins the predicate toselected_dot_kernel()per architecture.Validation
Relationship to #1027
Independent. #1027 fixes
accuracy_level = 0(borrowed path pre-empting MLAS CompFp32); this fixesaccuracy_level = 4M=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 == falsethere 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 constantfalsegate 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_pathlocks that in, andmatmulnbits_try_mlas_gates_decode_by_m_thresholdnow 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_kernelcross-checked the predicate againstselected_dot_kernel().uses_vnni_int4_direct()— the implementation restated. Fixed in the same commit: it now checks CPUID directly (avx512vnni || avxvnni), and skips whenONNX_GENAI_CPU_DOT_KERNELoverrides the selection, since there the predicate must follow what actually executes rather than what the hardware advertises.CI
mainis currently red for reasons that predate and are untouched by this PR — verified by checking out unmodifiedorigin/mainin this worktree, wherecargo fmt --all -- --checkalready flagscrates/onnx-genai-ort/src/lib.rsandcrates/onnx-runtime-ep-cuda/src/kernels/matmul_nbits.rs. The last 8CIruns onmainall concludefailure(independentCLI ORTbuild,CUDA compileinventory/clippy andRust (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:
Verified pre-existing CI failing set
Re-checked against the unmodified baseline commit
0b872ed2f(CI run 31902589284).
Exactly these eight jobs fail on
mainwith 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)andRust qualityboth fail in theirCheck formattingstep:cargo fmt --all -- --checkreports the identical 7diffs on this branch and on unmodified
origin/main, all inonnx-genai-ortand
onnx-runtime-ep-cuda, none in a file this PR touches.CI baseline
Rebuilt on top of
main@400fbe246(a plaingit merge origin/main, no history rewrite).Unmodified
main@400fbe246fails exactly these 6 jobs(run 31914831964):
mainCLI ORT (Linux x86_64)CLI ORT (Windows x86_64)CUDA compile (Linux x86_64)CUDA compile (Windows x86_64)Rust (Windows ARM64)Rust coverage (macOS arm64)None are touched by this PR.
Fast (Linux x86_64)andRust qualitypreviously failed onmaintoo (a repo-widecargo fmtdrift, fixed onmainby #1043); after merging currentmaininto this branch both are green here, which confirms those earlier reds were never mine.The jobs this PR is actually accountable for --
Fast (Linux x86_64),Rust quality,EP conformance (Linux x86_64),Rust coverage (Linux x86_64),Miri unsafe-crate soundness,auditandcodecov-- are green.Ratio convention (added post-merge for clarity)
Columns
before/ORTandafter/ORTare ours/ORT: >1 means we are slower.gainis 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.