use std::sync::OnceLock;
use rayon::iter::{
IndexedParallelIterator, IntoParallelIterator, IntoParallelRefIterator, ParallelIterator,
};
use rayon::slice::{ParallelSlice, ParallelSliceMut};
use rayon::{ThreadPool, ThreadPoolBuilder};
pub const MIN_WORK: usize = 65_536;
pub const WIDE_ITEM: usize = 256;
fn configured_threads() -> usize {
static N: OnceLock<usize> = OnceLock::new();
*N.get_or_init(|| {
match std::env::var("LIBJAY_THREADS").ok().and_then(|v| v.trim().parse::<usize>().ok()) {
Some(n) if n > 0 => n,
_ => std::thread::available_parallelism().map(|n| n.get()).unwrap_or(1),
}
})
}
fn pool() -> &'static ThreadPool {
static POOL: OnceLock<ThreadPool> = OnceLock::new();
POOL.get_or_init(|| {
ThreadPoolBuilder::new()
.num_threads(configured_threads())
.thread_name(|i| format!("libjay-{i}"))
.build()
.expect("building the libjay thread pool")
})
}
fn parallelism() -> usize {
match rayon::current_thread_index() {
Some(_) => rayon::current_num_threads(),
None => configured_threads(),
}
}
fn in_pool<R: Send>(f: impl FnOnce() -> R + Send) -> R {
match rayon::current_thread_index() {
Some(_) => f(),
None => pool().install(f),
}
}
#[cfg(test)]
pub fn with_threads<R: Send>(threads: usize, f: impl FnOnce() -> R + Send) -> R {
ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.expect("building a thread pool")
.install(f)
}
pub fn worth_it(n: usize) -> bool {
n >= MIN_WORK && parallelism() > 1
}
fn chunk_len(n: usize) -> usize {
n.div_ceil(parallelism() * 4).max(4096)
}
pub fn chunks(n: usize, work: usize) -> usize {
if !worth_it(work) {
return 1;
}
n.min(parallelism() * 4).max(1)
}
fn fill_chunks<U, F>(n: usize, chunk: usize, parallel: bool, f: F) -> (Vec<U>, bool)
where
U: Copy + Default + Send,
F: Fn(usize, &mut [U]) -> bool + Sync + Send,
{
let mut out = vec![U::default(); n];
let ok = if parallel {
in_pool(|| {
out.par_chunks_mut(chunk)
.enumerate()
.map(|(k, part)| f(k * chunk, part))
.reduce(|| true, |a, b| a && b)
})
} else {
f(0, &mut out)
};
(out, ok)
}
pub fn fill<U, F>(n: usize, f: F) -> (Vec<U>, bool)
where
U: Copy + Default + Send,
F: Fn(usize, &mut [U]) -> bool + Sync + Send,
{
fill_chunks(n, chunk_len(n), worth_it(n), f)
}
pub fn fill_wide<U, F>(n: usize, work: usize, f: F) -> (Vec<U>, bool)
where
U: Copy + Default + Send,
F: Fn(usize, &mut [U]) -> bool + Sync + Send,
{
let threads = parallelism();
let parallel = worth_it(work) && n >= threads;
fill_chunks(n, n.div_ceil(threads), parallel, f)
}
pub fn fill_rows<U, F>(rows: usize, width: usize, work: usize, f: F) -> Vec<U>
where
U: Copy + Default + Send,
F: Fn(usize, &mut [U]) + Sync + Send,
{
let mut out = vec![U::default(); rows * width];
let threads = parallelism();
if width > 0 && rows >= threads && worth_it(work) {
let per = rows.div_ceil(threads);
in_pool(|| {
out.par_chunks_mut(per * width)
.enumerate()
.for_each(|(k, part)| f(k * per, part));
});
} else if rows > 0 {
f(0, &mut out);
}
out
}
pub fn try_fill<U, E, F>(n: usize, f: F) -> Result<Vec<U>, E>
where
U: Copy + Default + Send,
E: Send,
F: Fn(usize, &mut [U]) -> Result<(), E> + Sync + Send,
{
let mut out = vec![U::default(); n];
if worth_it(n) {
let chunk = chunk_len(n);
in_pool(|| {
out.par_chunks_mut(chunk)
.enumerate()
.try_for_each(|(k, part)| f(k * chunk, part))
})?;
} else {
f(0, &mut out)?;
}
Ok(out)
}
pub fn map<T, U, F>(src: &[T], f: F) -> Vec<U>
where
T: Sync,
U: Send,
F: Fn(&T) -> U + Sync + Send,
{
if worth_it(src.len()) {
in_pool(|| src.par_iter().map(&f).collect())
} else {
src.iter().map(&f).collect()
}
}
pub fn try_map<T, U, F>(src: &[T], f: F) -> Option<Vec<U>>
where
T: Copy + Sync,
U: Copy + Default + Send,
F: Fn(T) -> Option<U> + Sync + Send,
{
let (out, ok) = fill(src.len(), |start, part| {
for (k, slot) in part.iter_mut().enumerate() {
match f(src[start + k]) {
Some(v) => *slot = v,
None => return false,
}
}
true
});
ok.then_some(out)
}
pub fn any<T, F>(v: &[T], f: F) -> bool
where
T: Sync,
F: Fn(&T) -> bool + Sync + Send,
{
let scan = |part: &[T]| part.iter().fold(false, |a, x| a | f(x));
if worth_it(v.len()) {
in_pool(|| v.par_chunks(chunk_len(v.len())).map(scan).reduce(|| false, |a, b| a | b))
} else {
scan(v)
}
}
pub fn map_indexed<U, F>(n: usize, f: F) -> Vec<U>
where
U: Send,
F: Fn(usize) -> U + Sync + Send,
{
in_pool(|| (0..n).into_par_iter().map(&f).collect())
}
pub fn try_fold_chunks<S, T, C, F>(v: &[S], seq: C, step: F) -> Option<T>
where
S: Copy + Send + Sync,
T: Copy + Send + Sync,
C: Fn(&[S]) -> Option<T> + Sync + Send,
F: Fn(T, T) -> Option<T> + Sync + Send,
{
if !worth_it(v.len()) {
return seq(v);
}
let chunk = chunk_len(v.len());
let parts: Vec<Option<T>> =
in_pool(|| v.par_chunks(chunk).map(&seq).collect::<Vec<Option<T>>>());
let parts: Option<Vec<T>> = parts.into_iter().collect();
let parts = parts?;
let mut acc = parts[parts.len() - 1];
for &x in parts[..parts.len() - 1].iter().rev() {
acc = step(x, acc)?;
}
Some(acc)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_small_pass_stays_on_one_thread() {
assert!(!worth_it(MIN_WORK - 1));
}
#[test]
fn one_thread_never_splits() {
assert!(with_threads(1, || !worth_it(MIN_WORK * 16)));
assert!(with_threads(4, || worth_it(MIN_WORK * 16)));
}
#[test]
fn a_split_fill_writes_every_element_once() {
let n = MIN_WORK * 4;
let run = |threads: usize| {
with_threads(threads, || {
fill(n, |start, part: &mut [i64]| {
for (k, slot) in part.iter_mut().enumerate() {
*slot = (start + k) as i64;
}
true
})
})
};
let (a, ok_a) = run(1);
let (b, ok_b) = run(4);
assert!(ok_a && ok_b);
assert_eq!(a, b);
assert_eq!(b[n - 1], (n - 1) as i64);
}
#[test]
fn a_failure_in_one_chunk_fails_the_whole_fill() {
let n = MIN_WORK * 4;
let (_, ok) = with_threads(4, || {
fill(n, |start, part: &mut [i64]| {
!(start..start + part.len()).contains(&0)
})
});
assert!(!ok);
}
}