use super::DisjointMutPtr;
use crate::policy::{ExecutionPolicy, Parallel};
use moirai_core::error::{ExecutorError, ExecutorResult};
use moirai_executor::{global, HybridExecutor, SchedulerScope, SyncTask};
use std::sync::Mutex;
enum Branch<F, R> {
Pending(F),
Claimed,
Done(R),
}
impl<F, R> Branch<F, R>
where
F: FnOnce() -> R,
{
fn claim(&mut self) -> Option<F> {
match std::mem::replace(self, Self::Claimed) {
Self::Pending(branch) => Some(branch),
other => {
*self = other;
None
}
}
}
fn complete(&mut self, result: R) {
*self = Self::Done(result);
}
fn run_shared(slot: &Mutex<Self>) {
let Some(branch) = lock(slot).claim() else {
return;
};
let result = branch();
lock(slot).complete(result);
}
fn run_here(&mut self) {
let Some(branch) = self.claim() else {
return;
};
let result = branch();
self.complete(result);
}
fn into_result(self) -> R {
match self {
Self::Done(result) => result,
_ => panic!("invariant: a join branch neither ran nor reported failure"),
}
}
}
fn lock<T>(slot: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
slot.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn lock_owned<T>(slot: Mutex<T>) -> T {
slot.into_inner()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
pub fn join_with<P, A, B, RA, RB>(left: A, right: B) -> (RA, RB)
where
P: ExecutionPolicy,
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB,
RA: Send,
{
if !P::parallelize_pair() {
return (left(), right());
}
join_on(global(), left, right)
}
pub(crate) fn join_on<A, B, RA, RB>(executor: &HybridExecutor, left: A, right: B) -> (RA, RB)
where
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB,
RA: Send,
{
let left_slot = Mutex::new(Branch::Pending(left));
let mut right_slot = Branch::Pending(right);
let forked = executor.scope::<SyncTask, _>(|scope| {
scope.spawn(|_| Branch::run_shared(&left_slot))?;
scope.flush()?;
right_slot.run_here();
Ok(())
});
match forked {
Ok(()) => {}
Err(ExecutorError::ShuttingDown | ExecutorError::ResourceExhausted(_)) => {
Branch::run_shared(&left_slot);
right_slot.run_here();
}
Err(error) => panic!("invariant: scheduled join branch failed ({error})"),
}
(
lock_owned(left_slot).into_result(),
right_slot.into_result(),
)
}
pub fn join<A, B, RA, RB>(left: A, right: B) -> (RA, RB)
where
A: FnOnce() -> RA + Send,
B: FnOnce() -> RB,
RA: Send,
{
join_with::<crate::Adaptive, _, _, _, _>(left, right)
}
pub struct Scope<'scope> {
inner: &'scope SchedulerScope<'scope, SyncTask>,
}
impl<'scope> Scope<'scope> {
#[inline]
pub fn spawn<F>(&self, task: F)
where
F: FnOnce() + Send + 'scope,
{
self.inner
.spawn(move |_| task())
.expect("moirai global executor: scope spawn");
}
}
#[inline]
pub fn scope<F, R>(body: F) -> R
where
F: for<'scope> FnOnce(&Scope<'scope>) -> R,
R: Send,
{
let mut result = None;
global()
.scope::<SyncTask, _>(|inner| {
let scope = Scope { inner };
result = Some(body(&scope));
ExecutorResult::Ok(())
})
.expect("moirai global executor: scope");
result.expect("scoped body must complete")
}
pub fn for_each_with<P, T, F>(data: &[T], f: F)
where
P: ExecutionPolicy,
T: Sync,
F: Fn(&T) + Send + Sync,
{
let n = data.len();
if n == 0 {
return;
}
if !P::parallelize(n) {
data.iter().for_each(f);
return;
}
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(n, move |i| f(&data[i]))
.expect("moirai global executor: for_each_with");
}
pub fn for_each_mut_with<P, T, F>(data: &mut [T], f: F)
where
P: ExecutionPolicy,
T: Send,
F: Fn(&mut T) + Send + Sync,
{
let n = data.len();
if n == 0 {
return;
}
if !P::parallelize(n) {
data.iter_mut().for_each(f);
return;
}
let base = DisjointMutPtr(data.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(n, move |i| {
f(unsafe { base.get_mut(i) });
})
.expect("moirai global executor: for_each_mut_with");
}
pub fn enumerate_with<P, T, F>(data: &[T], f: F)
where
P: ExecutionPolicy,
T: Sync,
F: Fn(usize, &T) + Send + Sync,
{
let n = data.len();
if n == 0 {
return;
}
if !P::parallelize(n) {
data.iter().enumerate().for_each(|(i, x)| f(i, x));
return;
}
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(n, move |i| f(i, &data[i]))
.expect("moirai global executor: enumerate_with");
}
pub fn enumerate_mut_with<P, T, F>(data: &mut [T], f: F)
where
P: ExecutionPolicy,
T: Send,
F: Fn(usize, &mut T) + Send + Sync,
{
let n = data.len();
if n == 0 {
return;
}
if !P::parallelize(n) {
data.iter_mut().enumerate().for_each(|(i, x)| f(i, x));
return;
}
let base = DisjointMutPtr(data.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(n, move |i| {
f(i, unsafe { base.get_mut(i) });
})
.expect("moirai global executor: enumerate_mut_with");
}
pub fn for_each_index_with<P, F>(len: usize, f: F)
where
P: ExecutionPolicy,
F: Fn(usize) + Send + Sync,
{
if len == 0 {
return;
}
if !P::parallelize(len) {
(0..len).for_each(f);
return;
}
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(len, f)
.expect("moirai global executor: for_each_index_with");
}
pub fn for_each_chunk_mut_with<P, T, F>(data: &mut [T], chunk_size: usize, f: F)
where
P: ExecutionPolicy,
T: Send,
F: Fn(&mut [T]) + Send + Sync,
{
let n = data.len();
if n == 0 || chunk_size == 0 {
return;
}
let num_chunks = n.div_ceil(chunk_size);
if !P::parallelize(n) || num_chunks <= 1 {
data.chunks_mut(chunk_size).for_each(&f);
return;
}
let base = DisjointMutPtr(data.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
let start = c * chunk_size;
if start >= n {
return;
}
let end = (start + chunk_size).min(n);
let chunk =
unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
f(chunk);
})
.expect("moirai global executor: for_each_chunk_mut_with");
}
pub fn for_each_chunk_mut_with_state<P, T, S, Init, F>(
data: &mut [T],
chunk_size: usize,
init: Init,
f: F,
) where
P: ExecutionPolicy,
T: Send,
S: Send,
Init: Fn() -> S + Send + Sync,
F: Fn(&mut S, &mut [T]) + Send + Sync,
{
let n = data.len();
if n == 0 || chunk_size == 0 {
return;
}
let num_chunks = n.div_ceil(chunk_size);
if !P::parallelize(n) || num_chunks <= 1 {
let mut state = init();
for chunk in data.chunks_mut(chunk_size) {
f(&mut state, chunk);
}
return;
}
let workers = std::thread::available_parallelism()
.map(|count| count.get())
.unwrap_or(1)
.min(num_chunks)
.max(1);
let chunks_per_worker = num_chunks.div_ceil(workers);
let base = DisjointMutPtr(data.as_mut_ptr());
let init = &init;
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(workers, move |worker| {
let first_chunk = worker * chunks_per_worker;
let last_chunk = ((worker + 1) * chunks_per_worker).min(num_chunks);
if first_chunk >= last_chunk {
return;
}
let mut state = init();
for chunk_index in first_chunk..last_chunk {
let start = chunk_index * chunk_size;
let end = (start + chunk_size).min(n);
let chunk =
unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
f(&mut state, chunk);
}
})
.expect("moirai global executor: for_each_chunk_mut_with_state");
}
pub fn for_each_chunk_pair_mut_enumerated_with<P, A, B, F>(
a: &mut [A],
b: &mut [B],
chunk_size: usize,
f: F,
) where
P: ExecutionPolicy,
A: Send,
B: Send,
F: Fn(usize, &mut [A], &mut [B]) + Send + Sync,
{
let na = a.len();
let nb = b.len();
if chunk_size == 0 || na == 0 {
return;
}
let num_chunks = na.div_ceil(chunk_size);
if !P::parallelize(na) || num_chunks <= 1 {
a.chunks_mut(chunk_size)
.zip(b.chunks_mut(chunk_size))
.enumerate()
.for_each(|(i, (ca, cb))| f(i, ca, cb));
return;
}
let abase = DisjointMutPtr(a.as_mut_ptr());
let bbase = DisjointMutPtr(b.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
let start = c * chunk_size;
if start >= na || start >= nb {
return;
}
let ea = (start + chunk_size).min(na);
let eb = (start + chunk_size).min(nb);
let ca =
unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
let cb =
unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
f(c, ca, cb);
})
.expect("moirai global executor: for_each_chunk_pair_mut_enumerated_with");
}
pub fn for_each_chunk_quad_mut_enumerated_with<P, A, B, C, D, F>(
a: &mut [A],
b: &mut [B],
c: &mut [C],
d: &mut [D],
chunk_size: usize,
f: F,
) where
P: ExecutionPolicy,
A: Send,
B: Send,
C: Send,
D: Send,
F: Fn(usize, &mut [A], &mut [B], &mut [C], &mut [D]) + Send + Sync,
{
let na = a.len();
let nb = b.len();
let nc = c.len();
let nd = d.len();
assert_eq!(na, nb, "quad chunk buffers must have equal lengths");
assert_eq!(na, nc, "quad chunk buffers must have equal lengths");
assert_eq!(na, nd, "quad chunk buffers must have equal lengths");
if chunk_size == 0 || na == 0 {
return;
}
let num_chunks = na.div_ceil(chunk_size);
if !P::parallelize(na) || num_chunks <= 1 {
a.chunks_mut(chunk_size)
.zip(b.chunks_mut(chunk_size))
.zip(c.chunks_mut(chunk_size))
.zip(d.chunks_mut(chunk_size))
.enumerate()
.for_each(|(i, (((ca, cb), cc), cd))| f(i, ca, cb, cc, cd));
return;
}
let abase = DisjointMutPtr(a.as_mut_ptr());
let bbase = DisjointMutPtr(b.as_mut_ptr());
let cbase = DisjointMutPtr(c.as_mut_ptr());
let dbase = DisjointMutPtr(d.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(num_chunks, move |chunk_index| {
let start = chunk_index * chunk_size;
if start >= na || start >= nb || start >= nc || start >= nd {
return;
}
let ea = (start + chunk_size).min(na);
let eb = (start + chunk_size).min(nb);
let ec = (start + chunk_size).min(nc);
let ed = (start + chunk_size).min(nd);
let ca =
unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
let cb =
unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
let cc =
unsafe { core::slice::from_raw_parts_mut(cbase.base().add(start), ec - start) };
let cd =
unsafe { core::slice::from_raw_parts_mut(dbase.base().add(start), ed - start) };
f(chunk_index, ca, cb, cc, cd);
})
.expect("moirai global executor: for_each_chunk_quad_mut_enumerated_with");
}
pub fn for_each_chunk_triple_mut_enumerated_with<P, A, B, C, F>(
a: &mut [A],
b: &mut [B],
c: &mut [C],
chunk_size: usize,
f: F,
) where
P: ExecutionPolicy,
A: Send,
B: Send,
C: Send,
F: Fn(usize, &mut [A], &mut [B], &mut [C]) + Send + Sync,
{
let na = a.len();
let nb = b.len();
let nc = c.len();
assert_eq!(na, nb, "triple chunk buffers must have equal lengths");
assert_eq!(na, nc, "triple chunk buffers must have equal lengths");
if chunk_size == 0 || na == 0 {
return;
}
let num_chunks = na.div_ceil(chunk_size);
if !P::parallelize(na) || num_chunks <= 1 {
a.chunks_mut(chunk_size)
.zip(b.chunks_mut(chunk_size))
.zip(c.chunks_mut(chunk_size))
.enumerate()
.for_each(|(i, ((ca, cb), cc))| f(i, ca, cb, cc));
return;
}
let abase = DisjointMutPtr(a.as_mut_ptr());
let bbase = DisjointMutPtr(b.as_mut_ptr());
let cbase = DisjointMutPtr(c.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(num_chunks, move |chunk_index| {
let start = chunk_index * chunk_size;
if start >= na || start >= nb || start >= nc {
return;
}
let ea = (start + chunk_size).min(na);
let eb = (start + chunk_size).min(nb);
let ec = (start + chunk_size).min(nc);
let ca =
unsafe { core::slice::from_raw_parts_mut(abase.base().add(start), ea - start) };
let cb =
unsafe { core::slice::from_raw_parts_mut(bbase.base().add(start), eb - start) };
let cc =
unsafe { core::slice::from_raw_parts_mut(cbase.base().add(start), ec - start) };
f(chunk_index, ca, cb, cc);
})
.expect("moirai global executor: for_each_chunk_triple_mut_enumerated_with");
}
pub fn for_each_chunk_mut_enumerated_with<P, T, F>(data: &mut [T], chunk_size: usize, f: F)
where
P: ExecutionPolicy,
T: Send,
F: Fn(usize, &mut [T]) + Send + Sync,
{
let n = data.len();
if n == 0 || chunk_size == 0 {
return;
}
let num_chunks = n.div_ceil(chunk_size);
if !P::parallelize(n) || num_chunks <= 1 {
data.chunks_mut(chunk_size)
.enumerate()
.for_each(|(i, c)| f(i, c));
return;
}
let base = DisjointMutPtr(data.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(num_chunks, move |c| {
let start = c * chunk_size;
if start >= n {
return;
}
let end = (start + chunk_size).min(n);
let chunk =
unsafe { core::slice::from_raw_parts_mut(base.base().add(start), end - start) };
f(c, chunk);
})
.expect("moirai global executor: for_each_chunk_mut_enumerated_with");
}
pub fn map_collect_with<P, T, R, F>(data: &[T], f: F) -> Vec<R>
where
P: ExecutionPolicy,
T: Sync,
R: Send,
F: Fn(&T) -> R + Send + Sync,
{
let n = data.len();
if !P::parallelize(n) {
return data.iter().map(f).collect();
}
let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(n);
unsafe {
out.set_len(n);
}
enumerate_mut_with::<Parallel, _, _>(&mut out, |i, slot| {
slot.write(f(&data[i]));
});
let mut out = core::mem::ManuallyDrop::new(out);
unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), n, out.capacity()) }
}
pub fn map_reduce_with<P, T, R, M, Rd>(data: &[T], identity: R, map: M, reduce: Rd) -> R
where
P: ExecutionPolicy,
T: Sync,
R: Send + Sync + Clone,
M: Fn(&T) -> R + Send + Sync,
Rd: Fn(R, R) -> R + Send + Sync,
{
let n = data.len();
if n == 0 || !P::parallelize(n) {
let mut acc = identity;
for item in data {
acc = reduce(acc, map(item));
}
return acc;
}
let map = ↦
let reduce = &reduce;
global()
.map_reduce_indexed::<SyncTask, _, _, _>(n, identity, move |i| map(&data[i]), reduce)
.expect("moirai global executor: map_reduce_with")
}
pub fn fold_reduce_with<P, A, Init, Fold, Red>(len: usize, init: Init, fold: Fold, reduce: Red) -> A
where
P: ExecutionPolicy,
A: Send,
Init: Fn() -> A + Send + Sync,
Fold: Fn(A, usize) -> A + Send + Sync,
Red: Fn(A, A) -> A,
{
if len == 0 {
return init();
}
if !P::parallelize(len) {
let mut acc = init();
for i in 0..len {
acc = fold(acc, i);
}
return acc;
}
let workers = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
let chunks = workers.min(len).max(1);
let chunk = len.div_ceil(chunks);
let mut slots: Vec<Option<A>> = (0..chunks).map(|_| None).collect();
let base = DisjointMutPtr(slots.as_mut_ptr());
let init_ref = &init;
let fold_ref = &fold;
global()
.for_each_indexed::<SyncTask, _>(chunks, move |ci| {
let start = ci * chunk;
if start >= len {
return;
}
let end = (start + chunk).min(len);
let mut acc = init_ref();
for i in start..end {
acc = fold_ref(acc, i);
}
unsafe {
*base.get_mut(ci) = Some(acc);
}
})
.expect("moirai global executor: fold_reduce_with");
slots
.into_iter()
.flatten()
.reduce(reduce)
.unwrap_or_else(init)
}
pub fn map_collect_index_with<P, R, Map>(len: usize, map: Map) -> Vec<R>
where
P: ExecutionPolicy,
R: Send,
Map: Fn(usize) -> R + Send + Sync,
{
if !P::parallelize(len) {
return (0..len).map(map).collect();
}
let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(len);
unsafe {
out.set_len(len);
}
enumerate_mut_with::<Parallel, _, _>(&mut out, |i, slot| {
slot.write(map(i));
});
let mut out = core::mem::ManuallyDrop::new(out);
unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), len, out.capacity()) }
}
pub fn map_collect_mut_with<P, T, R, F>(data: &mut [T], f: F) -> Vec<R>
where
P: ExecutionPolicy,
T: Send,
R: Send,
F: Fn(usize, &mut T) -> R + Send + Sync,
{
let n = data.len();
if !P::parallelize(n) {
return data.iter_mut().enumerate().map(|(i, x)| f(i, x)).collect();
}
let mut out: Vec<core::mem::MaybeUninit<R>> = Vec::with_capacity(n);
unsafe {
out.set_len(n);
}
let data_ptr = DisjointMutPtr(data.as_mut_ptr());
let out_ptr = DisjointMutPtr(out.as_mut_ptr());
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(n, move |i| {
let elem = unsafe { data_ptr.get_mut(i) };
let result = f(i, elem);
unsafe { out_ptr.get_mut(i).write(result) };
})
.expect("moirai global executor: map_collect_mut_with");
let mut out = core::mem::ManuallyDrop::new(out);
unsafe { Vec::from_raw_parts(out.as_mut_ptr().cast::<R>(), n, out.capacity()) }
}
pub fn reduce_index_with<P, R, Map, Red>(len: usize, identity: R, map: Map, reduce: Red) -> R
where
P: ExecutionPolicy,
R: Send + Sync + Clone,
Map: Fn(usize) -> R + Send + Sync,
Red: Fn(R, R) -> R + Send + Sync,
{
if len == 0 || !P::parallelize(len) {
let mut acc = identity;
for i in 0..len {
acc = reduce(acc, map(i));
}
return acc;
}
global()
.map_reduce_indexed::<SyncTask, _, _, _>(len, identity, map, reduce)
.expect("moirai global executor: reduce_index_with")
}