Repository navigation
perf(cpu): break the serial f32 reduction chain in the int4 decode GEMV (5.75x t=1) - #1667
Conversation
…MV (5.75x t=1)
`gemv_nk` reduced with `.map(|(&a, &b)| a * b).sum()`. f32 addition is not
associative, so LLVM must keep a single serial accumulator: the loop cannot
vectorize and issues one FMA per FMA *latency* rather than one per issue
slot. `dot_u8_f32`, twenty lines below it, already carries the comment
explaining exactly this -- "a plain `iter().map().sum()` keeps a single
serial `f32` reduction chain and stays scalar, which dominates 8-bit decode"
-- and was fixed with sixteen accumulators. The int4 route was not.
Sixteen accumulators, matching that precedent exactly. No new unsafe, no
intrinsics, portable to aarch64.
I went looking for a different bug. The hypothesis was memory traffic: this
path reads an expanded f32 [N, K] weight, 4 bytes per weight where the
packed form is 0.5, and at m=1 every weight is read once for a single FMA.
The attribution bench added here disproves it. The serial loop achieves
3.8-4.0 GB/s against a measured ~31-36 GB/s per-CCX ceiling -- it was never
close to bandwidth-bound. Breaking the chain alone, at identical traffic and
identical layout, is worth 5.3x-9.9x and lifts the same loop to 20-39 GB/s,
which is at the roofline. The bench keeps the two packed-nibble arms as a
recorded negative result: on-the-fly scalar dequant is 0.08x-0.15x, because
unpacking nibbles scalar-side costs far more than the traffic it saves.
End-to-end steady decode, min-of-6 interleaved, taskset to physical cores:
threads base fixed speedup
1 228.946 39.811 5.75x
4 57.883 23.292 2.49x
8 29.072 23.259 1.25x
16 23.517 23.163 1.02x
The fixed arm is flat from t=4 onward: it reaches a shared ceiling at
~23.2 ms/token and stops caring about threads. The baseline only reaches
that same floor at t=16, by spending 4x the cores to brute-force an
issue-bound loop. So the t=16 "wash" is not an absence of improvement, it is
the same throughput for a quarter of the machine.
That shows up as a large win as soon as the cores are contended. Aggregate
tokens/s, higher is better:
t=4, 2 sessions 18.2 -> 57.4 3.15x
t=4, 4 sessions 22.9 -> 68.3 2.98x
t=8, 2 sessions 33.0 -> 53.4 1.62x
t=8, 4 sessions 36.6 -> 67.9 1.86x
Scope, stated precisely because it is narrower than it first looks: this is
*not* the production default route. Default int4 decode (accuracy_level=0)
borrows the packed weight in place via `borrowed_affine_int4_matmul_nblock`
and never builds the f32 cache. `gemv_nk` is reached by accuracy_level=1,
by grouped quantization (g_idx), and by non-contiguous weights. Verified by
instrumenting `gemv_nk` directly: zero calls at accuracy_level 0 and 4,
calls at 1. My first A/B measured accuracy_level=0, found an exact wash, and
that wash is why the route was checked instead of the result being believed.
Numerics: reassociation changes rounding, not precision class. The compute
type stays f32, no operand is quantized. Two tests prove the direction --
against an f64 reference the reassociated order is at least as accurate as
the serial chain at k=1024/4096/14336, because sixteen accumulators form a
shallow pairwise tree while a serial chain adds every term at a
progressively worse exponent. A second test brackets the 16-lane tail at
k=0,1,2,15,16,17,31,32,33,47,63,64,65.
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 #1667 +/- ##
==========================================
- Coverage 81.41% 81.00% -0.42%
==========================================
Files 384 384
Lines 180713 180771 +58
Branches 180713 180771 +58
==========================================
- Hits 147134 146435 -699
- Misses 28637 29391 +754
- Partials 4942 4945 +3
Flags with carried forward coverage won't be shown. Click here to find out more.
🚀 New features to boost your workflow:
|
Local gate matrix: PASS=21 FAIL=0 SKIP=0On this branch merged to current Covered:
No-MLAS and fallback postureUnchanged by construction: this PR edits one reduction inside an existing native kernel and adds a bench. It links no new symbol, adds no aarch64The fix is portable by choice. Sixteen scalar accumulators, no intrinsics, so aarch64 gets the same autovectorized reduction rather than a second code path to keep in sync. Clippy-aarch64 and QEMU execution both green. MiriNot applicable: no |
Investigating the
|
#1674) `CUDA compile` fails on current `main`, on both the Linux and Windows lanes: ``` error: accessing first element with `matmul.inputs.get(0)` --> crates/onnx-runtime-ep-cuda/src/optimizer.rs:2035:12 | 2035 | if matmul.inputs.get(0).copied().flatten().is_none() | ^^^^^^^^^^^^^^^^^^^^ help: try: `matmul.inputs.first()` | = note: `-D clippy::get-first` implied by `-D warnings` error: could not compile `onnx-runtime-ep-cuda` (lib) due to 1 previous error ``` `clippy::get_first` is denied via `-D warnings`, so this is a hard compile failure, not a lint suggestion. The neighbouring `get(1)` / `get(2)` calls are correct as written — only index 0 has a dedicated accessor, which is exactly why this is easy to miss when writing the three lines as a block. Verified with `cargo clippy -p onnx-runtime-ep-cuda --features cuda -- -D warnings`: reproduces before the change, clean after. Also `cargo fmt --all -- --check` clean. ## Context This is the **eighth** time `main` has gone red in about a day. I hit it while merging latest `main` into #1667 for final validation — a PR that touches only the CPU int4 GEMV. The structural cause I flagged on #1653, #1655 and #1666 is unchanged: **required checks run on each PR's merge ref and never on the resulting `main`.** A PR can be green against its own merge base and still break `main`, and nothing re-checks `main` afterwards, so the breakage is discovered by whoever rebases next rather than by the PR that caused it. Two fixes would each close most of this class: 1. a **merge queue**, which tests the actual post-merge state; or 2. running `Rust quality` / `CUDA compile` **on `main` post-merge**, which at least surfaces breakage in minutes instead of at the next contributor's local build. I would rather stop paying this tax one PR at a time. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
…t=1/4/8 (#1852) ## The 1.84x acc0 gap is stale. Re-measured, it is ~1.12x. The published acc0 (`accuracy_level = 0`, the production default) int4 decode gap against ORT — **1.84x at t=1** — is what made acc0 the top remaining CPU MatMulNBits target in the ledger. It dates from `e9754e7ef` (#1628) and **eight merges have landed since**, three of them direct acc0 kernel work. On current main the gap measures **~1.12x**. | width | native tok/s | ORT tok/s | **gap** | gap range | cells (trusted/taken) | A/A range | |---:|---:|---:|---:|---:|---:|---:| | 1 | 27.9 | 31.2 | **1.120x** | 1.112–1.128 | 3/3, **2 retained** | 1.025–1.036 | | 4 | 107.2 | 122.3 | ~1.15x | 1.087–1.284 | 4/6 | **0.868–1.150** | | 8 | 211.0 | 238.0 | **1.120x** | 1.089–1.145 | 3/3 | 0.997–1.028 | | 16 | — | — | ~1.64x, **does not resolve** | 1.456–1.831 | 2/3 | **0.969–1.295** | Both arms are `tokens_s_total`, paired within each launch and then medianed. *Trusted/taken* is the harness's verdict; *retained* is editorial — all three `t=1` cells passed the guard and one was dropped afterwards by me, disclosed below. Only `t=1` and `t=8` resolve. `t=4` sits inside its own A/A null (0.868–1.150). **`t=16` reads ~1.64x and is the open row** — see below; a second revision of this PR corrects an earlier claim that it was wholly contaminated. **acc0 is no longer the top CPU target — conditional on `t=16`.** At the two widths that resolve, the remaining ~12% is a kernel efficiency difference that sits below several other open items. `t=16` is the width closest to an unconfined production process, and a confirmed 1.64x there would reverse that. ## Why the movement is kernel — measured, not inferred The first revision of this PR argued from a control: ORT re-measures at 31.99 ms, within 4.4% of its published 30.632, therefore the harness is comparable and *"the movement is entirely on our side."* **Review objected that this does not follow, and review was right.** The ORT arm reproducing shows the *ORT* ruler did not move. It says nothing about the native ruler, which sits in a different binary and changed repeatedly over the same window — `81e611c03` (#1722) is literally titled *"make the acc0 native and ORT arms measure one quantity"*. So the inference was replaced with a measurement. `e9754e7ef`'s tree is checked out in a second worktree, its `int4_decode_loop_ab` rebuilt, and run **beside** current main's on the same host, same environment, `PROBE_REPS=1` on both so neither gets a rep loop the other lacks, arms interleaved and the launch order alternated: | width | kernel-only, measured | published pair implies | verdict | |---:|---:|---:|---| | 1 | **1.64x** [1.61–1.88], 12 paired cells | 1.59x | apparent movement **is** kernel | | 8 | **1.82x** [1.78–1.89], 6 paired cells | 3.08x | **3.08x retracted** | ### Retracting the t=8 3.08x Both old figures reproduce today to within 0.4% — but only **unpinned**: | published | rebuilt `e9754e7ef` today, unpinned | delta | |---|---:|---:| | `56.307 ms` (t=1) | 56.519 (56.402 / 56.519 / 56.878) | +0.4% | | `14.091 ms` (t=8) | 14.115 (14.105 / 14.115 / 14.196) | +0.2% | The old bench never called `EpFactory::initialize`, so it never ran `bound_process_to_decode_budget()` and its process was never confined. That function — physical-core `select_budget_cpus` included — **already existed at `e9754e7ef`**, and production always called it; only the bench was missing the call, which #1766 `11cb8e5f3` added. The old `t=8` row therefore measured eight decode workers scattered over 32 logical CPUs onto SMT siblings: **a topology no served session ever ran in.** | binary at t=8 | 8 physical cores | 4 cores + SMT siblings | unpinned | |---|---:|---:|---:| | `e9754e7ef` | 8.430 ms | 16.121 ms | **14.115 ms** | | current main | 4.664 ms | 7.988 ms | **4.619 ms** | **1.67x of the claimed 3.08x was placement, not kernel work.** Today's binary is pin-insensitive (0.99x) because it confines itself. This is exactly the effect `docs/benchmarks/2026-08-21-decode-worker-cpu-placement.md` (#1680, ledger §24) already recorded — landing on a number I quoted two days later. ### Two corrections that look right and are not Recorded so they are not re-applied at this site: - **The ~11% warmup/spawn handicap of §27 is in `tokens_s_total`.** Both published figures are **`ms_token`** — the old `ort_matmulnbits_baseline.py` docstring names *"the native harness's `steady` column-2 median"* as its comparand, and the reproductions above land on it to 0.4%. Deducting 11% from `56.307` yields a number no run of either tree produces. - **The statistic asymmetry that *is* real points the other way.** Old ORT took `min` over reps of a per-`Run` median; old native was single-shot with no rep loop. Best-of-N against single-shot flatters ORT, so it made the old gap look *worse*. Calling the two arms "the same statistic", as the first revision did, was wrong. ## Headline-table defects fixed in this revision - **Mixed statistics.** The first table printed native *median latency* beside ORT *throughput-equivalent* and called the ratio a gap, so its columns did not yield its own gap figure. Both sides are now `tokens_s_total`; the mixed variant is shown, labelled, and noted to decline (1.113 → 1.098 → 1.085) rather than be flat. - **False precision at t=4.** `1.148x` quoted against an A/A null of 0.868–1.150. Now "~1.15x, does not resolve". - **Undisclosed post-hoc discard.** "three independent launches per width" was false (3 / 6 / 3 / 3 cells across two invocations), and the `t=1` 1.4% spread depended on discarding a cell after seeing it. Retained-cell figures are now published beside the headline (**1.112x [0.927–1.128], 18.1%, n=3** — the median barely moves, the *precision* does not survive), and the discard rule is stated prospectively for next time. - **Reproducibility.** `acc0_gap_matrix.py` gains `--launches`, a per-width `--tokens 1:64,4:192,8:384` map, and a `gap` column in ORT÷native orientation beside `ratio`, so the Reproduce block names a command that produces the published table. ## Method preconditions added 1. **The two arms were not getting the same machine.** `ONNX_GENAI_CPU_DECODE_THREADS=w` confines the *whole native process* to `w` CPUs; the script pinned ORT to all 16 even CPUs at every width. Measured effect on a quiet host: **1–2%** — real, small, and now data rather than argument. 2. **The realized width is read back and checked** (`decode_width requested=4 realized=4 as_requested`). Timings cannot detect a vacuous sweep. 3. **`LoadWatch` samples the runnable count *during* every arm**, refusing above `width + slack`. A pre-check cannot see a competitor that arrives mid-cell — one did, and four cells were discarded because of it. 4. **The wide-pin arm turned out to be a contention detector.** One `t=1` cell passed every host-level guard while CPU 0 alone was busy: both matched-pin arms ~2x slow, the roaming arm normal. A single-CPU pin is the most fragile cell in any width sweep, and it is what every speedup is quoted against. ## Still unresolved — and a correction to the first revision **`t=16`, and it is the row that matters.** The first revision of this PR wrote the width off as "every cell contaminated". **That was wrong, and wrong in the direction that flattered the conclusion.** Two of the three `t=16` cells passed the load guard cleanly (runnable 6, no competitor recorded): | `t=16` | gap | A/A | native spread | ORT spread | |---|---:|---:|---:|---:| | launch 1 | 1.831 | 1.295 | 17.8% | **55.4%** | | launch 2 | 1.456 | 0.969 | 27.7% | **19.6%** | | **median** | **1.643** | — | — | — | So it is **~1.64x from two accepted cells**, not "no data". It still does not resolve, but on the correct ground: the A/A null spans 0.969–1.295 (±30%, against 3.6% at `t=1` and 2.8% at `t=8`) and both arms are unstable at this width. The cell the guard *did* refuse reads 1.585 — between the two retained — so the discard is not load-bearing either way. This is the one open cell that could reverse the re-ranking, and it needs a dedicated quiet-host study with launch distributions and a pre-registered A/A acceptance threshold. ## Second-review fixes (this revision) An adversarial review returned MERGE AFTER FIXES with all six core claims surviving falsification and seven defects. All are fixed: 1. **`t=16` mischaracterised** — the headline fix above, propagated to all four sites that carried the re-ranking claim. 2. **Cell counts** — `t=4` is 4 trusted of **6** taken; the published table came from **two** script invocations, not three. 3. **Column semantics** — harness-trust and editorial retention were conflated in one column; now separated, which makes the `t=1` discard visible in the table rather than only in prose. 4. **`t=1` placement probe samples published**, including a `114.94 ms` outlier on one pinned rep of three (the old binary's bimodality — it is why the `t=1` A/B range reaches 1.88x). Placement is worth 1.9% at this width, against 1.67x at `t=8`, as expected for a single-threaded process with no SMT sibling to hit. 5. **ORT-side spread disclosed at `t=16`** (55.4%) — the denominator is no better behaved than the numerator there. 6. **Stale checksum constants** in `int4_decode_loop_ab`'s module doc corrected, with a note that they drift under reduction reassociation (#1667, #1783) and that the *pattern* — block 16 moves under `ONNX_GENAI_CPU_MM_INT4_GEBP=0`, block 32 does not — is the route evidence, not the digits. 7. **`--tokens` map footgun** — a map missing a `--threads` width died with a bare `KeyError` after the first cell had already waited out the load guard. It now refuses at parse time, naming the missing widths. ## Validation Docs plus one benchmark harness; no library code. The harness was stub-validated end to end (launch loop, per-width token map, `gap` column, paired per-width summary) and both binaries were run for the numbers above. Required CI (`Fast (Linux x86_64)`, `Rust quality`) must be green before merge — no admin bypass. --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: Roy <roy@squad.local>
The defect
gemv_nkreduced with.map(|(&a, &b)| a * b).sum().f32addition is not associative, so LLVM must keep a single serial accumulator: the loop cannot vectorize and issues one FMA per FMA latency rather than one per issue slot.dot_u8_f32, twenty lines below it in the same file, already carries the comment explaining precisely this — "a plainiter().map().sum()keeps a single serialf32reduction chain and stays scalar, which dominates 8-bit decode" — and was fixed with sixteen accumulators. The int4 route was not. This change applies the same fix, matching that precedent exactly: no newunsafe, no intrinsics, portable to aarch64.I was looking for a different bug
My hypothesis was memory traffic: this path reads an expanded f32
[N, K]weight — 4 bytes per weight where the packed form is 0.5 — and at m=1 every weight is read once for a single FMA, so time should track bytes.The attribution bench added here disproves that, which is why it is included rather than deleted:
The serial loop achieves 3.8–4.0 GB/s against a measured ~31–36 GB/s per-CCX ceiling. It was never anywhere near bandwidth-bound. Breaking the chain alone — identical traffic, identical layout — is worth 5.3x–9.9x and lifts the same loop to 20–39 GB/s, which is at the roofline.
The bench keeps its two packed-nibble arms as a recorded negative result: on-the-fly scalar dequant measures 0.08x–0.15x, i.e. 6–12x slower, because unpacking nibbles scalar-side costs far more than the traffic it saves. Those arms are deliberately naive references — they bound the idea in that form, they do not prove a SIMD unpack could not win.
End-to-end
Steady decode, min-of-6 interleaved,
tasksetto physical cores, two binaries built from the same tree:The t=16 wash is the interesting cell, and it is not an absence of improvement. The fixed arm is flat from t=4 onward — it reaches a shared ceiling at ~23.2 ms/token and stops caring about thread count. The baseline only reaches that same floor at t=16, by spending 4x the cores to brute-force an issue-bound loop. Same throughput, a quarter of the machine.
Which means the win reappears as soon as cores are contended. Aggregate tokens/s, higher is better:
Scope — narrower than it first looks
This is not the production default route, and I want that on the record rather than discovered later. Default int4 decode (
accuracy_level=0) borrows the packed weight in place viaborrowed_affine_int4_matmul_nblockand never builds the f32 cache.gemv_nkis reached byaccuracy_level=1(a legal ONNX value, fp32 compute by this kernel's own 0/1 convention), by grouped quantization (g_idx), and by non-contiguous weights.Verified by instrumenting
gemv_nkdirectly and counting calls: zero ataccuracy_level0 and 4, non-zero at 1.I found this the honest way. My first A/B measured
accuracy_level=0, showed an exact wash (56.341 → 56.266, 0.13%), and rather than explain the wash away I instrumented the route and found my patch was on a path that workload never takes.Production default
accuracy_level=0is unaffected by this PR and remains 1.84x behind ORT. That is a separate, still-open problem.Numerics
Reassociation changes rounding, not precision class: the compute type stays
f32throughout and no operand is quantized. This is the same tradedot_u8_f32already makes on theaccuracy_level=08-bit decode route, so it is not a new precedent.Two tests prove the direction rather than asserting it:
reassociated_dot_is_at_least_as_accurate_as_the_serial_chain— against an f64 reference at k=1024/4096/14336, the reassociated order is at least as accurate. It should be: sixteen accumulators form a shallow pairwise tree, while a serial chain adds every term at a progressively worse exponent and errs linearly in k. The test also asserts the reference is non-degenerate so it cannot pass vacuously.reassociated_dot_handles_every_tail_length— brackets the 16-lane body from every side at k=0,1,2,15,16,17,31,32,33,47,63,64,65, since losing or double-counting thek % 16remainder is the obvious failure mode and a large-k accuracy test would hide it inside its tolerance.Validation
Rebased on current
main(016f0612c, i.e. after theverify_execfix) before validating.cargo fmt --all -- --check— cleancargo test -p onnx-runtime-ep-cpu --lib— 1576 passed (1574 + the 2 new)