use std::cell::UnsafeCell;
use std::sync::Arc;
#[cfg(any(target_os = "android", target_os = "linux"))]
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
pub static FORCED_THREADS: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
pub static WORKER_TIDS: std::sync::Mutex<Vec<i32>> = std::sync::Mutex::new(Vec::new());
#[derive(Clone, Copy)]
struct TaskPtr(*const (dyn Fn(usize, usize) + Sync));
unsafe impl Send for TaskPtr {}
struct Inner {
epoch: AtomicUsize,
remaining: AtomicUsize,
slot: UnsafeCell<Option<(TaskPtr, usize, usize, usize)>>,
shutdown: AtomicBool,
spin_budget: AtomicUsize,
parked: Box<[AtomicBool]>,
#[cfg(any(target_os = "android", target_os = "linux"))]
registered: AtomicUsize,
#[cfg(any(target_os = "android", target_os = "linux"))]
worker_tids: Mutex<Vec<i32>>,
}
unsafe impl Sync for Inner {}
static DISPATCHES: AtomicUsize = AtomicUsize::new(0);
pub fn dispatch_count() -> usize {
DISPATCHES.load(Ordering::Relaxed)
}
pub struct Pool {
inner: Arc<Inner>,
threads: Vec<std::thread::Thread>,
joins: Vec<std::thread::JoinHandle<()>>,
}
fn spin_budget_from_env() -> usize {
std::env::var("CMF_POOL_SPIN")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(4000)
}
pub(crate) fn grain_for(rows: usize, workers: usize) -> usize {
if rows == 0 || workers <= 1 {
return rows.max(1);
}
let balanced = (rows / (workers * 8)).max(32);
balanced.min(rows.div_ceil(workers)).max(1)
}
impl Pool {
pub fn new(n_workers: usize) -> Self {
Self::with_spin(n_workers, spin_budget_from_env())
}
pub fn with_spin(n_workers: usize, spin_budget: usize) -> Self {
let inner = Arc::new(Inner {
epoch: AtomicUsize::new(0),
remaining: AtomicUsize::new(0),
slot: UnsafeCell::new(None),
shutdown: AtomicBool::new(false),
spin_budget: AtomicUsize::new(spin_budget),
parked: (0..n_workers).map(|_| AtomicBool::new(false)).collect(),
#[cfg(any(target_os = "android", target_os = "linux"))]
registered: AtomicUsize::new(0),
#[cfg(any(target_os = "android", target_os = "linux"))]
worker_tids: Mutex::new(Vec::with_capacity(n_workers)),
});
let mut joins = Vec::with_capacity(n_workers);
for w in 0..n_workers {
let inner = inner.clone();
let h = std::thread::Builder::new()
.name(format!("cmf-pool-{w}"))
.spawn(move || {
#[cfg(any(target_os = "android", target_os = "linux"))]
{
let tid = unsafe { libc::gettid() } as i32;
if let Ok(mut tids) = inner.worker_tids.lock() {
tids.push(tid);
}
inner.registered.fetch_add(1, Ordering::Release);
}
worker_loop(&inner, w)
})
.expect("spawn pool worker");
joins.push(h);
}
#[cfg(any(target_os = "android", target_os = "linux"))]
while inner.registered.load(Ordering::Acquire) < n_workers {
std::thread::yield_now();
}
#[cfg(any(target_os = "android", target_os = "linux"))]
if let (Ok(mut global), Ok(local)) = (WORKER_TIDS.lock(), inner.worker_tids.lock()) {
*global = local.clone();
}
let threads = joins.iter().map(|h| h.thread().clone()).collect();
Self {
inner,
threads,
joins,
}
}
#[cfg(all(
target_arch = "aarch64",
any(target_os = "linux", target_os = "android")
))]
fn big_cores() -> Option<usize> {
Self::cores_from_capacities(&core_capacities())
}
#[cfg_attr(
not(all(
target_arch = "aarch64",
any(target_os = "linux", target_os = "android")
)),
allow(dead_code)
)]
fn cores_from_capacities(caps: &[u64]) -> Option<usize> {
let max = *caps.iter().max()?;
let min = *caps.iter().min()?;
if caps.len() < 2 || max == min {
return None;
}
Some(caps.iter().filter(|&&c| c * 8 >= max * 5).count())
}
#[cfg(target_os = "macos")]
fn big_cores() -> Option<usize> {
if true {
return None;
}
#[allow(unreachable_code)]
unsafe extern "C" {
fn sysctlbyname(
name: *const std::ffi::c_char,
oldp: *mut std::ffi::c_void,
oldlenp: *mut usize,
newp: *mut std::ffi::c_void,
newlen: usize,
) -> std::ffi::c_int;
}
unsafe {
let name = std::ffi::CString::new("hw.perflevel0.physicalcpu").ok()?;
let mut count: i32 = 0;
let mut size = std::mem::size_of::<i32>();
let ret = sysctlbyname(
name.as_ptr(),
&mut count as *mut i32 as *mut std::ffi::c_void,
&mut size,
std::ptr::null_mut(),
0,
);
if ret == 0 && count > 0 {
Some(count as usize)
} else {
None
}
}
}
#[cfg(not(any(
all(
target_arch = "aarch64",
any(target_os = "linux", target_os = "android")
),
target_os = "macos"
)))]
fn big_cores() -> Option<usize> {
None
}
pub fn effective_threads() -> usize {
let forced = FORCED_THREADS.load(std::sync::atomic::Ordering::Relaxed);
if forced > 0 {
return forced;
}
match std::env::var("CMF_THREADS") {
Ok(v) => v.parse::<usize>().unwrap_or(0),
Err(_) => match Self::big_cores() {
Some(big) => big,
None => {
let avail = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
avail.saturating_sub(1).min(32)
}
},
}
}
pub fn from_env() -> Option<Arc<Self>> {
let n = Self::effective_threads();
if n <= 1 {
None
} else {
Some(Arc::new(Self::new(n)))
}
}
pub fn n_workers(&self) -> usize {
self.threads.len()
}
pub fn bind_numa(&self, regions: &[&[u8]]) {
#[cfg(target_os = "linux")]
{
let Some((node, cpus)) = numa::choose(regions, self.threads.len() + 1) else {
return;
};
let mut applied = 0usize;
if let Ok(tids) = self.inner.worker_tids.lock() {
for &tid in tids.iter() {
if numa::set_affinity(tid, &cpus) {
applied += 1;
}
}
}
numa::publish(cpus.clone());
numa::adopt_caller();
tracing::info!(
"numa: pool bound to node {node} ({} cpus, {applied}/{} workers)",
cpus.len(),
self.threads.len()
);
if std::env::var("CMF_NUMA_TRACE").is_ok_and(|v| v != "0") {
eprintln!(
"numa: pool bound to node {node}: {} cpus, {applied}/{} workers",
cpus.len(),
self.threads.len()
);
}
}
#[cfg(not(target_os = "linux"))]
let _ = regions;
}
pub(crate) fn set_spin_budget(&self, spins: usize) {
self.inner.spin_budget.store(spins, Ordering::Relaxed);
}
pub fn run_rows(&self, rows: usize, f: &(dyn Fn(usize, usize) + Sync)) {
let grain = grain_for(rows, self.threads.len() + 1);
let chunks = rows.div_ceil(grain.max(1));
let next = AtomicUsize::new(0);
self.run_limited(chunks, &|_w, _n| loop {
let start = next.fetch_add(grain, Ordering::Relaxed);
if start >= rows {
break;
}
f(start, (start + grain).min(rows));
});
}
fn run_limited(&self, max_workers: usize, f: &(dyn Fn(usize, usize) + Sync)) {
#[cfg(target_os = "linux")]
numa::adopt_caller();
let nw = self.threads.len().min(max_workers);
if nw == self.threads.len() {
return self.run(f);
}
DISPATCHES.fetch_add(1, Ordering::Relaxed);
let ptr: *const (dyn Fn(usize, usize) + Sync) = f;
let ptr: *const (dyn Fn(usize, usize) + Sync + 'static) =
unsafe { std::mem::transmute(ptr) };
let dev = crate::gpu::current_device();
unsafe { *self.inner.slot.get() = Some((TaskPtr(ptr), nw + 1, dev, nw)) };
self.inner.remaining.store(nw, Ordering::Relaxed);
self.inner.epoch.fetch_add(1, Ordering::SeqCst);
for (i, t) in self.threads.iter().enumerate().take(nw) {
if self.inner.parked[i].load(Ordering::SeqCst) {
t.unpark();
}
}
f(nw, nw + 1);
let mut spins = 0usize;
while self.inner.remaining.load(Ordering::Acquire) != 0 {
spins += 1;
if spins < 10_000 {
std::hint::spin_loop();
} else {
std::thread::yield_now();
}
}
}
pub fn run_many(&self, parts: &[(usize, &(dyn Fn(usize, usize) + Sync))]) {
let total: usize = parts.iter().map(|p| p.0).sum();
if total == 0 {
return;
}
let grain = grain_for(total, self.threads.len() + 1);
let chunks = total.div_ceil(grain.max(1));
let next = AtomicUsize::new(0);
self.run_limited(chunks, &|_w, _n| loop {
let s = next.fetch_add(grain, Ordering::Relaxed);
if s >= total {
break;
}
let e = (s + grain).min(total);
let mut base = 0usize;
for &(rows, f) in parts {
let a = s.max(base);
let b = e.min(base + rows);
if a < b {
f(a - base, b - base);
}
base += rows;
if base >= e {
break;
}
}
});
}
pub fn run(&self, f: &(dyn Fn(usize, usize) + Sync)) {
#[cfg(target_os = "linux")]
numa::adopt_caller();
DISPATCHES.fetch_add(1, Ordering::Relaxed);
let nw = self.threads.len();
let n = nw + 1; let ptr: *const (dyn Fn(usize, usize) + Sync) = f;
let ptr: *const (dyn Fn(usize, usize) + Sync + 'static) =
unsafe { std::mem::transmute(ptr) };
let dev = crate::gpu::current_device();
unsafe { *self.inner.slot.get() = Some((TaskPtr(ptr), n, dev, nw)) };
self.inner.remaining.store(nw, Ordering::Relaxed);
self.inner.epoch.fetch_add(1, Ordering::SeqCst);
for (i, t) in self.threads.iter().enumerate() {
if self.inner.parked[i].load(Ordering::SeqCst) {
t.unpark();
}
}
f(nw, n);
let mut spins = 0usize;
while self.inner.remaining.load(Ordering::Acquire) != 0 {
spins += 1;
if spins < 10_000 {
std::hint::spin_loop();
} else {
std::thread::yield_now();
}
}
}
}
impl Drop for Pool {
fn drop(&mut self) {
self.inner.shutdown.store(true, Ordering::SeqCst);
for t in &self.threads {
t.unpark();
}
for h in self.joins.drain(..) {
let _ = h.join();
}
}
}
#[cfg(target_os = "linux")]
mod numa {
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
static MASK: Mutex<Vec<usize>> = Mutex::new(Vec::new());
static EPOCH: AtomicUsize = AtomicUsize::new(0);
thread_local! {
static SEEN: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
pub(super) fn publish(cpus: Vec<usize>) {
if let Ok(mut m) = MASK.lock() {
*m = cpus;
}
EPOCH.fetch_add(1, Ordering::Release);
}
#[inline]
pub(super) fn adopt_caller() {
let e = EPOCH.load(Ordering::Acquire);
if e == 0 || SEEN.with(|c| c.get()) == e {
return;
}
SEEN.with(|c| c.set(e));
if let Ok(m) = MASK.lock() {
if !m.is_empty() {
set_affinity(0, &m);
}
}
}
pub(super) fn parse_list(s: &str) -> Vec<usize> {
let mut out = Vec::new();
for part in s.trim().split(',') {
let part = part.trim();
if part.is_empty() {
continue;
}
match part.split_once('-') {
Some((a, b)) => {
if let (Ok(a), Ok(b)) = (a.parse::<usize>(), b.parse::<usize>()) {
out.extend(a..=b);
}
}
None => {
if let Ok(a) = part.parse() {
out.push(a);
}
}
}
}
out
}
fn nodes() -> Vec<(usize, Vec<usize>)> {
let mut v = Vec::new();
let Ok(rd) = std::fs::read_dir("/sys/devices/system/node") else {
return v;
};
for e in rd.flatten() {
let name = e.file_name().to_string_lossy().to_string();
let Some(id) = name.strip_prefix("node").and_then(|x| x.parse::<usize>().ok()) else {
continue;
};
if let Ok(l) = std::fs::read_to_string(e.path().join("cpulist")) {
let cpus = parse_list(&l);
if !cpus.is_empty() {
v.push((id, cpus));
}
}
}
v.sort();
v
}
fn allowed() -> Vec<usize> {
unsafe {
let mut set: libc::cpu_set_t = std::mem::zeroed();
if libc::sched_getaffinity(0, std::mem::size_of::<libc::cpu_set_t>(), &mut set) != 0 {
return Vec::new();
}
(0..libc::CPU_SETSIZE as usize)
.filter(|&c| libc::CPU_ISSET(c, &set))
.collect()
}
}
fn primary(cpu: usize) -> bool {
let p = format!("/sys/devices/system/cpu/cpu{cpu}/topology/thread_siblings_list");
match std::fs::read_to_string(p) {
Ok(l) => parse_list(&l).first().is_none_or(|&f| f == cpu),
Err(_) => true,
}
}
fn page_nodes(regions: &[&[u8]]) -> (usize, usize, Vec<usize>) {
const PAGE: usize = 4096;
let total: usize = regions.iter().map(|r| r.len() / PAGE).sum();
if total == 0 {
return (0, 0, Vec::new());
}
let stride = total.div_ceil(4096).max(1);
let mut pages: Vec<*mut libc::c_void> = Vec::new();
for r in regions {
let base = (r.as_ptr() as usize).div_ceil(PAGE) * PAGE;
let end = r.as_ptr() as usize + r.len();
let mut a = base;
while a + PAGE <= end {
pages.push(a as *mut libc::c_void);
a += PAGE * stride;
}
}
let mut incore = 0usize;
for &p in &pages {
let mut vec = 0u8;
let cached = unsafe { libc::mincore(p, PAGE, &mut vec) } == 0 && vec & 1 == 1;
if cached {
incore += 1;
unsafe { std::ptr::read_volatile(p as *const u8) };
}
}
let mut status = vec![-1i32; pages.len()];
let rc = unsafe {
libc::syscall(
libc::SYS_move_pages,
0,
pages.len() as libc::c_ulong,
pages.as_mut_ptr(),
std::ptr::null::<libc::c_int>(),
status.as_mut_ptr(),
0,
)
};
let mut by = Vec::new();
if rc == 0 {
for &st in &status {
if st >= 0 {
let n = st as usize;
if by.len() <= n {
by.resize(n + 1, 0);
}
by[n] += 1;
}
}
} else {
by = numa_maps_nodes(regions);
}
(pages.len(), incore, by)
}
fn numa_maps_nodes(regions: &[&[u8]]) -> Vec<usize> {
let (Ok(maps), Ok(nm)) = (
std::fs::read_to_string("/proc/self/maps"),
std::fs::read_to_string("/proc/self/numa_maps"),
) else {
return Vec::new();
};
let spans: Vec<(usize, usize)> = regions
.iter()
.map(|r| (r.as_ptr() as usize, r.as_ptr() as usize + r.len()))
.collect();
let mut starts = std::collections::HashSet::new();
for line in maps.lines() {
let Some((range, _)) = line.split_once(' ') else {
continue;
};
let Some((a, b)) = range.split_once('-') else {
continue;
};
let (Ok(a), Ok(b)) = (usize::from_str_radix(a, 16), usize::from_str_radix(b, 16)) else {
continue;
};
if spans.iter().any(|&(s, e)| s < b && a < e) {
starts.insert(a);
}
}
let mut by = Vec::new();
for line in nm.lines() {
let mut it = line.split_whitespace();
let Some(a) = it.next().and_then(|a| usize::from_str_radix(a, 16).ok()) else {
continue;
};
if !starts.contains(&a) {
continue;
}
for f in it {
let Some((k, v)) = f.split_once('=') else {
continue;
};
let (Some(n), Ok(v)) = (
k.strip_prefix('N').and_then(|n| n.parse::<usize>().ok()),
v.parse::<usize>(),
) else {
continue;
};
if by.len() <= n {
by.resize(n + 1, 0);
}
by[n] += v;
}
}
by
}
pub(super) fn choose(regions: &[&[u8]], threads: usize) -> Option<(usize, Vec<usize>)> {
let env = std::env::var("CMF_NUMA").ok();
if matches!(env.as_deref(), Some("0") | Some("off")) {
return None;
}
let trace = std::env::var("CMF_NUMA_TRACE").is_ok_and(|v| v != "0");
let nodes = nodes();
if nodes.len() < 2 {
if trace {
eprintln!("numa: {} node(s) visible — nothing to bind", nodes.len());
}
return None;
}
let forced = env
.as_deref()
.and_then(|v| v.strip_prefix("node:"))
.and_then(|v| v.parse::<usize>().ok());
let node = match forced {
Some(n) => n,
None => {
let (sampled, incore, by) = page_nodes(regions);
let resident: usize = by.iter().sum();
if trace {
eprintln!(
"numa: sampled {sampled} weight pages, {incore} cached, mapped by node {by:?}"
);
}
if sampled == 0 || incore * 2 < sampled || resident == 0 {
return None;
}
let (n, &cnt) = by.iter().enumerate().max_by_key(|(_, c)| **c)?;
if cnt * 10 < resident * 9 {
return None;
}
n
}
};
let cpus = &nodes.iter().find(|(id, _)| *id == node)?.1;
let allowed = allowed();
let usable: Vec<usize> = cpus.iter().copied().filter(|c| allowed.contains(c)).collect();
let prim: Vec<usize> = usable.iter().copied().filter(|&c| primary(c)).collect();
if prim.len() >= threads {
Some((node, prim))
} else if usable.len() >= threads {
Some((node, usable))
} else {
None
}
}
pub(super) fn set_affinity(tid: i32, cpus: &[usize]) -> bool {
unsafe {
let mut set: libc::cpu_set_t = std::mem::zeroed();
for &c in cpus {
if c < libc::CPU_SETSIZE as usize {
libc::CPU_SET(c, &mut set);
}
}
libc::sched_setaffinity(tid, std::mem::size_of::<libc::cpu_set_t>(), &set) == 0
}
}
}
#[cfg(any(
target_os = "android",
all(target_arch = "aarch64", target_os = "linux")
))]
fn core_capacities() -> Vec<u64> {
let read_all = |leaf: &str| -> Vec<u64> {
let mut vals = Vec::new();
for cpu in 0.. {
let path = format!("/sys/devices/system/cpu/cpu{cpu}/{leaf}");
match std::fs::read_to_string(&path) {
Ok(v) => match v.trim().parse() {
Ok(x) => vals.push(x),
Err(_) => break,
},
Err(_) => break,
}
}
vals
};
let caps = read_all("cpu_capacity");
if caps.len() >= 2 {
return caps;
}
read_all("cpufreq/cpuinfo_max_freq")
}
#[cfg(target_os = "android")]
fn pin_thread_to_big_cores() {
use std::mem;
let caps = core_capacities();
let max = caps.iter().copied().max().unwrap_or(0);
let min = caps.iter().copied().min().unwrap_or(0);
if caps.len() < 2 || max == min {
return;
}
unsafe {
let mut set: libc::cpu_set_t = mem::zeroed();
for (i, &c) in caps.iter().enumerate() {
if c * 8 >= max * 5 {
libc::CPU_SET(i, &mut set);
}
}
libc::sched_setaffinity(0, mem::size_of::<libc::cpu_set_t>(), &set);
}
}
fn worker_loop(inner: &Inner, idx: usize) {
#[cfg(target_os = "android")]
pin_thread_to_big_cores();
#[cfg(target_os = "macos")]
unsafe {
libc::pthread_set_qos_class_self_np(libc::qos_class_t::QOS_CLASS_USER_INITIATED, 0);
}
let mut seen = 0usize;
loop {
let mut spins = 0usize;
loop {
let e = inner.epoch.load(Ordering::Acquire);
if e != seen {
seen = e;
break;
}
if inner.shutdown.load(Ordering::Relaxed) {
return;
}
if spins < inner.spin_budget.load(Ordering::Relaxed) {
spins += 1;
std::hint::spin_loop();
} else {
inner.parked[idx].store(true, Ordering::SeqCst);
if inner.epoch.load(Ordering::SeqCst) == seen
&& !inner.shutdown.load(Ordering::Relaxed)
{
std::thread::park();
}
inner.parked[idx].store(false, Ordering::SeqCst);
}
}
let (task, n, dev, limit) =
unsafe { (*inner.slot.get()).expect("job published with epoch") };
if idx >= limit {
continue;
}
let f = unsafe { &*task.0 };
crate::gpu::set_current_device(dev);
f(idx, n);
inner.remaining.fetch_sub(1, Ordering::AcqRel);
}
}
pub fn matvec_rows(pool: Option<&Pool>, w: &[f32], x: &[f32], out: &mut [f32]) {
let in_dim = x.len();
let out_dim = out.len();
debug_assert!(w.len() >= out_dim * in_dim);
let row_dot = |o: usize| -> f32 {
let row = &w[o * in_dim..(o + 1) * in_dim];
let mut sum = 0.0f32;
for j in 0..in_dim {
sum += row[j] * x[j];
}
sum
};
match pool {
Some(pool) if out_dim >= 256 => {
let out_addr = SendMut(out.as_mut_ptr());
let run_range = move |start: usize, end: usize| {
for o in start..end {
unsafe { *out_addr.at(o) = row_dot(o) };
}
};
pool.run_rows(out_dim, &run_range);
}
_ => {
for (o, dst) in out.iter_mut().enumerate() {
*dst = row_dot(o);
}
}
}
}
pub fn matvec_rows2(
pool: Option<&Pool>,
w: &[f32],
x1: &[f32],
x2: &[f32],
out1: &mut [f32],
out2: &mut [f32],
) {
let in_dim = x1.len();
debug_assert_eq!(x2.len(), in_dim);
let out_dim = out1.len();
debug_assert_eq!(out2.len(), out_dim);
debug_assert!(w.len() >= out_dim * in_dim);
let row_dots = |o: usize| -> (f32, f32) {
let row = &w[o * in_dim..(o + 1) * in_dim];
let (mut s1, mut s2) = (0.0f32, 0.0f32);
for j in 0..in_dim {
s1 += row[j] * x1[j];
s2 += row[j] * x2[j];
}
(s1, s2)
};
match pool {
Some(pool) if out_dim >= 256 => {
let o1 = SendMut(out1.as_mut_ptr());
let o2 = SendMut(out2.as_mut_ptr());
let run_range = move |start: usize, end: usize| {
for o in start..end {
let (s1, s2) = row_dots(o);
unsafe {
*o1.at(o) = s1;
*o2.at(o) = s2;
}
}
};
pool.run_rows(out_dim, &run_range);
}
_ => {
for o in 0..out_dim {
let (s1, s2) = row_dots(o);
out1[o] = s1;
out2[o] = s2;
}
}
}
}
pub(crate) struct SendMutT<T>(*mut T);
unsafe impl<T> Send for SendMutT<T> {}
unsafe impl<T> Sync for SendMutT<T> {}
impl<T> Clone for SendMutT<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for SendMutT<T> {}
impl<T> SendMutT<T> {
#[inline]
pub(crate) fn new(p: *mut T) -> Self {
Self(p)
}
#[inline]
pub(crate) fn at(self, i: usize) -> *mut T {
unsafe { self.0.add(i) }
}
}
#[derive(Clone, Copy)]
pub(crate) struct SendMut(*mut f32);
unsafe impl Send for SendMut {}
unsafe impl Sync for SendMut {}
impl SendMut {
#[inline]
pub(crate) fn new(p: *mut f32) -> Self {
Self(p)
}
#[inline]
pub(crate) fn at(self, i: usize) -> *mut f32 {
unsafe { self.0.add(i) }
}
}
#[cfg(test)]
mod tests {
#[test]
#[cfg(target_os = "linux")]
fn numa_cpulist_parses_ranges_and_singles() {
assert_eq!(super::numa::parse_list("0-3,8,10-11\n"), vec![0, 1, 2, 3, 8, 10, 11]);
assert_eq!(super::numa::parse_list(""), Vec::<usize>::new());
}
#[test]
#[cfg(any(target_os = "android", target_os = "linux"))]
fn worker_tids_registered_before_new_returns() {
use std::collections::HashSet;
let p = super::Pool::new(3);
let local: Vec<_> = p.inner.worker_tids.lock().unwrap().clone();
let registered = p.inner.registered.load(Ordering::Acquire);
let unique: HashSet<_> = local.iter().copied().collect();
assert!(
registered == 3
&& local.len() == 3
&& unique.len() == 3
&& local.iter().all(|&tid| tid > 0),
"all worker tids must be privately registered before new returns \
(registered {registered}, local {}, unique {})",
local.len(),
unique.len()
);
}
#[test]
fn forced_threads_overrides_env_and_topology() {
use std::sync::atomic::Ordering;
super::FORCED_THREADS.store(3, Ordering::Relaxed);
let pool = super::Pool::from_env().expect("forced 3 → pool");
assert_eq!(pool.n_workers(), 3);
super::FORCED_THREADS.store(1, Ordering::Relaxed);
assert!(super::Pool::from_env().is_none(), "forced 1 → serial");
super::FORCED_THREADS.store(0, Ordering::Relaxed);
}
#[test]
#[cfg(any(target_os = "android", target_os = "linux"))]
fn concurrent_pool_constructors_complete_without_registration_race() {
use std::sync::{Barrier, mpsc};
use std::time::Duration;
for round in 0..16 {
let start = Arc::new(Barrier::new(3));
let (done_tx, done_rx) = mpsc::channel();
let mut joins = Vec::new();
for workers in [8usize, 1usize] {
let start = start.clone();
let done_tx = done_tx.clone();
joins.push(std::thread::spawn(move || {
start.wait();
let pool = Pool::with_spin(workers, 0);
done_tx.send(pool.n_workers()).unwrap();
}));
}
drop(done_tx);
start.wait();
let mut sizes = Vec::with_capacity(2);
for _ in 0..2 {
sizes.push(
done_rx
.recv_timeout(Duration::from_secs(10))
.unwrap_or_else(|_| panic!("pool constructor stalled in round {round}")),
);
}
sizes.sort_unstable();
assert_eq!(sizes, [1, 8]);
for join in joins {
join.join().unwrap();
}
}
}
#[test]
fn capacity_split_clock_bins_vs_microarch() {
type P = super::Pool;
assert_eq!(
P::cores_from_capacities(&[1024, 1024, 1024, 1024, 768, 768, 768, 768]),
Some(8)
);
assert_eq!(
P::cores_from_capacities(&[1024, 1024, 1024, 1024, 350, 350, 350, 350]),
Some(4)
);
assert_eq!(
P::cores_from_capacities(&[1024, 800, 800, 800, 800, 300, 300, 300]),
Some(5)
);
assert_eq!(P::cores_from_capacities(&[1024; 8]), None);
assert_eq!(P::cores_from_capacities(&[]), None);
}
use super::*;
#[test]
fn parallel_matvec_equals_serial_bitexact() {
let (out_dim, in_dim) = (512, 64);
let w: Vec<f32> = (0..out_dim * in_dim)
.map(|i| (i as f32 * 0.013).sin())
.collect();
let x: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.07).cos()).collect();
let mut serial = vec![0.0f32; out_dim];
matvec_rows(None, &w, &x, &mut serial);
let pool = Pool::new(4);
let mut parallel = vec![0.0f32; out_dim];
matvec_rows(Some(&pool), &w, &x, &mut parallel);
assert_eq!(serial, parallel, "row-parallel must be bit-identical");
}
#[test]
fn fused_pair_equals_two_singles_bitexact() {
let (out_dim, in_dim) = (300, 48);
let w: Vec<f32> = (0..out_dim * in_dim)
.map(|i| (i as f32 * 0.011).sin())
.collect();
let x1: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.03).cos()).collect();
let x2: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.09).sin()).collect();
let mut a1 = vec![0.0f32; out_dim];
let mut a2 = vec![0.0f32; out_dim];
matvec_rows(None, &w, &x1, &mut a1);
matvec_rows(None, &w, &x2, &mut a2);
for pool in [None, Some(Pool::new(3))] {
let mut b1 = vec![0.0f32; out_dim];
let mut b2 = vec![0.0f32; out_dim];
matvec_rows2(pool.as_ref(), &w, &x1, &x2, &mut b1, &mut b2);
assert_eq!(a1, b1, "fused lane 1 must be bit-identical");
assert_eq!(a2, b2, "fused lane 2 must be bit-identical");
}
}
#[test]
fn pool_survives_many_runs() {
let pool = Pool::new(3);
let counter = AtomicUsize::new(0);
for _ in 0..100 {
pool.run(&|_, _| {
counter.fetch_add(1, Ordering::Relaxed);
});
}
assert_eq!(counter.load(Ordering::Relaxed), 400);
}
#[test]
fn pool_wakes_after_park() {
let pool = Pool::with_spin(2, 0);
let counter = AtomicUsize::new(0);
for _ in 0..50 {
pool.run(&|_, _| {
counter.fetch_add(1, Ordering::Relaxed);
});
std::thread::sleep(std::time::Duration::from_micros(200));
}
assert_eq!(counter.load(Ordering::Relaxed), 150);
}
#[test]
fn worker_indices_are_distinct_and_cover_range() {
let pool = Pool::new(3);
let hits: Vec<AtomicUsize> = (0..4).map(|_| AtomicUsize::new(0)).collect();
for _ in 0..20 {
pool.run(&|widx, n| {
assert_eq!(n, 4);
hits[widx].fetch_add(1, Ordering::Relaxed);
});
}
for (i, h) in hits.iter().enumerate() {
assert_eq!(h.load(Ordering::Relaxed), 20, "participant {i} missed runs");
}
}
}
#[cfg(test)]
mod grain_tests {
use super::grain_for;
#[test]
fn a_short_job_still_reaches_every_worker() {
assert_eq!(grain_for(24, 49), 1);
assert_eq!(grain_for(4096, 49), 32);
assert_eq!(grain_for(32768, 49), 83);
assert_eq!(grain_for(0, 49), 1);
assert_eq!(grain_for(7, 1), 7);
assert!(grain_for(1, 49) >= 1);
}
}