Repository navigation
perf(cpu): pack transposed B along its stored axis in the f16/bf16 GEMM (16.0x -> 7.0x vs ORT) - #1446
Conversation
`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>
Opus adversarial review — dispositionVerdict: APPROVE WITH NITS, no blocking findings. The reviewer independently re-implemented All five nits fixed in 1. Headline mixed two invocations. Fair, and it is precisely the error the harness README warns about. 2. Undisclosed benign behaviour change. Correct that this was not stated. On x86, transposed 3. 4. Dead zero-init of the 2 KiB scratch. Confirmed dead — every lane is fully overwritten by 5. Sizing invariant only Two reviewer observations I want to endorse rather than paper over
Re-verified after the fixes: |
🔴 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
|
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ 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
Flags with carried forward coverage won't be shown. Click here to find out more. 🚀 New features to boost your workflow:
|
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>
Pre-merge validation on the merged headRe-merged
1583 tests pass. 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 ( Numerics: 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 ( |
…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>
What
half_gemm.rspacks theBoperand into panels before the micro-kernel runs. For an ONNXGemmwith
transB = 1,Bis stored[n, k]and the layout isMatrixLayout::transposed(k)—row_stride = 1, column_stride = k. The general branch ofpack_bwalkedcolumnin its innerloop, so on transposed
Bit read with stridek: one cache line and one scalarto_f32perelement, never reaching the F16C/NEON widening the contiguous path already had.
The fix is in the packer, not the micro-kernel.
pack_b_transposedreads along the stored(contiguous) direction, so each logical column is a contiguous run through the same
T::pack_contiguoushelper the row-major path uses. Filling a[depth][column]panel fromcolumn-major input is a transpose, so it processes
TRANSPOSED_PACK_GROUP = 4columns at a time tokeep 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
mainmoved 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.rsis reverted for thebasearm, so the harness is byte-identicalin both arms. Native-only
p50, 5 interleaved trials,null= the base binary under a second namein 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.
Full 27-cell table, the
nnpre-transposed negative control (row-majorBmust not move — seven ofnine 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 = 8goes 14.0x -> 7.0x within one invocation.M = 1is unaffected, and a null-too-tight artifact worth recordingAt
M = 1withtransB = 1theGemmpath takesgemv_f16_nkand never reachespack_b, so nocode this PR touches runs. At 5 trials x 15 runs, three
M = 1cells nevertheless read asregressions 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 = 1cells 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 = 1regression is claimed, and none is supported.Numerics
Bit-identical for
f16: sweeping all 65536 patterns on this host, scalarf16::to_f32and F16C_mm256_cvtph_psagree on every pattern, NaNs included. The one divergence isbf16, on the126 signalling NaNs — scalar
bf16::to_f32quiets them (0x7f81 -> 0x7fc1_0000) where the<< 16widening preserves the signalling bit. Both are NaN, no numeric contract is affected, andthe change makes transposed
Bagree with row-majorB, which always used the shift.transposed_b_is_bit_identical_to_pre_transposed_row_majorassertsto_bits()equality (not atolerance) across shapes that straddle
TRANSPOSED_PACK_GROUP,NR,NCandKC— trailingpartial 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_correctlypins the general branch so the fast pathcannot quietly become the only column-strided route. End-to-end
parity=PASSon all 738 A/B rows.Validation
20-step local matrix on the merged head: 19/20. The only failure is
H check-win-arm64, whichis pre-existing on clean
main—onnx-genai-ort-sysbindgen 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-artifactspasses: the defaultartifact still has 0 MLAS symbols and never defers to the ORT CPU EP. 1583 tests pass.