use std::cell::UnsafeCell;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use crate::decode_affinity::{NodeShard, NumaTopology};
use crate::kernels::matmul_nbits::output_chunk_len;
pub const PERSISTENT_POOL_ENV: &str = "ONNX_GENAI_CPU_DECODE_PERSISTENT_POOL";
const SPIN_BEFORE_YIELD: u32 = 1 << 12;
const YIELD_BEFORE_PARK: u32 = 1 << 6;
const PARK_TIMEOUT: std::time::Duration = std::time::Duration::from_millis(50);
#[repr(align(128))]
struct Padded<T>(T);
#[derive(Clone, Copy)]
struct Job {
data: *const (),
call: unsafe fn(*const (), usize),
}
struct SharedState {
sequence: Padded<AtomicUsize>,
job: UnsafeCell<Option<Job>>,
node_pending: Vec<Padded<AtomicUsize>>,
parked: Vec<Padded<AtomicBool>>,
worker_node: Vec<usize>,
ready: AtomicUsize,
poisoned_worker: AtomicUsize,
shutdown: AtomicBool,
}
unsafe impl Sync for SharedState {}
unsafe impl Send for SharedState {}
impl SharedState {
fn publish(&self, job: Job, counts: &[usize], threads: &[thread::Thread]) {
unsafe {
*self.job.get() = Some(job);
}
for (counter, &count) in self.node_pending.iter().zip(counts) {
counter.0.store(count, Ordering::Release);
}
self.sequence.0.fetch_add(1, Ordering::SeqCst);
for (index, parked) in self.parked.iter().enumerate() {
if parked.0.load(Ordering::SeqCst) {
threads[index].unpark();
}
}
}
fn wait(&self) {
let mut spins = 0u32;
loop {
let done = self
.node_pending
.iter()
.all(|counter| counter.0.load(Ordering::Acquire) == 0);
if done {
return;
}
std::hint::spin_loop();
spins = spins.wrapping_add(1);
if spins >= SPIN_BEFORE_YIELD {
thread::yield_now();
}
}
}
fn panic_if_poisoned(&self) {
let poisoned = self.poisoned_worker.load(Ordering::Acquire);
if poisoned != 0 {
let worker = poisoned - 1;
panic!(
"persistent SPMD decode worker {worker} panicked while executing a decode op; \
the pool is poisoned and cannot continue. Disable \
ONNX_GENAI_CPU_DECODE_PERSISTENT_POOL or restart the process"
);
}
}
}
pub struct SpmdDecodePools {
shared: Arc<SharedState>,
worker_threads: Vec<thread::Thread>,
join_handles: Mutex<Vec<JoinHandle<()>>>,
node_worker_counts: Vec<usize>,
total_workers: usize,
}
impl SpmdDecodePools {
fn build(shards: &[NodeShard]) -> Self {
let node_count = shards.len();
let mut worker_node = Vec::new();
let mut node_worker_counts = Vec::with_capacity(node_count);
let mut assignment: Vec<(usize, Option<usize>)> = Vec::new();
for (node_position, shard) in shards.iter().enumerate() {
node_worker_counts.push(shard.workers);
for worker in 0..shard.workers {
worker_node.push(node_position);
let cpu = shard.cpus.get(worker % shard.cpus.len().max(1)).copied();
assignment.push((node_position, cpu));
}
}
let total_workers = assignment.len();
let shared = Arc::new(SharedState {
sequence: Padded(AtomicUsize::new(0)),
job: UnsafeCell::new(None),
node_pending: (0..node_count)
.map(|_| Padded(AtomicUsize::new(0)))
.collect(),
parked: (0..total_workers)
.map(|_| Padded(AtomicBool::new(false)))
.collect(),
worker_node,
ready: AtomicUsize::new(0),
poisoned_worker: AtomicUsize::new(0),
shutdown: AtomicBool::new(false),
});
let mut handles = Vec::with_capacity(total_workers);
for (global_index, (node_position, cpu)) in assignment.into_iter().enumerate() {
let shared = Arc::clone(&shared);
let handle = thread::Builder::new()
.name(format!("onnx-genai-spmd-n{node_position}-{global_index}"))
.spawn(move || {
if let Some(cpu) = cpu
&& let Err(message) = crate::decode_affinity::pin_current_thread_to_cpu(cpu)
{
report_spmd_fallback(&format!(
"worker {global_index} could not pin to cpu {cpu}: {message}"
));
}
worker_loop(shared, global_index);
})
.expect("spawn persistent SPMD decode worker");
handles.push(handle);
}
while shared.ready.load(Ordering::Acquire) < total_workers {
std::hint::spin_loop();
}
let worker_threads = handles.iter().map(|h| h.thread().clone()).collect();
Self {
shared,
worker_threads,
join_handles: Mutex::new(handles),
node_worker_counts,
total_workers,
}
}
pub fn total_workers(&self) -> usize {
self.total_workers
}
pub fn node_count(&self) -> usize {
self.node_worker_counts.len()
}
pub fn shutdown(&self) {
let handles: Vec<JoinHandle<()>> = {
let mut guard = self
.join_handles
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
guard.drain(..).collect()
};
if handles.is_empty() {
return;
}
self.shared.shutdown.store(true, Ordering::SeqCst);
self.shared.sequence.0.fetch_add(1, Ordering::SeqCst);
for thread in &self.worker_threads {
thread.unpark();
}
for handle in handles {
let _ = handle.join();
}
}
fn dispatch<F>(&self, job: &F)
where
F: Fn(usize) + Sync,
{
self.shared.panic_if_poisoned();
unsafe fn call<F>(data: *const (), global_index: usize)
where
F: Fn(usize) + Sync,
{
let job = unsafe { &*data.cast::<F>() };
job(global_index);
}
let job = Job {
data: std::ptr::from_ref(job).cast(),
call: call::<F>,
};
self.shared
.publish(job, &self.node_worker_counts, &self.worker_threads);
self.shared.wait();
self.shared.panic_if_poisoned();
}
fn node_row_lengths(&self, n: usize) -> Vec<usize> {
let node_count = self.node_worker_counts.len();
let mut lengths = Vec::with_capacity(node_count);
let mut assigned = 0;
for (position, &node_workers) in self.node_worker_counts.iter().enumerate() {
let rows = if position + 1 == node_count {
n - assigned
} else {
n.saturating_mul(node_workers) / self.total_workers
};
assigned += rows;
lengths.push(rows);
}
lengths
}
fn worker_row_segments(&self, n: usize) -> Vec<(usize, usize)> {
let node_lengths = self.node_row_lengths(n);
let mut segments = Vec::with_capacity(self.total_workers);
let mut node_start = 0;
for (&node_len, &node_workers) in node_lengths.iter().zip(&self.node_worker_counts) {
let base = node_len / node_workers;
let remainder = node_len % node_workers;
let mut offset = node_start;
for worker in 0..node_workers {
let len = base + usize::from(worker < remainder);
segments.push((offset, len));
offset += len;
}
node_start += node_len;
}
segments
}
fn worker_row_segments_aligned(&self, n: usize, align: usize) -> Vec<(usize, usize)> {
let base = self.worker_row_segments(n);
if align <= 1 {
return base;
}
let mut segments = Vec::with_capacity(base.len());
let mut prev_boundary = 0;
let mut cumulative = 0;
let last = base.len().saturating_sub(1);
for (index, &(_, len)) in base.iter().enumerate() {
cumulative += len;
let boundary = if index == last {
n
} else {
let rounded = ((cumulative + align / 2) / align) * align;
rounded.clamp(prev_boundary, n)
};
segments.push((prev_boundary, boundary - prev_boundary));
prev_boundary = boundary;
}
segments
}
pub fn dispatch_output_rows<F>(&self, result: &mut [f32], k: usize, compute: &F)
where
F: Fn(usize, &mut [f32]) + Sync,
{
let n = result.len();
if self.total_workers <= 1 || output_chunk_len(n, k) >= n {
compute(0, result);
return;
}
self.dispatch_rows_across_workers(result, &compute);
}
pub fn output_column_segments(&self, n: usize, align: usize) -> Vec<(usize, usize)> {
self.worker_row_segments_aligned(n, align)
}
pub fn dispatch_output_rows_indexed<F>(&self, result: &mut [f32], align: usize, compute: &F)
where
F: Fn(usize, usize, &mut [f32]) + Sync,
{
let n = result.len();
let segments = self.worker_row_segments_aligned(n, align);
let table = RowTable {
base: result.as_mut_ptr(),
segments: &segments,
};
let table = &table;
let job = move |global_index: usize| {
let (start, len) = table.segments[global_index];
if len == 0 {
return;
}
let outputs = unsafe { std::slice::from_raw_parts_mut(table.base.add(start), len) };
compute(global_index, start, outputs);
};
self.dispatch(&job);
}
pub fn dispatch_output_row_blocks<F>(
&self,
result: &mut [f32],
row_len: usize,
num_rows: usize,
compute: &F,
) where
F: Fn(usize, &mut [f32]) + Sync,
{
debug_assert_eq!(result.len(), row_len.saturating_mul(num_rows));
if self.total_workers <= 1 || num_rows <= 1 || row_len == 0 {
for row in 0..num_rows {
compute(row, &mut result[row * row_len..(row + 1) * row_len]);
}
return;
}
let segments = self.worker_row_segments(num_rows);
let table = RowBlockTable {
base: result.as_mut_ptr(),
row_len,
segments: &segments,
};
let table = &table;
let job = move |global_index: usize| {
let (start, len) = table.segments[global_index];
for row in start..start + len {
let slice = unsafe {
std::slice::from_raw_parts_mut(
table.base.add(row * table.row_len),
table.row_len,
)
};
compute(row, slice);
}
};
self.dispatch(&job);
}
pub fn dispatch_index_tasks<F>(&self, num_tasks: usize, compute: &F)
where
F: Fn(usize) + Sync,
{
if self.total_workers <= 1 || num_tasks <= 1 {
for task in 0..num_tasks {
compute(task);
}
return;
}
let segments = self.worker_row_segments(num_tasks);
let segments = &segments;
let job = move |global_index: usize| {
let (start, len) = segments[global_index];
for task in start..start + len {
compute(task);
}
};
self.dispatch(&job);
}
fn dispatch_rows_across_workers<F>(&self, result: &mut [f32], compute: &F)
where
F: Fn(usize, &mut [f32]) + Sync,
{
let n = result.len();
let segments = self.worker_row_segments(n);
let table = RowTable {
base: result.as_mut_ptr(),
segments: &segments,
};
let table = &table;
let job = move |global_index: usize| {
let (start, len) = table.segments[global_index];
if len == 0 {
return;
}
let outputs = unsafe { std::slice::from_raw_parts_mut(table.base.add(start), len) };
compute(start, outputs);
};
self.dispatch(&job);
}
pub fn place_rows<T: Copy + Send + Sync>(&self, src: &[T], n: usize) -> Vec<T> {
if n == 0 || src.is_empty() || self.total_workers <= 1 {
return src.to_vec();
}
let stride = src.len() / n;
debug_assert_eq!(stride * n, src.len());
let mut dst: Vec<T> = Vec::with_capacity(src.len());
#[allow(clippy::uninit_vec)]
unsafe {
dst.set_len(src.len());
}
let segments = self.worker_row_segments(n);
let table = CopyTable {
dst: dst.as_mut_ptr(),
src: src.as_ptr(),
stride,
segments: &segments,
};
let table = &table;
let job = move |global_index: usize| {
let (start, len) = table.segments[global_index];
if len == 0 {
return;
}
unsafe {
let dst = table.dst.add(start * table.stride);
let src = table.src.add(start * table.stride);
std::ptr::copy_nonoverlapping(src, dst, len * table.stride);
}
};
self.dispatch(&job);
dst
}
}
impl Drop for SpmdDecodePools {
fn drop(&mut self) {
self.shutdown();
}
}
struct RowTable<'a> {
base: *mut f32,
segments: &'a [(usize, usize)],
}
unsafe impl Sync for RowTable<'_> {}
struct RowBlockTable<'a> {
base: *mut f32,
row_len: usize,
segments: &'a [(usize, usize)],
}
unsafe impl Sync for RowBlockTable<'_> {}
struct CopyTable<'a, T> {
dst: *mut T,
src: *const T,
stride: usize,
segments: &'a [(usize, usize)],
}
unsafe impl<T: Send + Sync> Sync for CopyTable<'_, T> {}
struct WorkerCompletion<'a> {
shared: &'a SharedState,
node: usize,
global_index: usize,
}
impl WorkerCompletion<'_> {
fn complete(self) {
self.shared.node_pending[self.node]
.0
.fetch_sub(1, Ordering::AcqRel);
std::mem::forget(self);
}
}
impl Drop for WorkerCompletion<'_> {
fn drop(&mut self) {
self.shared
.poisoned_worker
.compare_exchange(
0,
self.global_index + 1,
Ordering::Release,
Ordering::Relaxed,
)
.ok();
self.shared.node_pending[self.node]
.0
.fetch_sub(1, Ordering::AcqRel);
}
}
fn worker_loop(shared: Arc<SharedState>, global_index: usize) {
let node = shared.worker_node[global_index];
let mut local_seq = 0usize;
shared.ready.fetch_add(1, Ordering::AcqRel);
loop {
let mut spins = 0u32;
let mut yields = 0u32;
let new_seq = loop {
if shared.shutdown.load(Ordering::Acquire) {
return;
}
let seq = shared.sequence.0.load(Ordering::Acquire);
if seq != local_seq {
break seq;
}
std::hint::spin_loop();
spins = spins.wrapping_add(1);
if spins >= SPIN_BEFORE_YIELD {
yields = yields.wrapping_add(1);
if yields >= YIELD_BEFORE_PARK {
shared.parked[global_index].0.store(true, Ordering::SeqCst);
while shared.sequence.0.load(Ordering::SeqCst) == local_seq
&& !shared.shutdown.load(Ordering::SeqCst)
{
thread::park_timeout(PARK_TIMEOUT);
}
shared.parked[global_index].0.store(false, Ordering::SeqCst);
yields = 0;
}
spins = 0;
}
};
local_seq = new_seq;
if shared.shutdown.load(Ordering::Acquire) {
return;
}
let job = unsafe { (*shared.job.get()).expect("published SPMD job") };
let completion = WorkerCompletion {
shared: &shared,
node,
global_index,
};
unsafe { (job.call)(job.data, global_index) };
completion.complete();
}
}
static POOLS: OnceLock<Option<SpmdDecodePools>> = OnceLock::new();
pub fn pools() -> Option<&'static SpmdDecodePools> {
POOLS
.get_or_init(|| build_from_env(default_threads()))
.as_ref()
}
pub fn shutdown_pools() {
if let Some(Some(pool)) = POOLS.get() {
pool.shutdown();
}
}
fn default_threads() -> Option<usize> {
crate::kernels::matmul_nbits::configured_persistent_decode_threads()
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum PersistenceMode {
Off,
Auto,
Forced,
}
pub(crate) fn persistence_mode_from_raw(raw: Option<&str>) -> PersistenceMode {
match raw.map(str::trim) {
Some("0") => PersistenceMode::Off,
Some("1") => PersistenceMode::Forced,
_ => PersistenceMode::Auto,
}
}
fn persistence_mode() -> PersistenceMode {
persistence_mode_from_raw(std::env::var(PERSISTENT_POOL_ENV).ok().as_deref())
}
fn pool_mode_builds(mode: PersistenceMode) -> bool {
matches!(mode, PersistenceMode::Forced | PersistenceMode::Auto)
}
fn pool_mode_forces(mode: PersistenceMode) -> bool {
matches!(mode, PersistenceMode::Forced)
}
pub(crate) fn is_forced() -> bool {
pool_mode_forces(persistence_mode())
}
pub fn build_from_env(threads: Option<usize>) -> Option<SpmdDecodePools> {
let mode = persistence_mode();
if !pool_mode_builds(mode) {
return None;
}
if matches!(mode, PersistenceMode::Auto) && explicit_decode_affinity_set() {
return None;
}
let Some(total) = threads else {
report_spmd_fallback(
"ONNX_GENAI_CPU_DECODE_THREADS=0 opts out of the bounded pool; the persistent \
SPMD pool needs a bounded worker count -- leaving the decode path unchanged",
);
return None;
};
if total == 0 {
return None;
}
if let Some(allowed) = crate::decode_affinity::allowed_cpus()
&& allowed.len() == 1
{
report_spmd_fallback(
"the process is confined to a single CPU (cpuset/taskset), which leaves no core \
for the inline dispatcher alongside a spinning worker -- leaving decode on the \
flat path instead of starving the persistent SPMD pool",
);
return None;
}
report_pool_built(mode);
let shards = node_shards(total);
Some(SpmdDecodePools::build(&shards))
}
fn explicit_decode_affinity_set() -> bool {
std::env::var(crate::decode_affinity::DECODE_AFFINITY_ENV)
.ok()
.is_some_and(|value| !value.trim().is_empty())
}
fn node_shards(total: usize) -> Vec<NodeShard> {
let allowed = crate::decode_affinity::allowed_cpus();
if let Some(topology) = NumaTopology::detect() {
let topology = topology.restrict_to_allowed(allowed.as_deref());
if let Some(mut shards) = topology.split_workers(total) {
reserve_split_headroom(&mut shards);
return shards;
}
}
let cpus = allowed.unwrap_or_default();
let workers = reserve_single_group_headroom(total, cpus.len());
vec![NodeShard {
index: 0,
cpus,
workers,
}]
}
const DISPATCHER_RESERVED_CPUS: usize = 1;
fn reserve_single_group_headroom(total: usize, allowed_count: usize) -> usize {
if allowed_count == 0 || total < allowed_count {
return total;
}
allowed_count
.saturating_sub(DISPATCHER_RESERVED_CPUS)
.max(1)
}
fn reserve_split_headroom(shards: &mut [NodeShard]) {
for shard in shards.iter_mut() {
let cap = shard
.cpus
.len()
.saturating_sub(DISPATCHER_RESERVED_CPUS)
.max(1);
shard.workers = shard.workers.min(cap);
}
}
fn report_spmd_fallback(message: &str) {
static REPORTED: OnceLock<()> = OnceLock::new();
if REPORTED.set(()).is_ok() {
eprintln!("onnx-genai: persistent SPMD decode pool: {message}");
}
}
fn report_pool_built(mode: PersistenceMode) {
static REPORTED: OnceLock<()> = OnceLock::new();
if REPORTED.set(()).is_ok() {
match mode {
PersistenceMode::Forced => eprintln!(
"onnx-genai: persistent SPMD decode pool forced on via \
ONNX_GENAI_CPU_DECODE_PERSISTENT_POOL=1 (always dispatches to the pool)"
),
_ => eprintln!(
"onnx-genai: persistent SPMD decode pool built for auto-calibration \
(ONNX_GENAI_CPU_DECODE_PERSISTENT_POOL unset); each decode step is timed \
both ways and the faster path is kept -- the flat path stays committed \
under load. Set =0 to force flat, =1 to force the pool"
),
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum AutoPath {
Pool,
Flat,
}
const CALIB_WARMUP_STEPS: u64 = 2;
const CALIB_PROBE_SAMPLES: usize = 5;
const CALIB_RECAL_PERIOD: u64 = 600;
const CALIB_SWITCH_MARGIN_PCT: u64 = 8;
const CALIB_PROBE_DISCARD: usize = 1;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum CalibPhase {
Warmup,
ProbeFlat,
ProbePool,
Committed,
}
struct Calibrator {
phase: CalibPhase,
warmup_left: u64,
discard_left: usize,
pool_ns: Vec<u64>,
flat_ns: Vec<u64>,
committed: AutoPath,
committed_left: u64,
}
impl Calibrator {
fn new() -> Self {
Self {
phase: CalibPhase::Warmup,
warmup_left: CALIB_WARMUP_STEPS,
discard_left: 0,
pool_ns: Vec::with_capacity(CALIB_PROBE_SAMPLES),
flat_ns: Vec::with_capacity(CALIB_PROBE_SAMPLES),
committed: AutoPath::Flat,
committed_left: 0,
}
}
fn choose(&self) -> AutoPath {
match self.phase {
CalibPhase::Warmup | CalibPhase::ProbePool => AutoPath::Pool,
CalibPhase::ProbeFlat => AutoPath::Flat,
CalibPhase::Committed => self.committed,
}
}
fn record(&mut self, path: AutoPath, ns: u64) {
match self.phase {
CalibPhase::Warmup => {
self.warmup_left = self.warmup_left.saturating_sub(1);
if self.warmup_left == 0 {
self.enter_flat_probe();
}
}
CalibPhase::ProbeFlat => {
if path == AutoPath::Flat {
self.push_sample_or_discard(ns, true);
}
if self.flat_ns.len() >= CALIB_PROBE_SAMPLES {
self.enter_pool_probe();
}
}
CalibPhase::ProbePool => {
if path == AutoPath::Pool {
self.push_sample_or_discard(ns, false);
}
if self.pool_ns.len() >= CALIB_PROBE_SAMPLES {
self.commit_from_samples();
}
}
CalibPhase::Committed => {
self.committed_left = self.committed_left.saturating_sub(1);
if self.committed_left == 0 {
self.enter_flat_probe();
}
}
}
}
fn push_sample_or_discard(&mut self, ns: u64, flat: bool) {
if self.discard_left > 0 {
self.discard_left -= 1;
return;
}
let block = if flat {
&mut self.flat_ns
} else {
&mut self.pool_ns
};
if block.len() < CALIB_PROBE_SAMPLES {
block.push(ns);
}
}
fn enter_flat_probe(&mut self) {
self.phase = CalibPhase::ProbeFlat;
self.flat_ns.clear();
self.pool_ns.clear();
self.discard_left = CALIB_PROBE_DISCARD;
}
fn enter_pool_probe(&mut self) {
self.phase = CalibPhase::ProbePool;
self.discard_left = CALIB_PROBE_DISCARD;
}
fn commit_from_samples(&mut self) {
let pool = median_ns(&mut self.pool_ns);
let flat = median_ns(&mut self.flat_ns);
let pool_scaled = u128::from(pool) * 100;
let flat_scaled = u128::from(flat) * u128::from(100 - CALIB_SWITCH_MARGIN_PCT);
self.committed = if pool_scaled <= flat_scaled {
AutoPath::Pool
} else {
AutoPath::Flat
};
self.phase = CalibPhase::Committed;
self.committed_left = CALIB_RECAL_PERIOD;
self.pool_ns.clear();
self.flat_ns.clear();
}
}
fn median_ns(samples: &mut [u64]) -> u64 {
if samples.is_empty() {
return u64::MAX;
}
samples.sort_unstable();
samples[samples.len() / 2]
}
fn calibrator() -> &'static Mutex<Calibrator> {
static CALIBRATOR: OnceLock<Mutex<Calibrator>> = OnceLock::new();
CALIBRATOR.get_or_init(|| Mutex::new(Calibrator::new()))
}
pub(crate) fn auto_choose_path() -> AutoPath {
calibrator()
.lock()
.map(|calib| calib.choose())
.unwrap_or(AutoPath::Flat)
}
pub(crate) fn auto_record_sample(path: AutoPath, elapsed: Duration) {
let ns = u64::try_from(elapsed.as_nanos()).unwrap_or(u64::MAX);
if let Ok(mut calib) = calibrator().lock() {
calib.record(path, ns);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn two_group_pool() -> SpmdDecodePools {
let shards = vec![
NodeShard {
index: 0,
cpus: vec![],
workers: 2,
},
NodeShard {
index: 1,
cpus: vec![],
workers: 2,
},
];
SpmdDecodePools::build(&shards)
}
fn single_group_pool(workers: usize) -> SpmdDecodePools {
let shards = vec![NodeShard {
index: 0,
cpus: vec![],
workers,
}];
SpmdDecodePools::build(&shards)
}
#[test]
fn reserve_single_group_headroom_frees_a_dispatcher_cpu_when_fully_subscribed() {
assert_eq!(reserve_single_group_headroom(32, 32), 31);
assert_eq!(reserve_single_group_headroom(40, 32), 31);
assert_eq!(reserve_single_group_headroom(2, 2), 1);
}
#[test]
fn reserve_single_group_headroom_is_a_noop_when_headroom_exists_or_affinity_unknown() {
assert_eq!(reserve_single_group_headroom(16, 32), 16);
assert_eq!(reserve_single_group_headroom(31, 32), 31);
assert_eq!(reserve_single_group_headroom(32, 0), 32);
assert_eq!(reserve_single_group_headroom(1, 0), 1);
}
#[test]
fn reserve_split_headroom_reserves_one_cpu_per_node_only_when_fully_subscribed() {
let mut shards = vec![
NodeShard {
index: 0,
cpus: (0..16).collect(),
workers: 16,
},
NodeShard {
index: 1,
cpus: (16..32).collect(),
workers: 10,
},
];
reserve_split_headroom(&mut shards);
assert_eq!(shards[0].workers, 15);
assert_eq!(shards[1].workers, 10);
assert_eq!(shards[0].cpus.len(), 16);
}
#[test]
fn reserve_split_headroom_floors_at_one_worker_per_node() {
let mut shards = vec![NodeShard {
index: 0,
cpus: vec![7],
workers: 1,
}];
reserve_split_headroom(&mut shards);
assert_eq!(shards[0].workers, 1);
}
#[test]
fn node_row_lengths_split_proportionally_and_cover_all_rows() {
let pool = two_group_pool();
assert_eq!(pool.node_row_lengths(100), vec![50, 50]);
assert_eq!(pool.node_row_lengths(101), vec![50, 51]);
assert_eq!(pool.node_row_lengths(1), vec![0, 1]);
assert_eq!(pool.node_row_lengths(0), vec![0, 0]);
}
#[test]
fn worker_row_segments_are_disjoint_and_cover_every_row() {
let pool = two_group_pool();
let n = 37usize;
let segments = pool.worker_row_segments(n);
assert_eq!(segments.len(), pool.total_workers());
let mut expected_start = 0;
for (start, len) in &segments {
assert_eq!(*start, expected_start);
expected_start += len;
}
assert_eq!(expected_start, n);
}
#[test]
fn worker_row_segments_aligned_snaps_interior_boundaries_and_covers_every_row() {
let pool = single_group_pool(3);
for &align in &[4usize, 16] {
for &n in &[97usize, 128, 151936, 1, 0, 5, 17] {
let segments = pool.worker_row_segments_aligned(n, align);
assert_eq!(segments.len(), pool.total_workers());
let mut expected_start = 0;
for (index, &(start, len)) in segments.iter().enumerate() {
assert_eq!(start, expected_start, "n={n} align={align} seg {index}");
assert_eq!(
start % align,
0,
"n={n} align={align}: segment start {start} not aligned"
);
expected_start += len;
}
assert_eq!(expected_start, n, "n={n} align={align}: must cover 0..n");
}
}
}
#[test]
fn worker_row_segments_aligned_is_identity_for_align_one() {
let pool = two_group_pool();
for &n in &[0usize, 1, 37, 100, 101] {
assert_eq!(
pool.worker_row_segments_aligned(n, 1),
pool.worker_row_segments(n),
"align=1 must reproduce the unaligned split (n={n})"
);
}
}
#[test]
fn dispatch_output_rows_matches_flat_computation() {
let pool = two_group_pool();
let n = 101usize;
let compute = |output_start: usize, outputs: &mut [f32]| {
for (offset, out) in outputs.iter_mut().enumerate() {
*out = (output_start + offset) as f32 * 2.5 - 3.0;
}
};
let mut sharded = vec![0.0f32; n];
pool.dispatch_rows_across_workers(&mut sharded, &compute);
let mut flat = vec![0.0f32; n];
compute(0, &mut flat);
assert_eq!(sharded, flat);
}
#[test]
fn dispatch_preserves_per_row_reduction_bit_for_bit() {
let pool = two_group_pool();
let n = 257usize;
let k = 320usize;
let activation: Vec<f32> = (0..k)
.map(|i| ((i * 37 % 101) as f32 - 50.0) * 0.031_25)
.collect();
let weight = |row: usize, col: usize| -> f32 {
(((row * 131 + col * 17) % 251) as f32 - 125.0) * 0.007_812_5
};
let compute = |output_start: usize, outputs: &mut [f32]| {
for (offset, out) in outputs.iter_mut().enumerate() {
let row = output_start + offset;
let mut acc = 0.0f32;
for (col, &a) in activation.iter().enumerate() {
acc += a * weight(row, col);
}
*out = acc;
}
};
let mut sharded = vec![0.0f32; n];
pool.dispatch_rows_across_workers(&mut sharded, &compute);
let mut reference = vec![0.0f32; n];
compute(0, &mut reference);
assert_eq!(
sharded.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
reference.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"row-sharded dispatch must be bit-identical to the serial reference"
);
}
#[test]
fn dispatch_output_row_blocks_matches_flat_computation() {
for (num_rows, row_len) in [
(28usize, 128usize),
(3, 128),
(1, 64),
(5, 3),
(37, 1),
(0, 8),
] {
let pool = two_group_pool();
let compute = |row_index: usize, row: &mut [f32]| {
for (offset, out) in row.iter_mut().enumerate() {
let mut acc = 0.0f32;
for step in 0..=offset {
acc += (row_index * 7 + step) as f32 * 0.015_625 - 1.0;
}
*out = acc;
}
};
let mut sharded = vec![0.0f32; num_rows * row_len];
pool.dispatch_output_row_blocks(&mut sharded, row_len, num_rows, &compute);
let mut reference = vec![0.0f32; num_rows * row_len];
for row in 0..num_rows {
compute(row, &mut reference[row * row_len..(row + 1) * row_len]);
}
assert_eq!(
sharded.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
reference.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"row-block dispatch must be bit-identical to the serial reference \
(num_rows={num_rows}, row_len={row_len})"
);
}
}
#[test]
fn dispatch_is_reusable_across_many_ops() {
let pool = single_group_pool(4);
for round in 0..200usize {
let n = 53usize;
let compute = move |output_start: usize, outputs: &mut [f32]| {
for (offset, out) in outputs.iter_mut().enumerate() {
*out = (round * 1000 + output_start + offset) as f32;
}
};
let mut got = vec![0.0f32; n];
pool.dispatch_rows_across_workers(&mut got, &compute);
let mut want = vec![0.0f32; n];
compute(0, &mut want);
assert_eq!(got, want, "round {round}");
}
}
#[test]
fn build_then_immediate_dispatch_never_hangs() {
for _ in 0..40usize {
let pool = single_group_pool(6);
let n = 61usize;
let compute = |output_start: usize, outputs: &mut [f32]| {
for (offset, out) in outputs.iter_mut().enumerate() {
*out = (output_start + offset) as f32;
}
};
let mut got = vec![-1.0f32; n];
pool.dispatch_rows_across_workers(&mut got, &compute);
let mut want = vec![0.0f32; n];
compute(0, &mut want);
assert_eq!(got, want);
}
}
#[test]
fn place_rows_preserves_bytes() {
let pool = two_group_pool();
let n = 7usize;
let stride = 4usize;
let src: Vec<u8> = (0..(n * stride) as u8).collect();
assert_eq!(pool.place_rows(&src, n), src);
let scales: Vec<f32> = (0..n).map(|row| row as f32 * 0.5).collect();
assert_eq!(pool.place_rows(&scales, n), scales);
}
#[test]
fn tiny_ops_run_serially_but_correctly() {
let pool = single_group_pool(8);
let n = 3usize;
let compute = |output_start: usize, outputs: &mut [f32]| {
for (offset, out) in outputs.iter_mut().enumerate() {
*out = (output_start + offset) as f32;
}
};
let mut got = vec![0.0f32; n];
pool.dispatch_output_rows(&mut got, 4096, &compute);
assert_eq!(got, vec![0.0, 1.0, 2.0]);
}
#[test]
fn panicking_worker_poison_is_reported_without_hanging() {
let pool = single_group_pool(4);
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
pool.dispatch(&|worker| {
assert_ne!(worker, 2, "intentional SPMD worker panic");
});
}));
let panic = result.expect_err("dispatcher must report a worker panic");
let message = panic
.downcast_ref::<String>()
.map(String::as_str)
.or_else(|| panic.downcast_ref::<&str>().copied())
.unwrap_or("");
assert!(
message.contains("persistent SPMD decode worker 2 panicked")
&& message.contains("pool is poisoned"),
"unexpected dispatcher diagnostic: {message}"
);
}
#[test]
fn persistence_mode_parses_env_values() {
assert_eq!(persistence_mode_from_raw(None), PersistenceMode::Auto);
assert_eq!(persistence_mode_from_raw(Some("")), PersistenceMode::Auto);
assert_eq!(
persistence_mode_from_raw(Some(" ")),
PersistenceMode::Auto
);
assert_eq!(persistence_mode_from_raw(Some("0")), PersistenceMode::Off);
assert_eq!(persistence_mode_from_raw(Some(" 0 ")), PersistenceMode::Off);
assert_eq!(
persistence_mode_from_raw(Some("1")),
PersistenceMode::Forced
);
assert_eq!(
persistence_mode_from_raw(Some(" 1 ")),
PersistenceMode::Forced
);
assert_eq!(
persistence_mode_from_raw(Some("true")),
PersistenceMode::Auto
);
assert_eq!(persistence_mode_from_raw(Some("2")), PersistenceMode::Auto);
assert!(pool_mode_builds(persistence_mode_from_raw(Some("1"))));
assert!(pool_mode_builds(persistence_mode_from_raw(None)));
assert!(!pool_mode_builds(persistence_mode_from_raw(Some("0"))));
assert!(pool_mode_forces(persistence_mode_from_raw(Some("1"))));
assert!(!pool_mode_forces(persistence_mode_from_raw(None)));
}
#[test]
fn auto_and_forced_build_the_pool_but_only_forced_dispatches_unconditionally() {
assert!(pool_mode_builds(PersistenceMode::Auto));
assert!(pool_mode_builds(PersistenceMode::Forced));
assert!(!pool_mode_builds(PersistenceMode::Off));
assert!(pool_mode_forces(PersistenceMode::Forced));
assert!(!pool_mode_forces(PersistenceMode::Auto));
assert!(!pool_mode_forces(PersistenceMode::Off));
assert!(pool_mode_builds(persistence_mode_from_raw(None)));
assert!(!pool_mode_forces(persistence_mode_from_raw(None)));
assert!(!pool_mode_builds(persistence_mode_from_raw(Some("0"))));
assert!(pool_mode_builds(persistence_mode_from_raw(Some("2"))));
assert!(pool_mode_forces(persistence_mode_from_raw(Some("1"))));
}
fn drive_to_commit(calib: &mut Calibrator, pool_ns: u64, flat_ns: u64) {
for _ in 0..100_000 {
if calib.phase == CalibPhase::Committed {
return;
}
let path = calib.choose();
let ns = match path {
AutoPath::Pool => pool_ns,
AutoPath::Flat => flat_ns,
};
calib.record(path, ns);
}
panic!("calibrator never reached the committed phase");
}
fn run_one_probe(pool_ns: u64, flat_ns: u64) -> Calibrator {
let mut calib = Calibrator::new();
drive_to_commit(&mut calib, pool_ns, flat_ns);
calib
}
#[test]
fn calibrator_defaults_to_flat_before_any_measurement() {
let calib = Calibrator::new();
assert_eq!(calib.committed, AutoPath::Flat);
}
#[test]
fn calibrator_probe_measures_flat_block_before_pool_block() {
let mut calib = Calibrator::new();
for _ in 0..CALIB_WARMUP_STEPS {
assert_eq!(calib.choose(), AutoPath::Pool);
assert_eq!(calib.phase, CalibPhase::Warmup);
calib.record(AutoPath::Pool, 1_000);
}
assert_eq!(calib.phase, CalibPhase::ProbeFlat);
while calib.phase == CalibPhase::ProbeFlat {
assert_eq!(calib.choose(), AutoPath::Flat);
calib.record(AutoPath::Flat, 100);
}
assert_eq!(calib.phase, CalibPhase::ProbePool);
while calib.phase == CalibPhase::ProbePool {
assert_eq!(calib.choose(), AutoPath::Pool);
calib.record(AutoPath::Pool, 100);
}
assert_eq!(calib.phase, CalibPhase::Committed);
}
#[test]
fn calibrator_probe_discards_the_transition_sample() {
let mut calib = Calibrator::new();
for _ in 0..CALIB_WARMUP_STEPS {
calib.record(AutoPath::Pool, 1);
}
assert_eq!(calib.phase, CalibPhase::ProbeFlat);
assert_eq!(calib.discard_left, CALIB_PROBE_DISCARD);
while calib.phase == CalibPhase::ProbeFlat {
calib.record(AutoPath::Flat, 100);
}
assert_eq!(calib.flat_ns.len(), CALIB_PROBE_SAMPLES);
}
#[test]
fn calibrator_commits_pool_only_when_clearly_faster() {
let calib = run_one_probe(80, 100);
assert_eq!(calib.phase, CalibPhase::Committed);
assert_eq!(calib.committed, AutoPath::Pool);
assert_eq!(calib.choose(), AutoPath::Pool);
}
#[test]
fn calibrator_stays_flat_when_pool_slower_simulating_contention() {
let calib = run_one_probe(200, 100);
assert_eq!(calib.committed, AutoPath::Flat);
assert_eq!(calib.choose(), AutoPath::Flat);
}
#[test]
fn calibrator_stays_flat_within_the_hysteresis_margin() {
let calib = run_one_probe(95, 100);
assert_eq!(calib.committed, AutoPath::Flat);
}
#[test]
fn calibrator_reprobes_after_the_recal_period_and_can_fall_back() {
let mut calib = run_one_probe(80, 100);
assert_eq!(calib.committed, AutoPath::Pool);
for _ in 0..CALIB_RECAL_PERIOD {
assert_eq!(calib.choose(), AutoPath::Pool);
calib.record(AutoPath::Pool, 80);
}
assert_eq!(calib.phase, CalibPhase::ProbeFlat);
drive_to_commit(&mut calib, 300, 100);
assert_eq!(calib.committed, AutoPath::Flat);
}
#[test]
fn calibrator_probe_median_rejects_a_single_load_spike() {
let mut calib = Calibrator::new();
for _ in 0..CALIB_WARMUP_STEPS {
calib.record(AutoPath::Pool, 80);
}
while calib.phase == CalibPhase::ProbeFlat {
calib.record(AutoPath::Flat, 100);
}
let mut pool_samples = [80u64, 80, 80, 80, 80, 100_000].into_iter();
while calib.phase == CalibPhase::ProbePool {
calib.record(AutoPath::Pool, pool_samples.next().unwrap_or(80));
}
assert_eq!(calib.committed, AutoPath::Pool);
}
#[test]
fn median_ns_picks_the_middle_and_guards_empty() {
assert_eq!(median_ns(&mut []), u64::MAX);
assert_eq!(median_ns(&mut [5]), 5);
assert_eq!(median_ns(&mut [30, 10, 20]), 20);
assert_eq!(median_ns(&mut [10, 40, 20, 30]), 30);
}
}