diff --git a/crates/onnx-runtime-ep-cpu/src/kernels/half_gemm.rs b/crates/onnx-runtime-ep-cpu/src/kernels/half_gemm.rs index a76d140bd9..d5572a408d 100644 --- a/crates/onnx-runtime-ep-cpu/src/kernels/half_gemm.rs +++ b/crates/onnx-runtime-ep-cpu/src/kernels/half_gemm.rs @@ -329,9 +329,11 @@ fn gemm_impl( c.fill(0.0); let threads = rayon::current_num_threads(); - let mc = if threads <= 1 { - MAX_MC.min(m) - } else { + // Splitting is not free: the fork costs, and the workers keep spinning + // after the region ends, which competes with whatever runs next. Below + // `PARALLEL_MIN_WORK` that costs more than the split saves, so stay serial + // -- the same guard `half_gemv` and `accelerate_gemm` already apply. + let split_mc = { let rows = m.div_ceil(threads.saturating_mul(2)).clamp(1, MAX_MC); if rows == 1 { 1 @@ -339,16 +341,143 @@ fn gemm_impl( rows.div_ceil(MR).saturating_mul(MR).min(MAX_MC) } }; + let parallel = forced_route().unwrap_or_else(|| { + threads > 1 + && m.saturating_mul(k).saturating_mul(n) >= PARALLEL_MIN_WORK + // A split into a single block cannot use more than one thread, so + // it is pure fork overhead. `m == 1` always lands here. + && m.div_ceil(split_mc) > 1 + }); + let mc = if parallel { split_mc } else { MAX_MC.min(m) }; + + let block = |block_index: usize, c_block: &mut [f32]| { + let first_row = block_index * mc; + let rows = c_block.len() / n; + gemm_block::( + a, a_layout, b, b_layout, c_block, first_row, rows, k, n, path, + ); + }; + + if !parallel { + count_serial_gemm(); + for (block_index, c_block) in c.chunks_mut(mc * n).enumerate() { + block(block_index, c_block); + } + return; + } c.par_chunks_mut(mc * n) .enumerate() - .for_each(|(block_index, c_block)| { - let first_row = block_index * mc; - let rows = c_block.len() / n; - gemm_block::( - a, a_layout, b, b_layout, c_block, first_row, rows, k, n, path, - ); - }); + .for_each(|(block_index, c_block)| block(block_index, c_block)); +} + +/// Minimum multiply-accumulate count (`m*k*n`) before splitting a half GEMM +/// across the pool. +/// +/// `gemm_impl` used to fork unconditionally, so it spent 0.14 ms of pool +/// overhead to multiply an 8x64 by a 64x64 -- work a single core finishes in +/// 0.046 ms. That is the common case now: since the f16 prefill gate landed, +/// this kernel only ever sees the *small* shapes that decline widening, so the +/// unguarded fork was mis-sized for every shape it still serves. +/// +/// Measured `serial/parallel` (>1 = splitting wins) at `m = 8`, interleaved +/// rep-by-rep, `p50` of 9, pinned to 16 physical cores, two independent runs +/// per thread count (`bench_half_gemm_parallel_threshold`): +/// +/// | m*k*n | T=2 | T=4 | T=8 | T=16 | T=32 | +/// |-----------|-----------|-----------|-----------|-----------|-----------| +/// | 32_768 | 0.88/0.86 | 0.51/0.55 | 0.63/0.75 | 0.36/0.31 | 0.21/0.19 | +/// | 65_536 | 0.79/0.78 | 0.63/0.63 | 0.85/0.88 | 0.53/0.53 | 0.41/0.33 | +/// | 131_072 | 1.02/1.02 | 0.79/0.80 | 1.14/1.18 | 0.76/0.60 | 0.63/0.54 | +/// | 262_144 | 1.20/1.20 | 0.92/0.92 | 1.16/1.51 | 1.03/1.02 | 0.86/0.80 | +/// | 393_216 | 1.27/1.28 | 0.91/0.97 | 1.27/1.64 | 1.04/1.14 | 0.99/1.04 | +/// | 524_288 | 1.34/1.30 | 1.01/1.02 | 1.29/1.76 | 1.20/1.08 | 1.19/1.12 | +/// | 1_048_576 | 1.41/1.37 | 1.08/1.04 | 1.47/1.95 | 1.34/1.30 | 1.31/1.31 | +/// +/// Re-measured after `codegen-units = 1` was pinned for this crate, which makes +/// the *serial* route materially faster and therefore moves the crossover up. +/// Same harness, two runs per thread count: +/// +/// | m*k*n | T=2 | T=4 | T=8 | T=16 | +/// |-----------|-----------|-----------|-----------|-----------| +/// | 262_144 | 1.20 | 0.92/0.91 | 1.52 | 0.64/0.95 | +/// | 393_216 | 1.28 | 0.99/0.96 | 1.66 | 0.80/1.10 | +/// | 524_288 | 1.32 | 0.99/0.98 | 1.81 | 1.03/1.19 | +/// | 786_432 | 1.33 | 1.03/0.92 | 1.35 | 0.97/1.25 | +/// | 1_048_576 | 1.37 | 1.05/1.05 | 1.46 | 1.31/1.34 | +/// +/// `1_048_576` is now the smallest size that wins at every measured thread count +/// in every run (worst case 1.05x). `524_288` no longer qualifies: it is a wash +/// at T=4 (0.99/0.98) and `786_432` dips to 0.92 there and 0.97 at T=16. The +/// rule is unchanged -- smallest size that wins everywhere -- only the build it +/// is measured against is faster, so the answer moved. +/// +/// Below the threshold the picture is not merely mixed but bad: 0.32-0.37 at +/// T=16 for the smallest shapes, which is the 3x the guard exists to stop. +/// +/// Deliberately a single constant rather than a function of thread count: the +/// measured crossover is *not* monotone in threads (T=4 wants the highest +/// threshold, T=8 the lowest), because for small `m` the block count is pinned +/// by `m` rather than by the pool, so a thread-scaled formula would be fitting +/// noise. +/// +/// `m*k*n` is a proxy, not a cost model. The split is over *rows of C*, so the +/// achievable speedup is bounded by the block count, and the parallel route +/// re-packs `B` per row-block -- for small `m` (`mc == 1`) that is `m` +/// redundant packs. That is why the ratios plateau near 1.1x at T=4 rather than +/// scaling with the pool, and why some shapes above the threshold still gain +/// little. Sibling kernels apply the same kind of guard +/// (`half_gemv::PARALLEL_MIN_WORK`, `accelerate_gemm`'s `k*n` bound). +const PARALLEL_MIN_WORK: usize = 1_048_576; + +#[cfg(test)] +thread_local! { + /// Test-only override letting a benchmark time both routes in one process. + static FORCED_ROUTE: std::cell::Cell> = const { std::cell::Cell::new(None) }; +} + +#[inline] +fn forced_route() -> Option { + #[cfg(test)] + { + FORCED_ROUTE.with(std::cell::Cell::get) + } + #[cfg(not(test))] + { + None + } +} + +#[cfg(test)] +fn with_forced_route(parallel: bool, op: impl FnOnce() -> R) -> R { + FORCED_ROUTE.with(|c| c.set(Some(parallel))); + let out = op(); + FORCED_ROUTE.with(|c| c.set(None)); + out +} + +#[cfg(test)] +thread_local! { + /// Counts half GEMMs that declined to split, so a test can prove the guard + /// fired rather than inferring it from a timing. + static SERIAL_GEMMS: std::cell::Cell = const { std::cell::Cell::new(0) }; +} + +#[inline] +fn count_serial_gemm() { + #[cfg(test)] + SERIAL_GEMMS.with(|c| c.set(c.get() + 1)); +} + +/// Number of unsplit half GEMMs on this thread since the last reset. +#[cfg(test)] +fn serial_gemms() -> usize { + SERIAL_GEMMS.with(std::cell::Cell::get) +} + +#[cfg(test)] +fn reset_serial_gemms() { + SERIAL_GEMMS.with(|c| c.set(0)); } #[allow(clippy::too_many_arguments)] @@ -1087,3 +1216,290 @@ mod tests { } } } + +#[cfg(test)] +mod half_gemm_par_bench { + use super::*; + use std::time::Instant; + + fn operands(m: usize, k: usize, n: usize) -> (Vec, Vec, Vec) { + let a = (0..m * k) + .map(|i| half::f16::from_f32(((i % 13) as f32 - 6.0) * 0.1).to_bits()) + .collect(); + let b = (0..k * n) + .map(|i| half::f16::from_f32(((i % 7) as f32 - 3.0) * 0.1).to_bits()) + .collect(); + (a, b, vec![0.0f32; m * n]) + } + + fn run(a: &[u16], b: &[u16], c: &mut [f32], m: usize, k: usize, n: usize, parallel: bool) { + with_forced_route(parallel, || { + gemm_impl::( + a, + MatrixLayout::row_major(k), + b, + MatrixLayout::row_major(n), + c, + m, + k, + n, + ExecutionPath::Scalar, + ) + }); + } + + /// Sites [`PARALLEL_MIN_WORK`]: serial vs pool-split half GEMM, interleaved + /// rep-by-rep so both routes see the same machine load. Pin the run + /// (`taskset -c 0-15`) and set `RAYON_NUM_THREADS`. + #[test] + #[ignore = "microbench: run explicitly with --ignored --nocapture"] + fn bench_half_gemm_parallel_threshold() { + println!("threads={}", rayon::current_num_threads()); + println!("m,k,n,work,serial_ms,par_ms,serial/par"); + for &(m, k, n) in &[ + (8usize, 64usize, 64usize), // 32_768 + (8, 128, 64), // 65_536 + (8, 128, 128), // 131_072 + (8, 192, 128), // 196_608 + (8, 256, 128), // 262_144 + (8, 256, 192), // 393_216 + (8, 256, 256), // 524_288 + (8, 384, 256), // 786_432 + (8, 512, 256), // 1_048_576 + (8, 512, 384), // 1_572_864 + (8, 512, 512), // 2_097_152 + (8, 768, 512), // 3_145_728 + ] { + let (a, b, mut c) = operands(m, k, n); + run(&a, &b, &mut c, m, k, n, false); + let expected = c.clone(); + run(&a, &b, &mut c, m, k, n, true); + assert_eq!(c, expected, "{m}x{k}x{n}: split changed the result"); + + let (mut sv, mut pv) = (Vec::new(), Vec::new()); + for _ in 0..9 { + let t = Instant::now(); + run(&a, &b, &mut c, m, k, n, false); + sv.push(t.elapsed().as_secs_f64() * 1e3); + let t = Instant::now(); + run(&a, &b, &mut c, m, k, n, true); + pv.push(t.elapsed().as_secs_f64() * 1e3); + } + sv.sort_by(f64::total_cmp); + pv.sort_by(f64::total_cmp); + println!( + "{m},{k},{n},{},{:.4},{:.4},{:.2}", + m * k * n, + sv[4], + pv[4], + sv[4] / pv[4] + ); + } + } +} + +#[cfg(test)] +mod half_gemm_guard_tests { + use super::*; + + fn operands(m: usize, k: usize, n: usize) -> (Vec, Vec, Vec) { + let a = (0..m * k) + .map(|i| half::f16::from_f32(((i % 23) as f32 - 11.0) * 0.07).to_bits()) + .collect(); + let b = (0..k * n) + .map(|i| half::f16::from_f32(((i % 17) as f32 - 8.0) * 0.05).to_bits()) + .collect(); + (a, b, vec![0.0f32; m * n]) + } + + fn gemm_as(a: &[u16], b: &[u16], c: &mut [f32], m: usize, k: usize, n: usize) { + gemm_impl::( + a, + MatrixLayout::row_major(k), + b, + MatrixLayout::row_major(n), + c, + m, + k, + n, + ExecutionPath::Scalar, + ); + } + + fn gemm(a: &[u16], b: &[u16], c: &mut [f32], m: usize, k: usize, n: usize) { + gemm_impl::( + a, + MatrixLayout::row_major(k), + b, + MatrixLayout::row_major(n), + c, + m, + k, + n, + ExecutionPath::Scalar, + ); + } + + /// The threshold is a measured crossover, not a guess -- an earlier guess + /// sat 8x too low, inside the range the sweep shows losing 0.3x-0.9x. Pin + /// the value so a silent edit has to come back through the sweep, which + /// spans five thread counts precisely because the crossover moves with the + /// pool. + #[test] + fn half_gemm_parallel_threshold_matches_the_measured_crossover() { + assert_eq!( + PARALLEL_MIN_WORK, 1_048_576, + "smallest m*k*n measured to win at every thread count; re-run \ + bench_half_gemm_parallel_threshold before changing it" + ); + } + + /// Falsifies the guard itself: work below the crossover must stay on one + /// thread, and work above it must still split. Without the first half, an + /// 8x64x64 GEMM forks and takes 3.1x longer than doing it serially. + #[test] + fn half_gemm_declines_to_split_work_below_the_crossover() { + if rayon::current_num_threads() < 2 { + eprintln!("skipping: needs a multi-thread pool"); + return; + } + // 8*64*64 = 32_768, the shape measured at 0.32x when split. + let (a, b, mut c) = operands(8, 64, 64); + reset_serial_gemms(); + gemm(&a, &b, &mut c, 8, 64, 64); + assert_eq!( + serial_gemms(), + 1, + "work below PARALLEL_MIN_WORK must not fork the pool" + ); + + // 8*512*384 = 1_572_864, measured 1.33x-1.51x in favour of splitting. + let (a, b, mut c) = operands(8, 512, 384); + reset_serial_gemms(); + gemm(&a, &b, &mut c, 8, 512, 384); + assert_eq!( + serial_gemms(), + 0, + "work above PARALLEL_MIN_WORK must still be split" + ); + } + + /// Pins the comparison at *exactly* the threshold. Without this, swapping + /// `>=` for `>` in the guard changes behaviour at the boundary and every + /// other test still passes. + #[test] + fn the_crossover_itself_splits_and_one_mac_below_it_does_not() { + if rayon::current_num_threads() < 2 { + eprintln!("skipping: needs a multi-thread pool"); + return; + } + // 8 * 512 * 256 == 1_048_576 == PARALLEL_MIN_WORK exactly. + let (m, k, n) = (8usize, 512usize, 256usize); + assert_eq!( + m * k * n, + PARALLEL_MIN_WORK, + "shape must sit on the boundary" + ); + let (a, b, mut c) = operands(m, k, n); + reset_serial_gemms(); + gemm(&a, &b, &mut c, m, k, n); + assert_eq!( + serial_gemms(), + 0, + "work equal to PARALLEL_MIN_WORK must split (the bound is inclusive)" + ); + + // One MAC below the threshold: k*n is one column short. + let n1 = n - 1; + assert!(m * k * n1 < PARALLEL_MIN_WORK); + let (a, b, mut c) = operands(m, k, n1); + reset_serial_gemms(); + gemm(&a, &b, &mut c, m, k, n1); + assert_eq!( + serial_gemms(), + 1, + "work below PARALLEL_MIN_WORK must stay serial" + ); + } + + /// A split into one block cannot use more than one thread, so forking for + /// it is pure overhead however large the operand is. `m == 1` is the case + /// that matters: a 1x2048 by 2048x2048 GEMV is 4.2M MACs, far above the + /// threshold, yet yields exactly one block. + #[test] + fn a_single_block_is_never_split_however_large_the_operand() { + if rayon::current_num_threads() < 2 { + eprintln!("skipping: needs a multi-thread pool"); + return; + } + let (m, k, n) = (1usize, 1024usize, 2048usize); + assert!(m * k * n > PARALLEL_MIN_WORK, "must clear the work bound"); + let (a, b, mut c) = operands(m, k, n); + reset_serial_gemms(); + gemm(&a, &b, &mut c, m, k, n); + assert_eq!( + serial_gemms(), + 1, + "a one-block split is pure fork overhead and must be declined" + ); + } + + /// The routing is element-type agnostic, so bf16 must get the same + /// bit-for-bit guarantee as f16 rather than inheriting it by assumption. + #[test] + fn bf16_routes_agree_bit_for_bit() { + for &(m, k, n) in &[(8usize, 64usize, 64usize), (13, 129, 67), (130, 65, 67)] { + let a: Vec = (0..m * k) + .map(|i| half::bf16::from_f32(((i % 23) as f32 - 11.0) * 0.07).to_bits()) + .collect(); + let b: Vec = (0..k * n) + .map(|i| half::bf16::from_f32(((i % 17) as f32 - 8.0) * 0.05).to_bits()) + .collect(); + let mut c = vec![0.0f32; m * n]; + with_forced_route(false, || gemm_as::(&a, &b, &mut c, m, k, n)); + let serial = c.clone(); + with_forced_route(true, || gemm_as::(&a, &b, &mut c, m, k, n)); + let sb: Vec = serial.iter().map(|v| v.to_bits()).collect(); + let pb: Vec = c.iter().map(|v| v.to_bits()).collect(); + assert!( + sb == pb, + "{m}x{k}x{n}: split and serial bf16 half GEMM disagree bit-for-bit" + ); + } + } + + /// The guard changes *scheduling*, never arithmetic: both routes must agree + /// bit-for-bit, including on shapes whose row count is not a multiple of + /// the register-block height so the final block is a tail. + /// + /// Deliberately spans `m > MAX_MC`: the serial route blocks at + /// `MAX_MC.min(m)`, so every shape with `m <= 64` yields a *single* block + /// and never exercises the loop's block indexing. Without a `m > 64` shape + /// here, corrupting the serial route's `first_row` goes undetected. + #[test] + fn both_routes_agree_bit_for_bit() { + assert_eq!(MAX_MC, 64, "shape list below is chosen to straddle MAX_MC"); + for &(m, k, n) in &[ + (1usize, 64usize, 64usize), + (3, 65, 33), + (8, 64, 64), + (13, 129, 67), + (15, 256, 256), + (16, 257, 129), + (65, 33, 49), // just over MAX_MC: two serial blocks, second a tail + (130, 65, 67), // several serial blocks with a tail + (192, 17, 23), // exact multiple of MAX_MC: no tail + ] { + let (a, b, mut c) = operands(m, k, n); + with_forced_route(false, || gemm(&a, &b, &mut c, m, k, n)); + let serial = c.clone(); + with_forced_route(true, || gemm(&a, &b, &mut c, m, k, n)); + let sb: Vec = serial.iter().map(|v| v.to_bits()).collect(); + let pb: Vec = c.iter().map(|v| v.to_bits()).collect(); + assert!( + sb == pb, + "{m}x{k}x{n}: split and serial half GEMM disagree bit-for-bit" + ); + } + } +}