use rayon::prelude::*;
use crate::cpu_pool::CpuPool;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Backend {
Rayon,
Spin,
}
const TASKS_PER_THREAD: usize = 8;
pub fn backend() -> Backend {
use std::sync::OnceLock;
static BACKEND: OnceLock<Backend> = OnceLock::new();
*BACKEND.get_or_init(|| {
match std::env::var("FERROX_CPU_POOL")
.ok()
.map(|v| v.trim().to_ascii_lowercase())
.as_deref()
{
Some("spin") | Some("persistent") | Some("1") | Some("on") | Some("true") => {
Backend::Spin
}
_ => Backend::Rayon,
}
})
}
fn pool() -> &'static CpuPool {
use std::sync::OnceLock;
static POOL: OnceLock<CpuPool> = OnceLock::new();
POOL.get_or_init(|| CpuPool::new(crate::threads::resolve_cpu_threads()))
}
pub fn num_threads() -> usize {
match backend() {
Backend::Rayon => rayon::current_num_threads().max(1),
Backend::Spin => pool().num_threads(),
}
}
pub fn task_count(n_items: usize) -> usize {
if n_items == 0 {
return 0;
}
n_items.min(num_threads().saturating_mul(TASKS_PER_THREAD).max(1))
}
fn split(n_items: usize) -> (usize, usize) {
let n_tasks = task_count(n_items);
if n_tasks == 0 {
return (0, 0);
}
(n_items.div_ceil(n_tasks), n_tasks)
}
struct SendPtr<T>(*mut T);
impl<T> Clone for SendPtr<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for SendPtr<T> {}
impl<T> SendPtr<T> {
unsafe fn at(self, offset: usize) -> *mut T {
unsafe { self.0.add(offset) }
}
}
unsafe impl<T: Send> Send for SendPtr<T> {}
unsafe impl<T: Send> Sync for SendPtr<T> {}
pub fn indices<F>(n: usize, min_len: usize, f: F)
where
F: Fn(usize) + Send + Sync,
{
if n == 0 {
return;
}
if backend() == Backend::Spin {
let (per, n_tasks) = split(n);
let task = |t: usize| {
let lo = t * per;
let hi = ((t + 1) * per).min(n);
for i in lo..hi {
f(i);
}
};
if pool().run(n_tasks, &task) {
return;
}
}
(0..n)
.into_par_iter()
.with_min_len(min_len.max(1))
.for_each(&f);
}
pub fn indices_init<S, I, F>(n: usize, min_len: usize, init: I, f: F)
where
S: Send,
I: Fn() -> S + Send + Sync,
F: Fn(&mut S, usize) + Send + Sync,
{
if n == 0 {
return;
}
if backend() == Backend::Spin {
let (per, n_tasks) = split(n);
let task = |t: usize| {
let lo = t * per;
let hi = ((t + 1) * per).min(n);
if lo >= hi {
return;
}
let mut state = init();
for i in lo..hi {
f(&mut state, i);
}
};
if pool().run(n_tasks, &task) {
return;
}
}
(0..n)
.into_par_iter()
.with_min_len(min_len.max(1))
.for_each_init(&init, |state, i| f(state, i));
}
pub fn items_mut<T, F>(data: &mut [T], min_len: usize, f: F)
where
T: Send,
F: Fn(usize, &mut T) + Send + Sync,
{
let n = data.len();
if n == 0 {
return;
}
if backend() == Backend::Spin {
let base = SendPtr(data.as_mut_ptr());
let (per, n_tasks) = split(n);
let task = |t: usize| {
let lo = t * per;
let hi = ((t + 1) * per).min(n);
for i in lo..hi {
f(i, unsafe { &mut *base.at(i) });
}
};
if pool().run(n_tasks, &task) {
return;
}
}
data.par_iter_mut()
.with_min_len(min_len.max(1))
.enumerate()
.for_each(|(i, slot)| f(i, slot));
}
pub fn chunks_mut2<T, U, F>(a: &mut [T], b: &mut [U], chunk_len: usize, min_len: usize, f: F)
where
T: Send,
U: Send,
F: Fn(usize, &mut [T], &mut [U]) + Send + Sync,
{
assert!(chunk_len > 0, "chunk length must be positive");
assert_eq!(a.len(), b.len(), "zipped slices must be the same length");
let len = a.len();
if len == 0 {
return;
}
let n_chunks = len.div_ceil(chunk_len);
if backend() == Backend::Spin {
let base_a = SendPtr(a.as_mut_ptr());
let base_b = SendPtr(b.as_mut_ptr());
let (per, n_tasks) = split(n_chunks);
let task = |t: usize| {
let lo = t * per;
let hi = ((t + 1) * per).min(n_chunks);
for c in lo..hi {
unsafe {
f(
c,
chunk_of(base_a, len, chunk_len, c),
chunk_of(base_b, len, chunk_len, c),
);
}
}
};
if pool().run(n_tasks, &task) {
return;
}
}
a.par_chunks_mut(chunk_len)
.zip(b.par_chunks_mut(chunk_len))
.with_min_len(min_len.max(1))
.enumerate()
.for_each(|(c, (ca, cb))| f(c, ca, cb));
}
pub fn chunks_mut<T, F>(data: &mut [T], chunk_len: usize, min_len: usize, f: F)
where
T: Send,
F: Fn(usize, &mut [T]) + Send + Sync,
{
assert!(chunk_len > 0, "chunk length must be positive");
let len = data.len();
if len == 0 {
return;
}
let n_chunks = len.div_ceil(chunk_len);
if backend() == Backend::Spin {
let base = SendPtr(data.as_mut_ptr());
let (per, n_tasks) = split(n_chunks);
let task = |t: usize| {
let lo = t * per;
let hi = ((t + 1) * per).min(n_chunks);
for c in lo..hi {
f(c, unsafe { chunk_of(base, len, chunk_len, c) });
}
};
if pool().run(n_tasks, &task) {
return;
}
}
data.par_chunks_mut(chunk_len)
.with_min_len(min_len.max(1))
.enumerate()
.for_each(|(c, chunk)| f(c, chunk));
}
unsafe fn chunk_of<'a, T>(base: SendPtr<T>, len: usize, chunk_len: usize, c: usize) -> &'a mut [T] {
let start = c * chunk_len;
let end = ((c + 1) * chunk_len).min(len);
debug_assert!(start < end && end <= len);
unsafe { std::slice::from_raw_parts_mut(base.at(start), end - start) }
}
pub fn chunks_mut_init<T, S, I, F>(data: &mut [T], chunk_len: usize, min_len: usize, init: I, f: F)
where
T: Send,
S: Send,
I: Fn() -> S + Send + Sync,
F: Fn(&mut S, usize, &mut [T]) + Send + Sync,
{
assert!(chunk_len > 0, "chunk length must be positive");
let len = data.len();
if len == 0 {
return;
}
let n_chunks = len.div_ceil(chunk_len);
if backend() == Backend::Spin {
let base = SendPtr(data.as_mut_ptr());
let (per, n_tasks) = split(n_chunks);
let task = |t: usize| {
let lo = t * per;
let hi = ((t + 1) * per).min(n_chunks);
if lo >= hi {
return;
}
let mut state = init();
for c in lo..hi {
f(&mut state, c, unsafe { chunk_of(base, len, chunk_len, c) });
}
};
if pool().run(n_tasks, &task) {
return;
}
}
data.par_chunks_mut(chunk_len)
.with_min_len(min_len.max(1))
.enumerate()
.for_each_init(&init, |state, (c, chunk)| f(state, c, chunk));
}
pub fn join2<A, B, RA, RB>(a: A, b: B) -> (RA, RB)
where
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB + Send,
RA: Send,
RB: Send,
{
match backend() {
Backend::Rayon => rayon::join(a, b),
Backend::Spin => (a(), b()),
}
}
pub fn join3<A, B, C, RA, RB, RC>(a: A, b: B, c: C) -> (RA, RB, RC)
where
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB + Send,
C: FnOnce() -> RC + Send,
RA: Send,
RB: Send,
RC: Send,
{
match backend() {
Backend::Rayon => {
let (ra, (rb, rc)) = rayon::join(a, || rayon::join(b, c));
(ra, rb, rc)
}
Backend::Spin => (a(), b(), c()),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU32, Ordering};
#[test]
fn the_backend_is_read_from_one_env_var_and_defaults_to_rayon() {
if std::env::var_os("FERROX_CPU_POOL").is_none() {
assert_eq!(backend(), Backend::Rayon);
}
}
#[test]
fn the_spin_arm_chunks_by_pool_width_with_no_work_threshold() {
assert_eq!(task_count(0), 0);
assert_eq!(task_count(1), 1);
assert_eq!(task_count(3), 3);
let wide = task_count(1_000_000);
assert_eq!(wide, num_threads() * TASKS_PER_THREAD);
assert_eq!(task_count(4096), task_count(4096));
let (per, n) = split(1000);
assert_eq!(n, task_count(1000));
assert!(per * n >= 1000 && (per - 1) * n < 1000);
}
#[test]
fn indices_visits_every_index_exactly_once() {
for n in [0usize, 1, 7, 64, 5000] {
let hits: Vec<AtomicU32> = (0..n).map(|_| AtomicU32::new(0)).collect();
indices(n, 8, |i| {
hits[i].fetch_add(1, Ordering::Relaxed);
});
assert!(hits.iter().all(|h| h.load(Ordering::Relaxed) == 1), "n={n}");
}
}
#[test]
fn items_mut_writes_every_slot_with_its_own_index() {
for n in [0usize, 1, 9, 257] {
let mut data = vec![0u32; n];
items_mut(&mut data, 4, |i, slot| *slot = i as u32 + 1);
assert_eq!(data, (1..=n as u32).collect::<Vec<_>>(), "n={n}");
}
}
#[test]
fn chunks_mut_delivers_a_short_trailing_chunk() {
let mut data = vec![0u32; 10];
let seen: std::sync::Mutex<Vec<(usize, usize)>> = std::sync::Mutex::new(Vec::new());
chunks_mut(&mut data, 4, 1, |c, chunk| {
seen.lock().unwrap().push((c, chunk.len()));
for (i, slot) in chunk.iter_mut().enumerate() {
*slot = (c * 4 + i) as u32;
}
});
let mut seen = seen.into_inner().unwrap();
seen.sort_unstable();
assert_eq!(seen, vec![(0, 4), (1, 4), (2, 2)]);
assert_eq!(data, (0..10).collect::<Vec<u32>>());
}
#[test]
fn chunks_mut_init_gives_each_task_its_own_scratch() {
let mut data = vec![0u64; 512];
chunks_mut_init(
&mut data,
8,
1,
|| Vec::<u64>::with_capacity(8),
|scratch: &mut Vec<u64>, c, chunk| {
scratch.clear();
scratch.extend(chunk.iter().map(|_| c as u64));
chunk.copy_from_slice(scratch);
},
);
for (c, chunk) in data.chunks(8).enumerate() {
assert!(chunk.iter().all(|&v| v == c as u64));
}
}
#[test]
fn joins_return_every_result_in_order() {
assert_eq!(join2(|| 1u8, || 2u8), (1, 2));
assert_eq!(join3(|| 1u8, || 2u8, || 3u8), (1, 2, 3));
}
#[test]
fn the_spin_arm_and_the_rayon_arm_produce_identical_results() {
let pool = CpuPool::new(4);
for n in [1usize, 5, 63, 1024] {
let mut spun = vec![0f32; n];
let base = SendPtr(spun.as_mut_ptr());
let (per, n_tasks) = split(n);
let task = |t: usize| {
let lo = t * per;
let hi = ((t + 1) * per).min(n);
for i in lo..hi {
unsafe { *base.at(i) = (i as f32) * 0.5 + 1.0 };
}
};
assert!(pool.run(n_tasks, &task));
let mut forked = vec![0f32; n];
forked
.par_iter_mut()
.with_min_len(8)
.enumerate()
.for_each(|(i, slot)| *slot = (i as f32) * 0.5 + 1.0);
assert_eq!(spun, forked, "n={n}");
}
}
}