From 24292d677bd0e2366a901e69eb88ff5fca4d4069 Mon Sep 17 00:00:00 2001 From: Deckard Date: Mon, 17 Aug 2026 20:41:56 +0000 Subject: [PATCH 1/5] perf(cpu-ep): run the elementwise split on ORT's pool, not a second one Our elementwise kernels split long slices across a rayon pool. That is right for the native executor, which owns the machine, and wrong inside an ORT session: ORT already has an intra-op pool and its workers spin, so ours is a second pool on the same cores. Measured at `intra_op_num_threads = 16`, 1 Mi f32, p50 (serial vs our rayon split, same binary): | op | serial | split | |----------|--------|---------| | Sqrt | 252us | 777us | | Sigmoid | 521us | 993us | | FastGelu | 1040us | 1479us | Every op lost by parallelising. Raising the length threshold does not fix it, because at `intra_op = 1` the same split is a 2-5x *win* from 1 Mi upwards -- the variable that matters is how much of the machine the host is already using, and only the host knows that. So ask it. `OrtApi::KernelContext_ParallelFor` runs a callback on the session's own intra-op pool. This adds a `host_parallel` seam in `onnx-runtime-ep-api` (the only crate both the kernels and the plugin already depend on), an ORT-backed implementation in the plugin, and installs it for the dynamic extent of each `compute_execute`. One pool, sized by whatever the user configured, and the oversubscription is gone by construction rather than by tuning. The native executor installs nothing and keeps its rayon path. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../onnx-runtime-ep-api/src/host_parallel.rs | 323 +++++++++++++++ crates/onnx-runtime-ep-api/src/lib.rs | 3 + .../src/kernels/simd_activations.rs | 392 +++++++++++++++++- crates/onnx-runtime-ep-plugin/src/compute.rs | 10 + .../onnx-runtime-ep-plugin/src/host_pool.rs | 335 +++++++++++++++ crates/onnx-runtime-ep-plugin/src/lib.rs | 1 + 6 files changed, 1063 insertions(+), 1 deletion(-) create mode 100644 crates/onnx-runtime-ep-api/src/host_parallel.rs create mode 100644 crates/onnx-runtime-ep-plugin/src/host_pool.rs diff --git a/crates/onnx-runtime-ep-api/src/host_parallel.rs b/crates/onnx-runtime-ep-api/src/host_parallel.rs new file mode 100644 index 0000000000..872a5f4ef9 --- /dev/null +++ b/crates/onnx-runtime-ep-api/src/host_parallel.rs @@ -0,0 +1,323 @@ +//! The host runtime's own parallel-for, borrowed for the duration of a kernel. +//! +//! # Why this exists +//! +//! Our CPU kernels split long elementwise slices across a `rayon` pool. That is +//! the right thing to do when we own the machine — the native executor does — +//! but it is the *wrong* thing to do inside an ORT session, because ORT already +//! has an intra-op pool of its own and that pool **spins**. Running our sixteen +//! rayon workers alongside ORT's sixteen spinning workers puts thirty-two +//! runnable threads on sixteen cores, and the result is not a small tax: +//! +//! | op, 1 Mi f32, `intra_op = 16` | serial | our rayon split | +//! |---|---|---| +//! | `Sqrt` | 252 us | 777 us | +//! | `Sigmoid` | 521 us | 993 us | +//! | `FastGelu` | 1040 us | 1479 us | +//! +//! Every one of those is a *loss* from parallelising, and it gets worse the +//! more threads ORT was given. Raising the length threshold until the split +//! only fires on very long slices trades one wrong answer for another: with +//! `intra_op = 1` there is nothing to contend with and the same split is a +//! 2-5x win from 1 Mi upwards. No constant can be right for both, because the +//! variable that actually matters is *how much of the machine the host is +//! already using*, and only the host knows that. +//! +//! So we stop guessing and ask. When ORT calls into our compute function it +//! hands us an `OrtKernelContext`, and `OrtApi::KernelContext_ParallelFor` +//! runs a callback on the session's *own* intra-op pool. Routing our chunk +//! split through it means there is exactly one pool on the machine, sized by +//! whatever the user asked for, and oversubscription disappears by +//! construction rather than by tuning. +//! +//! # Shape of the seam +//! +//! This crate is the only thing both `onnx-runtime-ep-cpu` (which has the +//! kernels) and `onnx-runtime-ep-plugin` (which has the `OrtKernelContext`) +//! already depend on, so the seam lives here and neither of them needs to +//! learn about the other. +//! +//! The plugin installs a [`HostParallel`] for the dynamic extent of one +//! compute call with [`scope`]; kernels ask for it with [`current`]. Outside a +//! plugin compute — the native executor, unit tests, anything that never +//! installs one — [`current`] is `None` and callers keep their existing +//! behaviour. +//! +//! # Threading contract +//! +//! * The installed value is **thread-local**. It is only visible on the thread +//! ORT called us on, which is the only thread that may legally touch the +//! `OrtKernelContext` it was built from. +//! * [`HostParallel::run`] is **blocking**: it returns only once every index +//! has run. That is what makes borrowing a `&dyn Fn` from the caller's frame +//! sound. +//! * Bodies run with [`in_host_task`] set, so a kernel that reaches a second +//! parallel split from inside a task can see that it is already inside the +//! host's pool and stay serial instead of nesting. + +use core::ffi::c_void; + +/// Runs `body(i)` for every `i` in `0..total` on the host's threads. +/// +/// # Safety +/// +/// The implementation receives the `host` pointer that was paired with it in +/// [`HostParallel::new`], and may assume it is still valid — which is what +/// [`scope`]'s dynamic extent guarantees. It must invoke `body` exactly once +/// per index in `0..total`, must not let a Rust panic escape into foreign +/// frames, and must not return until every invocation has finished. +pub type HostParallelForFn = + unsafe fn(host: *mut c_void, total: usize, body: &(dyn Fn(usize) + Sync)); + +/// A borrowed handle to the host runtime's thread pool. +/// +/// Deliberately `Copy` and pointer-sized: it is read out of a thread-local on +/// the hot path, and a clone there would be pure overhead. +#[derive(Clone, Copy)] +pub struct HostParallel { + host: *mut c_void, + run: HostParallelForFn, +} + +impl core::fmt::Debug for HostParallel { + fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { + f.debug_struct("HostParallel") + .field("host", &self.host) + .finish_non_exhaustive() + } +} + +impl HostParallel { + /// Pairs a host context pointer with the function that drives its pool. + /// + /// # Safety + /// + /// `host` must stay valid for as long as this handle is installed, and + /// `run` must honour the contract on [`HostParallelForFn`] for that + /// pointer. Installing it with [`scope`] is what bounds the first half; + /// the second half is on the implementer. + pub const unsafe fn new(host: *mut c_void, run: HostParallelForFn) -> Self { + Self { host, run } + } + + /// Runs `body(0..total)` on the host's pool and waits for all of it. + /// + /// `total` of zero is a no-op, and one is run inline rather than handed to + /// the pool — a single index cannot be split, so a round trip through the + /// host would be pure overhead. + pub fn run(&self, total: usize, body: &(dyn Fn(usize) + Sync)) { + match total { + 0 => (), + 1 => { + let _task = TaskGuard::enter(); + body(0); + } + _ => { + // Mark the *task*, not the dispatch: the guard has to be taken + // on whichever of the host's threads ends up running the body, + // which is why it is inside the closure rather than around the + // call. + let marked = |index: usize| { + let _task = TaskGuard::enter(); + body(index); + }; + // SAFETY: `host` and `run` were paired in `new`, whose safety + // contract puts the validity of `host` on the installer, and + // `scope` bounds it to the extent the handle is reachable. + unsafe { (self.run)(self.host, total, &marked) } + } + } + } +} + +thread_local! { + /// The host pool this thread may borrow, if it is inside a compute call. + static CURRENT: core::cell::Cell> = + const { core::cell::Cell::new(None) }; + + /// Set while this thread is running a body handed to [`HostParallel::run`]. + static IN_TASK: core::cell::Cell = const { core::cell::Cell::new(false) }; +} + +/// A [`HostParallel`] installed on this thread until this value is dropped. +/// +/// The closure form ([`scope`]) is the one to reach for. This exists for the +/// FFI entry points, whose bodies are hundreds of lines long and would have to +/// be re-indented wholesale to take a closure. +/// +/// Not `Send`, because [`HostParallel`] holds raw pointers — which is exactly +/// the property that keeps the handle on the thread ORT called us on. +pub struct Installed { + prev: Option, +} + +impl Installed { + /// Installs `host` on the calling thread. + pub fn new(host: HostParallel) -> Self { + Self { + prev: CURRENT.with(|c| c.replace(Some(host))), + } + } +} + +impl Drop for Installed { + fn drop(&mut self) { + CURRENT.with(|c| c.set(self.prev)); + } +} + +/// Installs `host` for the duration of `f` on the calling thread. +/// +/// Restores the previous value on unwind as well as on return. That is not +/// tidiness: the handle borrows an `OrtKernelContext` that ORT frees when the +/// compute call returns, so a leaked handle would be a dangling pointer the +/// next kernel on this thread would happily dispatch through. +pub fn scope(host: HostParallel, f: impl FnOnce() -> T) -> T { + let _installed = Installed::new(host); + f() +} + +/// Runs `f` with no host pool installed, restoring it afterwards. +/// +/// For the paths that have to reach a kernel from somewhere the host context +/// is no longer valid, and for tests that need the un-installed behaviour. +pub fn without(f: impl FnOnce() -> T) -> T { + struct Restore(Option); + impl Drop for Restore { + fn drop(&mut self) { + CURRENT.with(|c| c.set(self.0)); + } + } + let _restore = Restore(CURRENT.with(|c| c.replace(None))); + f() +} + +/// The host pool installed on this thread, if any. +#[inline] +pub fn current() -> Option { + CURRENT.with(core::cell::Cell::get) +} + +/// Whether this thread is currently running a [`HostParallel::run`] body. +/// +/// A kernel that reaches a parallel split from inside one is already occupying +/// a host worker; splitting again would nest a second pool inside the first. +#[inline] +pub fn in_host_task() -> bool { + IN_TASK.with(core::cell::Cell::get) +} + +struct TaskGuard(bool); + +impl TaskGuard { + fn enter() -> Self { + Self(IN_TASK.with(|c| c.replace(true))) + } +} + +impl Drop for TaskGuard { + fn drop(&mut self) { + IN_TASK.with(|c| c.set(self.0)); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicUsize, Ordering}; + + /// A host that runs everything inline, in order, on the calling thread. + /// + /// # Safety + /// + /// Ignores `host`, so any pointer is valid for it. + unsafe fn serial_host(_host: *mut c_void, total: usize, body: &(dyn Fn(usize) + Sync)) { + for index in 0..total { + body(index); + } + } + + fn serial() -> HostParallel { + // SAFETY: `serial_host` never dereferences its `host` argument. + unsafe { HostParallel::new(core::ptr::null_mut(), serial_host) } + } + + #[test] + fn no_host_is_installed_by_default() { + assert!(current().is_none()); + assert!(!in_host_task()); + } + + #[test] + fn scope_installs_and_restores() { + scope(serial(), || { + assert!(current().is_some()); + without(|| assert!(current().is_none())); + assert!(current().is_some()); + }); + assert!(current().is_none()); + } + + #[test] + fn scope_restores_on_unwind() { + let unwound = std::panic::catch_unwind(|| { + scope(serial(), || panic!("kernel failed")); + }); + assert!(unwound.is_err()); + assert!(current().is_none(), "a leaked handle would dangle"); + } + + #[test] + fn run_covers_every_index_exactly_once() { + let seen: Vec = (0..7).map(|_| AtomicUsize::new(0)).collect(); + serial().run(seen.len(), &|index| { + seen[index].fetch_add(1, Ordering::Relaxed); + }); + assert!(seen.iter().all(|c| c.load(Ordering::Relaxed) == 1)); + } + + #[test] + fn empty_total_runs_nothing() { + let calls = AtomicUsize::new(0); + serial().run(0, &|_| { + calls.fetch_add(1, Ordering::Relaxed); + }); + assert_eq!(calls.load(Ordering::Relaxed), 0); + } + + #[test] + fn a_single_index_runs_inline_and_is_still_marked() { + let marked = AtomicUsize::new(0); + serial().run(1, &|_| { + marked.fetch_add(usize::from(in_host_task()), Ordering::Relaxed); + }); + assert_eq!(marked.load(Ordering::Relaxed), 1); + } + + #[test] + fn bodies_run_marked_and_the_mark_does_not_leak() { + assert!(!in_host_task()); + let marked = AtomicUsize::new(0); + serial().run(4, &|_| { + marked.fetch_add(usize::from(in_host_task()), Ordering::Relaxed); + }); + assert_eq!(marked.load(Ordering::Relaxed), 4); + assert!(!in_host_task()); + } + + #[test] + fn the_mark_survives_a_panicking_body() { + let unwound = std::panic::catch_unwind(|| { + serial().run(4, &|index| assert_ne!(index, 2, "body failed")); + }); + assert!(unwound.is_err()); + assert!(!in_host_task()); + } + + #[test] + fn a_handle_is_not_visible_from_another_thread() { + scope(serial(), || { + assert!(std::thread::spawn(|| current().is_none()).join().unwrap()); + }); + } +} diff --git a/crates/onnx-runtime-ep-api/src/lib.rs b/crates/onnx-runtime-ep-api/src/lib.rs index 585c73d5de..1248e3027e 100644 --- a/crates/onnx-runtime-ep-api/src/lib.rs +++ b/crates/onnx-runtime-ep-api/src/lib.rs @@ -18,6 +18,7 @@ //! * [`tensor`] — [`TensorView`] / [`TensorMut`] zero-copy device views. //! * [`weight`] — capability-negotiated lazy [`WeightHandle`] delivery. //! * [`abi`] — ORT graph ABI bridge for legacy plugin EPs (Phase 2). +//! * [`host_parallel`] — the host runtime's own thread pool, borrowed per compute call. // This crate hands raw pointers across an FFI boundary, so keeping a Rust // value alive past its scope is a recurring temptation. It is almost never the @@ -34,6 +35,7 @@ pub mod abi; pub mod epcontext; +pub mod host_parallel; pub mod kernel; pub mod provider; pub mod registry; @@ -43,6 +45,7 @@ pub mod weight; pub use abi::{LegacyOrtEp, PluginCompiledKernel, PluginExecutionPlan, SubgraphClaim}; pub use epcontext::{EpContext, EpContextRegistry, build_ep_context_registry}; pub use error::{EpError, Result}; +pub use host_parallel::HostParallel; pub use kernel::{ ARG_BYTES, ARG_DEVICE, ARG_FLOPS, ARG_KERNEL_VARIANT, ARG_KERNEL_VARIANT_REASON, CAT_KERNEL_WORKER, CaptureSupport, ClaimPreference, Cost, Kernel, KernelInput, KernelMatch, diff --git a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs index 931e29f8a5..7e9492f8ac 100644 --- a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs +++ b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs @@ -290,6 +290,130 @@ fn par_chunk_len(len: usize, threads: usize) -> Option { (chunk < len).then_some(chunk) } +/// Chunks to cut `len` into when the *host* runtime owns the pool. +/// +/// The rayon split asks the pool how wide it is and cuts exactly that many +/// pieces, because a rayon `par_chunks` costs a wake-up per piece and handing +/// it more pieces than threads is wasted scheduling. The host pool is the +/// opposite: it is already spinning, we cannot ask ORT how many intra-op +/// threads it was given, and `TrySimpleParallelFor` claims indices +/// dynamically. So we cut by *size* rather than by thread count and let the +/// host decide how many of its threads to point at the result — which also +/// makes the split independent of the machine, so the chunk boundaries (and +/// therefore the bits) do not move when the session is reconfigured. +/// +/// [`PAR_MIN_CHUNK`] is still the floor, so a chunk is never too short to be +/// worth dispatching or to stay on the vector path. The cap keeps a very long +/// slice from turning into thousands of tiny tasks. +const MAX_HOST_CHUNKS: usize = 64; + +/// Chunk length and count for the host-pool split, or `None` to stay serial. +#[inline] +fn host_chunk_len(len: usize) -> Option<(usize, usize)> { + let chunk = len + .div_ceil(MAX_HOST_CHUNKS) + .max(PAR_MIN_CHUNK) + .next_multiple_of(8); + (chunk < len).then(|| (chunk, len.div_ceil(chunk))) +} + +/// [`host_chunk_len`] for the bias-fused kernels: whole multiples of `width`. +#[inline] +fn host_chunk_len_rows(len: usize, width: usize) -> Option<(usize, usize)> { + if width == 0 { + return None; + } + let rows = len + .div_ceil(width) + .div_ceil(MAX_HOST_CHUNKS) + .max(PAR_MIN_CHUNK.div_ceil(width)); + let chunk = rows.checked_mul(width)?; + (chunk < len).then(|| (chunk, len.div_ceil(chunk))) +} + +/// `input`/`output` as raw pointers, so the host's threads can share them. +/// +/// Rayon proves disjointness with `par_chunks_mut`; the host pool has no such +/// API, so the split is done by index arithmetic and the disjointness argument +/// moves into [`run_on_host`]. +struct HostChunks { + input: *const f32, + output: *mut f32, + len: usize, + chunk: usize, +} + +// SAFETY: the only thing shared across threads is a pair of pointers plus the +// two lengths needed to derive a chunk from an index. `run_on_host` gives each +// index a disjoint half-open range of both slices, and `HostParallel::run` +// promises each index runs exactly once, so no two threads ever hold +// overlapping references to the output. +unsafe impl Sync for HostChunks {} + +impl HostChunks { + /// Runs `body` on the `index`-th chunk of the input and output slices. + /// + /// Takes the body rather than returning the two slices so that the + /// caller's closure captures the whole `&HostChunks` — which is `Sync` — + /// instead of the raw pointer fields individually, which are not. + /// + /// # Safety + /// + /// `index` must be below the chunk count this was built for, and no two + /// live calls may share one: the `&mut` handed to `body` is only unique + /// because distinct indices give disjoint ranges. + unsafe fn run_chunk(&self, index: usize, body: &F) + where + F: Fn(&[f32], &mut [f32]), + { + let start = index * self.chunk; + // The last chunk is short whenever `chunk` does not divide `len`. + let this = self.chunk.min(self.len - start); + // SAFETY: `start + this <= len` by construction, and the caller + // guarantees no other thread holds this range. + let (input, output) = unsafe { + ( + core::slice::from_raw_parts(self.input.add(start), this), + core::slice::from_raw_parts_mut(self.output.add(start), this), + ) + }; + body(input, output); + } +} + +/// Splits `input`/`output` into `count` chunks of `chunk` and runs `body` on +/// each one on the host runtime's pool. +/// +/// Bit-identical to calling `body` on the whole slice, for the same reason the +/// rayon path is: the kernels are elementwise, every chunk is a multiple of +/// eight lanes, and none is shorter than [`SIMD_MIN_LEN`], so no chunk takes a +/// different code path from its neighbours. +fn run_on_host( + host: onnx_runtime_ep_api::HostParallel, + input: &[f32], + output: &mut [f32], + chunk: usize, + count: usize, + body: &F, +) where + F: Fn(&[f32], &mut [f32]) + Sync, +{ + debug_assert_eq!(input.len(), output.len()); + debug_assert!(chunk > 0 && count > 0); + let shared = HostChunks { + input: input.as_ptr(), + output: output.as_mut_ptr(), + len: input.len(), + chunk, + }; + let shared = &shared; + host.run(count, &|index| { + // SAFETY: `HostParallel::run` invokes every index in `0..count` + // exactly once, so this call owns its range for its whole duration. + unsafe { shared.run_chunk(index, body) }; + }); +} + /// Chunk length for [`run_chunked_rows`], always a whole multiple of `width`. #[inline] fn par_chunk_len_rows(len: usize, width: usize, threads: usize) -> Option { @@ -334,6 +458,24 @@ where body(input, output); return; } + // Inside a host task we are already occupying one of the host's threads; + // splitting again would nest a pool inside the pool we were handed. + if onnx_runtime_ep_api::host_parallel::in_host_task() { + body(input, output); + return; + } + // Prefer the host's pool over ours whenever the host offered one. Ours + // would be a *second* pool on the same cores: measured 1.5-3.1x slower + // than staying serial at 1 Mi under an `intra_op = 16` session. + if let Some(host) = onnx_runtime_ep_api::host_parallel::current() { + if let Some((chunk, count)) = host_chunk_len(len) { + note_parallel_dispatch(); + run_on_host(host, input, output, chunk, count, &body); + return; + } + body(input, output); + return; + } // Nesting a `par_chunks` inside an outer parallel region only adds // scheduling overhead: the outer region is already keeping the pool busy. if rayon::current_thread_index().is_some() { @@ -391,7 +533,11 @@ pub(crate) fn run_chunked_fn(input: &[f32], output: &mut [f32], body: fn(&[f32], pub(crate) fn clip_chunked(input: &[f32], output: &mut [f32], minimum: f32, maximum: f32) { // Take the serial decision here, so the common short case is a direct call // instead of one through a closure the optimiser can no longer see into. - if input.len() < PAR_MIN_LEN || force_serial() || rayon::current_thread_index().is_some() { + if input.len() < PAR_MIN_LEN + || force_serial() + || onnx_runtime_ep_api::host_parallel::in_host_task() + || rayon::current_thread_index().is_some() + { mlas_sys::compute_clip(input, output, minimum, maximum); return; } @@ -422,6 +568,18 @@ where body(input, output); return; } + if onnx_runtime_ep_api::host_parallel::in_host_task() { + body(input, output); + return; + } + if let Some(host) = onnx_runtime_ep_api::host_parallel::current() { + if let Some((chunk, count)) = host_chunk_len_rows(len, width) { + run_on_host(host, input, output, chunk, count, &body); + return; + } + body(input, output); + return; + } if rayon::current_thread_index().is_some() { body(input, output); return; @@ -3938,3 +4096,235 @@ mod chunking_instantiation_is_local { ); } } + +/// The host-pool split: what happens when ORT lends us its intra-op threads. +/// +/// The production host is ORT's `KernelContext_ParallelFor`, which needs a +/// live session to exercise. These tests stand a real thread pool in its +/// place, so everything on our side of the seam — the chunk policy, the +/// disjointness of the ranges, the suppression of the rayon path and of +/// nesting — is covered without one. +#[cfg(test)] +mod host_pool_split { + use super::*; + use onnx_runtime_ep_api::HostParallel; + use onnx_runtime_ep_api::host_parallel; + use std::ffi::c_void; + use std::sync::atomic::{AtomicUsize, Ordering}; + + thread_local! { + /// Indices dispatched through the fake host from *this* thread since + /// the last reset. + /// + /// Thread-local rather than a global atomic for the same reason + /// `PARALLEL_DISPATCHES` is: the test binary runs these tests + /// concurrently, and a global counter would let them bump each + /// other's. The increment happens on whichever thread called + /// `threaded_host`, which is also the thread that reads it back -- + /// including inside a task, which is what makes the nesting test + /// work. + static HOST_INDICES: std::cell::Cell = const { std::cell::Cell::new(0) }; + } + + fn dispatched() -> usize { + HOST_INDICES.with(std::cell::Cell::get) + } + + fn reset_dispatched() { + HOST_INDICES.with(|c| c.set(0)); + } + + /// Stands in for ORT: runs the indices on four real threads. + /// + /// Genuinely concurrent on purpose. A serial stand-in would still prove + /// the arithmetic but not that the ranges are disjoint, which is the part + /// that would corrupt an output tensor if it were wrong. + /// + /// # Safety + /// + /// Ignores `host`, so any pointer is valid for it. + unsafe fn threaded_host(_host: *mut c_void, total: usize, body: &(dyn Fn(usize) + Sync)) { + HOST_INDICES.with(|c| c.set(c.get() + total)); + let next = AtomicUsize::new(0); + std::thread::scope(|scope| { + for _ in 0..4 { + scope.spawn(|| { + loop { + let index = next.fetch_add(1, Ordering::Relaxed); + if index >= total { + break; + } + body(index); + } + }); + } + }); + } + + fn fake_host() -> HostParallel { + // SAFETY: `threaded_host` never dereferences its `host` argument. + unsafe { HostParallel::new(core::ptr::null_mut(), threaded_host) } + } + + /// Long enough to be split, and deliberately not a multiple of the chunk + /// size, so the final chunk is short. + const N: usize = 3 * PAR_MIN_LEN + 37; + + fn probe(len: usize) -> Vec { + (0..len) + .map(|i| (i as f32 / 991.0).sin() * 24.0 + (i % 7) as f32 * 1e-7 - 3.0) + .collect() + } + + /// Asserts the host split computes exactly what one call would. + fn assert_same(name: &str, run: impl Fn(&[f32], &mut [f32]) + Sync) { + let x = probe(N); + let mut want = vec![0.0f32; N]; + serial_scope(|| run(&x, &mut want)); + + reset_dispatched(); + let mut got = vec![0.0f32; N]; + host_parallel::scope(fake_host(), || run(&x, &mut got)); + assert!( + dispatched() > 1, + "{name}: the host pool was never asked to split anything" + ); + + for (i, (&w, &g)) in want.iter().zip(&got).enumerate() { + assert_eq!( + w.to_bits(), + g.to_bits(), + "{name}: element {i} differs (whole {w:e}, host-split {g:e})" + ); + } + } + + #[test] + fn unary_kernels_match_the_unsplit_result() { + assert_same("tanh", tanh_f32_slice); + assert_same("sqrt", sqrt_f32_slice); + assert_same("sigmoid", sigmoid_f32_slice); + assert_same("exp", exp_f32_slice); + assert_same("erf", erf_f32_slice); + assert_same("erf_gelu", erf_gelu_f32_slice); + assert_same("tanh_gelu", tanh_gelu_f32_slice); + assert_same("quick_gelu", |x, y| quick_gelu_f32_slice(x, y, 1.702)); + } + + #[test] + fn bias_kernels_match_the_unsplit_result() { + for width in [1usize, 3, 11, 4096, 4099] { + let bias: Vec = (0..width).map(|i| (i as f32) * 0.013 - 0.4).collect(); + assert_same("tanh_gelu_bias", |x, y| { + tanh_gelu_bias_f32_slice(x, &bias, width, y) + }); + assert_same("erf_gelu_bias", |x, y| { + erf_gelu_bias_f32_slice(x, &bias, width, y) + }); + } + } + + /// The point of the whole change: with a host pool installed we must not + /// also start our own. + #[test] + fn the_rayon_pool_is_not_used_when_a_host_is_installed() { + if rayon::current_num_threads() < 2 { + eprintln!("skipped: single-threaded rayon pool cannot show the difference"); + return; + } + let x = probe(N); + let mut y = vec![0.0f32; N]; + + reset_dispatched(); + host_parallel::scope(fake_host(), || tanh_f32_slice(&x, &mut y)); + let host_indices = dispatched(); + + assert!(host_indices > 1, "the host pool was not used"); + assert_eq!( + host_indices, + host_chunk_len(N).expect("N is long enough to split").1, + "every chunk should have been dispatched exactly once" + ); + } + + /// A kernel reached from inside a host task is already on a host thread. + /// Splitting again would nest a pool inside the pool we were handed. + #[test] + fn a_nested_split_stays_serial() { + let x = probe(N); + let mut y = vec![0.0f32; N]; + host_parallel::scope(fake_host(), || { + fake_host().run(2, &|_| { + assert!(host_parallel::in_host_task()); + let before = dispatched(); + let mut inner = vec![0.0f32; N]; + tanh_f32_slice(&x, &mut inner); + assert_eq!( + dispatched(), + before, + "a kernel inside a host task dispatched again" + ); + }); + tanh_f32_slice(&x, &mut y); + }); + } + + /// `serial_scope` has to keep suppressing the split on the host path too, + /// or the f16/bf16 sandwich picks its measured regression back up. + #[test] + fn serial_scope_still_suppresses_the_split() { + let x = probe(N); + let mut y = vec![0.0f32; N]; + reset_dispatched(); + host_parallel::scope(fake_host(), || { + serial_scope(|| tanh_f32_slice(&x, &mut y)); + }); + assert_eq!(dispatched(), 0); + } + + /// The host chunk policy is pure, so sweep it over lengths and widths this + /// machine's memory could not hold, and assert the invariants the kernels + /// depend on: whole vectors, never below the vector threshold, and — for + /// the bias kernels — never a cut through the middle of a row. + #[test] + fn host_chunk_policy_holds_across_lengths() { + for len in [ + 0, + 1, + 8, + PAR_MIN_CHUNK - 1, + PAR_MIN_CHUNK, + PAR_MIN_CHUNK + 1, + PAR_MIN_LEN - 1, + PAR_MIN_LEN, + PAR_MIN_LEN + 1, + N, + 1 << 26, + (1 << 26) + 13, + usize::MAX / 2, + ] { + if let Some((chunk, count)) = host_chunk_len(len) { + assert!(chunk < len, "chunk {chunk} !< len {len}"); + assert_eq!(chunk % 8, 0, "chunk {chunk} is not a whole vector"); + assert!(chunk >= PAR_MIN_CHUNK, "chunk {chunk} is too short"); + assert!(chunk >= SIMD_MIN_LEN, "chunk {chunk} would go scalar"); + assert_eq!(count, len.div_ceil(chunk)); + assert!( + (count - 1) * chunk < len, + "chunk {chunk} x {count} would dispatch an empty range" + ); + assert!(count <= MAX_HOST_CHUNKS, "{count} chunks is past the cap"); + } + for width in [1usize, 3, 8, 64, 4099, PAR_MIN_LEN] { + if let Some((chunk, count)) = host_chunk_len_rows(len, width) { + assert!(chunk < len, "rows: chunk {chunk} !< len {len}"); + assert_eq!(chunk % width, 0, "rows: chunk {chunk} cuts width {width}"); + assert!(chunk >= PAR_MIN_CHUNK.min(len), "rows: chunk {chunk} short"); + assert_eq!(count, len.div_ceil(chunk)); + assert!((count - 1) * chunk < len, "rows: empty final range"); + } + } + assert_eq!(host_chunk_len_rows(len, 0), None, "width 0 must not split"); + } + } +} diff --git a/crates/onnx-runtime-ep-plugin/src/compute.rs b/crates/onnx-runtime-ep-plugin/src/compute.rs index 2d4f675319..72acc2f78a 100644 --- a/crates/onnx-runtime-ep-plugin/src/compute.rs +++ b/crates/onnx-runtime-ep-plugin/src/compute.rs @@ -2075,6 +2075,16 @@ unsafe extern "C" fn compute_execute( } let api_ref = unsafe { &*api }; + // Lend ORT's intra-op pool to the kernels for this call. Ours would be + // a second pool on the same cores, and ORT's workers spin: at + // `intra_op = 16`, splitting a 1 Mi `Sqrt` across our rayon pool cost + // 252 -> 777 us against staying serial. Dropped at the end of the + // call, before `kernel_context` goes away. + // + // SAFETY: `kernel_context` is the context ORT handed this call and + // stays valid until it returns, which is after the guard is dropped. + let _host_pool = unsafe { crate::host_pool::install(api_ref, kernel_context) }; + // Memory info for intermediate scratch. On a device EP this is device // memory, so multi-node intermediates stay on the GPU (a host buffer // would make the next kernel dereference a host pointer as device → diff --git a/crates/onnx-runtime-ep-plugin/src/host_pool.rs b/crates/onnx-runtime-ep-plugin/src/host_pool.rs new file mode 100644 index 0000000000..1f4a60ecc2 --- /dev/null +++ b/crates/onnx-runtime-ep-plugin/src/host_pool.rs @@ -0,0 +1,335 @@ +//! ORT's intra-op thread pool, exposed to our kernels for one compute call. +//! +//! Our elementwise CPU kernels split long slices across a `rayon` pool. Under +//! an ORT session that pool is a *second* pool on the same cores, and ORT's +//! intra-op workers spin, so the two fight: at `intra_op_num_threads = 16` a +//! 1 Mi `Sqrt` went 252 us serial to 777 us split, and `FastGelu` 1040 us to +//! 1479 us. Parallelising made every op slower, and the more threads the user +//! asked for the worse it got. +//! +//! `OrtApi::KernelContext_ParallelFor` runs a callback on the session's own +//! intra-op pool. Pointing our chunk split at it leaves exactly one pool on +//! the machine, sized by whatever the user configured — so the contention is +//! gone by construction, and we stop spending threads the caller never +//! offered us. This module builds the +//! [`HostParallel`](onnx_runtime_ep_api::HostParallel) that does it, and +//! [`scope`] installs it for the dynamic extent of one compute call. +//! +//! Sessions whose ORT is older than 1.17 have a null `KernelContext_ParallelFor` +//! and get no handle at all, which leaves the kernels on their existing rayon +//! path. + +use core::ffi::c_void; +use onnx_genai_ort_sys as ort; +use onnx_runtime_ep_api::HostParallel; +use onnx_runtime_ep_api::host_parallel; +use std::panic::AssertUnwindSafe; +use std::sync::atomic::{AtomicBool, Ordering}; + +/// What [`ort_parallel_for`] needs to reach ORT, behind one `void*`. +struct HostPool { + parallel_for: unsafe extern "C" fn( + *const ort::OrtKernelContext, + Option, + usize, + usize, + *mut c_void, + ) -> *mut ort::OrtStatus, + release_status: Option, + ctx: *mut ort::OrtKernelContext, +} + +/// The closure a dispatch is running, plus somewhere to record a panic. +struct Task<'a> { + body: &'a (dyn Fn(usize) + Sync), + panicked: AtomicBool, +} + +/// Trampoline handed to ORT: one index of one dispatch. +/// +/// # Safety +/// +/// `usr_data` must be the `*mut Task` that [`ort_parallel_for`] passed to +/// `KernelContext_ParallelFor`, and must outlive the call — which it does, +/// because that call is blocking. +unsafe extern "C" fn run_index(usr_data: *mut c_void, index: usize) { + // A Rust panic must not unwind into ORT's frames: that is undefined + // behaviour, and this callback is called from C++. Catch it here, record + // it, and let `ort_parallel_for` re-raise on the calling thread once every + // worker has finished. + let task = unsafe { &*(usr_data.cast::>()) }; + if std::panic::catch_unwind(AssertUnwindSafe(|| (task.body)(index))).is_err() { + task.panicked.store(true, Ordering::Relaxed); + } +} + +/// Runs `body(0..total)` on the ORT session's intra-op pool. +/// +/// # Safety +/// +/// `host` must be the `*mut HostPool` that [`scope`] built, still valid — i.e. +/// this must be reached from inside the compute call that installed it. +unsafe fn ort_parallel_for(host: *mut c_void, total: usize, body: &(dyn Fn(usize) + Sync)) { + let pool = unsafe { &*(host.cast::()) }; + let task = Task { + body, + panicked: AtomicBool::new(false), + }; + // `num_batch = 0` means "no limit": ORT gives every index its own task and + // its workers claim them dynamically. That is what we want, because we cut + // the slice by size rather than by thread count — we cannot ask ORT how + // wide its pool is, so we hand it enough pieces to fill any pool and let + // it decide how many to run at once. + let status = unsafe { + (pool.parallel_for)( + pool.ctx, + Some(run_index), + total, + 0, + (&raw const task).cast::().cast_mut(), + ) + }; + if !status.is_null() { + // The dispatch itself failed. Every index still has to run or the + // output tensor keeps whatever was in the buffer, so fall back to + // running them here rather than silently returning short. + if let Some(release) = pool.release_status { + unsafe { release(status) }; + } + for index in 0..total { + body(index); + } + return; + } + assert!( + !task.panicked.load(Ordering::Relaxed), + "a kernel panicked on an ORT intra-op thread" + ); +} + +/// ORT's pool, installed on this thread until the guard is dropped. +/// +/// Field order is the safety argument: `installed` is dropped first, so the +/// thread-local handle is gone before `pool` — the allocation it points at — +/// is freed. +pub struct Guard { + installed: Option, + pool: Option>, +} + +impl Guard { + /// A guard that installs nothing, for hosts that offer no pool. + const fn inert() -> Self { + Self { + installed: None, + pool: None, + } + } + + /// Whether a handle is actually installed, for tests and diagnostics. + #[must_use] + pub const fn is_installed(&self) -> bool { + self.installed.is_some() + } +} + +impl Drop for Guard { + fn drop(&mut self) { + // Explicit, and in this order, so a later field reshuffle cannot turn + // the handle into a dangling pointer without failing to compile. + drop(self.installed.take()); + drop(self.pool.take()); + } +} + +/// Installs ORT's pool on the calling thread until the returned guard drops. +/// +/// Returns an inert guard when this ORT has no `KernelContext_ParallelFor` +/// (pre-1.17) or the context is null, which leaves the kernels on their +/// existing rayon path. +/// +/// # Safety +/// +/// `api` must be a valid `OrtApi`, and `ctx` a valid `OrtKernelContext` that +/// stays valid for as long as the guard is alive. Because the handle is +/// thread-local and the guard is confined to the compute call's frame, it +/// cannot be reached from a later call whose context has been freed. +#[must_use = "dropping the guard immediately uninstalls the pool"] +pub unsafe fn install(api: &ort::OrtApi, ctx: *mut ort::OrtKernelContext) -> Guard { + let (Some(parallel_for), false) = (api.KernelContext_ParallelFor, ctx.is_null()) else { + return Guard::inert(); + }; + let pool = Box::new(HostPool { + parallel_for, + release_status: api.ReleaseStatus, + ctx, + }); + // SAFETY: the box outlives the handle (see `Guard`'s drop order), and + // `ort_parallel_for` only ever reads the pointer back as a `*mut HostPool`. + let handle = unsafe { + HostParallel::new( + (&raw const *pool).cast::().cast_mut(), + ort_parallel_for, + ) + }; + Guard { + installed: Some(host_parallel::Installed::new(handle)), + pool: Some(pool), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::AtomicUsize; + + /// Stands in for ORT: runs every index inline and reports success. + /// + /// # Safety + /// + /// Matches the ABI ORT expects; `usr_data` is passed straight through. + unsafe extern "C" fn inline_parallel_for( + _ctx: *const ort::OrtKernelContext, + body: Option, + total: usize, + _num_batch: usize, + usr_data: *mut c_void, + ) -> *mut ort::OrtStatus { + let body = body.expect("ORT is always given a callback"); + for index in 0..total { + unsafe { body(usr_data, index) }; + } + core::ptr::null_mut() + } + + /// Stands in for an ORT that refuses the dispatch. + /// + /// # Safety + /// + /// Returns a non-null status that is never dereferenced, because the + /// `release_status` hook in these tests is `None`. + unsafe extern "C" fn failing_parallel_for( + _ctx: *const ort::OrtKernelContext, + _body: Option, + _total: usize, + _num_batch: usize, + _usr_data: *mut c_void, + ) -> *mut ort::OrtStatus { + core::ptr::dangling_mut() + } + + fn pool( + parallel_for: unsafe extern "C" fn( + *const ort::OrtKernelContext, + Option, + usize, + usize, + *mut c_void, + ) -> *mut ort::OrtStatus, + ) -> HostPool { + HostPool { + parallel_for, + release_status: None, + ctx: core::ptr::null_mut(), + } + } + + #[test] + fn every_index_runs_once() { + let mut pool = pool(inline_parallel_for); + let seen: Vec = (0..5).map(|_| AtomicUsize::new(0)).collect(); + unsafe { + ort_parallel_for((&raw mut pool).cast::(), seen.len(), &|index| { + seen[index].fetch_add(1, Ordering::Relaxed); + }); + } + assert!(seen.iter().all(|c| c.load(Ordering::Relaxed) == 1)); + } + + #[test] + fn a_refused_dispatch_still_runs_every_index() { + let mut pool = pool(failing_parallel_for); + let seen: Vec = (0..5).map(|_| AtomicUsize::new(0)).collect(); + unsafe { + ort_parallel_for((&raw mut pool).cast::(), seen.len(), &|index| { + seen[index].fetch_add(1, Ordering::Relaxed); + }); + } + assert!( + seen.iter().all(|c| c.load(Ordering::Relaxed) == 1), + "a short write would leave the output tensor uninitialised" + ); + } + + #[test] + fn a_panicking_body_does_not_unwind_into_ort() { + let mut pool = pool(inline_parallel_for); + let ran = AtomicUsize::new(0); + let unwound = std::panic::catch_unwind(AssertUnwindSafe(|| unsafe { + ort_parallel_for((&raw mut pool).cast::(), 4, &|index| { + ran.fetch_add(1, Ordering::Relaxed); + assert_ne!(index, 1, "kernel failed"); + }); + })); + // The panic surfaces on the calling thread, *after* the dispatch has + // drained -- every index still ran. + assert!(unwound.is_err()); + assert_eq!(ran.load(Ordering::Relaxed), 4); + } + + #[test] + fn an_ort_without_parallel_for_installs_nothing() { + let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; + api.KernelContext_ParallelFor = None; + let guard = unsafe { install(&api, core::ptr::dangling_mut()) }; + assert!(!guard.is_installed()); + assert!(host_parallel::current().is_none()); + } + + #[test] + fn a_null_context_installs_nothing() { + let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; + api.KernelContext_ParallelFor = Some(inline_parallel_for); + let guard = unsafe { install(&api, core::ptr::null_mut()) }; + assert!(!guard.is_installed()); + assert!(host_parallel::current().is_none()); + } + + #[test] + fn install_gives_a_working_handle_and_removes_it_on_drop() { + let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; + api.KernelContext_ParallelFor = Some(inline_parallel_for); + let seen: Vec = (0..6).map(|_| AtomicUsize::new(0)).collect(); + { + let guard = unsafe { install(&api, core::ptr::dangling_mut()) }; + assert!(guard.is_installed()); + let host = host_parallel::current().expect("install publishes a handle"); + host.run(seen.len(), &|index| { + seen[index].fetch_add(1, Ordering::Relaxed); + }); + } + assert!(seen.iter().all(|c| c.load(Ordering::Relaxed) == 1)); + assert!( + host_parallel::current().is_none(), + "the handle must not outlive the compute call" + ); + } + + #[test] + fn a_nested_install_restores_the_outer_handle() { + let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; + api.KernelContext_ParallelFor = Some(inline_parallel_for); + let outer = unsafe { install(&api, core::ptr::dangling_mut()) }; + assert!(outer.is_installed()); + { + let _inner = unsafe { install(&api, core::ptr::dangling_mut()) }; + assert!(host_parallel::current().is_some()); + } + assert!( + host_parallel::current().is_some(), + "the inner guard must not uninstall the outer one" + ); + drop(outer); + assert!(host_parallel::current().is_none()); + } +} diff --git a/crates/onnx-runtime-ep-plugin/src/lib.rs b/crates/onnx-runtime-ep-plugin/src/lib.rs index 18987e4007..2ddec2f467 100644 --- a/crates/onnx-runtime-ep-plugin/src/lib.rs +++ b/crates/onnx-runtime-ep-plugin/src/lib.rs @@ -32,6 +32,7 @@ pub mod device; pub mod ep; pub mod factory; pub mod graph_reader; +pub mod host_pool; pub mod kernel_ctx; pub mod status; pub mod transfer; From d77eccc21bc44993f26771af79507112747bbc9a Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Mon, 17 Aug 2026 21:31:53 +0000 Subject: [PATCH 2/5] perf(cpu-ep): only borrow the host's pool when it really has workers The previous commit always preferred `KernelContext_ParallelFor` over our own rayon pool. That is right when ORT's intra-op pool is wide -- ours would be a second pool fighting it for the same cores -- but wrong when the session was built with `intra_op = 1`. There the host is not using the machine at all, and borrowing its single thread ran 2-9x slower over 1-4 Mi than splitting across rayon, which is what main did. We cannot ask ORT how wide its pool is, so observe it: the first dispatch with more than one index records whether any index ran on a thread other than the caller's, and stores the verdict in a cell owned by the fused node (not a global -- one process may hold both a 1-thread and a 16-thread session, and the right answer is the opposite for each). Until that verdict exists the host path is taken, which costs at most one dispatch on a serial host. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../onnx-runtime-ep-api/src/host_parallel.rs | 98 ++++++++- crates/onnx-runtime-ep-api/src/lib.rs | 2 +- .../src/kernels/simd_activations.rs | 195 +++++++++++++++--- crates/onnx-runtime-ep-plugin/src/compute.rs | 15 +- .../onnx-runtime-ep-plugin/src/host_pool.rs | 152 +++++++++++++- 5 files changed, 414 insertions(+), 48 deletions(-) diff --git a/crates/onnx-runtime-ep-api/src/host_parallel.rs b/crates/onnx-runtime-ep-api/src/host_parallel.rs index 872a5f4ef9..795d8c05ed 100644 --- a/crates/onnx-runtime-ep-api/src/host_parallel.rs +++ b/crates/onnx-runtime-ep-api/src/host_parallel.rs @@ -56,6 +56,31 @@ //! host's pool and stay serial instead of nesting. use core::ffi::c_void; +use core::sync::atomic::{AtomicU8, Ordering}; + +/// How many threads the host's pool turned out to have. +/// +/// The host runtime does not tell us, and it matters: a pool of one is not +/// using the machine, so ours may. Rather than probe — which would cost a +/// dispatch per session to answer a question the next real dispatch answers +/// for free — the implementation *observes* its own first dispatch and records +/// the result here. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum HostWidth { + /// No dispatch has completed yet. + Unknown, + /// Every body ran on the dispatching thread: the host has no workers. + Serial, + /// At least one body ran on another thread. + Parallel, +} + +/// Nothing observed yet. +pub const WIDTH_UNKNOWN: u8 = 0; +/// The host pool ran everything inline. +pub const WIDTH_SERIAL: u8 = 1; +/// The host pool used at least one other thread. +pub const WIDTH_PARALLEL: u8 = 2; /// Runs `body(i)` for every `i` in `0..total` on the host's threads. /// @@ -77,6 +102,9 @@ pub type HostParallelForFn = pub struct HostParallel { host: *mut c_void, run: HostParallelForFn, + /// Where the implementation records [`HostWidth`], or null if it never + /// will — in which case the pool is assumed parallel. + width: *const AtomicU8, } impl core::fmt::Debug for HostParallel { @@ -96,8 +124,36 @@ impl HostParallel { /// `run` must honour the contract on [`HostParallelForFn`] for that /// pointer. Installing it with [`scope`] is what bounds the first half; /// the second half is on the implementer. - pub const unsafe fn new(host: *mut c_void, run: HostParallelForFn) -> Self { - Self { host, run } + /// + /// `width` must be null or point at a cell that outlives this handle, and + /// that only ever holds [`WIDTH_UNKNOWN`], [`WIDTH_SERIAL`] or + /// [`WIDTH_PARALLEL`]. + pub const unsafe fn new( + host: *mut c_void, + run: HostParallelForFn, + width: *const AtomicU8, + ) -> Self { + Self { host, run, width } + } + + /// What the last completed dispatch said about the host pool's width. + /// + /// [`HostWidth::Serial`] is the interesting answer: it means the host's + /// pool has no workers of its own, so borrowing it would serialise us for + /// nothing and our own pool is free to use the machine instead. + #[must_use] + pub fn width(&self) -> HostWidth { + if self.width.is_null() { + return HostWidth::Parallel; + } + // SAFETY: `new`'s contract puts the validity of this pointer on the + // installer, and `scope` bounds it to the extent the handle is + // reachable. + match unsafe { &*self.width }.load(Ordering::Relaxed) { + WIDTH_SERIAL => HostWidth::Serial, + WIDTH_PARALLEL => HostWidth::Parallel, + _ => HostWidth::Unknown, + } } /// Runs `body(0..total)` on the host's pool and waits for all of it. @@ -238,8 +294,42 @@ mod tests { } fn serial() -> HostParallel { - // SAFETY: `serial_host` never dereferences its `host` argument. - unsafe { HostParallel::new(core::ptr::null_mut(), serial_host) } + // SAFETY: `serial_host` never dereferences its `host` argument, and a + // null width cell is explicitly allowed. + unsafe { HostParallel::new(core::ptr::null_mut(), serial_host, core::ptr::null()) } + } + + fn with_width(cell: &AtomicU8) -> HostParallel { + // SAFETY: as `serial`, plus `cell` outlives every use below. + unsafe { + HostParallel::new( + core::ptr::null_mut(), + serial_host, + core::ptr::from_ref(cell), + ) + } + } + + #[test] + fn a_handle_without_a_width_cell_reads_as_parallel() { + assert_eq!(serial().width(), HostWidth::Parallel); + } + + #[test] + fn width_reflects_the_cell() { + let cell = AtomicU8::new(WIDTH_UNKNOWN); + let host = with_width(&cell); + assert_eq!(host.width(), HostWidth::Unknown); + cell.store(WIDTH_SERIAL, Ordering::Relaxed); + assert_eq!(host.width(), HostWidth::Serial); + cell.store(WIDTH_PARALLEL, Ordering::Relaxed); + assert_eq!(host.width(), HostWidth::Parallel); + cell.store(200, Ordering::Relaxed); + assert_eq!( + host.width(), + HostWidth::Unknown, + "an unknown code is not a promise" + ); } #[test] diff --git a/crates/onnx-runtime-ep-api/src/lib.rs b/crates/onnx-runtime-ep-api/src/lib.rs index 1248e3027e..a51a84ffd1 100644 --- a/crates/onnx-runtime-ep-api/src/lib.rs +++ b/crates/onnx-runtime-ep-api/src/lib.rs @@ -45,7 +45,7 @@ pub mod weight; pub use abi::{LegacyOrtEp, PluginCompiledKernel, PluginExecutionPlan, SubgraphClaim}; pub use epcontext::{EpContext, EpContextRegistry, build_ep_context_registry}; pub use error::{EpError, Result}; -pub use host_parallel::HostParallel; +pub use host_parallel::{HostParallel, HostWidth}; pub use kernel::{ ARG_BYTES, ARG_DEVICE, ARG_FLOPS, ARG_KERNEL_VARIANT, ARG_KERNEL_VARIANT_REASON, CAT_KERNEL_WORKER, CaptureSupport, ClaimPreference, Cost, Kernel, KernelInput, KernelMatch, diff --git a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs index 7e9492f8ac..645491b05e 100644 --- a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs +++ b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs @@ -302,17 +302,46 @@ fn par_chunk_len(len: usize, threads: usize) -> Option { /// makes the split independent of the machine, so the chunk boundaries (and /// therefore the bits) do not move when the session is reconfigured. /// -/// [`PAR_MIN_CHUNK`] is still the floor, so a chunk is never too short to be -/// worth dispatching or to stay on the vector path. The cap keeps a very long -/// slice from turning into thousands of tiny tasks. +/// [`HOST_MIN_CHUNK`] is the floor, so a chunk is never too short to be worth +/// dispatching or to stay on the vector path; this cap is the other end, and +/// keeps a very long slice from turning into thousands of tiny tasks. Raising +/// it to 256 or 1024 measured within noise of 64 at every size swept, so 64 — +/// four tasks per thread on a sixteen-thread session, enough for the host to +/// balance with — is the one that dispatches least. const MAX_HOST_CHUNKS: usize = 64; -/// Chunk length and count for the host-pool split, or `None` to stay serial. +/// Length below which the host split is not worth dispatching. +/// +/// Sixteen times lower than [`PAR_MIN_LEN`], and for a concrete reason: that +/// constant pays for waking rayon's workers, and ORT's intra-op workers are +/// already awake — they spin. Measured at `intra_op = 16` against ORT's own +/// CPU EP, splitting a 256 Ki slice on the host pool instead of leaving it +/// whole moved `Gelu` from 0.15x to 0.94x, `FastGelu` 0.11x to 0.64x, `Erf` +/// 0.24x to 1.33x and `Sqrt` 0.19x to 0.91x. At 64 Ki the split still helps +/// (`Tanh` 0.40x to 1.36x, `Relu` 0.61x to 1.12x); at 16 Ki it stopped +/// mattering, every gate measuring within noise of every other, so this is the +/// last length where the evidence is one-sided. +const HOST_MIN_LEN: usize = 65_536; + +/// Shortest chunk worth handing to one of the host's threads. +/// +/// [`PAR_MIN_CHUNK`] is 256 Ki because a rayon wake-up costs ~50 us. A host +/// task costs a fraction of that, and the sweep is monotone in this direction: +/// at `intra_op = 16` over 64 Ki - 1 Mi, dropping the floor from 256 Ki to +/// 64 Ki to 16 Ki to 4 Ki improved almost every op at almost every size +/// (`Erf` at 256 Ki: 0.18x, 0.33x, 0.79x, 1.02x against ORT; `Gelu`: 0.16x, +/// 0.31x, 0.90x, 0.98x). 4 Ki is where it flattened. +const HOST_MIN_CHUNK: usize = 4_096; + +/// Chunk length and count for the host-pool split, or `None` to stay whole. #[inline] fn host_chunk_len(len: usize) -> Option<(usize, usize)> { + if len < HOST_MIN_LEN { + return None; + } let chunk = len .div_ceil(MAX_HOST_CHUNKS) - .max(PAR_MIN_CHUNK) + .max(HOST_MIN_CHUNK) .next_multiple_of(8); (chunk < len).then(|| (chunk, len.div_ceil(chunk))) } @@ -320,13 +349,13 @@ fn host_chunk_len(len: usize) -> Option<(usize, usize)> { /// [`host_chunk_len`] for the bias-fused kernels: whole multiples of `width`. #[inline] fn host_chunk_len_rows(len: usize, width: usize) -> Option<(usize, usize)> { - if width == 0 { + if width == 0 || len < HOST_MIN_LEN { return None; } let rows = len .div_ceil(width) .div_ceil(MAX_HOST_CHUNKS) - .max(PAR_MIN_CHUNK.div_ceil(width)); + .max(HOST_MIN_CHUNK.div_ceil(width)); let chunk = rows.checked_mul(width)?; (chunk < len).then(|| (chunk, len.div_ceil(chunk))) } @@ -454,7 +483,7 @@ where // a 4096-element decode call cost ~1.4 us, which is more than the whole // call. Measured as a uniform 0.6-0.8x regression on every short case // until this check moved above it. - if len < PAR_MIN_LEN || force_serial() { + if len < HOST_MIN_LEN || force_serial() { body(input, output); return; } @@ -464,15 +493,28 @@ where body(input, output); return; } - // Prefer the host's pool over ours whenever the host offered one. Ours - // would be a *second* pool on the same cores: measured 1.5-3.1x slower - // than staying serial at 1 Mi under an `intra_op = 16` session. - if let Some(host) = onnx_runtime_ep_api::host_parallel::current() { - if let Some((chunk, count)) = host_chunk_len(len) { - note_parallel_dispatch(); - run_on_host(host, input, output, chunk, count, &body); - return; + // Prefer the host's pool over ours whenever the host has one. Ours would + // be a *second* pool on the same cores: measured 1.5-3.1x slower than + // staying serial at 1 Mi under an `intra_op = 16` session. + // + // A host pool of *one* thread is the opposite case. It is not using the + // machine, so there is nothing to contend with and our own pool is the + // right tool: at `intra_op = 1`, rayon beat borrowing ORT's single thread + // by 2-9x over 1-4 Mi. `HostWidth` is observed rather than assumed, so + // each session gets the answer that is true for it. + if let Some(host) = onnx_runtime_ep_api::host_parallel::current() + && host.width() != onnx_runtime_ep_api::HostWidth::Serial + { + match host_chunk_len(len) { + Some((chunk, count)) => { + note_parallel_dispatch(); + run_on_host(host, input, output, chunk, count, &body); + } + None => body(input, output), } + return; + } + if len < PAR_MIN_LEN { body(input, output); return; } @@ -533,11 +575,13 @@ pub(crate) fn run_chunked_fn(input: &[f32], output: &mut [f32], body: fn(&[f32], pub(crate) fn clip_chunked(input: &[f32], output: &mut [f32], minimum: f32, maximum: f32) { // Take the serial decision here, so the common short case is a direct call // instead of one through a closure the optimiser can no longer see into. - if input.len() < PAR_MIN_LEN + if input.len() < HOST_MIN_LEN || force_serial() || onnx_runtime_ep_api::host_parallel::in_host_task() || rayon::current_thread_index().is_some() { + // Short, or already on a worker: one direct call, with no closure for + // the optimiser to lose sight of. mlas_sys::compute_clip(input, output, minimum, maximum); return; } @@ -564,7 +608,7 @@ where F: Fn(&[f32], &mut [f32]) + Sync + Send, { let len = input.len(); - if len < PAR_MIN_LEN || width == 0 || force_serial() { + if len < HOST_MIN_LEN || width == 0 || force_serial() { body(input, output); return; } @@ -572,11 +616,16 @@ where body(input, output); return; } - if let Some(host) = onnx_runtime_ep_api::host_parallel::current() { - if let Some((chunk, count)) = host_chunk_len_rows(len, width) { - run_on_host(host, input, output, chunk, count, &body); - return; + if let Some(host) = onnx_runtime_ep_api::host_parallel::current() + && host.width() != onnx_runtime_ep_api::HostWidth::Serial + { + match host_chunk_len_rows(len, width) { + Some((chunk, count)) => run_on_host(host, input, output, chunk, count, &body), + None => body(input, output), } + return; + } + if len < PAR_MIN_LEN { body(input, output); return; } @@ -4161,14 +4210,41 @@ mod host_pool_split { }); } + /// A host that has really been seen to use more than one thread. + static PARALLEL_WIDTH: std::sync::atomic::AtomicU8 = + std::sync::atomic::AtomicU8::new(host_parallel::WIDTH_PARALLEL); + + /// A host whose pool turned out to have no workers of its own. + static SERIAL_WIDTH: std::sync::atomic::AtomicU8 = + std::sync::atomic::AtomicU8::new(host_parallel::WIDTH_SERIAL); + fn fake_host() -> HostParallel { - // SAFETY: `threaded_host` never dereferences its `host` argument. - unsafe { HostParallel::new(core::ptr::null_mut(), threaded_host) } + // SAFETY: `threaded_host` never dereferences its `host` argument, and + // the width cell is a `static`. + unsafe { + HostParallel::new( + core::ptr::null_mut(), + threaded_host, + core::ptr::from_ref(&PARALLEL_WIDTH), + ) + } + } + + /// The same host, but known to have a single thread. + fn serial_host_handle() -> HostParallel { + // SAFETY: as `fake_host`. + unsafe { + HostParallel::new( + core::ptr::null_mut(), + threaded_host, + core::ptr::from_ref(&SERIAL_WIDTH), + ) + } } /// Long enough to be split, and deliberately not a multiple of the chunk /// size, so the final chunk is short. - const N: usize = 3 * PAR_MIN_LEN + 37; + const N: usize = 5 * HOST_MIN_LEN + 37; fn probe(len: usize) -> Vec { (0..len) @@ -4269,6 +4345,53 @@ mod host_pool_split { }); } + /// Below the gate the slice stays whole even with a host installed: the + /// dispatch would cost more than the work it hands out. + #[test] + fn a_short_slice_is_not_dispatched() { + let n = HOST_MIN_LEN - 8; + let x = probe(n); + let mut y = vec![0.0f32; n]; + reset_dispatched(); + host_parallel::scope(fake_host(), || tanh_f32_slice(&x, &mut y)); + assert_eq!(dispatched(), 0, "{n} elements should have stayed whole"); + + let mut want = vec![0.0f32; n]; + serial_scope(|| tanh_f32_slice(&x, &mut want)); + assert!(want.iter().zip(&y).all(|(a, b)| a.to_bits() == b.to_bits())); + } + + /// A host pool with no workers is not using the machine, so ours may. + /// Borrowing its single thread instead measured 2-9x slower over 1-4 Mi at + /// `intra_op = 1`, so this is worth an assertion. + #[test] + fn a_serial_host_is_not_borrowed() { + if rayon::current_num_threads() < 2 { + eprintln!("skipped: single-threaded rayon pool cannot show the difference"); + return; + } + let n = 2 * PAR_MIN_LEN; + let x = probe(n); + let mut y = vec![0.0f32; n]; + reset_dispatched(); + let before = parallel_dispatches(); + host_parallel::scope(serial_host_handle(), || tanh_f32_slice(&x, &mut y)); + assert_eq!( + dispatched(), + 0, + "a one-thread host pool was borrowed anyway" + ); + assert_eq!( + parallel_dispatches() - before, + 1, + "the work should have gone to our own pool instead" + ); + + let mut want = vec![0.0f32; n]; + serial_scope(|| tanh_f32_slice(&x, &mut want)); + assert!(want.iter().zip(&y).all(|(a, b)| a.to_bits() == b.to_bits())); + } + /// `serial_scope` has to keep suppressing the split on the host path too, /// or the f16/bf16 sandwich picks its measured regression back up. #[test] @@ -4292,12 +4415,14 @@ mod host_pool_split { 0, 1, 8, - PAR_MIN_CHUNK - 1, + HOST_MIN_CHUNK - 1, + HOST_MIN_CHUNK, + HOST_MIN_CHUNK + 1, + HOST_MIN_LEN - 1, + HOST_MIN_LEN, + HOST_MIN_LEN + 1, PAR_MIN_CHUNK, - PAR_MIN_CHUNK + 1, - PAR_MIN_LEN - 1, PAR_MIN_LEN, - PAR_MIN_LEN + 1, N, 1 << 26, (1 << 26) + 13, @@ -4306,7 +4431,8 @@ mod host_pool_split { if let Some((chunk, count)) = host_chunk_len(len) { assert!(chunk < len, "chunk {chunk} !< len {len}"); assert_eq!(chunk % 8, 0, "chunk {chunk} is not a whole vector"); - assert!(chunk >= PAR_MIN_CHUNK, "chunk {chunk} is too short"); + assert!(len >= HOST_MIN_LEN, "split below the gate"); + assert!(chunk >= HOST_MIN_CHUNK, "chunk {chunk} is too short"); assert!(chunk >= SIMD_MIN_LEN, "chunk {chunk} would go scalar"); assert_eq!(count, len.div_ceil(chunk)); assert!( @@ -4319,12 +4445,19 @@ mod host_pool_split { if let Some((chunk, count)) = host_chunk_len_rows(len, width) { assert!(chunk < len, "rows: chunk {chunk} !< len {len}"); assert_eq!(chunk % width, 0, "rows: chunk {chunk} cuts width {width}"); - assert!(chunk >= PAR_MIN_CHUNK.min(len), "rows: chunk {chunk} short"); + assert!(len >= HOST_MIN_LEN, "rows: split below the gate"); + assert!( + chunk >= HOST_MIN_CHUNK.min(len), + "rows: chunk {chunk} short" + ); assert_eq!(count, len.div_ceil(chunk)); assert!((count - 1) * chunk < len, "rows: empty final range"); } } assert_eq!(host_chunk_len_rows(len, 0), None, "width 0 must not split"); + if len < HOST_MIN_LEN { + assert_eq!(host_chunk_len(len), None, "{len} is below the gate"); + } } } } diff --git a/crates/onnx-runtime-ep-plugin/src/compute.rs b/crates/onnx-runtime-ep-plugin/src/compute.rs index 72acc2f78a..195473631c 100644 --- a/crates/onnx-runtime-ep-plugin/src/compute.rs +++ b/crates/onnx-runtime-ep-plugin/src/compute.rs @@ -1036,6 +1036,14 @@ pub struct ExportedComputeInfo { /// CPU EP (and any host EP), which uses its inputs verbatim exactly as /// before. device_staging: Option, + /// How wide ORT's intra-op pool turned out to be for this session. + /// + /// Lives here, rather than in a process-global, because a process may hold + /// one session at `intra_op = 1` and another at `intra_op = 16`, and the + /// right answer -- whether to borrow ORT's threads or use our own -- is + /// the opposite for each. Written once, by the first dispatch that has + /// more than one index; see `host_pool`. + host_pool_width: std::sync::atomic::AtomicU8, } /// Everything `Compute` needs to stage host-resident boundary inputs onto the @@ -1093,6 +1101,9 @@ impl ExportedComputeInfo { routing: None, workspace_plans, device_staging: None, + host_pool_width: std::sync::atomic::AtomicU8::new( + onnx_runtime_ep_api::host_parallel::WIDTH_UNKNOWN, + ), } } @@ -2083,7 +2094,9 @@ unsafe extern "C" fn compute_execute( // // SAFETY: `kernel_context` is the context ORT handed this call and // stays valid until it returns, which is after the guard is dropped. - let _host_pool = unsafe { crate::host_pool::install(api_ref, kernel_context) }; + let _host_pool = unsafe { + crate::host_pool::install(api_ref, kernel_context, &exported.host_pool_width) + }; // Memory info for intermediate scratch. On a device EP this is device // memory, so multi-node intermediates stay on the GPU (a host buffer diff --git a/crates/onnx-runtime-ep-plugin/src/host_pool.rs b/crates/onnx-runtime-ep-plugin/src/host_pool.rs index 1f4a60ecc2..46a32562a1 100644 --- a/crates/onnx-runtime-ep-plugin/src/host_pool.rs +++ b/crates/onnx-runtime-ep-plugin/src/host_pool.rs @@ -23,8 +23,9 @@ use core::ffi::c_void; use onnx_genai_ort_sys as ort; use onnx_runtime_ep_api::HostParallel; use onnx_runtime_ep_api::host_parallel; +use onnx_runtime_ep_api::host_parallel::{WIDTH_PARALLEL, WIDTH_SERIAL, WIDTH_UNKNOWN}; use std::panic::AssertUnwindSafe; -use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU8, Ordering}; /// What [`ort_parallel_for`] needs to reach ORT, behind one `void*`. struct HostPool { @@ -37,12 +38,20 @@ struct HostPool { ) -> *mut ort::OrtStatus, release_status: Option, ctx: *mut ort::OrtKernelContext, + /// Where this session records how wide ORT's pool turned out to be. + width: *const AtomicU8, } -/// The closure a dispatch is running, plus somewhere to record a panic. +/// The closure a dispatch is running, plus somewhere to record what happened. struct Task<'a> { body: &'a (dyn Fn(usize) + Sync), panicked: AtomicBool, + /// The thread that called `KernelContext_ParallelFor`, while we still do + /// not know whether ORT has any workers. `None` once the answer is known, + /// which turns the observation below into a single relaxed load. + caller: Option, + /// Set when a body ran somewhere other than `caller`. + saw_another_thread: AtomicBool, } /// Trampoline handed to ORT: one index of one dispatch. @@ -58,6 +67,11 @@ unsafe extern "C" fn run_index(usr_data: *mut c_void, index: usize) { // it, and let `ort_parallel_for` re-raise on the calling thread once every // worker has finished. let task = unsafe { &*(usr_data.cast::>()) }; + if let Some(caller) = task.caller + && std::thread::current().id() != caller + { + task.saw_another_thread.store(true, Ordering::Relaxed); + } if std::panic::catch_unwind(AssertUnwindSafe(|| (task.body)(index))).is_err() { task.panicked.store(true, Ordering::Relaxed); } @@ -71,9 +85,15 @@ unsafe extern "C" fn run_index(usr_data: *mut c_void, index: usize) { /// this must be reached from inside the compute call that installed it. unsafe fn ort_parallel_for(host: *mut c_void, total: usize, body: &(dyn Fn(usize) + Sync)) { let pool = unsafe { &*(host.cast::()) }; + // Only look at thread identities until the answer is in. A session runs + // this on every dispatch, and `thread::current()` is not free. + let observing = total > 1 + && unsafe { pool.width() }.is_some_and(|w| w.load(Ordering::Relaxed) == WIDTH_UNKNOWN); let task = Task { body, panicked: AtomicBool::new(false), + caller: observing.then(|| std::thread::current().id()), + saw_another_thread: AtomicBool::new(false), }; // `num_batch = 0` means "no limit": ORT gives every index its own task and // its workers claim them dynamically. That is what we want, because we cut @@ -89,6 +109,20 @@ unsafe fn ort_parallel_for(host: *mut c_void, total: usize, body: &(dyn Fn(usize (&raw const task).cast::().cast_mut(), ) }; + if observing && status.is_null() { + // ORT ran every index; if none of them landed on another thread, its + // intra-op pool has no workers. Recorded per session, so a process + // with both a one-thread and a sixteen-thread session gets the right + // answer for each. + let observed = if task.saw_another_thread.load(Ordering::Relaxed) { + WIDTH_PARALLEL + } else { + WIDTH_SERIAL + }; + if let Some(width) = unsafe { pool.width() } { + width.store(observed, Ordering::Relaxed); + } + } if !status.is_null() { // The dispatch itself failed. Every index still has to run or the // output tensor keeps whatever was in the buffer, so fall back to @@ -107,6 +141,18 @@ unsafe fn ort_parallel_for(host: *mut c_void, total: usize, body: &(dyn Fn(usize ); } +impl HostPool { + /// The session's width cell, if it has one. + /// + /// # Safety + /// + /// The pointer must still be valid, which `install`'s contract requires + /// for as long as the guard is alive. + unsafe fn width(&self) -> Option<&AtomicU8> { + (!self.width.is_null()).then(|| unsafe { &*self.width }) + } +} + /// ORT's pool, installed on this thread until the guard is dropped. /// /// Field order is the safety argument: `installed` is dropped first, so the @@ -155,7 +201,11 @@ impl Drop for Guard { /// thread-local and the guard is confined to the compute call's frame, it /// cannot be reached from a later call whose context has been freed. #[must_use = "dropping the guard immediately uninstalls the pool"] -pub unsafe fn install(api: &ort::OrtApi, ctx: *mut ort::OrtKernelContext) -> Guard { +pub unsafe fn install( + api: &ort::OrtApi, + ctx: *mut ort::OrtKernelContext, + width: &AtomicU8, +) -> Guard { let (Some(parallel_for), false) = (api.KernelContext_ParallelFor, ctx.is_null()) else { return Guard::inert(); }; @@ -163,6 +213,7 @@ pub unsafe fn install(api: &ort::OrtApi, ctx: *mut ort::OrtKernelContext) -> Gua parallel_for, release_status: api.ReleaseStatus, ctx, + width: core::ptr::from_ref(width), }); // SAFETY: the box outlives the handle (see `Guard`'s drop order), and // `ort_parallel_for` only ever reads the pointer back as a `*mut HostPool`. @@ -170,6 +221,7 @@ pub unsafe fn install(api: &ort::OrtApi, ctx: *mut ort::OrtKernelContext) -> Gua HostParallel::new( (&raw const *pool).cast::().cast_mut(), ort_parallel_for, + core::ptr::from_ref(width), ) }; Guard { @@ -218,6 +270,28 @@ mod tests { core::ptr::dangling_mut() } + /// Stands in for an ORT whose intra-op pool has real workers. + /// + /// # Safety + /// + /// Matches the ABI ORT expects; `usr_data` is passed straight through. + unsafe extern "C" fn threaded_parallel_for( + _ctx: *const ort::OrtKernelContext, + body: Option, + total: usize, + _num_batch: usize, + usr_data: *mut c_void, + ) -> *mut ort::OrtStatus { + let body = body.expect("ORT is always given a callback"); + let usr = usr_data as usize; + std::thread::scope(|scope| { + for index in 0..total { + scope.spawn(move || unsafe { body(usr as *mut c_void, index) }); + } + }); + core::ptr::null_mut() + } + fn pool( parallel_for: unsafe extern "C" fn( *const ort::OrtKernelContext, @@ -226,17 +300,20 @@ mod tests { usize, *mut c_void, ) -> *mut ort::OrtStatus, + width: &AtomicU8, ) -> HostPool { HostPool { parallel_for, release_status: None, ctx: core::ptr::null_mut(), + width: core::ptr::from_ref(width), } } #[test] fn every_index_runs_once() { - let mut pool = pool(inline_parallel_for); + let width = AtomicU8::new(WIDTH_UNKNOWN); + let mut pool = pool(inline_parallel_for, &width); let seen: Vec = (0..5).map(|_| AtomicUsize::new(0)).collect(); unsafe { ort_parallel_for((&raw mut pool).cast::(), seen.len(), &|index| { @@ -246,9 +323,57 @@ mod tests { assert!(seen.iter().all(|c| c.load(Ordering::Relaxed) == 1)); } + /// An ORT that runs everything inline has no workers, and our own pool is + /// then free to use the machine. Getting this wrong in either direction + /// costs multiples: our pool alongside ORT's is 1.5-3.1x slower than + /// serial, and serial where ORT has no workers is 2-9x slower than ours. + #[test] + fn an_inline_host_is_observed_as_serial() { + let width = AtomicU8::new(WIDTH_UNKNOWN); + let mut pool = pool(inline_parallel_for, &width); + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + assert_eq!(width.load(Ordering::Relaxed), WIDTH_SERIAL); + } + + #[test] + fn a_threaded_host_is_observed_as_parallel() { + let width = AtomicU8::new(WIDTH_UNKNOWN); + let mut pool = pool(threaded_parallel_for, &width); + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + assert_eq!(width.load(Ordering::Relaxed), WIDTH_PARALLEL); + } + + /// One index proves nothing: even a sixteen-thread pool runs it inline. + #[test] + fn a_single_index_dispatch_does_not_decide_the_width() { + let width = AtomicU8::new(WIDTH_UNKNOWN); + let mut pool = pool(inline_parallel_for, &width); + unsafe { ort_parallel_for((&raw mut pool).cast::(), 1, &|_| {}) }; + assert_eq!(width.load(Ordering::Relaxed), WIDTH_UNKNOWN); + } + + /// A refused dispatch says nothing about the pool either. + #[test] + fn a_refused_dispatch_does_not_decide_the_width() { + let width = AtomicU8::new(WIDTH_UNKNOWN); + let mut pool = pool(failing_parallel_for, &width); + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + assert_eq!(width.load(Ordering::Relaxed), WIDTH_UNKNOWN); + } + + /// Once decided, later dispatches stop paying for the observation. + #[test] + fn an_answered_width_is_not_revisited() { + let width = AtomicU8::new(WIDTH_PARALLEL); + let mut pool = pool(inline_parallel_for, &width); + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + assert_eq!(width.load(Ordering::Relaxed), WIDTH_PARALLEL); + } + #[test] fn a_refused_dispatch_still_runs_every_index() { - let mut pool = pool(failing_parallel_for); + let width = AtomicU8::new(WIDTH_UNKNOWN); + let mut pool = pool(failing_parallel_for, &width); let seen: Vec = (0..5).map(|_| AtomicUsize::new(0)).collect(); unsafe { ort_parallel_for((&raw mut pool).cast::(), seen.len(), &|index| { @@ -263,7 +388,8 @@ mod tests { #[test] fn a_panicking_body_does_not_unwind_into_ort() { - let mut pool = pool(inline_parallel_for); + let width = AtomicU8::new(WIDTH_UNKNOWN); + let mut pool = pool(inline_parallel_for, &width); let ran = AtomicUsize::new(0); let unwound = std::panic::catch_unwind(AssertUnwindSafe(|| unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|index| { @@ -281,7 +407,8 @@ mod tests { fn an_ort_without_parallel_for_installs_nothing() { let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; api.KernelContext_ParallelFor = None; - let guard = unsafe { install(&api, core::ptr::dangling_mut()) }; + let width = AtomicU8::new(WIDTH_UNKNOWN); + let guard = unsafe { install(&api, core::ptr::dangling_mut(), &width) }; assert!(!guard.is_installed()); assert!(host_parallel::current().is_none()); } @@ -290,7 +417,8 @@ mod tests { fn a_null_context_installs_nothing() { let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; api.KernelContext_ParallelFor = Some(inline_parallel_for); - let guard = unsafe { install(&api, core::ptr::null_mut()) }; + let width = AtomicU8::new(WIDTH_UNKNOWN); + let guard = unsafe { install(&api, core::ptr::null_mut(), &width) }; assert!(!guard.is_installed()); assert!(host_parallel::current().is_none()); } @@ -301,7 +429,8 @@ mod tests { api.KernelContext_ParallelFor = Some(inline_parallel_for); let seen: Vec = (0..6).map(|_| AtomicUsize::new(0)).collect(); { - let guard = unsafe { install(&api, core::ptr::dangling_mut()) }; + let width = AtomicU8::new(WIDTH_UNKNOWN); + let guard = unsafe { install(&api, core::ptr::dangling_mut(), &width) }; assert!(guard.is_installed()); let host = host_parallel::current().expect("install publishes a handle"); host.run(seen.len(), &|index| { @@ -319,10 +448,11 @@ mod tests { fn a_nested_install_restores_the_outer_handle() { let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; api.KernelContext_ParallelFor = Some(inline_parallel_for); - let outer = unsafe { install(&api, core::ptr::dangling_mut()) }; + let width = AtomicU8::new(WIDTH_UNKNOWN); + let outer = unsafe { install(&api, core::ptr::dangling_mut(), &width) }; assert!(outer.is_installed()); { - let _inner = unsafe { install(&api, core::ptr::dangling_mut()) }; + let _inner = unsafe { install(&api, core::ptr::dangling_mut(), &width) }; assert!(host_parallel::current().is_some()); } assert!( From 6783f6196a9b1249969d33a6dc79c9ff52710a6b Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Mon, 17 Aug 2026 23:00:37 +0000 Subject: [PATCH 3/5] perf(cpu-ep): decide by evidence which pool to split across The previous commit inferred the host pool's width from one dispatch: if every body ran on the calling thread, ORT had no workers. Measurement says that inference is unsound. On a 16-thread session 35% of dispatches ran entirely on the calling thread -- in runs of up to 60 -- because ORT hands its indices out dynamically and an unstalled caller drains them before a worker wakes. Acting on that would start our pool alongside ORT's sixteen, the 3-10x pathology this seam exists to remove. So stop inferring and require positive evidence: a body seen running on a thread that was not the one that dispatched it. Only a pool with workers can produce that, so the verdict is permanent and cannot be faked by a serial session. Until it arrives the kernels stay on their own pool -- exactly what they did before this seam existed -- except on probe dispatches, which hold the caller's first index open for 100 us so a worker that exists has time to claim another. Probes run back to back for the first 32 dispatches of a session, then back off geometrically to one in a thousand, so a session whose pool really is serial pays almost nothing and still recovers if the opening burst was unlucky. Measured at 16 threads (ORT/ours, 4 interleaved rounds, taskset 0-15): every session now latches, and 1 Mi goes from 0.07-0.14 to 0.44-1.35. At intra_op=1 the branch matches main within noise (0.45-2.79 vs 0.46-2.99), i.e. the rayon split is kept exactly where it was winning. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../onnx-runtime-ep-api/src/host_parallel.rs | 243 ++++++++++++---- crates/onnx-runtime-ep-api/src/lib.rs | 2 +- .../src/kernels/matmul_nbits.rs | 24 +- .../src/kernels/simd_activations.rs | 50 ++-- crates/onnx-runtime-ep-plugin/src/compute.rs | 14 +- .../onnx-runtime-ep-plugin/src/host_pool.rs | 275 +++++++++++++----- 6 files changed, 439 insertions(+), 169 deletions(-) diff --git a/crates/onnx-runtime-ep-api/src/host_parallel.rs b/crates/onnx-runtime-ep-api/src/host_parallel.rs index 795d8c05ed..e9af3d3a21 100644 --- a/crates/onnx-runtime-ep-api/src/host_parallel.rs +++ b/crates/onnx-runtime-ep-api/src/host_parallel.rs @@ -56,31 +56,34 @@ //! host's pool and stay serial instead of nesting. use core::ffi::c_void; -use core::sync::atomic::{AtomicU8, Ordering}; +use core::sync::atomic::{AtomicU32, Ordering}; -/// How many threads the host's pool turned out to have. +/// Sentinel for "the host's pool has been seen doing our work". /// -/// The host runtime does not tell us, and it matters: a pool of one is not -/// using the machine, so ours may. Rather than probe — which would cost a -/// dispatch per session to answer a question the next real dispatch answers -/// for free — the implementation *observes* its own first dispatch and records -/// the result here. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum HostWidth { - /// No dispatch has completed yet. - Unknown, - /// Every body ran on the dispatching thread: the host has no workers. - Serial, - /// At least one body ran on another thread. - Parallel, -} +/// The only *permanent* state, and it takes positive evidence to reach: a body +/// ran on a thread that was not the one that dispatched it, which a pool with +/// no workers can never produce. Nothing ever clears it. +pub const HOST_HELPED: u32 = u32::MAX; -/// Nothing observed yet. -pub const WIDTH_UNKNOWN: u8 = 0; -/// The host pool ran everything inline. -pub const WIDTH_SERIAL: u8 = 1; -/// The host pool used at least one other thread. -pub const WIDTH_PARALLEL: u8 = 2; +/// Dispatches a session probes back to back before it starts backing off. +/// +/// A probe is deliberately made hard to fail — the implementation holds the +/// calling thread's first index open so a worker that exists has time to claim +/// another — but even so, a 16-thread ORT session answers only about half of +/// them, because it hands its indices out dynamically and a caller that is not +/// stalled drains them first. A burst of 32 turns that coin-flip into a +/// certainty; getting it wrong the other way is the 3-10x mistake. +pub const PROBE_BURST: u32 = 32; + +/// Gap between probes when the burst has produced nothing, before backoff. +pub const PROBE_MIN: u32 = 16; + +/// Longest gap between probes once the host has never been seen to help. +/// +/// Bounds the steady-state cost of asking on a session whose pool really is +/// serial: one dispatch in a thousand runs on the host's single thread instead +/// of ours. +pub const PROBE_MAX: u32 = 1024; /// Runs `body(i)` for every `i` in `0..total` on the host's threads. /// @@ -102,9 +105,9 @@ pub type HostParallelForFn = pub struct HostParallel { host: *mut c_void, run: HostParallelForFn, - /// Where the implementation records [`HostWidth`], or null if it never - /// will — in which case the pool is assumed parallel. - width: *const AtomicU8, + /// Where this session records whether the host's pool has ever helped, + /// and when to ask again. Null means "never ask, always use the host". + probe: *const AtomicU32, } impl core::fmt::Debug for HostParallel { @@ -125,35 +128,90 @@ impl HostParallel { /// pointer. Installing it with [`scope`] is what bounds the first half; /// the second half is on the implementer. /// - /// `width` must be null or point at a cell that outlives this handle, and - /// that only ever holds [`WIDTH_UNKNOWN`], [`WIDTH_SERIAL`] or - /// [`WIDTH_PARALLEL`]. + /// `probe` must be null or point at a cell that outlives this handle and + /// is shared by every handle for the same host pool. Start it at zero; + /// [`HOST_HELPED`] is the implementation's way of saying the pool has been + /// seen running our bodies on its own threads. This crate owns every other + /// value in it. pub const unsafe fn new( host: *mut c_void, run: HostParallelForFn, - width: *const AtomicU8, + probe: *const AtomicU32, ) -> Self { - Self { host, run, width } + Self { host, run, probe } } - /// What the last completed dispatch said about the host pool's width. + /// Whether this dispatch should go to the host's pool rather than ours. /// - /// [`HostWidth::Serial`] is the interesting answer: it means the host's - /// pool has no workers of its own, so borrowing it would serialise us for - /// nothing and our own pool is free to use the machine instead. - #[must_use] - pub fn width(&self) -> HostWidth { - if self.width.is_null() { - return HostWidth::Parallel; + /// Answering it needs a fact the host runtime will not tell us: whether + /// its pool has any workers. A session built with `intra_op = 1` has none, + /// and borrowing its single thread would serialise work our own pool could + /// spread over the machine — measured 2-9x slower over 1-4 Mi. A session + /// with a wide pool is the exact opposite: ours would be a *second* pool + /// on the same cores, and that measured 3-10x slower than borrowing. + /// + /// So ask the only question that has a trustworthy answer: *has the host's + /// pool ever actually run one of our bodies on another thread?* Only a + /// pool with workers can, so a "yes" is permanent and cannot be faked. + /// Until then this returns `false` — keeping the caller on its own pool, + /// which is what it did before this seam existed — except on the + /// occasional probe dispatch that gives the host a chance to answer. + /// + /// The first [`PROBE_BURST`] dispatches all ask, because at 16 threads a + /// single dispatch answering "nobody helped" is uninformative — it happens + /// 35% of the time — while a burst that long is answered almost surely. + /// After that, probes back off geometrically from [`PROBE_MIN`] to + /// [`PROBE_MAX`], so a session that really is serial settles at one probe + /// per thousand dispatches and still recovers if the burst was unlucky. + /// + /// Mutates the cell, so call it once per dispatch decision. + pub fn prefer_host(&self) -> bool { + let Some(cell) = self.probe_cell() else { + return true; + }; + let mut seen = cell.load(Ordering::Relaxed); + loop { + if seen == HOST_HELPED { + return true; + } + let period = seen >> 16; + let countdown = seen & 0xFFFF; + let (next, probe) = if period == 0 { + // Still in the opening burst: ask every time, and count how + // many times asking has told us nothing. + let asked = countdown + 1; + let next = if asked >= PROBE_BURST { + (PROBE_MIN << 16) | PROBE_MIN + } else { + asked + }; + (next, true) + } else if countdown == 0 { + // Probe now, and put the next one twice as far out. + let period = period.saturating_mul(2).min(PROBE_MAX); + ((period << 16) | period, true) + } else { + ((period << 16) | (countdown - 1), false) + }; + match cell.compare_exchange_weak(seen, next, Ordering::Relaxed, Ordering::Relaxed) { + Ok(_) => return probe, + Err(current) => seen = current, + } } + } + + /// Whether the host's pool has been seen running our work. + #[must_use] + pub fn helped(&self) -> bool { + self.probe_cell() + .is_none_or(|cell| cell.load(Ordering::Relaxed) == HOST_HELPED) + } + + fn probe_cell(&self) -> Option<&AtomicU32> { // SAFETY: `new`'s contract puts the validity of this pointer on the // installer, and `scope` bounds it to the extent the handle is // reachable. - match unsafe { &*self.width }.load(Ordering::Relaxed) { - WIDTH_SERIAL => HostWidth::Serial, - WIDTH_PARALLEL => HostWidth::Parallel, - _ => HostWidth::Unknown, - } + (!self.probe.is_null()).then(|| unsafe { &*self.probe }) } /// Runs `body(0..total)` on the host's pool and waits for all of it. @@ -299,7 +357,7 @@ mod tests { unsafe { HostParallel::new(core::ptr::null_mut(), serial_host, core::ptr::null()) } } - fn with_width(cell: &AtomicU8) -> HostParallel { + fn with_probe(cell: &AtomicU32) -> HostParallel { // SAFETY: as `serial`, plus `cell` outlives every use below. unsafe { HostParallel::new( @@ -311,25 +369,94 @@ mod tests { } #[test] - fn a_handle_without_a_width_cell_reads_as_parallel() { - assert_eq!(serial().width(), HostWidth::Parallel); + fn a_handle_without_a_probe_cell_always_uses_the_host() { + assert!(serial().prefer_host()); + assert!(serial().helped()); } + /// The opening burst asks every time: one silent dispatch is not evidence + /// of a serial pool, a run of 64 is. #[test] - fn width_reflects_the_cell() { - let cell = AtomicU8::new(WIDTH_UNKNOWN); - let host = with_width(&cell); - assert_eq!(host.width(), HostWidth::Unknown); - cell.store(WIDTH_SERIAL, Ordering::Relaxed); - assert_eq!(host.width(), HostWidth::Serial); - cell.store(WIDTH_PARALLEL, Ordering::Relaxed); - assert_eq!(host.width(), HostWidth::Parallel); - cell.store(200, Ordering::Relaxed); + fn the_opening_burst_probes_every_dispatch() { + let cell = AtomicU32::new(0); + let host = with_probe(&cell); + for step in 0..PROBE_BURST { + assert!(host.prefer_host(), "dispatch {step} of the burst"); + } + assert!(!host.prefer_host(), "the burst has to end somewhere"); + } + + /// Until the host has been seen helping, work stays on our pool -- which + /// is what the caller did before this seam existed. + #[test] + fn an_unhelpful_host_is_asked_less_and_less_often() { + let cell = AtomicU32::new(0); + let host = with_probe(&cell); + let mut gaps = Vec::new(); + let mut gap = 0u32; + for _ in 0..8400 { + if host.prefer_host() { + gaps.push(gap); + gap = 0; + } else { + gap += 1; + } + } + let burst = usize::try_from(PROBE_BURST).unwrap(); + assert!( + gaps[..burst].iter().all(|&g| g == 0), + "the opening burst asks on every dispatch" + ); assert_eq!( - host.width(), - HostWidth::Unknown, - "an unknown code is not a promise" + &gaps[burst..burst + 4], + &[PROBE_MIN, PROBE_MIN * 2, PROBE_MIN * 4, PROBE_MIN * 8], + "probes should back off geometrically once the burst is over" ); + assert!( + gaps.iter().all(|&g| g <= PROBE_MAX), + "the gap must stay bounded so a wrong guess still self-corrects" + ); + assert_eq!( + *gaps.last().unwrap(), + PROBE_MAX, + "and it should settle at the cap" + ); + } + + /// One sighting of a worker thread is permanent: only a pool with workers + /// can produce it, so no later evidence can argue with it. + #[test] + fn a_helping_host_is_used_from_then_on() { + let cell = AtomicU32::new(0); + let host = with_probe(&cell); + assert!(!host.helped()); + cell.store(HOST_HELPED, Ordering::Relaxed); + assert!(host.helped()); + for _ in 0..1000 { + assert!(host.prefer_host()); + } + assert_eq!(cell.load(Ordering::Relaxed), HOST_HELPED); + } + + /// Two threads running the same session must not lose or double-count a + /// probe, and must never corrupt the cell into the sentinel. + #[test] + fn concurrent_dispatches_keep_the_cell_sane() { + let cell = AtomicU32::new(0); + std::thread::scope(|scope| { + for _ in 0..4 { + scope.spawn(|| { + let host = with_probe(&cell); + for _ in 0..5000 { + host.prefer_host(); + } + }); + } + }); + let seen = cell.load(Ordering::Relaxed); + assert_ne!(seen, HOST_HELPED, "no thread may invent the sentinel"); + assert!((seen >> 16) <= PROBE_MAX); + assert!((seen & 0xFFFF) <= PROBE_MAX); } #[test] diff --git a/crates/onnx-runtime-ep-api/src/lib.rs b/crates/onnx-runtime-ep-api/src/lib.rs index a51a84ffd1..1248e3027e 100644 --- a/crates/onnx-runtime-ep-api/src/lib.rs +++ b/crates/onnx-runtime-ep-api/src/lib.rs @@ -45,7 +45,7 @@ pub mod weight; pub use abi::{LegacyOrtEp, PluginCompiledKernel, PluginExecutionPlan, SubgraphClaim}; pub use epcontext::{EpContext, EpContextRegistry, build_ep_context_registry}; pub use error::{EpError, Result}; -pub use host_parallel::{HostParallel, HostWidth}; +pub use host_parallel::HostParallel; pub use kernel::{ ARG_BYTES, ARG_DEVICE, ARG_FLOPS, ARG_KERNEL_VARIANT, ARG_KERNEL_VARIANT_REASON, CAT_KERNEL_WORKER, CaptureSupport, ClaimPreference, Cost, Kernel, KernelInput, KernelMatch, diff --git a/crates/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs b/crates/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs index afc1287e87..97c7806285 100644 --- a/crates/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs +++ b/crates/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs @@ -16526,7 +16526,7 @@ mod tests { let pb = &packed_bytes[start * k_blocks * blob..(start + len) * k_blocks * blob]; let sc = &scales[start * k_blocks..(start + len) * k_blocks]; let pack = mlas_sys::SQNBitPackedB::new(len, k, 4, block_size, comp, pb, sc, None) - .expect("MLAS SQNBit int4 shard must pack"); + .expect("MLAS SQNBit int4 shard must pack"); (start, len, pack) }) .collect(); @@ -16579,9 +16579,7 @@ mod tests { .max() .unwrap_or(0); let full_identical = bits(&o_serial) == bits(&o_full); - eprintln!( - "#1138 full_width vs sharded: byte_identical={full_identical} max_ulp={max_ulp}" - ); + eprintln!("#1138 full_width vs sharded: byte_identical={full_identical} max_ulp={max_ulp}"); let pool = rayon::ThreadPoolBuilder::new() .num_threads(hw) @@ -16599,21 +16597,21 @@ mod tests { pool.install(|| { let mut out = vec![0.0f32; n]; for _ in 0..50 { - run(&mut out); + run(&mut out); } let mut samples = [0.0f64; 5]; for s in samples.iter_mut() { - let iters = 500u32; - let start = Instant::now(); - for _ in 0..iters { - run(&mut out); - } - *s = start.elapsed().as_secs_f64() * 1e6 / iters as f64; + let iters = 500u32; + let start = Instant::now(); + for _ in 0..iters { + run(&mut out); + } + *s = start.elapsed().as_secs_f64() * 1e6 / iters as f64; } samples.sort_by(|a, b| a.partial_cmp(b).unwrap()); eprintln!( - "#1138 int4 M=1 K={k} N={n} shards={shards} hw={hw}t {label:16}: p50={:.3}ms", - samples[2] / 1000.0, + "#1138 int4 M=1 K={k} N={n} shards={shards} hw={hw}t {label:16}: p50={:.3}ms", + samples[2] / 1000.0, ); }); } diff --git a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs index 645491b05e..6ea3f07d64 100644 --- a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs +++ b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs @@ -500,10 +500,10 @@ where // A host pool of *one* thread is the opposite case. It is not using the // machine, so there is nothing to contend with and our own pool is the // right tool: at `intra_op = 1`, rayon beat borrowing ORT's single thread - // by 2-9x over 1-4 Mi. `HostWidth` is observed rather than assumed, so - // each session gets the answer that is true for it. + // by 2-9x over 1-4 Mi. `prefer_host` decides which of the two a session is + // by watching whether the host's pool ever actually runs our chunks. if let Some(host) = onnx_runtime_ep_api::host_parallel::current() - && host.width() != onnx_runtime_ep_api::HostWidth::Serial + && host.prefer_host() { match host_chunk_len(len) { Some((chunk, count)) => { @@ -617,7 +617,7 @@ where return; } if let Some(host) = onnx_runtime_ep_api::host_parallel::current() - && host.width() != onnx_runtime_ep_api::HostWidth::Serial + && host.prefer_host() { match host_chunk_len_rows(len, width) { Some((chunk, count)) => run_on_host(host, input, output, chunk, count, &body), @@ -4210,34 +4210,36 @@ mod host_pool_split { }); } - /// A host that has really been seen to use more than one thread. - static PARALLEL_WIDTH: std::sync::atomic::AtomicU8 = - std::sync::atomic::AtomicU8::new(host_parallel::WIDTH_PARALLEL); + /// A host pool that has been seen running our chunks on its own threads. + static HELPED: std::sync::atomic::AtomicU32 = + std::sync::atomic::AtomicU32::new(host_parallel::HOST_HELPED); - /// A host whose pool turned out to have no workers of its own. - static SERIAL_WIDTH: std::sync::atomic::AtomicU8 = - std::sync::atomic::AtomicU8::new(host_parallel::WIDTH_SERIAL); + /// A host pool that never has, past its opening burst of probes: the + /// kernels should keep the work on their own pool. + static NEVER_HELPED: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new( + (host_parallel::PROBE_MIN << 16) | host_parallel::PROBE_MIN, + ); fn fake_host() -> HostParallel { // SAFETY: `threaded_host` never dereferences its `host` argument, and - // the width cell is a `static`. + // the probe cell is a `static`. unsafe { HostParallel::new( core::ptr::null_mut(), threaded_host, - core::ptr::from_ref(&PARALLEL_WIDTH), + core::ptr::from_ref(&HELPED), ) } } - /// The same host, but known to have a single thread. - fn serial_host_handle() -> HostParallel { + /// The same host, but one that has never been seen to help. + fn unhelpful_host() -> HostParallel { // SAFETY: as `fake_host`. unsafe { HostParallel::new( core::ptr::null_mut(), threaded_host, - core::ptr::from_ref(&SERIAL_WIDTH), + core::ptr::from_ref(&NEVER_HELPED), ) } } @@ -4361,11 +4363,11 @@ mod host_pool_split { assert!(want.iter().zip(&y).all(|(a, b)| a.to_bits() == b.to_bits())); } - /// A host pool with no workers is not using the machine, so ours may. - /// Borrowing its single thread instead measured 2-9x slower over 1-4 Mi at - /// `intra_op = 1`, so this is worth an assertion. + /// A host pool that has never been seen doing our work is not using the + /// machine, so ours may. Borrowing a single ORT thread instead measured + /// 2-9x slower over 1-4 Mi at `intra_op = 1`, so this is worth asserting. #[test] - fn a_serial_host_is_not_borrowed() { + fn an_unhelpful_host_is_not_borrowed() { if rayon::current_num_threads() < 2 { eprintln!("skipped: single-threaded rayon pool cannot show the difference"); return; @@ -4375,11 +4377,17 @@ mod host_pool_split { let mut y = vec![0.0f32; n]; reset_dispatched(); let before = parallel_dispatches(); - host_parallel::scope(serial_host_handle(), || tanh_f32_slice(&x, &mut y)); + // Past the opening burst with nothing to show for it, which is the + // steady state for a session whose host pool has no workers. + NEVER_HELPED.store( + (host_parallel::PROBE_MIN << 16) | host_parallel::PROBE_MIN, + std::sync::atomic::Ordering::Relaxed, + ); + host_parallel::scope(unhelpful_host(), || tanh_f32_slice(&x, &mut y)); assert_eq!( dispatched(), 0, - "a one-thread host pool was borrowed anyway" + "an unhelpful host pool was borrowed anyway" ); assert_eq!( parallel_dispatches() - before, diff --git a/crates/onnx-runtime-ep-plugin/src/compute.rs b/crates/onnx-runtime-ep-plugin/src/compute.rs index 195473631c..463262190b 100644 --- a/crates/onnx-runtime-ep-plugin/src/compute.rs +++ b/crates/onnx-runtime-ep-plugin/src/compute.rs @@ -1036,14 +1036,14 @@ pub struct ExportedComputeInfo { /// CPU EP (and any host EP), which uses its inputs verbatim exactly as /// before. device_staging: Option, - /// How wide ORT's intra-op pool turned out to be for this session. + /// Whether ORT's intra-op pool has ever been seen running our elementwise + /// chunks, and when to ask again if not. /// /// Lives here, rather than in a process-global, because a process may hold /// one session at `intra_op = 1` and another at `intra_op = 16`, and the /// right answer -- whether to borrow ORT's threads or use our own -- is - /// the opposite for each. Written once, by the first dispatch that has - /// more than one index; see `host_pool`. - host_pool_width: std::sync::atomic::AtomicU8, + /// the opposite for each. See `onnx_runtime_ep_api::host_parallel`. + host_pool_probe: std::sync::atomic::AtomicU32, } /// Everything `Compute` needs to stage host-resident boundary inputs onto the @@ -1101,9 +1101,7 @@ impl ExportedComputeInfo { routing: None, workspace_plans, device_staging: None, - host_pool_width: std::sync::atomic::AtomicU8::new( - onnx_runtime_ep_api::host_parallel::WIDTH_UNKNOWN, - ), + host_pool_probe: std::sync::atomic::AtomicU32::new(0), } } @@ -2095,7 +2093,7 @@ unsafe extern "C" fn compute_execute( // SAFETY: `kernel_context` is the context ORT handed this call and // stays valid until it returns, which is after the guard is dropped. let _host_pool = unsafe { - crate::host_pool::install(api_ref, kernel_context, &exported.host_pool_width) + crate::host_pool::install(api_ref, kernel_context, &exported.host_pool_probe) }; // Memory info for intermediate scratch. On a device EP this is device diff --git a/crates/onnx-runtime-ep-plugin/src/host_pool.rs b/crates/onnx-runtime-ep-plugin/src/host_pool.rs index 46a32562a1..7ab7750d99 100644 --- a/crates/onnx-runtime-ep-plugin/src/host_pool.rs +++ b/crates/onnx-runtime-ep-plugin/src/host_pool.rs @@ -23,9 +23,9 @@ use core::ffi::c_void; use onnx_genai_ort_sys as ort; use onnx_runtime_ep_api::HostParallel; use onnx_runtime_ep_api::host_parallel; -use onnx_runtime_ep_api::host_parallel::{WIDTH_PARALLEL, WIDTH_SERIAL, WIDTH_UNKNOWN}; +use onnx_runtime_ep_api::host_parallel::HOST_HELPED; use std::panic::AssertUnwindSafe; -use std::sync::atomic::{AtomicBool, AtomicU8, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU32, Ordering}; /// What [`ort_parallel_for`] needs to reach ORT, behind one `void*`. struct HostPool { @@ -38,8 +38,9 @@ struct HostPool { ) -> *mut ort::OrtStatus, release_status: Option, ctx: *mut ort::OrtKernelContext, - /// Where this session records how wide ORT's pool turned out to be. - width: *const AtomicU8, + /// Where this session records whether ORT's pool has ever helped, and + /// when to ask again. + probe: *const AtomicU32, } /// The closure a dispatch is running, plus somewhere to record what happened. @@ -52,6 +53,9 @@ struct Task<'a> { caller: Option, /// Set when a body ran somewhere other than `caller`. saw_another_thread: AtomicBool, + /// Set once the calling thread has held its first index open, to give + /// ORT's workers a chance to claim one of the others. + stalled: AtomicBool, } /// Trampoline handed to ORT: one index of one dispatch. @@ -67,16 +71,44 @@ unsafe extern "C" fn run_index(usr_data: *mut c_void, index: usize) { // it, and let `ort_parallel_for` re-raise on the calling thread once every // worker has finished. let task = unsafe { &*(usr_data.cast::>()) }; - if let Some(caller) = task.caller - && std::thread::current().id() != caller - { - task.saw_another_thread.store(true, Ordering::Relaxed); + if let Some(caller) = task.caller { + if std::thread::current().id() == caller { + // A probe that the calling thread simply drains tells us nothing: + // ORT hands its indices out dynamically, so a pool of sixteen + // looks exactly like a pool of one when the chunks are small. + // Hold the first index this thread takes for long enough that a + // worker which is going to help has time to claim another. Paid + // only on probe dispatches, and only until one of them answers. + if !task.stalled.swap(true, Ordering::Relaxed) { + stall(PROBE_STALL); + } + } else { + task.saw_another_thread.store(true, Ordering::Relaxed); + } } if std::panic::catch_unwind(AssertUnwindSafe(|| (task.body)(index))).is_err() { task.panicked.store(true, Ordering::Relaxed); } } +/// How long the calling thread holds its first index open on a probe. +/// +/// Long enough to cover a sleeping worker's wake-up on this machine, short +/// enough that a session whose pool really is serial barely notices: probes +/// back off to one dispatch in a thousand, so this is ~0.1 us amortised. +const PROBE_STALL: core::time::Duration = core::time::Duration::from_micros(100); + +/// Busy-waits for `how_long`. +/// +/// Deliberately not a sleep: this runs on ORT's calling thread, and parking it +/// would hand the core to the very workers whose presence is being tested. +fn stall(how_long: core::time::Duration) { + let until = std::time::Instant::now() + how_long; + while std::time::Instant::now() < until { + std::hint::spin_loop(); + } +} + /// Runs `body(0..total)` on the ORT session's intra-op pool. /// /// # Safety @@ -87,13 +119,16 @@ unsafe fn ort_parallel_for(host: *mut c_void, total: usize, body: &(dyn Fn(usize let pool = unsafe { &*(host.cast::()) }; // Only look at thread identities until the answer is in. A session runs // this on every dispatch, and `thread::current()` is not free. - let observing = total > 1 - && unsafe { pool.width() }.is_some_and(|w| w.load(Ordering::Relaxed) == WIDTH_UNKNOWN); + // Watch which threads run the bodies until one of them is not ours. That + // sighting is the whole answer -- see `HostParallel::prefer_host` -- and + // once it is in, stop paying for `thread::current()` on every index. + let observing = total > 1 && !unsafe { pool.helped() }; let task = Task { body, panicked: AtomicBool::new(false), caller: observing.then(|| std::thread::current().id()), saw_another_thread: AtomicBool::new(false), + stalled: AtomicBool::new(false), }; // `num_batch = 0` means "no limit": ORT gives every index its own task and // its workers claim them dynamically. That is what we want, because we cut @@ -109,19 +144,14 @@ unsafe fn ort_parallel_for(host: *mut c_void, total: usize, body: &(dyn Fn(usize (&raw const task).cast::().cast_mut(), ) }; - if observing && status.is_null() { - // ORT ran every index; if none of them landed on another thread, its - // intra-op pool has no workers. Recorded per session, so a process - // with both a one-thread and a sixteen-thread session gets the right - // answer for each. - let observed = if task.saw_another_thread.load(Ordering::Relaxed) { - WIDTH_PARALLEL - } else { - WIDTH_SERIAL - }; - if let Some(width) = unsafe { pool.width() } { - width.store(observed, Ordering::Relaxed); - } + if observing + && status.is_null() + && task.saw_another_thread.load(Ordering::Relaxed) + && let Some(probe) = unsafe { pool.probe() } + { + // Recorded per session, so a process holding both a one-thread and a + // sixteen-thread session gets the right answer for each. + probe.store(HOST_HELPED, Ordering::Relaxed); } if !status.is_null() { // The dispatch itself failed. Every index still has to run or the @@ -142,14 +172,23 @@ unsafe fn ort_parallel_for(host: *mut c_void, total: usize, body: &(dyn Fn(usize } impl HostPool { - /// The session's width cell, if it has one. + /// The session's probe cell, if it has one. /// /// # Safety /// /// The pointer must still be valid, which `install`'s contract requires /// for as long as the guard is alive. - unsafe fn width(&self) -> Option<&AtomicU8> { - (!self.width.is_null()).then(|| unsafe { &*self.width }) + unsafe fn probe(&self) -> Option<&AtomicU32> { + (!self.probe.is_null()).then(|| unsafe { &*self.probe }) + } + + /// Whether this session has already seen ORT run a body on its own thread. + /// + /// # Safety + /// + /// As [`HostPool::probe`]. + unsafe fn helped(&self) -> bool { + unsafe { self.probe() }.is_none_or(|cell| cell.load(Ordering::Relaxed) == HOST_HELPED) } } @@ -204,7 +243,7 @@ impl Drop for Guard { pub unsafe fn install( api: &ort::OrtApi, ctx: *mut ort::OrtKernelContext, - width: &AtomicU8, + probe: &AtomicU32, ) -> Guard { let (Some(parallel_for), false) = (api.KernelContext_ParallelFor, ctx.is_null()) else { return Guard::inert(); @@ -213,7 +252,7 @@ pub unsafe fn install( parallel_for, release_status: api.ReleaseStatus, ctx, - width: core::ptr::from_ref(width), + probe: core::ptr::from_ref(probe), }); // SAFETY: the box outlives the handle (see `Guard`'s drop order), and // `ort_parallel_for` only ever reads the pointer back as a `*mut HostPool`. @@ -221,7 +260,7 @@ pub unsafe fn install( HostParallel::new( (&raw const *pool).cast::().cast_mut(), ort_parallel_for, - core::ptr::from_ref(width), + core::ptr::from_ref(probe), ) }; Guard { @@ -233,6 +272,7 @@ pub unsafe fn install( #[cfg(test)] mod tests { use super::*; + use onnx_runtime_ep_api::host_parallel::{PROBE_MAX, PROBE_MIN}; use std::sync::atomic::AtomicUsize; /// Stands in for ORT: runs every index inline and reports success. @@ -300,19 +340,26 @@ mod tests { usize, *mut c_void, ) -> *mut ort::OrtStatus, - width: &AtomicU8, + probe: &AtomicU32, ) -> HostPool { HostPool { parallel_for, release_status: None, ctx: core::ptr::null_mut(), - width: core::ptr::from_ref(width), + probe: core::ptr::from_ref(probe), } } + /// Only ever used to read a probe cell back through the public API. + /// + /// # Safety + /// + /// Trivially sound: it touches neither argument. + unsafe fn never_runs(_host: *mut c_void, _total: usize, _body: &(dyn Fn(usize) + Sync)) {} + #[test] fn every_index_runs_once() { - let width = AtomicU8::new(WIDTH_UNKNOWN); + let width = AtomicU32::new(0); let mut pool = pool(inline_parallel_for, &width); let seen: Vec = (0..5).map(|_| AtomicUsize::new(0)).collect(); unsafe { @@ -323,56 +370,148 @@ mod tests { assert!(seen.iter().all(|c| c.load(Ordering::Relaxed) == 1)); } - /// An ORT that runs everything inline has no workers, and our own pool is - /// then free to use the machine. Getting this wrong in either direction - /// costs multiples: our pool alongside ORT's is 1.5-3.1x slower than - /// serial, and serial where ORT has no workers is 2-9x slower than ours. + /// The sighting that matters: a body ran somewhere other than the thread + /// that dispatched it, which only a pool with workers can do. #[test] - fn an_inline_host_is_observed_as_serial() { - let width = AtomicU8::new(WIDTH_UNKNOWN); - let mut pool = pool(inline_parallel_for, &width); + fn a_threaded_dispatch_latches_the_host_in() { + let probe = AtomicU32::new(0); + let mut pool = pool(threaded_parallel_for, &probe); unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; - assert_eq!(width.load(Ordering::Relaxed), WIDTH_SERIAL); + assert_eq!(probe.load(Ordering::Relaxed), HOST_HELPED); } + /// An ORT that ran everything inline has told us nothing. It happens on + /// 16-thread sessions too -- 35% of dispatches, in runs of up to 60 -- + /// because ORT hands its indices out dynamically and the calling thread + /// can drain the queue before a worker wakes. Concluding "no workers" from + /// that would start our pool alongside ORT's, the 3-10x pathology. #[test] - fn a_threaded_host_is_observed_as_parallel() { - let width = AtomicU8::new(WIDTH_UNKNOWN); - let mut pool = pool(threaded_parallel_for, &width); - unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; - assert_eq!(width.load(Ordering::Relaxed), WIDTH_PARALLEL); + fn inline_dispatches_never_decide_anything() { + let probe = AtomicU32::new(0); + let mut pool = pool(inline_parallel_for, &probe); + for _ in 0..256 { + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + } + assert_eq!(probe.load(Ordering::Relaxed), 0, "silence is not evidence"); } - /// One index proves nothing: even a sixteen-thread pool runs it inline. + /// Once latched, later dispatches stop paying for the observation, and + /// nothing can unlatch it. #[test] - fn a_single_index_dispatch_does_not_decide_the_width() { - let width = AtomicU8::new(WIDTH_UNKNOWN); - let mut pool = pool(inline_parallel_for, &width); + fn a_latched_host_is_not_revisited() { + let probe = AtomicU32::new(HOST_HELPED); + let mut pool = pool(inline_parallel_for, &probe); + for _ in 0..16 { + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + } + assert_eq!(probe.load(Ordering::Relaxed), HOST_HELPED); + } + + /// One index cannot be split, so it proves nothing either way. + #[test] + fn a_single_index_dispatch_is_not_evidence() { + let probe = AtomicU32::new(0); + let mut pool = pool(threaded_parallel_for, &probe); unsafe { ort_parallel_for((&raw mut pool).cast::(), 1, &|_| {}) }; - assert_eq!(width.load(Ordering::Relaxed), WIDTH_UNKNOWN); + assert_eq!(probe.load(Ordering::Relaxed), 0); } - /// A refused dispatch says nothing about the pool either. + /// The stall must not run once a session has latched: it is a probe cost, + /// not a per-dispatch one. #[test] - fn a_refused_dispatch_does_not_decide_the_width() { - let width = AtomicU8::new(WIDTH_UNKNOWN); - let mut pool = pool(failing_parallel_for, &width); - unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; - assert_eq!(width.load(Ordering::Relaxed), WIDTH_UNKNOWN); + fn a_latched_session_pays_no_stall() { + let probe = AtomicU32::new(HOST_HELPED); + let mut pool = pool(inline_parallel_for, &probe); + let started = std::time::Instant::now(); + for _ in 0..8 { + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + } + assert!( + started.elapsed() < PROBE_STALL, + "eight latched dispatches took longer than a single probe stall" + ); } - /// Once decided, later dispatches stop paying for the observation. + /// And it must run at most once per probe, however many indices there are. #[test] - fn an_answered_width_is_not_revisited() { - let width = AtomicU8::new(WIDTH_PARALLEL); - let mut pool = pool(inline_parallel_for, &width); - unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; - assert_eq!(width.load(Ordering::Relaxed), WIDTH_PARALLEL); + fn a_probe_stalls_once_not_once_per_index() { + let probe = AtomicU32::new(0); + let mut pool = pool(inline_parallel_for, &probe); + let started = std::time::Instant::now(); + unsafe { ort_parallel_for((&raw mut pool).cast::(), 64, &|_| {}) }; + let elapsed = started.elapsed(); + assert!( + elapsed >= PROBE_STALL, + "a probe has to give workers a chance" + ); + assert!( + elapsed < PROBE_STALL * 4, + "the stall is per dispatch, not per index: {elapsed:?}" + ); + } + + /// A refused dispatch ran on our fallback path, not ORT's pool. + #[test] + fn a_refused_dispatch_is_not_evidence() { + let probe = AtomicU32::new(0); + let mut pool = pool(failing_parallel_for, &probe); + for _ in 0..8 { + unsafe { ort_parallel_for((&raw mut pool).cast::(), 4, &|_| {}) }; + } + assert_eq!(probe.load(Ordering::Relaxed), 0); + } + + /// End to end through the public seam: a session whose pool helps ends up + /// preferring the host, and one whose pool never does keeps asking, but + /// rarely. + #[test] + fn the_probe_and_the_latch_agree() { + let probe = AtomicU32::new(0); + let mut threaded = pool(threaded_parallel_for, &probe); + // SAFETY: `never_runs` dereferences neither argument. + let handle = unsafe { + HostParallel::new( + core::ptr::null_mut(), + never_runs, + core::ptr::from_ref(&probe), + ) + }; + assert!(handle.prefer_host(), "the first dispatch always asks"); + unsafe { ort_parallel_for((&raw mut threaded).cast::(), 4, &|_| {}) }; + assert!(handle.helped()); + assert!(handle.prefer_host()); + + let quiet = AtomicU32::new(0); + let mut inline = pool(inline_parallel_for, &quiet); + // SAFETY: as above. + let handle = unsafe { + HostParallel::new( + core::ptr::null_mut(), + never_runs, + core::ptr::from_ref(&quiet), + ) + }; + let mut asks = 0; + for _ in 0..4096 { + if handle.prefer_host() { + asks += 1; + unsafe { ort_parallel_for((&raw mut inline).cast::(), 4, &|_| {}) }; + } + } + assert!(!handle.helped()); + assert!( + (2..=4096 / usize::try_from(PROBE_MIN).unwrap()).contains(&asks), + "asked {asks} times in 4096 dispatches" + ); + assert!( + asks * usize::try_from(PROBE_MAX).unwrap() >= 4096, + "a serial-looking session must still re-ask often enough to recover" + ); } #[test] fn a_refused_dispatch_still_runs_every_index() { - let width = AtomicU8::new(WIDTH_UNKNOWN); + let width = AtomicU32::new(0); let mut pool = pool(failing_parallel_for, &width); let seen: Vec = (0..5).map(|_| AtomicUsize::new(0)).collect(); unsafe { @@ -388,7 +527,7 @@ mod tests { #[test] fn a_panicking_body_does_not_unwind_into_ort() { - let width = AtomicU8::new(WIDTH_UNKNOWN); + let width = AtomicU32::new(0); let mut pool = pool(inline_parallel_for, &width); let ran = AtomicUsize::new(0); let unwound = std::panic::catch_unwind(AssertUnwindSafe(|| unsafe { @@ -407,7 +546,7 @@ mod tests { fn an_ort_without_parallel_for_installs_nothing() { let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; api.KernelContext_ParallelFor = None; - let width = AtomicU8::new(WIDTH_UNKNOWN); + let width = AtomicU32::new(0); let guard = unsafe { install(&api, core::ptr::dangling_mut(), &width) }; assert!(!guard.is_installed()); assert!(host_parallel::current().is_none()); @@ -417,7 +556,7 @@ mod tests { fn a_null_context_installs_nothing() { let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; api.KernelContext_ParallelFor = Some(inline_parallel_for); - let width = AtomicU8::new(WIDTH_UNKNOWN); + let width = AtomicU32::new(0); let guard = unsafe { install(&api, core::ptr::null_mut(), &width) }; assert!(!guard.is_installed()); assert!(host_parallel::current().is_none()); @@ -429,7 +568,7 @@ mod tests { api.KernelContext_ParallelFor = Some(inline_parallel_for); let seen: Vec = (0..6).map(|_| AtomicUsize::new(0)).collect(); { - let width = AtomicU8::new(WIDTH_UNKNOWN); + let width = AtomicU32::new(0); let guard = unsafe { install(&api, core::ptr::dangling_mut(), &width) }; assert!(guard.is_installed()); let host = host_parallel::current().expect("install publishes a handle"); @@ -448,7 +587,7 @@ mod tests { fn a_nested_install_restores_the_outer_handle() { let mut api: ort::OrtApi = unsafe { core::mem::zeroed() }; api.KernelContext_ParallelFor = Some(inline_parallel_for); - let width = AtomicU8::new(WIDTH_UNKNOWN); + let width = AtomicU32::new(0); let outer = unsafe { install(&api, core::ptr::dangling_mut(), &width) }; assert!(outer.is_installed()); { From 5fe0c69af1a63bf59a865bb58b07f44d9f91ba78 Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Mon, 17 Aug 2026 23:25:03 +0000 Subject: [PATCH 4/5] perf(cpu-ep): outline the host-pool branch out of run_chunked Writing the host decision inline in `run_chunked` cost `Relu` 34% at 1 Mi on one thread -- 236 -> 315 us, reproducible to the microsecond across runs -- with no change to the path it actually took: with `intra_op = 1` nothing ever latches, so every one of those dispatches went down the same serial call as before. Adding an `eprintln!` to diagnose it moved `Relu` back to 237 us and pushed `Tanh` from 440 to 1110, which is the tell: this is the codegen-unit repartitioning already documented for `clip_chunked`, not a runtime effect. The branch has no business being inline anyway. It decides with one relaxed load and the split behind it only happens inside a session whose pool has proved parallel, while `run_chunked`'s callers are the hottest elementwise kernels in the crate. `#[inline(never)]` on `try_host` and `try_host_rows`, with a test that keeps it there. Measured at intra_op=1, rayon=1, 4 interleaved rounds (branch vs main, us at 1 Mi): Relu 236.4/236.3, Clip 268.1/314.2, Gelu 1260.8/1303.9, Tanh 443.7/472.9, Erf 894.7/894.8 -- i.e. the regression is gone and Clip gained. The 16-thread win is unchanged: Clip 1 Mi 697.5 -> 85.5 us, Relu 693.7 -> 79.0, Erf 1398.0 -> 165.9, Gelu 1937.4 -> 240.6. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../src/kernels/simd_activations.rs | 102 +++++++++++++++--- 1 file changed, 85 insertions(+), 17 deletions(-) diff --git a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs index 6ea3f07d64..1f7d7b8c99 100644 --- a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs +++ b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs @@ -410,6 +410,60 @@ impl HostChunks { } } +/// Runs `body` on the host runtime's pool if this session has one worth using. +/// +/// Returns whether `body` was run at all; a `false` return leaves `output` +/// untouched and the caller free to pick another split. +/// +/// `#[inline(never)]` on purpose. Everything here is cold -- the decision is a +/// relaxed load, and the split it guards only happens inside an ORT session +/// whose pool has been proven parallel -- while `run_chunked`'s callers are the +/// hottest elementwise kernels in the crate. Inlining it grew `run_chunked` +/// enough to repartition codegen units: at `intra_op = 1` it cost `Relu` 34% at +/// 1 Mi (236 -> 315 us) with no path change at all, the same failure mode as +/// the instantiation note on `clip_chunked`. +#[inline(never)] +fn try_host(input: &[f32], output: &mut [f32], body: &F) -> bool +where + F: Fn(&[f32], &mut [f32]) + Sync + Send, +{ + let Some(host) = onnx_runtime_ep_api::host_parallel::current() else { + return false; + }; + if !host.prefer_host() { + return false; + } + match host_chunk_len(input.len()) { + Some((chunk, count)) => { + note_parallel_dispatch(); + run_on_host(host, input, output, chunk, count, body); + } + // Too short to split across the host's threads. Our own pool is not + // the answer either: it would be a second pool on the same cores. + None => body(input, output), + } + true +} + +/// [`try_host`] for the row-shaped split. Outlined for the same reason. +#[inline(never)] +fn try_host_rows(input: &[f32], output: &mut [f32], width: usize, body: &F) -> bool +where + F: Fn(&[f32], &mut [f32]) + Sync + Send, +{ + let Some(host) = onnx_runtime_ep_api::host_parallel::current() else { + return false; + }; + if !host.prefer_host() { + return false; + } + match host_chunk_len_rows(input.len(), width) { + Some((chunk, count)) => run_on_host(host, input, output, chunk, count, body), + None => body(input, output), + } + true +} + /// Splits `input`/`output` into `count` chunks of `chunk` and runs `body` on /// each one on the host runtime's pool. /// @@ -502,16 +556,7 @@ where // right tool: at `intra_op = 1`, rayon beat borrowing ORT's single thread // by 2-9x over 1-4 Mi. `prefer_host` decides which of the two a session is // by watching whether the host's pool ever actually runs our chunks. - if let Some(host) = onnx_runtime_ep_api::host_parallel::current() - && host.prefer_host() - { - match host_chunk_len(len) { - Some((chunk, count)) => { - note_parallel_dispatch(); - run_on_host(host, input, output, chunk, count, &body); - } - None => body(input, output), - } + if try_host(input, &mut *output, &body) { return; } if len < PAR_MIN_LEN { @@ -616,13 +661,7 @@ where body(input, output); return; } - if let Some(host) = onnx_runtime_ep_api::host_parallel::current() - && host.prefer_host() - { - match host_chunk_len_rows(len, width) { - Some((chunk, count)) => run_on_host(host, input, output, chunk, count, &body), - None => body(input, output), - } + if try_host_rows(input, &mut *output, width, &body) { return; } if len < PAR_MIN_LEN { @@ -4144,6 +4183,35 @@ mod chunking_instantiation_is_local { Offending call sites: {offenders:?}" ); } + + /// The host branch is cold, and inlining it into `run_chunked` costs 34%. + /// + /// `try_host`/`try_host_rows` decide with one relaxed load, and the split + /// they guard only runs inside an ORT session whose intra-op pool has been + /// proven parallel. `run_chunked`'s callers, meanwhile, are the hottest + /// elementwise kernels in the crate. When the branch was written inline, + /// `Relu` at 1 Mi and one thread went 236 -> 315 us with no runtime path + /// change at all -- the same codegen-unit repartitioning as above. Adding + /// an `eprintln!` for diagnosis moved it back to 237 us, which is how the + /// cause was identified. `#[inline(never)]` pins it. + #[test] + fn the_host_branch_stays_outlined() { + let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("src/kernels/simd_activations.rs"); + let text = std::fs::read_to_string(&path).expect("read source"); + for name in ["fn try_host", "fn try_host_rows"] { + let at = text.find(name).unwrap_or_else(|| panic!("{name} exists")); + assert!( + text[..at] + .lines() + .next_back() + .is_some_and(|l| l.trim() == "#[inline(never)]"), + "{name} must be preceded by #[inline(never)]: inlining the host \ + branch into run_chunked repartitioned codegen units and cost \ + Relu 34% at 1 Mi (236 -> 315 us) with no path change" + ); + } + } } /// The host-pool split: what happens when ORT lends us its intra-op threads. From 17ec3d5ae104736198ef737993a21679b277f0dd Mon Sep 17 00:00:00 2001 From: justinchuby <223556219+Copilot@users.noreply.github.com> Date: Mon, 17 Aug 2026 16:53:22 -0700 Subject: [PATCH 5/5] test(cpu-ep): three-arm interleaved bench for the host-pool split Add a gated (EP_BENCH=1, #[ignore]d) three-arm microbenchmark in simd_activations::three_arm_bench that measures serial vs rayon split vs host-pool split in one interleaved process, driving the real kernel paths against a spinning stand-in pool that models ORT's always-hot intra-op pool. The host-pool split beats serial ~5.5-6.6x at intra_op=16 on 1 Mi f32; the rayon-split arm reproduces its known loss as a built-in control; the no-host fall-through still wins at intra_op=1. Also strengthen the MAX_HOST_CHUNKS comment: cutting by size rather than pool width keeps the chunk boundaries (and thus the bit pattern) independent of the session's thread count, which is a correctness property, not a tuning one. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../src/kernels/simd_activations.rs | 385 +++++++++++++++++- 1 file changed, 382 insertions(+), 3 deletions(-) diff --git a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs index 1f7d7b8c99..c4854c4e17 100644 --- a/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs +++ b/crates/onnx-runtime-ep-cpu/src/kernels/simd_activations.rs @@ -298,9 +298,16 @@ fn par_chunk_len(len: usize, threads: usize) -> Option { /// opposite: it is already spinning, we cannot ask ORT how many intra-op /// threads it was given, and `TrySimpleParallelFor` claims indices /// dynamically. So we cut by *size* rather than by thread count and let the -/// host decide how many of its threads to point at the result — which also -/// makes the split independent of the machine, so the chunk boundaries (and -/// therefore the bits) do not move when the session is reconfigured. +/// host decide how many of its threads to point at the result. +/// +/// Cutting by size rather than by pool width is also a **correctness** +/// property, not only a scheduling one: the chunk boundaries — and therefore +/// the exact bit pattern of the output — depend only on `len`, never on how +/// many intra-op threads the session happens to have. The same tensor produces +/// the same bits whether it is run at `intra_op = 1` or `16`, so a result can +/// never move when the session is reconfigured. Had we cut by pool width (as +/// the rayon path does), reducing `intra_op` would re-chunk the slice and +/// perturb the last-lane rounding of the split boundaries. /// /// [`HOST_MIN_CHUNK`] is the floor, so a chunk is never too short to be worth /// dispatching or to stay on the vector path; this cap is the other end, and @@ -4537,3 +4544,375 @@ mod host_pool_split { } } } + +/// Interleaved three-arm micro-measurement for the host-pool split. +/// +/// This is the arm that decides PR #1143: everything measured before compared +/// *our rayon split against staying serial*, and serial won every row. Nothing +/// showed that dispatching the split onto the **host's** pool — the change this +/// PR actually makes — beats serial. This harness runs all three arms in one +/// process, interleaved per rep, so between-arm drift on a shared box cannot +/// flip the conclusion. +/// +/// # Why a stand-in pool, and why it is trustworthy +/// +/// A real ORT session cannot switch arms within one process — its rayon width +/// is a process global and its host install is a code path — so the three arms +/// are driven directly against the real `simd_activations` code paths. The one +/// thing that has to be reproduced faithfully is *why* our rayon split loses +/// under an ORT session: ORT's intra-op pool **spins**, so a second (rayon) +/// pool alongside it oversubscribes the cores. The stand-in models exactly +/// that — a persistent pool of workers that spin while idle (as ORT's do) and +/// claim indices dynamically when dispatched (`num_batch = 0` semantics) — and +/// it latches [`host_parallel::HOST_HELPED`] the honest way, by having a worker +/// (not the dispatcher) run one of the chunks, so `prefer_host` reaches its +/// steady state through the real mechanism rather than by fiat. +/// +/// The `rayon split` arm is therefore a **built-in control**: its result is +/// already known from the PR's nine-row table (serial wins at `intra_op = 16`). +/// If this harness reproduces `rayon < serial`, the model is validated and its +/// `host-pool` number can be trusted. If it does not, the run is unusable and +/// we know it before drawing a conclusion. +/// +/// Ignored by the normal gate; run with: +/// `EP_BENCH=1 cargo test --release -p onnx-runtime-ep-cpu --lib three_arm_bench -- --nocapture --ignored --test-threads=1` +#[cfg(test)] +mod three_arm_bench { + use super::*; + use onnx_runtime_ep_api::HostParallel; + use onnx_runtime_ep_api::host_parallel; + use std::ffi::c_void; + use std::sync::Arc; + use std::sync::atomic::{AtomicBool, AtomicU32, AtomicUsize, Ordering}; + use std::time::{Duration, Instant}; + + /// State shared between the dispatcher and the stand-in pool's workers. + /// + /// A dispatch publishes the body (as its two raw fat-pointer words) and the + /// index count, then bumps `generation`; workers pick the new generation up + /// and claim indices off `cursor` until it is drained, recording progress + /// in `done`. Everything is lock-free so idle workers can *spin*, which is + /// the property that makes a coexisting rayon pool oversubscribe. + struct Shared { + generation: AtomicUsize, + total: AtomicUsize, + cursor: AtomicUsize, + done: AtomicUsize, + body_data: AtomicUsize, + body_vtable: AtomicUsize, + stop: AtomicBool, + /// The probe cell `prefer_host` reads: starts at zero (opening burst) + /// and reaches `HOST_HELPED` the first time a worker runs a chunk. + probe: AtomicU32, + /// Set by a worker (never the dispatcher) when it runs a chunk, so the + /// dispatcher can latch the probe cell only on real positive evidence. + worker_helped: AtomicBool, + } + + impl Shared { + fn new() -> Self { + Self { + generation: AtomicUsize::new(0), + total: AtomicUsize::new(0), + cursor: AtomicUsize::new(0), + done: AtomicUsize::new(0), + body_data: AtomicUsize::new(0), + body_vtable: AtomicUsize::new(0), + stop: AtomicBool::new(false), + probe: AtomicU32::new(0), + worker_helped: AtomicBool::new(false), + } + } + + /// Claims and runs indices off `cursor` for the current generation. + /// + /// Shared by the workers and the calling thread — ORT's own + /// `ParallelFor` runs tasks on the caller too, so the calling thread is + /// one of the `intra_op` threads, not a spectator. A worker flags + /// `worker_helped` so the dispatcher can latch the probe cell. + /// + /// # Safety + /// + /// The published body pointer must still be valid, which the blocking + /// dispatch guarantees for its whole duration. + unsafe fn drain(&self, total: usize, is_worker: bool) { + let data = self.body_data.load(Ordering::Relaxed); + let vtable = self.body_vtable.load(Ordering::Relaxed); + // SAFETY: `data`/`vtable` are the two words of a `&(dyn Fn(usize) + + // Sync)` published under the `generation` release/acquire, and the + // referent outlives the dispatch. + let body: *const (dyn Fn(usize) + Sync) = + unsafe { core::mem::transmute([data, vtable]) }; + let body = unsafe { &*body }; + let mut ran = false; + loop { + let index = self.cursor.fetch_add(1, Ordering::Relaxed); + if index >= total { + break; + } + body(index); + ran = true; + self.done.fetch_add(1, Ordering::Relaxed); + } + if is_worker && ran { + self.worker_helped.store(true, Ordering::Relaxed); + } + } + } + + /// A persistent pool of spinning workers standing in for ORT's intra-op + /// pool. Sized so that `workers + 1` (the calling thread) equals the + /// modelled `intra_op_num_threads`. + struct SpinPool { + shared: Arc, + handles: Vec>, + } + + impl SpinPool { + fn new(workers: usize) -> Self { + let shared = Arc::new(Shared::new()); + let mut handles = Vec::with_capacity(workers); + for _ in 0..workers { + let shared = Arc::clone(&shared); + handles.push(std::thread::spawn(move || { + let mut seen = 0usize; + loop { + if shared.stop.load(Ordering::Relaxed) { + break; + } + let g = shared.generation.load(Ordering::Acquire); + if g != seen { + seen = g; + let total = shared.total.load(Ordering::Relaxed); + // SAFETY: the dispatch that bumped `generation` is + // blocked until `done == total`, so the body it + // published is alive for this whole drain. + unsafe { shared.drain(total, true) }; + } else { + // Idle: burn the core, exactly as ORT's workers do + // between parallel regions. This is what makes the + // rayon arm oversubscribe. + std::hint::spin_loop(); + } + } + })); + } + Self { shared, handles } + } + + /// The `HostParallel` handle our kernels see. + fn handle(&self) -> HostParallel { + // SAFETY: `spin_dispatch` reads the host pointer back as + // `*const Shared`, and the probe pointer is the cell inside the same + // `Arc`. The `Arc` in `self` keeps both alive for as long as the + // handle is installed (the handle never escapes this struct). + unsafe { + HostParallel::new( + Arc::as_ptr(&self.shared).cast::().cast_mut(), + spin_dispatch, + core::ptr::from_ref(&self.shared.probe), + ) + } + } + } + + impl Drop for SpinPool { + fn drop(&mut self) { + self.shared.stop.store(true, Ordering::Relaxed); + for h in self.handles.drain(..) { + let _ = h.join(); + } + } + } + + /// Publishes one dispatch and drives it to completion on the pool + caller. + /// + /// # Safety + /// + /// `host` must be the `*const Shared` produced by [`SpinPool::handle`]. + unsafe fn spin_dispatch(host: *mut c_void, total: usize, body: &(dyn Fn(usize) + Sync)) { + let shared = unsafe { &*host.cast::() }; + let raw = body as *const (dyn Fn(usize) + Sync); + let [data, vtable]: [usize; 2] = unsafe { core::mem::transmute(raw) }; + shared.worker_helped.store(false, Ordering::Relaxed); + shared.total.store(total, Ordering::Relaxed); + shared.cursor.store(0, Ordering::Relaxed); + shared.done.store(0, Ordering::Relaxed); + shared.body_data.store(data, Ordering::Relaxed); + shared.body_vtable.store(vtable, Ordering::Relaxed); + // Release the body/total stores to the workers, then help drain. + shared.generation.fetch_add(1, Ordering::Release); + // SAFETY: we published the body above and block below until every index + // is done, so it stays valid for the whole call. + unsafe { shared.drain(total, false) }; + while shared.done.load(Ordering::Acquire) < total { + std::hint::spin_loop(); + } + // Latch the probe the honest way: only if a worker (not this calling + // thread) actually ran a chunk. A pool with no workers can never set + // this, which is exactly the fact `prefer_host` is trying to learn. + if shared.worker_helped.load(Ordering::Relaxed) { + shared + .probe + .store(host_parallel::HOST_HELPED, Ordering::Relaxed); + } + } + + /// A reproducible non-degenerate input in a range that exercises every + /// branch of the activation polynomials. + fn probe(len: usize) -> Vec { + (0..len) + .map(|i| (i as f32 / 977.0).sin() * 9.0 + (i % 13) as f32 * 1e-6 - 2.0) + .collect() + } + + fn pct(sorted: &[Duration], num: usize, den: usize) -> Duration { + sorted[(sorted.len() * num / den).min(sorted.len() - 1)] + } + + fn us(d: Duration) -> f64 { + d.as_secs_f64() * 1e6 + } + + type Kernel = fn(&[f32], &mut [f32]); + + fn kernels() -> [(&'static str, Kernel); 2] { + [ + ("Sqrt", sqrt_f32_slice as Kernel), + ("Gelu", erf_gelu_f32_slice as Kernel), + ] + } + + /// The three arms at `intra_op = 16`, interleaved per rep. + #[test] + #[ignore = "measurement harness; run explicitly with EP_BENCH=1"] + fn three_arm_intra_op_16() { + if std::env::var("EP_BENCH").is_err() { + return; + } + const N: usize = 1 << 20; // 1 Mi f32 + const REPS: usize = 201; + const WARMUP: usize = 25; + // The modelled intra-op width. Defaults to this box's physical core + // count (14) rather than the spec's 16, because the stand-in's spinning + // workers are *real* threads: on a 14-core box, 15 spinners + the caller + // already oversubscribe, which starves the serial baseline before the + // rayon arm even runs. Matching the core count gives every arm a clean + // core and keeps the rayon arm the only one that oversubscribes. + let intra_op: usize = std::env::var("EP_INTRA_OP") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(14); + + // (intra_op - 1) spinning workers + the calling thread == intra_op threads. + let pool = SpinPool::new(intra_op - 1); + let host = pool.handle(); + let rayon_pool = rayon::ThreadPoolBuilder::new() + .num_threads(intra_op) + .build() + .expect("rayon pool"); + + let x = probe(N); + let mut out = vec![0.0f32; N]; + + println!("\n=== three-arm, intra_op = {intra_op}, N = 1 Mi f32, p50 us (p10..p90) ==="); + println!( + "host stand-in: {} spinning workers + caller; rayon: {intra_op} threads", + intra_op - 1 + ); + for (name, kernel) in kernels() { + // Interleave the arms within one process: each rep samples all + // three back to back, so a busy neighbour hits every arm equally. + let arms = ["serial", "rayon split", "host-pool split"]; + let mut samples: [Vec; 3] = [Vec::new(), Vec::new(), Vec::new()]; + for _ in 0..WARMUP + REPS { + let mut run_arm = |arm: usize| { + let t = Instant::now(); + match arm { + 0 => serial_scope(|| kernel(&x, &mut out)), + 1 => rayon_pool.install(|| kernel(&x, &mut out)), + _ => host_parallel::scope(host, || kernel(&x, &mut out)), + } + t.elapsed() + }; + for (arm, bucket) in samples.iter_mut().enumerate() { + let d = run_arm(arm); + bucket.push(d); + } + } + // The host arm must have latched, or its numbers are the rayon + // fall-through, not the host path. + assert!( + host.helped(), + "{name}: the stand-in host was never seen to help; numbers invalid" + ); + print!("{name:>5}: "); + let mut p50s = [0.0f64; 3]; + for (i, arm) in arms.iter().enumerate() { + let s = &mut samples[i]; + s.drain(..WARMUP); + s.sort_unstable(); + p50s[i] = us(pct(s, 1, 2)); + print!( + "{arm} {:.0} ({:.0}..{:.0}) ", + us(pct(s, 1, 2)), + us(pct(s, 1, 10)), + us(pct(s, 9, 10)) + ); + } + println!(); + println!( + " -> rayon/serial {:.2}x host/serial {:.2}x (control: rayon must be > 1x)", + p50s[1] / p50s[0], + p50s[2] / p50s[0] + ); + } + } + + /// The no-host fall-through on a free machine: rayon split must keep its + /// large win over serial at `intra_op = 1`. No stand-in pool, so nothing + /// spins and the machine is the native executor's to use. + #[test] + #[ignore = "measurement harness; run explicitly with EP_BENCH=1"] + fn no_host_fall_through_intra_op_1() { + if std::env::var("EP_BENCH").is_err() { + return; + } + const REPS: usize = 151; + const WARMUP: usize = 20; + println!("\n=== no-host fall-through, free machine, p50 us (p10..p90) ==="); + for n_label in ["1 Mi", "4 Mi"] { + let n = if n_label == "1 Mi" { 1 << 20 } else { 1 << 22 }; + let x = probe(n); + let mut out = vec![0.0f32; n]; + for (name, kernel) in kernels() { + let mut ser = Vec::with_capacity(REPS); + let mut par = Vec::with_capacity(REPS); + for r in 0..WARMUP + REPS { + let t = Instant::now(); + serial_scope(|| kernel(&x, &mut out)); + let ds = t.elapsed(); + let t = Instant::now(); + // No host installed and not inside rayon: run_chunked takes + // the rayon fall-through on the global pool. + kernel(&x, &mut out); + let dp = t.elapsed(); + if r >= WARMUP { + ser.push(ds); + par.push(dp); + } + } + ser.sort_unstable(); + par.sort_unstable(); + let (s, prl) = (pct(&ser, 1, 2), pct(&par, 1, 2)); + println!( + "{n_label} {name:>5}: serial {:.0} rayon {:.0} -> {:.2}x faster", + us(s), + us(prl), + us(s) / us(prl) + ); + } + } + } +} +