use crate::ops::shards::{BufferArray, chunk_shards, drive_chunks};
use crate::policy::ExecutionPolicy;
use melinoe::MelinoeCell;
use melinoe::region::WriterShard;
use moirai_executor::{SyncTask, global};
#[cfg(test)]
mod tests;
chunk_shards! {
struct SingleShards { data: T }
}
chunk_shards! {
struct PairShards { a: A, b: B }
}
chunk_shards! {
struct TripleShards { a: A, b: B, c: C }
}
chunk_shards! {
struct QuadShards { a: A, b: B, c: C, d: D }
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum ChunkBuffersError {
LengthMismatch {
buffer_index: usize,
expected: usize,
actual: usize,
},
}
impl core::fmt::Display for ChunkBuffersError {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::LengthMismatch {
buffer_index,
expected,
actual,
} => write!(
formatter,
"chunk buffer {buffer_index} has length {actual}, expected {expected}"
),
}
}
}
impl std::error::Error for ChunkBuffersError {}
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,
{
drive_chunks::<P, _, _>(
SingleShards { data },
chunk_size,
"moirai global executor: for_each_chunk_mut_with",
|_chunk_index, (chunk,)| f(chunk),
);
}
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_chunks(n, num_chunks) || num_chunks <= 1 {
let mut state = init();
for chunk in data.chunks_mut(chunk_size) {
f(&mut state, chunk);
}
return;
}
let workers = moirai_core::executor::logical_parallelism()
.min(num_chunks)
.max(1);
let chunks_per_worker = num_chunks.div_ceil(workers);
let partitions = WriterShard::new(MelinoeCell::from_mut_slice(data)).par_chunks(chunk_size);
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 chunk = unsafe { partitions.get_unchecked_chunk(chunk_index) }.into_mut_slice();
f(&mut state, chunk);
}
})
.expect("moirai global executor: for_each_chunk_mut_with_state");
}
pub fn for_each_chunk_buffers_mut_enumerated_with<P, T, F, const N: usize>(
buffers: [&mut [T]; N],
chunk_size: usize,
f: F,
) -> Result<(), ChunkBuffersError>
where
P: ExecutionPolicy,
T: Send,
F: for<'chunk> Fn(usize, [&'chunk mut [T]; N]) + Send + Sync,
{
let length = buffers.first().map_or(0, |first| first.len());
if let Some((buffer_index, actual)) =
buffers
.iter()
.enumerate()
.skip(1)
.find_map(|(buffer_index, buffer)| {
(buffer.len() != length).then_some((buffer_index, buffer.len()))
})
{
return Err(ChunkBuffersError::LengthMismatch {
buffer_index,
expected: length,
actual,
});
}
drive_chunks::<P, _, _>(
BufferArray { buffers },
chunk_size,
"moirai global executor: for_each_chunk_buffers_mut_enumerated_with",
f,
);
Ok(())
}
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,
{
drive_chunks::<P, _, _>(
PairShards { a, b },
chunk_size,
"moirai global executor: for_each_chunk_pair_mut_enumerated_with",
|chunk_index, (run_a, run_b)| f(chunk_index, run_a, run_b),
);
}
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,
{
assert_eq!(
a.len(),
b.len(),
"quad chunk buffers must have equal lengths"
);
assert_eq!(
a.len(),
c.len(),
"quad chunk buffers must have equal lengths"
);
assert_eq!(
a.len(),
d.len(),
"quad chunk buffers must have equal lengths"
);
drive_chunks::<P, _, _>(
QuadShards { a, b, c, d },
chunk_size,
"moirai global executor: for_each_chunk_quad_mut_enumerated_with",
|chunk_index, (run_a, run_b, run_c, run_d)| f(chunk_index, run_a, run_b, run_c, run_d),
);
}
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,
{
assert_eq!(
a.len(),
b.len(),
"triple chunk buffers must have equal lengths"
);
assert_eq!(
a.len(),
c.len(),
"triple chunk buffers must have equal lengths"
);
drive_chunks::<P, _, _>(
TripleShards { a, b, c },
chunk_size,
"moirai global executor: for_each_chunk_triple_mut_enumerated_with",
|chunk_index, (run_a, run_b, run_c)| f(chunk_index, run_a, run_b, run_c),
);
}
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,
{
drive_chunks::<P, _, _>(
SingleShards { data },
chunk_size,
"moirai global executor: for_each_chunk_mut_enumerated_with",
|chunk_index, (chunk,)| f(chunk_index, chunk),
);
}