Skip to content

perf(cpu): pack transposed B along its stored axis in the f16/bf16 GEMM (16.0x -> 7.0x vs ORT) - #1446

Merged
justinchuby merged 6 commits into
mainfrom
squad/roy-f16-nt-gemm
Aug 20, 2026
Merged

justinchuby merged 6 commits into
mainfrom
squad/roy-f16-nt-gemm

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 19, 2026 •

Copy link
Copy Markdown
Owner

What

half_gemm.rs packs the B operand into panels before the micro-kernel runs. For an ONNX Gemm
with transB = 1, B is stored [n, k] and the layout is MatrixLayout::transposed(k) —
row_stride = 1, column_stride = k. The general branch of pack_b walked column in its inner
loop, so on transposed B it read with stride k: one cache line and one scalar to_f32 per
element, never reaching the F16C/NEON widening the contiguous path already had.

The fix is in the packer, not the micro-kernel. pack_b_transposed reads along the stored
(contiguous) direction, so each logical column is a contiguous run through the same
T::pack_contiguous helper the row-major path uses. Filling a [depth][column] panel from
column-major input is a transpose, so it processes TRANSPOSED_PACK_GROUP = 4 columns at a time to
keep the stores contiguous as well. The group width was swept, not guessed — the table is in the
constant's docs.

No new intrinsics and no new architecture gates. All widening is delegated to the existing
per-format dispatcher, so aarch64/NEON is served by the unchanged widen_f16_neon / widen_bf16_neon.

Measured on the merged head, immediately before landing

main moved a long way since this PR was opened, so the full sweep was re-run on the merged head.
Same-tree A/B — only half_gemm.rs is reverted for the base arm, so the harness is byte-identical
in both arms. Native-only p50, 5 interleaved trials, null = the base binary under a second name
in the same invocation.

26 of 27 prefill cells improve beyond their own noise floor, -24.2% to -65.0% (1.32x - 2.86x) —
stronger than the 1.29x - 2.66x this PR originally claimed. The 27th, qwen3_0p6b_m32 t=8, reads
-40.9% but its null moved 46.3%, so it is not claimed.

cell t base ms new ms delta speedup
llama3_8b_m32 8 57.18 20.01 -65.0% 2.86x
qwen3_8b_m32 16 106.15 41.58 -60.8% 2.55x
llama3_8b_m128 8 132.66 47.61 -64.1% 2.79x
qwen3_8b_m128 8 79.82 35.78 -55.2% 2.23x
llama3_8b_m512 16 203.15 98.97 -51.3% 2.05x
qwen3_8b_m512 1 654.63 496.33 -24.2% 1.32x

Full 27-cell table, the nn pre-transposed negative control (row-major B must not move — seven of
nine within 0.5%), and the roofline arithmetic are in
docs/benchmarks/2026-08-19-f16-nt-gemm-packing.md.

Headline against ORT: M = 128 / t = 8 goes 14.0x -> 7.0x within one invocation.

M = 1 is unaffected, and a null-too-tight artifact worth recording

At M = 1 with transB = 1 the Gemm path takes gemv_f16_nk and never reaches pack_b, so no
code this PR touches runs. At 5 trials x 15 runs, three M = 1 cells nevertheless read as
regressions above their null (+10.20% vs 2.04%, +6.10% vs 1.22%, +1.85% vs 0.12%). Re-measured at
11 trials x 40 runs, all six M = 1 cells are within noise (+6.82% vs 11.36%, -3.16% vs 11.58%,
+0.62% vs 0.74%, +0.00%, +0.00%, -2.21% vs 5.15%) — the longer run widened the nulls from 0.1-2% to
5-11.6% and the deltas did not follow. No M = 1 regression is claimed, and none is supported.

Numerics

Bit-identical for f16: sweeping all 65536 patterns on this host, scalar f16::to_f32 and F16C
_mm256_cvtph_ps agree on every pattern, NaNs included. The one divergence is bf16, on the
126 signalling NaNs — scalar bf16::to_f32 quiets them (0x7f81 -> 0x7fc1_0000) where the
<< 16 widening preserves the signalling bit. Both are NaN, no numeric contract is affected, and
the change makes transposed B agree with row-major B, which always used the shift.

transposed_b_is_bit_identical_to_pre_transposed_row_major asserts to_bits() equality (not a
tolerance) across shapes that straddle TRANSPOSED_PACK_GROUP, NR, NC and KC — trailing
partial groups of size 1, 2 and 3, partial depth panels, single-column panels — for both formats and
every execution path including Scalar.
a_layout_strided_in_both_directions_still_packs_correctly pins the general branch so the fast path
cannot quietly become the only column-strided route. End-to-end parity=PASS on all 738 A/B rows.

Validation

20-step local matrix on the merged head: 19/20. The only failure is H check-win-arm64, which
is pre-existing on clean main — onnx-genai-ort-sys bindgen needs the Windows SDK
('stdlib.h' file not found) and fails identically without this change. G clippy-aarch64-linux
(-D warnings) passes and is the meaningful ARM gate. K no-mlas-artifacts passes: the default
artifact still has 0 MLAS symbols and never defers to the ORT CPU EP. 1583 tests pass.

roy and others added 2 commits August 19, 2026 09:36
`Gemm` f16 with `transB = 1` -- the shape a fused QKV projection takes when
the weight is stored output-major -- measured 4.48x against ORT at one thread
and 15.99x at eight on `M = 128, K = N = 3584`. It was one of the few cells
that got dramatically *worse* as threads were added.

The cost was in `half_gemm::pack_b`, not in any kernel. `MatrixLayout::
transposed(k)` has `column_stride = k`, and the packer's inner loop walked
`column`, so it read `B` with stride `k` and did one scalar `to_f32` per
element, never reaching F16C. Transposed `B` stores each logical column
*contiguously*; the packer was walking the one axis that strides.

Localised with a control rather than a ratio: `gen_f16_nt.py` emits each shape
twice, `transB = 1` and `transB = 0` over the same array pre-transposed. Same
product, same numbers, same micro-kernel -- only the layout handed to `pack_b`
differs. The transposed spelling cost a flat ~62 ms more at *both* t=1 and
t=8, and a constant that will not shrink under 8x the cores is not arithmetic.
It is `pack_b` running once per row-block while `gemm_impl` scales the
row-block count with the thread count.

`pack_b_transposed` reads along the stored direction, so each column is a
contiguous run through the same `T::pack_contiguous` the row-major path uses
(F16C on x86, FP16/bf16 on NEON -- no new intrinsics, no new arch gates).
Columns go `TRANSPOSED_PACK_GROUP = 4` at a time so the transposing stores are
contiguous too; the group width is swept in the constant's docs.

27 cells across three geometries, 27 wins, 1.29x-2.66x. The headline cell,
M=128/t=8, goes 16.0x -> 7.0x against ORT.

Results are bit-identical: widening is exact and elementwise, so fill order
cannot change the packed panel. `transposed_b_is_bit_identical_to_pre_
transposed_row_major` asserts `to_bits()` equality across 10 shapes straddling
every blocking constant, both formats, and every execution path including
Scalar. `a_layout_strided_in_both_directions_still_packs_correctly` keeps the
generic branch reachable.

Bounded honestly: this removes two-thirds to three-quarters of the layout
penalty and nothing else. The pre-transposed control is itself still 2.6-3.3x
behind ORT, and `B` is still re-packed per row-block -- most of the 9-20 ms
residual. Both are recorded as open.

Also fixes `bench_generic --ort-only`, which refused any f16 graph: it
synthesized only f32/i32/i64 while the paired path already handled Float16,
Uint8 and Int8. That made the separate-arm method -- mandatory here, since
ORT's spin-waiting intra-op pool depresses a paired native arm -- unavailable
for exactly the half-precision kernels that most need it.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
- The headline before/after mixed two invocations: 16.0x came from the
  reproduction run, 7.0x from the results table. ORT's own time differs between
  those sessions (5.99 vs 5.70 ms), which is exactly the mistake the harness
  README warns about. Quote the same-invocation pair, 14.0x -> 7.0x, and say
  where the 15.99x came from.
- Disclose the one real behaviour change: on x86 transposed B now converts
  through F16C rather than scalar `to_f32`. `f16 -> f32` is exact for every
  finite value, subnormal and infinity, so the two can differ only in NaN
  payload bits -- and this makes the transposed path agree with the row-major
  path, which always used F16C. The bit-identity test's sin/cos data never
  generates a NaN, so note that it does not cover this.
- `gen_f16_nt.py` described `M = 1` as a negative control that "routes to
  half_gemv, not to this GEMM". That is the post-#1417 world; on the base
  commit `M = 1` still falls into the blocked GEMM, as the benchmark doc
  itself says.
- Tie the `widened` buffer size to the `panel_depth <= KC` invariant it
  depends on, and note the zero-init is dead but not worth `MaybeUninit`.

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

Copy link
Copy Markdown
Owner Author

Opus adversarial review — disposition

Verdict: APPROVE WITH NITS, no blocking findings. The reviewer independently re-implemented pack_b_transposed in a standalone harness and fuzzed it against the general branch (0 mismatches over 31 tiles, including k=1, n=1, panel_columns not a multiple of GROUP, and panel_depth < KC), mutation-tested the indices to confirm the new test actually catches a subtle swap/off-by-one, verified the fast path reads the identical index set to the pre-existing branch (so no new OOB is possible), recomputed six sampled rows of the derived tables, and confirmed scope was clean.

All five nits fixed in 186c184d8:

1. Headline mixed two invocations. Fair, and it is precisely the error the harness README warns about. 16.0x -> 7.0x took "before" from the reproduction run and "after" from the results table; ORT's own time differs between those sessions (5.99 vs 5.70 ms). Now quotes the same-invocation pair 14.0x -> 7.0x in both the benchmark doc and the work list, and says explicitly where the 15.99x came from and why the denominator moved. The size of the change is unaffected.

2. Undisclosed benign behaviour change. Correct that this was not stated. On x86, transposed B now converts through F16C _mm256_cvtph_ps instead of scalar half::f16::to_f32. f16 -> f32 is exact for every finite value, subnormal and infinity, so the two can differ only in NaN payload bits — and the change makes the transposed path agree with the row-major path, which always used F16C. Now documented both at pack_b_transposed and in the benchmark doc's numerics section, including the admission that the test's sin/cos data never generates a NaN so the test does not cover it.

3. gen_f16_nt.py comment described the post-#1417 world. Right — it called M = 1 a negative control that "routes to half_gemv, not to this GEMM", but on e13460af6 it still falls into the blocked GEMM, as the benchmark doc itself says two paragraphs later. Rewritten to match the base commit.

4. Dead zero-init of the 2 KiB scratch. Confirmed dead — every lane is fully overwritten by pack_contiguous before it is read. Left as-is with a comment: removing it means MaybeUninit, and buying back one 2 KiB memset per tile is not worth introducing unsafe into a packer.

5. Sizing invariant only debug_assert-checked. Added a comment tying the buffer size to panel_depth <= KC and naming the caller that guarantees it, so a future second caller has something to trip over.

Two reviewer observations I want to endorse rather than paper over

  • "'Arithmetic parallelises' is rhetorically loose — memory-bound arithmetic wouldn't either." Correct. The inference survives only because the transB = 0 control already removes all of the arithmetic and all of its memory traffic; the flat ~62 ms is what is left when the only surviving difference is pack_b's route. The conclusion stands on the control, not on that sentence.
  • The reviewer noted k == 1 makes transposed(1) have column_stride == 1, so it falls into the contiguous branch — safe, and coincidentally correct, since both spellings degenerate to the same thing.

Re-verified after the fixes: cargo fmt --check clean, 1449 tests pass, clippy clean, generator parses.

@github-actions

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
🔴 gather/small_f32_threads=1-internal/4096 639.1 ns 903.9 ns +41.4%
⚠️ sampling_latency/top_k_per_token 62.95 µs 81.18 µs +29.0%
⚠️ reduce_mean/large_f32_threads=1-internal/262144 957.28 µs 1.18 ms +23.7%
⚠️ gather/small_f16_threads=1-internal/4096 485.7 ns 574.1 ns +18.2%
⚠️ tokenization/decode_tokens_per_second 6.36 ms 7.44 ms +17.1%
⚠️ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.82 ms 4.43 ms +15.8%
⚠️ reduce_mean/medium_f32_threads=1-internal/65536 249.51 µs 287.20 µs +15.1%
✅ gather/large_f16_threads=1-internal/131072 13.49 µs 15.39 µs +14.0%
✅ matmul/large_generic_f32_threads=8/32x1024x1024 3.82 ms 4.34 ms +13.4%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 92.11 µs 100.64 µs +9.3%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 2.29 ms 2.49 ms +8.8%
✅ matmul/small_generic_f16_threads=1/1x256x256 31.46 µs 33.88 µs +7.7%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.71 ms 6.05 ms +6.0%
✅ reduce_mean/small_f32_threads=1-internal/4096 14.64 µs 15.51 µs +5.9%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 80.71 µs 84.97 µs +5.3%
✅ block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 109.41 µs 114.24 µs +4.4%
✅ sampling_latency/greedy_per_token 3.71 µs 3.86 µs +4.0%
✅ add/medium_bf16_threads=1-internal/262144 113.32 µs 115.48 µs +1.9%
✅ gather/small_bf16_threads=1-internal/4096 553.4 ns 562.3 ns +1.6%
✅ matmul/small_generic_f32_threads=8/1x256x256 40.47 µs 40.88 µs +1.0%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 405.16 µs 408.59 µs +0.8%
✅ matmul/large_generic_bf16_threads=8/32x1024x1024 1.28 ms 1.29 ms +0.8%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.08 ms 9.14 ms +0.6%
✅ add/large_f16_threads=1-internal/4194304 1.76 ms 1.76 ms +0.0%
✅ matmul/small_generic_f32_threads=1/1x256x256 42.67 µs 42.38 µs -0.7%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 568.96 µs 563.74 µs -0.9%
✅ qwen3_sampling_processors/top_k_top_p_fast 720.93 µs 708.77 µs -1.7%
✅ add/large_f32_threads=1-internal/4194304 855.26 µs 840.08 µs -1.8%
✅ matmul/medium_generic_bf16_threads=8/32x512x512 383.14 µs 373.54 µs -2.5%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 1.91 ms 1.86 ms -2.7%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 518.26 µs 502.65 µs -3.0%
✅ qwen3_sampling_processors/top_k_partial_selection 158.44 µs 153.51 µs -3.1%
✅ add/medium_f16_threads=1-internal/262144 128.59 µs 124.04 µs -3.5%
✅ gather/large_bf16_threads=1-internal/131072 14.48 µs 13.90 µs -4.0%
✅ sampling_latency/min_p_per_token 239.75 µs 228.68 µs -4.6%
✅ tokenization/encode_tokens_per_second 398.16 µs 378.28 µs -5.0%
✅ add/large_bf16_threads=1-internal/4194304 1.99 ms 1.87 ms -6.3%
✅ matmul/small_generic_bf16_threads=8/1x256x256 34.40 µs 31.79 µs -7.6%
✅ matmul/small_generic_bf16_threads=1/1x256x256 38.21 µs 34.83 µs -8.9%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.51 ms 2.27 ms -9.5%
✅ matmul/medium_generic_f16_threads=8/32x512x512 33.43 µs 30.02 µs -10.2%
✅ logit_processing/seven_processor_chain_per_step 389.37 µs 342.84 µs -11.9%
✅ kv_cache/alloc_dealloc_pages 46.68 µs 40.93 µs -12.3%
✅ block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 1.12 ms 977.35 µs -12.9%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 86.95 µs 75.71 µs -12.9%
✅ matmul/small_generic_f16_threads=8/1x256x256 34.58 µs 30.05 µs -13.1%
✅ sampling_latency/top_p_per_token 565.28 µs 489.19 µs -13.5%
🟢 matmul/medium_generic_f32_threads=8/32x512x512 1.09 ms 905.79 µs -16.6%
🟢 matmul/medium_generic_f16_threads=1/32x512x512 35.28 µs 28.96 µs -17.9%
🟢 gather/large_f32_threads=1-internal/131072 33.68 µs 27.56 µs -18.2%
🟢 block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 102.04 µs 79.76 µs -21.8%
🟢 add/small_bf16_threads=1-internal/1024 600.1 ns 455.0 ns -24.2%
🟢 add/medium_f32_threads=1-internal/262144 35.39 µs 26.57 µs -24.9%
🟢 add/small_f32_threads=1-internal/1024 289.7 ns 211.8 ns -26.9%
🟢 grammar_masking/llguidance_compute_mask/32 118.75 µs 84.21 µs -29.1%
🟢 gather/medium_f16_threads=1-internal/32768 4.07 µs 2.75 µs -32.6%
🟢 add/small_f16_threads=1-internal/1024 734.1 ns 470.9 ns -35.9%
🟢 gather/medium_f32_threads=1-internal/32768 7.12 µs 4.52 µs -36.5%
🟢 gather/medium_bf16_threads=1-internal/32768 5.25 µs 2.46 µs -53.2%

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: { 5.71 3.63 5.57 }
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)

@codecov

codecov Bot commented Aug 19, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 82.64%. Comparing base (3e2f21d) to head (390d207).
⚠️ Report is 41 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@             Coverage Diff             @@
##             main    #1446       +/-   ##
===========================================
+ Coverage   80.38%   82.64%    +2.26%     
===========================================
  Files         378       12      -366     
  Lines      167923     5475   -162448     
  Branches   167923     5475   -162448     
===========================================
- Hits       134980     4525   -130455     
+ Misses      28092      757    -27335     
+ Partials     4851      193     -4658     
Flag Coverage Δ
cli-ort-linux 82.60% <ø> (?)
cli-ort-windows 82.19% <ø> (ø)
offline ?

Flags with carried forward coverage won't be shown. Click here to find out more.
see 369 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.

justinchuby and others added 4 commits August 20, 2026 07:11
Renumber the ledger section 5 -> 8 behind the three that landed while this was
open, and close out section 6's forward reference to the NT packing problem.
Exhaustive sweep of all 65536 patterns: f16 scalar to_f32 and F16C
_mm256_cvtph_ps agree everywhere, so the f16 path is bit-identical
outright. The real divergence is 126 bf16 signalling NaNs, which scalar
to_f32 quiets and the << 16 widening preserves.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
26 of 27 prefill cells improve beyond their own null, -24.2% to -65.0%
(1.32x-2.86x), stronger than the 1.29x-2.66x claimed when the PR was
opened. The 27th is not claimed: its null moved 46.3%.

M=1 is unaffected and unclaimed (it takes gemv_f16_nk, never pack_b).
Three M=1 cells read as regressions above a very tight null at 5 trials;
at 11 trials x 40 runs all six were within noise, with the nulls widening
from 0.1-2% to 5-11.6%. Recorded as a null-too-tight artifact.

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

Copy link
Copy Markdown
Owner Author

Pre-merge validation on the merged head

Re-merged origin/main (3f09dd3c8, #1552) and re-ran everything. 19/20, the only failure pre-existing on clean main.

step result
A fmt-all / B build-offline / C test-ep-cpu / D clippy-offline PASS
E test-mlas-feature / F clippy-native-be / G clippy-aarch64-linux PASS
H check-win-arm64 FAIL — pre-existing, onnx-genai-ort-sys bindgen needs the Windows SDK ('stdlib.h' file not found); fails identically on unmodified main
I ep-cpu-no-default / J ep-cpu-all-features / K no-mlas-artifacts / L cross-compile-sh PASS
8 guard scripts (publish order, profile table, platform naming, dispatch reachability, dispatch manifest, feature-gate coverage, documented env vars, workspace test packages) PASS

1583 tests pass. K no-mlas-artifacts: default artifact has 0 MLAS symbols, no ORT CPU EP deferral.

Performance, re-measured on this head: 26 of 27 prefill cells improve beyond their own null, -24.2% to -65.0% (1.32x - 2.86x) — stronger than the 1.29x - 2.66x claimed when the PR was opened. The 27th (qwen3_0p6b_m32 t=8) is not claimed: its null moved 46.3%. M = 1 is unaffected and unclaimed; three cells that read as regressions at 5 trials were all within noise at 11 trials x 40 runs.

Numerics: parity=PASS on all 738 A/B rows. f16 is bit-identical (scalar and F16C agree on all 65536 patterns); bf16 differs only on the 126 signalling NaNs, where the scalar converter quiets and the << 16 widening preserves — both NaN, and it makes transposed B agree with row-major B.

Opus review: no BLOCKING findings. The one NON-BLOCKING finding was that my own comment and doc mis-attributed the NaN divergence to f16 payload bits; I brute-forced all 65536 patterns for both formats and rewrote both to the measured result (f16 0 differences, bf16 126 sNaN).

@justinchuby
justinchuby merged commit 583cafd into main Aug 20, 2026
6 checks passed
@justinchuby
justinchuby deleted the squad/roy-f16-nt-gemm branch August 20, 2026 07:59
justinchuby added a commit that referenced this pull request Aug 20, 2026
…row gate (1.31x on llama3_8b_qkv_t8) (#1556)

Closes #1471.

## The issue's premise was stale; the measured cost was real

#1471 proposed fusing dequantization into the L2 panel so the int4
prefill would stop
materializing a full f32 weight. **That fix had already landed** — in
#1356, the PR the issue
cites. `pack_b_quant` already dequantizes straight into the `KC x NR`
panel and
`quant_prefill_gebp` reuses each panel across all `m` rows. There is no
180 MB f32 panel on `main`
to remove.

The *fixed cost* the issue measured was real, though. Fitting `t = fixed
+ marginal * m` on the
GEBP arm at `4096x11008` gives **fixed = 4.80 ms**. Against the 22.5 MB
of packed weight that is
**4.7 GB/s** — 6% of this host's 75.8 GB/s roofline. It was never
bandwidth.

Reading the pack, the cause is the store pattern.
`Int4Weight::dequant_column` fills **one column**
of the panel at a time, writing `dst[p * NR + slot]` for consecutive
`p`: one f32 every 64 bytes,
so a separate cache line and a separate scalar store per element, on top
of fully scalar nibble
unpacking.

## The change

`BlockQuantWeight` gains a `dequant_panel` method whose **default
implementation is the old
per-column loop**, so `Int8Weight` is byte-for-byte unchanged.
`Int4Weight` overrides it with
`dequant_panel_avx2`, which dequantizes `DEQUANT_GROUP = 8` columns at
once into eight `__m256`,
transposes them in registers (`transpose8x8_ps`), and issues eight
contiguous 32-byte stores —
turning 8 scattered scalar stores into 1 vector store. This is the same
shape as #1446's
`TRANSPOSED_PACK_GROUP`; it is a recurring pattern in these packers.

The override is gated on `block_size % 8 == 0` and `has_simd_x86()`;
anything else keeps the
scalar default.

**The result is bit-identical**, not merely close. `_mm256_cvtepi32_ps`
on a zero-extended nibble
is exact, and the `sub` and `mul` are kept as separate intrinsics so
they cannot contract into an
FMA. `int4_dequant_panel_is_bit_identical_to_the_per_column_path`
asserts equality against
`dequant_column` across partial groups, `kc < 8`, `nr < 8`, both
zero-point modes, non-zero `pc`,
and the fallback block sizes. I mutation-tested it: perturbing the zero
point and swapping the
nibble order are both caught.

## The row gate had to move with it

`INT4_PREFILL_GEBP_MIN_ROWS` was measured against the *scalar* pack.
Halving the pack's fixed cost
makes that constant stale by construction, so I re-derived both
crossovers over 5 interleaved reps:

| regime | constant | old | new |
|---|---|---:|---:|
| non-resident | `INT4_PREFILL_GEBP_MIN_ROWS` | 12 | **5** |
| L2-resident | `INT4_PREFILL_GEBP_MIN_ROWS_L2_RESIDENT` | 24 | **12** |

`INT4_PREFILL_GEBP_MIN_ROWS_UNBLOCKED = 4` is untouched — block sizes
that are not a multiple of
32 have no better row kernel to cross over to.

Kernel-level effect at `4096x11008`: fixed **4.80 -> 2.24 ms**; `m=16`
6.05 -> 3.52 ms (1.72x);
`m=32` 7.30 -> 4.81 ms (1.52x).

## Production A/B (real ONNX models, `--native-only`, null control)

23 cells over `llama3_8b_qkv`, `llama3_8b_mlp`, `qwen3_0p6b_qkv` at t =
1/8/128/256/512 and
8/16/32 threads: **14 improved 1.02x - 1.31x, 9 within noise, 0
surviving regressions.**

| cell | threads | base ms | new ms | speedup |
|---|---:|---:|---:|---:|
| `llama3_8b_qkv_t8` | 8 | 3.808 | 2.897 | **1.31x** |
| `llama3_8b_mlp_t8` | 8 | 8.847 | 6.789 | **1.30x** |
| `llama3_8b_mlp_t128` | 8 | 23.71 | 18.19 | **1.30x** |
| `llama3_8b_qkv_t128` | 8 | 10.42 | 8.42 | **1.24x** |
| `llama3_8b_mlp_t8` | 16 | 5.933 | 4.877 | **1.22x** |

`llama3_8b_qkv_t8` is the exact cell the scheduler evidence flagged as
~10x behind ORT when
measured native-alone.

One cell (`qwen3_0p6b_qkv_t1`, 16 threads) read +7.69% against a 6.15%
null at 5 trials. It is
`m = 1`, which provably cannot reach any changed code — the lowest row
gate of any route is 4, and
the thresholds this PR moves do not apply at `m = 1` either (1 < 12
before, 1 < 5 after). It had to
be noise, and at 11 trials x 40 runs it is +1.82% against a 3.4% null.

## Versus ORT (paired, before/after from the same invocation)

| cell | before | after |
|---|---:|---:|
| `llama3_8b_qkv_t8` | 4.61x | **3.56x** |
| `llama3_8b_mlp_t8` | 3.42x | **2.73x** |
| `llama3_8b_qkv_t128` | 1.98x | **1.57x** |
| t512 shapes | ~1.2x | **~1.1x** |

## What this does *not* fix

- **8-bit prefill is untouched.** `Int8Weight` keeps the per-column
default. The same transpose
would help it; deliberately out of scope so the 8-bit path is provably
unchanged.
- **`m = 1` is unchanged.** Decode does not route here at all.
- **Still behind ORT at small `m`.** That gap is structural: MLAS
`SQNBitGemm` CompInt8 wants VNNI,
  which this host (AMD EPYC 9V74, AVX2+FMA+F16C) does not have.
- **`block_size % 8 != 0` falls back** to the scalar pack.
- Non-x86 targets are unaffected (scalar default).

## Validation

- 20-step matrix **19/20** — only `H check-win-arm64`, which fails
identically on clean `main`
(`onnx-genai-ort-sys` bindgen needs the Windows SDK: `'stdlib.h' file
not found`). Step G
  (clippy `aarch64-unknown-linux-gnu` `-D warnings`) is green.
- `cargo test -p onnx-runtime-ep-cpu --lib` — **1548 pass**, including
the new bit-identity test.
- Default-artifact no-MLAS guard green; production never defers to the
ORT CPU EP.
- Miri: **no coverage**, not a pass — `is_x86_feature_detected!("avx2")`
is false under Miri, so it
  exercises the scalar fallback rather than the new code.

Reviewed by Opus: no BLOCKING findings. It independently brute-forced
the nibble unpack against a
scalar oracle (0 mismatches over 40,960 byte quadruples) and verified
`transpose8x8_ps` with a
marker matrix. Its two NON-BLOCKING doc nits are fixed in `732868326`.

---------

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