use melinoe::MelinoeCell;
use melinoe::region::{ParChunks, WriterShard};
use moirai_executor::{SyncTask, global};
use crate::policy::ExecutionPolicy;
pub(super) trait ChunkShards {
type Element: Send;
type Chunk;
type Views: Send + Sync;
fn len(&self) -> usize;
fn task_count(&self, chunk_size: usize) -> usize;
fn split(self, chunk_size: usize) -> Self::Views;
unsafe fn chunk(views: &Self::Views, index: usize) -> Self::Chunk;
}
macro_rules! chunk_shards {
(
$(#[$attr:meta])*
struct $name:ident {
$first_field:ident : $first_ty:ident
$(, $field:ident : $ty:ident)* $(,)?
}
) => {
$(#[$attr])*
pub(in crate::ops) struct $name<'buffer, $first_ty $(, $ty)*> {
pub(in crate::ops) $first_field: &'buffer mut [$first_ty]
$(, pub(in crate::ops) $field: &'buffer mut [$ty])*
}
impl<'buffer, $first_ty: Send $(, $ty: Send)*> $crate::ops::shards::ChunkShards
for $name<'buffer, $first_ty $(, $ty)*>
{
type Element = $first_ty;
type Chunk = (&'buffer mut [$first_ty] $(, &'buffer mut [$ty])* ,);
type Views = (
melinoe::region::ParChunks<'buffer, 'buffer, $first_ty>
$(, melinoe::region::ParChunks<'buffer, 'buffer, $ty>)*
,
);
#[inline]
fn len(&self) -> usize {
self.$first_field.len()
}
#[inline]
fn task_count(&self, chunk_size: usize) -> usize {
[self.$first_field.len() $(, self.$field.len())*]
.into_iter()
.map(|len| len.div_ceil(chunk_size))
.min()
.unwrap_or(0)
}
#[inline]
fn split(self, chunk_size: usize) -> Self::Views {
let Self { $first_field $(, $field)* } = self;
(
melinoe::region::WriterShard::new(
melinoe::MelinoeCell::from_mut_slice($first_field),
)
.par_chunks(chunk_size)
$(
, melinoe::region::WriterShard::new(
melinoe::MelinoeCell::from_mut_slice($field),
)
.par_chunks(chunk_size)
)*
,
)
}
#[inline]
unsafe fn chunk(views: &Self::Views, index: usize) -> Self::Chunk {
let ($first_field $(, $field)* ,) = views;
(
unsafe { $first_field.get_unchecked_chunk(index) }.into_mut_slice()
$(
, unsafe { $field.get_unchecked_chunk(index) }.into_mut_slice()
)*
,
)
}
}
};
}
pub(super) use chunk_shards;
pub(super) struct BufferArray<'buffer, T, const N: usize> {
pub(super) buffers: [&'buffer mut [T]; N],
}
impl<'buffer, T: Send, const N: usize> ChunkShards for BufferArray<'buffer, T, N> {
type Element = T;
type Chunk = [&'buffer mut [T]; N];
type Views = [ParChunks<'buffer, 'buffer, T>; N];
#[inline]
fn len(&self) -> usize {
self.buffers.first().map_or(0, |buffer| buffer.len())
}
#[inline]
fn task_count(&self, chunk_size: usize) -> usize {
self.buffers
.iter()
.map(|buffer| buffer.len().div_ceil(chunk_size))
.min()
.unwrap_or(0)
}
#[inline]
fn split(self, chunk_size: usize) -> Self::Views {
self.buffers.map(|buffer| {
WriterShard::new(MelinoeCell::from_mut_slice(buffer)).par_chunks(chunk_size)
})
}
#[inline]
unsafe fn chunk(views: &Self::Views, index: usize) -> Self::Chunk {
core::array::from_fn(|buffer| {
unsafe { views[buffer].get_unchecked_chunk(index) }.into_mut_slice()
})
}
}
pub(super) fn drive_chunks<P, B, F>(buffers: B, chunk_size: usize, context: &'static str, f: F)
where
P: ExecutionPolicy,
B: ChunkShards,
F: Fn(usize, B::Chunk) + Send + Sync,
{
let len = buffers.len();
if len == 0 || chunk_size == 0 {
return;
}
let num_chunks = len.div_ceil(chunk_size);
let tasks = buffers.task_count(chunk_size);
let views = buffers.split(chunk_size);
if !P::parallelize_chunks(len, num_chunks) || num_chunks <= 1 {
for index in 0..tasks {
let chunk = unsafe { B::chunk(&views, index) };
f(index, chunk);
}
return;
}
let f = &f;
global()
.for_each_indexed::<SyncTask, _>(tasks, move |index| {
let chunk = unsafe { B::chunk(&views, index) };
f(index, chunk);
})
.expect(context);
}