Repository navigation
Keep the f16 decode GEMV accumulator in registers - #1436
Conversation
half_gemv's module doc states the design premise -- at M=1 "each weight element is touched exactly once, so the kernel is purely memory-bound". It was not. Against this host's measured 75.8 GB/s sustained read bandwidth (roofline_bandwidth, saturating by 4-8 threads) the kernel ran at 12-47 GB/s, under half the machine on every large cell, and it anti-scaled: l3_3584 went 0.568 ms at t=4 to 0.830 ms at t=32. There was no f16 GEMV benchmark cell, which is why a kernel whose entire premise is a bandwidth claim had nothing checking it. gen_f16_gemv.py adds five, sweeping the weight working set from 0.5 MB (inside one core's L2) to 134 MB (past any LLC here). That sweep is the diagnosis: the L2-resident cell moved the same GB/s as the DRAM-resident one. A memory-bound kernel would be far faster per byte on the small one, so the limit was per-core and the extra threads were only contending. The p loop was outermost, so each output's accumulator was live across the whole contraction but lived in acc. Three memory operations per 8-lane FMA -- load the weight, load the accumulator, store it back -- plus a store-to-load forward from the previous p, against one FMA. STRIPE = 512 keeps those accumulators in L1, which the old comment offered as reassurance; L1 is only cheap next to L3, not next to a register, and 16 ymm were sitting idle. Tiling the output into TILE = 64 columns and hoisting p inside leaves eight accumulators in registers for the whole contraction, stored once: one memory operation per FMA, the minimum this problem admits. No extra traffic -- the same k * STRIPE elements at the same stride, and a tile is exactly two cache lines with STRIPE a whole number of tiles. 14 of 15 cells win by 1.22x-2.60x above their null control, and the large cells go from ~46% of the memory roofline to 79-86%. l3_2048 at t=32 would not settle across three runs and is not claimed. TILE = 64 is measured, not asserted: 32/64/96/128 across all 15 cells. 64 wins 11 outright and ties a twelfth. 128 needs 18 of 16 architectural ymm and spills -- on the smallest cell it is slower than changing nothing. Bit-identical, which the int4 row-blocking change could not claim: tiling changes which register holds a partial sum, never the order it is built in, so every pre-existing oracle test passes unmodified. Two tests added for the new structure -- a contiguous width sweep across the tile boundary at four stripe offsets, and the divisibility facts the no-extra-traffic argument rests on. ab.py grows --native-only. Per sebastian-paired-harness-coresidency ORT's intra-op pool spin-waits; on these cells a paired run depressed the native median by up to 6x and drove the null control to 27%, larger than most of the effects here. The first paired attempt at this measurement produced a table with three sign errors in it. The flag belongs in the shared driver, not a private script. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Four nits from adversarial review. The module doc still asserted the exact non-sequitur this commit's parent disproves -- that touching each weight once makes the kernel "purely memory-bound". Touch-once sets the ceiling; it says nothing about the distance to it, which was a factor of two to six. Reworded to say that, with the measured numbers and a pointer to the benchmark record. ab.py's --native-only mode reuses the `ratio` column for native milliseconds, so a machine consumer could read ms as a dimensionless ratio and only notice from `ort` being NaN. The CSV now carries native_only, and the README documents both the flag and the column. The default paired path is unchanged: same regex, same printed line, same columns, verified by running both modes. TILE carried a redundant x86 cfg; the module is already x86-only and its sibling constants are ungated. stripe_widths_around_the_tile_boundary_are_exact skipped any width exceeding n rather than failing, so a future edit to n could have silently shrunk the sweep to nothing while still passing. It now asserts the bound holds and asserts its own combination count. Prose said the sweep ran to 2*TILE+8; it runs to 2*TILE+9, which is two whole tiles plus an 8-lane block plus a scalar -- all three paths in one call. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
|
Opus adversarial review: APPROVE WITH NITS — all four fixed in The review independently confirmed the load-bearing claims: the three loops partition 1. The module doc still asserted the thing this PR disproves. It read: each weight is touched exactly once, "so the kernel is purely memory-bound". That is the exact non-sequitur the commit message spends paragraphs refuting — touch-once sets the ceiling, it says nothing about the distance to it, and the distance was a factor of two to six. Fixing the kernel while leaving the false premise in place would have been the worse half of the job. Reworded with the measured numbers and a pointer to the benchmark record. 2. 3. 4. The width sweep could have silently shrunk to nothing. One review observation worth recording for whoever reviews this next: this worktree's local Gates 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❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #1436 +/- ##
==========================================
+ Coverage 80.38% 80.51% +0.13%
==========================================
Files 378 379 +1
Lines 167923 169694 +1771
Branches 167923 169694 +1771
==========================================
+ Hits 134980 136629 +1649
- Misses 28092 28212 +120
- Partials 4851 4853 +2
Flags with carried forward coverage won't be shown. Click here to find out more.
🚀 New features to boost your workflow:
|
Port the register tiling into main's per-format stripe_simd_fn! macro so it serves bf16 as well as f16; renumber the ledger section 5 -> 7.
The tiling now lives in `stripe_simd_fn!`, so it is instantiated once per widening kernel; each instance needs its own bit-identity check against its own scalar reference.
…es that reach it #1381 landed a decode handover while this was open: an M=1 half MatMul of 1,048,576 elements or more now takes the fused widen-pack GEBP, so four of the five MatMul cells no longer reach this kernel in a default build. Measure the MatMul cells with the GEBP switched off to isolate the kernel, and add Gemm transB=0 cells -- which have no weight gate -- to measure what a default build runs. 15/15 and 13/15 above their null controls. Also record that #1381's handover was placed against the untiled GEMV and partly inverts against the tiled one, and that retuning it needs its own sweep.
15/15 and 13/15 leaves two unclaimed cells, not four, and their nulls are 22.7% and 48.7%. Caught in review.
Local validation — latest
|
| target | step | result |
|---|---|---|
x86_64-unknown-linux-gnu |
A-F, I-L, tests | PASS |
aarch64-unknown-linux-gnu |
G clippy -D warnings on onnx-runtime-ep-cpu |
PASS |
aarch64-pc-windows-msvc |
H | blocked on Windows SDK (see above) |
| big-endian | F clippy-native-be |
PASS |
kernels::half_gemv is #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] at its mod declaration, so on either aarch64 target this diff compiles to nothing at all. Step G is therefore the meaningful ARM check, and it is green.
Also green: cargo fmt --all --check, clippy -D warnings (offline / native / big-endian), --no-default-features, --all-features, the no-MLAS artifact guard (K), cross-compile (L), and all 8 guard scripts including verify_documented_env_vars and check_dispatch_reachability. 1545 tests pass in onnx-runtime-ep-cpu --lib.
Miri — reported as a no-op gate, not as a pass
is_x86_feature_detected!("avx2") returns false under Miri, so every #[target_feature] path takes its scalar fallback. stripe_widths_around_the_tile_boundary_are_exact "passes" under Miri in 7.07s — because it continues past both formats and checks nothing. The AVX2 intrinsics added here get zero Miri coverage. Their evidence is instead:
- the native bit-identity sweep (1096 width x offset x format combinations, asserting its own combination count so it cannot silently shrink),
- the four pre-existing oracle tests, passing unmodified,
- the unchanged real
assert!(b.len() >= k*n)— a hard check rather than adebug_assert, precisely because the kernel does unchecked pointer arithmetic fromk,nandj0.
Two defects found and fixed on main while validating this
Neither was mine, and both were red on a clean origin/main:
- fix(ci): restore
cargo fmton main #1553 —cargo fmt --all --check, one of the two required checks, broken by perf(cuda): let the split-K decode GEMV store bf16 directly #1550. Mergedb3d46de70. - (earlier this session: fix(ci): account for the deleted ONNX_GENAI_GEMV_FOLDSCALE knob #1536
verify_documented_env_vars, fix(ci): restore cargo fmt on main #1546cargo fmt.)
Three in three days, same mechanism each time: two required checks plus strict_required_status_checks_policy=false lets a PR merge on a green run that predates the commit which breaks the gate.
Review
Opus review: one finding, non-blocking — the ledger said "four unclaimed cells … nulls of 16-49%" where 15/15 + 13/15 leaves exactly two, with nulls 22.7% and 48.7%. Fixed in 1b812456f. The review independently re-derived the exactly-once-write partition of [0, w), the worst-case pointer bound (k-1)·n + j0 + w - 1, the k == 0 safety of dropping acc.fill(0.0), and the bit-identity claim across all three loops.
half_gemv's module documentation states the design premise plainly: atM == 1"each weight element is touched exactly once, so the kernel is purely memory-bound". It was not.The gap
Against this host's measured sustained read bandwidth —
roofline_bandwidth --threads 1,2,4,8,16,32 --mib 1024, which reports 75.8 GB/s and saturates by 4-8 threads — the kernel ran at 12–47 GB/s. Under half the machine on every large cell. It also anti-scaled:l3_3584went 0.568 ms at t=4 to 0.830 ms at t=32.Why nobody noticed
There was no f16 GEMV benchmark cell.
gen_gemm.pycovers block-quantised and f32 dense GEMM, so the one kernel whose entire premise is a bandwidth claim had nothing measuring it.scripts/ort_ab/gen_f16_gemv.pyadds five, sweeping the weight working set — the only variable that matters for a kernel that reads each weight once:l2_512l2_1024l3_2048l3_3584dram_8192That sweep is the diagnosis. The L2-resident cell moved the same GB/s as the DRAM-resident one (19.4 vs 34.9 at t=32). A memory-bound kernel would be far faster per byte on the small cell. So the limit was per-core, and the extra threads were only contending for it.
Root cause
The
p(contraction) loop was outermost, so each output's accumulator was live across the whole contraction but lived inacc:Three memory operations per 8-lane FMA — load the weight, load the accumulator, store it back — plus a store-to-load forwarding round trip from the same address one
pearlier, all to feed one arithmetic op. The load/store ports set the rate, not the FMA units.STRIPE = 512puts the accumulators in L1, which the existing comment offers as reassurance. That was the problem rather than the comfort it reads as: L1 is only cheap next to L3. It is not cheap next to a register, and 16ymmwere sitting idle.The fix
Tile the output into
TILE = 64columns and hoistpinside, so eight accumulators stay inymmfor the entire contraction and reach memory exactly once — one memory operation per FMA, the minimum the problem admits.Tiling costs no extra traffic: the same
k * STRIPEelements, the samen-element stride between consecutivep, andTILEis exactly two 64-byte cache lines withSTRIPEa whole number of tiles, so no fetched line is ever partially consumed.a_tile_never_straddles_a_stripepins all three divisibility facts rather than leaving them to the comment.Which paths actually reach this kernel
This PR was opened before #1381, which landed a decode handover: an
M == 1halfMatMulof1,048,576 elements or more now goes to the fused widen-pack GEBP instead. So four of the five
MatMulcells above no longer reach this kernel in a default build, and quoting the originalnumbers as-is would be quoting a measurement of a path that no longer runs.
What still reaches it:
MatMulf16/bf16 decodeGemmf16 decode,transB=0ONNX_GENAI_CPU_MM_HALF_GEBP=0x86So everything below is re-measured on latest
main, in two sets:MatMulwith the GEBP switchedoff, which isolates the kernel over the whole weight range, and
Gemmwith no environment set atall, which is what a default build runs.
gen_f16_gemv.py --op gemmemits the second set.The merge also moved the change: main had since generalised the stripe kernel to bf16 via the
stripe_simd_fn!macro, so the tiling went into the macro rather than into a f16-only function.It now serves both formats, which the original revision did not.
Result
ab.py --native-only --null-control, 7 trials x 30 runs, medians.nullis the baseline binaryunder a second name; its delta is the host's noise floor for that cell.
MatMulcells, GEBP switched off — the kernel over its whole rangel2_512l2_512l2_512l2_1024l2_1024l2_1024l3_2048l3_2048l3_2048l3_3584l3_3584l3_3584dram_8192dram_8192dram_819215 of 15 above their null control, 1.13x-2.61x.
Gemmcells, shipped default — no environment setl2_512l2_512l2_512l2_1024l2_1024l2_1024l3_2048l3_2048l3_2048l3_3584l3_3584l3_3584dram_8192dram_8192dram_819213 of 15 above their null control, 1.25x-2.21x. Both unclaimed cells have unusable controls
(nulls of 22.7% and 48.7%), not small effects.
On the bandwidth numbers
Three cells now read above the host's 75.8 GB/s DRAM ceiling (up to 88.6). That is not an error:
at 8.4 and 25.7 MB those weights are L3-resident and never reach DRAM, so the DRAM roofline is the
wrong ceiling for them. Before this change they ran at 20-43 GB/s, far enough below it that the
distinction never surfaced — that it surfaces now is itself the result.
dram_8192is the only cell whose weight genuinely comes from DRAM, and it goes 33.0 → 63.1 GB/s,44% → 83% of roofline.
A consequence this PR deliberately does not act on
#1381 placed its handover using a sweep of the untiled GEMV. Re-running that same harness
(
bench half_decode_gemv_ab, 5 interleaved reps,ratio = GEMV/GEBP, below 1.00 favours the GEMV)against the tiled kernel:
K x NThe two largest shapes —
mlpandlm_head, the ones that dominate decode — invert. 2048x2048narrows from 1.78 to 1.26 without inverting, and there the two harnesses disagree by thread count
(
ab.pyreads -83.5% at t=4 and -61.8% at t=16 in the GEMV's favour, but only -2.13% within noiseat t=32, which is where
cargo benchruns).So the threshold is wrong at both ends and right in the middle, and a single weight cutoff cannot
express a thread-dependent crossover. Retuning it needs its own
k x n xthread sweep and its owncontrol; stacking a routing change on evidence that conflicts between two harnesses is the kind of
unproven change the ledger exists to refuse. Filed as a follow-up.
This PR changes no routing. Every shape takes the route it took before, only faster.
Numerics
Bit-identical. Tiling changes only which register holds a partial sum, never the order it is
built in: within any one output element the contraction still runs
p = 0 .. k-1with the same FMAat each step. Every pre-existing oracle test passes unmodified; nothing was weakened to a
tolerance. 1545 tests green.
stripe_widths_around_the_tile_boundary_are_exactsweeps every width from 1 to2 * TILE + 9atfour stripe offsets against the scalar stripe, bit for bit, and now does so for both formats,
asserting its own combination count so it cannot silently shrink.
a_tile_never_straddles_a_stripepins
STRIPE % TILE == 0,TILE % 8 == 0,TILE * 2 % 64 == 0.The macro no longer needs
acc.fill(0.0): the three loops write[0, w)exactly once between them,and
k == 0is already short-circuited by the caller.Validation
20-step local matrix: 19 PASS, 1 FAIL. The failure is step H
(
cargo check --target aarch64-pc-windows-msvc), which fails on unmodifiedmaintoo — bindgencannot find the Windows SDK headers in this container. Everything else green: fmt, clippy
-D warnings(offline/native/BE), Linux tests, cross-compile, feature combinations, no-MLASartifact guard, and all 8 guard scripts.
Miri is a no-op gate for this change and is reported as such, not as a pass.
is_x86_feature_detected!("avx2")returns false under Miri, so every#[target_feature]path takesits scalar fallback:
stripe_widths_around_the_tile_boundary_are_exactcompletes in 7s under Miribecause it skips both formats. The AVX2 intrinsics added here get zero Miri coverage. Their
safety evidence is the native bit-identity sweep plus the unchanged
assert!onb.len().