use super::{
ctrl::SUB_CONTEXT,
task::{ParTask, ParTaskHolder, TaskId},
};
use crate::global;
use rayon::iter::{
plumbing::{
Consumer, Folder, Producer, ProducerCallback, Reducer, UnindexedConsumer, UnindexedProducer,
},
IndexedParallelIterator, ParallelIterator,
};
use std::cell::Cell;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(transparent)]
pub(crate) struct FnContext {
migrated: bool,
}
impl FnContext {
pub(super) const MIGRATED: Self = Self { migrated: true };
pub(super) const NOT_MIGRATED: Self = Self { migrated: false };
}
#[derive(Debug, Clone, Copy)]
#[repr(transparent)]
struct Splitter {
splits: usize,
}
impl Splitter {
fn new() -> Option<Self> {
let ptr = SUB_CONTEXT.with(Cell::get);
(!ptr.is_dangling()).then(|| {
Self {
splits: unsafe { ptr.as_ref().get_comm().num_siblings() },
}
})
}
fn try_split(&mut self, migrated: bool) -> bool {
if !migrated {
self.splits /= 2;
}
self.splits > 0
}
}
#[derive(Debug, Clone, Copy)]
struct LengthSplitter {
inner: Splitter,
min: usize,
}
impl LengthSplitter {
fn new(min: usize, max: usize, len: usize) -> Option<Self> {
let mut inner = Splitter::new()?;
let min_splits = len / max.max(1);
inner.splits = inner.splits.max(min_splits);
Some(Self {
inner,
min: min.max(1),
})
}
fn try_split(&mut self, len: usize, migrated: bool) -> bool {
len / 2 >= self.min && self.inner.try_split(migrated)
}
}
fn bridge<I, C>(par_iter: I, consumer: C) -> C::Result
where
I: IndexedParallelIterator,
C: Consumer<I::Item>,
{
global::stat::increase_parallel_task_count();
let len = par_iter.len();
return par_iter.with_producer(Callback { len, consumer });
struct Callback<C> {
len: usize,
consumer: C,
}
impl<C, I> ProducerCallback<I> for Callback<C>
where
C: Consumer<I>,
{
type Output = C::Result;
fn callback<P>(self, producer: P) -> C::Result
where
P: Producer<Item = I>,
{
bridge_producer_consumer(self.len, producer, self.consumer)
}
}
}
fn bridge_producer_consumer<P, C>(len: usize, producer: P, consumer: C) -> C::Result
where
P: Producer,
C: Consumer<P::Item>,
{
const MIGRATED: bool = false;
let res =
if let Some(splitter) = LengthSplitter::new(producer.min_len(), producer.max_len(), len) {
helper(len, MIGRATED, splitter, producer, consumer)
} else {
helper_no_split(producer, consumer)
};
return res;
fn helper<P, C>(
len: usize,
migrated: bool,
mut splitter: LengthSplitter,
producer: P,
consumer: C,
) -> C::Result
where
P: Producer,
C: Consumer<P::Item>,
{
if consumer.full() {
consumer.into_folder().complete()
} else if splitter.try_split(len, migrated) {
let mid = len / 2;
let (l_producer, r_producer) = producer.split_at(mid);
let (l_consumer, r_consumer, reducer) = consumer.split_at(mid);
let (l_result, r_result) = join_context(
|f_cx: FnContext| helper(mid, f_cx.migrated, splitter, l_producer, l_consumer),
|f_cx: FnContext| {
helper(len - mid, f_cx.migrated, splitter, r_producer, r_consumer)
},
);
reducer.reduce(l_result, r_result)
} else {
producer.fold_with(consumer.into_folder()).complete()
}
}
fn helper_no_split<P, C>(producer: P, consumer: C) -> C::Result
where
P: Producer,
C: Consumer<P::Item>,
{
if consumer.full() {
consumer.into_folder().complete()
} else {
producer.fold_with(consumer.into_folder()).complete()
}
}
}
#[allow(dead_code)] fn bridge_unindexed<P, C>(producer: P, consumer: C) -> C::Result
where
P: UnindexedProducer,
C: UnindexedConsumer<P::Item>,
{
global::stat::increase_parallel_task_count();
let splitter = Splitter::new().unwrap();
bridge_unindexed_producer_consumer(false, splitter, producer, consumer)
}
fn bridge_unindexed_producer_consumer<P, C>(
migrated: bool,
mut splitter: Splitter,
producer: P,
consumer: C,
) -> C::Result
where
P: UnindexedProducer,
C: UnindexedConsumer<P::Item>,
{
if consumer.full() {
consumer.into_folder().complete()
} else if splitter.try_split(migrated) {
match producer.split() {
(l_producer, Some(r_producer)) => {
let (reducer, l_consumer, r_consumer) =
(consumer.to_reducer(), consumer.split_off_left(), consumer);
let bridge = bridge_unindexed_producer_consumer;
let (l_result, r_result) = join_context(
|f_cx: FnContext| bridge(f_cx.migrated, splitter, l_producer, l_consumer),
|f_cx: FnContext| bridge(f_cx.migrated, splitter, r_producer, r_consumer),
);
reducer.reduce(l_result, r_result)
}
(producer, None) => producer.fold_with(consumer.into_folder()).complete(),
}
} else {
producer.fold_with(consumer.into_folder()).complete()
}
}
fn join_context<L, R, Lr, Rr>(l_f: L, r_f: R) -> (Lr, Rr)
where
L: FnOnce(FnContext) -> Lr + Send,
R: FnOnce(FnContext) -> Rr + Send,
Lr: Send,
Rr: Send,
{
let cx = unsafe { SUB_CONTEXT.with(Cell::get).as_ref() };
let r_holder = ParTaskHolder::new(r_f);
let r_task = unsafe { ParTask::new(&r_holder) };
let r_task_id = TaskId::Parallel(r_task);
cx.get_comm().push_parallel_task(r_task);
cx.get_comm().signal().sub().notify_one();
#[cfg(not(target_arch = "wasm32"))]
let l_res = {
let executor = std::panic::AssertUnwindSafe(move || l_f(FnContext::NOT_MIGRATED));
match std::panic::catch_unwind(executor) {
Ok(l_res) => l_res,
Err(payload) => {
if let Some(task) = cx.get_comm().pop_local() {
debug_assert_eq!(task.id(), r_task_id);
} else {
while !r_holder.is_executed() {
std::thread::yield_now();
}
}
std::panic::resume_unwind(payload);
}
}
};
#[cfg(target_arch = "wasm32")]
let l_res = l_f(FnContext::NOT_MIGRATED);
if let Some(task) = cx.get_comm().pop_local() {
debug_assert_eq!(task.id(), r_task_id);
let wid = cx.get_comm().worker_id();
r_task.execute(wid, FnContext::NOT_MIGRATED);
} else {
while !r_holder.is_executed() {
let mut steal = cx.get_comm().search();
cx.work(&mut steal);
}
}
match unsafe { r_holder.return_or_panic_unchecked() } {
Ok(r_res) => (l_res, r_res),
Err(payload) => std::panic::resume_unwind(payload),
}
}
pub trait IntoEcsPar: ParallelIterator {
#[inline]
fn into_ecs_par(self) -> EcsPar<Self> {
EcsPar(self)
}
#[doc(hidden)]
fn drive_unindexed<C>(self, consumer: C) -> C::Result
where
C: UnindexedConsumer<Self::Item>;
}
impl<I: IndexedParallelIterator> IntoEcsPar for I {
#[inline]
fn drive_unindexed<C>(self, consumer: C) -> C::Result
where
C: UnindexedConsumer<Self::Item>,
{
bridge(self, consumer)
}
}
#[derive(Clone)]
#[repr(transparent)]
pub struct EcsPar<I>(pub I);
impl<I: IntoEcsPar> ParallelIterator for EcsPar<I> {
type Item = I::Item;
fn drive_unindexed<C>(self, consumer: C) -> C::Result
where
C: UnindexedConsumer<Self::Item>,
{
IntoEcsPar::drive_unindexed(self.0, consumer)
}
}
impl<I> IndexedParallelIterator for EcsPar<I>
where
I: IntoEcsPar + IndexedParallelIterator,
{
#[inline]
fn len(&self) -> usize {
self.0.len()
}
#[inline]
fn drive<C: Consumer<Self::Item>>(self, consumer: C) -> C::Result {
bridge(self, consumer)
}
#[inline]
fn with_producer<CB: ProducerCallback<Self::Item>>(self, callback: CB) -> CB::Output {
self.0.with_producer(callback)
}
}
#[cfg(test)]
mod tests {
#[test]
fn test_into_ecs_par() {
use super::*;
use rayon::iter::IntoParallelIterator;
let iter: rayon::array::IntoIter<i32, 2> = [0, 1].into_par_iter();
let _ecs_iter = iter.into_ecs_par();
let iter: rayon::range::Iter<i32> = (0..2).into_par_iter();
let _ecs_iter = iter.into_ecs_par();
let iter: rayon::slice::Iter<'_, i32> = [0, 1][..].into_par_iter();
let _ecs_iter = iter.into_ecs_par();
let range_iter0: rayon::range::Iter<i32> = (0..2).into_par_iter();
let range_iter1: rayon::range::Iter<i32> = (0..2).into_par_iter();
let zip_iter = range_iter0.zip(range_iter1);
let _ecs_iter = zip_iter.into_ecs_par();
}
}