use crate::policy::{ExecutionPolicy, Parallel};
use melinoe::MelinoeCell;
use melinoe::region::WriterShard;
use moirai_core::error::{ExecutorError, ExecutorResult};
use moirai_executor::{HybridExecutor, SchedulerScope, SyncTask, global};
use std::sync::Mutex;
mod chunks;
mod shards;
mod unit_tasks;
pub use chunks::{
ChunkBuffersError, for_each_chunk_buffers_mut_enumerated_with,
for_each_chunk_mut_enumerated_with, for_each_chunk_mut_with, for_each_chunk_mut_with_state,
for_each_chunk_pair_mut_enumerated_with, for_each_chunk_quad_mut_enumerated_with,
for_each_chunk_triple_mut_enumerated_with,
};
pub use unit_tasks::{
UNIT_TASK_BYTES, for_each_unit_task_many_mut_with, for_each_unit_task_mut_with,
for_each_unit_task_pair_mut_with, for_each_unit_task_range_with,
for_each_unit_task_triple_mut_with, units_per_task,
};
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, 'env: 'scope> {
inner: &'scope SchedulerScope<'env, SyncTask>,
}
impl<'scope, 'env: 'scope> Scope<'scope, 'env> {
#[inline]
pub fn spawn<F>(&self, task: F)
where
F: FnOnce() + Send + 'env,
{
self.inner
.spawn(move |_| task())
.expect("moirai global executor: scope spawn");
}
}
#[inline]
pub fn scope<'env, F, R>(body: F) -> R
where
F: for<'scope> FnOnce(&Scope<'scope, 'env>) -> 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");
}
fn worker_chunk_size(len: usize) -> usize {
let workers = moirai_core::executor::logical_parallelism().max(1);
len.div_ceil(workers).max(1)
}
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 partitions =
WriterShard::new(MelinoeCell::from_mut_slice(data)).par_chunks(worker_chunk_size(n));
let tasks = partitions.len();
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(tasks, move |c| {
let mut shard = unsafe { partitions.get_unchecked_chunk(c) };
for element in shard.iter_mut() {
f(element);
}
})
.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 chunk_size = worker_chunk_size(n);
let partitions = WriterShard::new(MelinoeCell::from_mut_slice(data)).par_chunks(chunk_size);
let tasks = partitions.len();
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(tasks, move |c| {
let start = c * chunk_size;
let mut shard = unsafe { partitions.get_unchecked_chunk(c) };
for (offset, element) in shard.iter_mut().enumerate() {
f(start + offset, element);
}
})
.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 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 = moirai_core::executor::logical_parallelism().max(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 partitions =
WriterShard::new(MelinoeCell::from_mut_slice(slots.as_mut_slice())).par_chunks(1);
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);
}
let mut shard = unsafe { partitions.get_unchecked_chunk(ci) };
if let Some(slot) = shard.get_mut(0) {
*slot = 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 chunk_size = worker_chunk_size(n);
let data_partitions =
WriterShard::new(MelinoeCell::from_mut_slice(data)).par_chunks(chunk_size);
let out_partitions =
WriterShard::new(MelinoeCell::from_mut_slice(out.as_mut_slice())).par_chunks(chunk_size);
let tasks = data_partitions.len();
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(tasks, move |c| {
let start = c * chunk_size;
let mut data_shard = unsafe { data_partitions.get_unchecked_chunk(c) };
let mut out_shard = unsafe { out_partitions.get_unchecked_chunk(c) };
for (offset, (element, slot)) in
data_shard.iter_mut().zip(out_shard.iter_mut()).enumerate()
{
slot.write(f(start + offset, element));
}
})
.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")
}