Skip to main content

moirai_parallel/ops/
chunks.rs

1//! Mutable chunk operators over one or more disjoint buffers.
2//!
3//! Each operator plans the pass, consults the policy, runs a sequential
4//! fallback, and otherwise dispatches the disjoint partitions through the global
5//! executor. That skeleton — and its single `SAFETY` argument — lives once in
6//! [`crate::ops::shards::drive_chunks`]; the operators below are thin wrappers
7//! that name their buffer set and adapt their closure.
8
9use crate::ops::shards::{BufferArray, chunk_shards, drive_chunks};
10use crate::policy::ExecutionPolicy;
11use melinoe::MelinoeCell;
12use melinoe::region::WriterShard;
13use moirai_executor::{SyncTask, global};
14
15#[cfg(test)]
16mod tests;
17
18chunk_shards! {
19    /// The single-buffer chunk operator's buffer set.
20    struct SingleShards { data: T }
21}
22
23chunk_shards! {
24    /// The paired chunk operator's buffer set.
25    struct PairShards { a: A, b: B }
26}
27
28chunk_shards! {
29    /// The triple chunk operator's buffer set.
30    struct TripleShards { a: A, b: B, c: C }
31}
32
33chunk_shards! {
34    /// The quad chunk operator's buffer set.
35    struct QuadShards { a: A, b: B, c: C, d: D }
36}
37
38/// Failure to partition a fixed set of mutable buffers into matching chunks.
39#[derive(Clone, Copy, Debug, Eq, PartialEq)]
40#[non_exhaustive]
41pub enum ChunkBuffersError {
42    /// A buffer does not have the same element count as the first buffer.
43    LengthMismatch {
44        /// Zero-based position of the mismatched buffer.
45        buffer_index: usize,
46        /// Required element count, taken from the first buffer.
47        expected: usize,
48        /// Actual element count of the mismatched buffer.
49        actual: usize,
50    },
51}
52
53impl core::fmt::Display for ChunkBuffersError {
54    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
55        match self {
56            Self::LengthMismatch {
57                buffer_index,
58                expected,
59                actual,
60            } => write!(
61                formatter,
62                "chunk buffer {buffer_index} has length {actual}, expected {expected}"
63            ),
64        }
65    }
66}
67
68impl std::error::Error for ChunkBuffersError {}
69
70/// Apply `f` to each consecutive `chunk_size`-element mutable chunk of `data` in
71/// parallel, scheduled by policy `P`. The final chunk may be shorter.
72///
73/// Synchronous equivalent of rayon's `data.par_chunks_mut(chunk_size).for_each(f)`
74/// — the natural shape for batched/lane-wise transforms.
75pub fn for_each_chunk_mut_with<P, T, F>(data: &mut [T], chunk_size: usize, f: F)
76where
77    P: ExecutionPolicy,
78    T: Send,
79    F: Fn(&mut [T]) + Send + Sync,
80{
81    drive_chunks::<P, _, _>(
82        SingleShards { data },
83        chunk_size,
84        "moirai global executor: for_each_chunk_mut_with",
85        |_chunk_index, (chunk,)| f(chunk),
86    );
87}
88
89/// Apply `f(state, chunk)` to each consecutive mutable chunk, creating one
90/// reusable state value per scheduled worker shard.
91///
92/// This is the scratch-buffer form of [`for_each_chunk_mut_with`]. It matches
93/// the allocation discipline of Rayon-style `for_each_init`/`for_each_with`
94/// loops: a worker shard initializes `S` once, then reuses it for every logical
95/// chunk assigned to that shard.
96pub fn for_each_chunk_mut_with_state<P, T, S, Init, F>(
97    data: &mut [T],
98    chunk_size: usize,
99    init: Init,
100    f: F,
101) where
102    P: ExecutionPolicy,
103    T: Send,
104    S: Send,
105    Init: Fn() -> S + Send + Sync,
106    F: Fn(&mut S, &mut [T]) + Send + Sync,
107{
108    let n = data.len();
109    if n == 0 || chunk_size == 0 {
110        return;
111    }
112    let num_chunks = n.div_ceil(chunk_size);
113    if !P::parallelize_chunks(n, num_chunks) || num_chunks <= 1 {
114        let mut state = init();
115        for chunk in data.chunks_mut(chunk_size) {
116            f(&mut state, chunk);
117        }
118        return;
119    }
120
121    let workers = moirai_core::executor::logical_parallelism()
122        .min(num_chunks)
123        .max(1);
124    let chunks_per_worker = num_chunks.div_ceil(workers);
125    let partitions = WriterShard::new(MelinoeCell::from_mut_slice(data)).par_chunks(chunk_size);
126    let init = &init;
127    let f = &f;
128    global()
129        .for_each_indexed::<SyncTask, _>(workers, move |worker| {
130            let first_chunk = worker * chunks_per_worker;
131            let last_chunk = ((worker + 1) * chunks_per_worker).min(num_chunks);
132            if first_chunk >= last_chunk {
133                return;
134            }
135            let mut state = init();
136            for chunk_index in first_chunk..last_chunk {
137                // SAFETY: the worker ranges `first_chunk..last_chunk` partition
138                // `0..num_chunks`, so each chunk index is visited by exactly one
139                // worker and distinct indices name disjoint element ranges.
140                let chunk = unsafe { partitions.get_unchecked_chunk(chunk_index) }.into_mut_slice();
141                f(&mut state, chunk);
142            }
143        })
144        .expect("moirai global executor: for_each_chunk_mut_with_state");
145}
146
147/// Apply `f(index, chunks)` to matching chunks from a fixed set of distinct
148/// mutable buffers, scheduled by policy `P`.
149///
150/// All buffers must have the same length. Validation completes before any
151/// buffer is mutated. The final chunk may be shorter than `chunk_size`; zero
152/// buffers, empty buffers, and a zero chunk size are no-ops.
153///
154/// The fixed-size array and chunk derivation add no heap allocation after the
155/// global executor is initialized. A first parallel call can allocate while
156/// constructing that process-wide executor and its worker pool. The operation
157/// lets callers fuse any homogeneous number of output-buffer passes without
158/// adding another arity-specific operator.
159///
160/// # Examples
161///
162/// ```
163/// use moirai_parallel::{
164///     for_each_chunk_buffers_mut_enumerated_with, ChunkBuffersError, Sequential,
165/// };
166///
167/// let mut left = [0_u32; 5];
168/// let mut right = [0_u32; 5];
169/// for_each_chunk_buffers_mut_enumerated_with::<Sequential, _, _, 2>(
170///     [&mut left, &mut right],
171///     2,
172///     |chunk_index, [left, right]| {
173///         left.fill(chunk_index as u32);
174///         right.fill((chunk_index as u32) + 10);
175///     },
176/// )?;
177///
178/// assert_eq!(left, [0, 0, 1, 1, 2]);
179/// assert_eq!(right, [10, 10, 11, 11, 12]);
180/// # Ok::<(), ChunkBuffersError>(())
181/// ```
182///
183/// # Errors
184///
185/// Returns [`ChunkBuffersError::LengthMismatch`] when a buffer length differs
186/// from the first buffer's length.
187pub fn for_each_chunk_buffers_mut_enumerated_with<P, T, F, const N: usize>(
188    buffers: [&mut [T]; N],
189    chunk_size: usize,
190    f: F,
191) -> Result<(), ChunkBuffersError>
192where
193    P: ExecutionPolicy,
194    T: Send,
195    F: for<'chunk> Fn(usize, [&'chunk mut [T]; N]) + Send + Sync,
196{
197    let length = buffers.first().map_or(0, |first| first.len());
198    if let Some((buffer_index, actual)) =
199        buffers
200            .iter()
201            .enumerate()
202            .skip(1)
203            .find_map(|(buffer_index, buffer)| {
204                (buffer.len() != length).then_some((buffer_index, buffer.len()))
205            })
206    {
207        return Err(ChunkBuffersError::LengthMismatch {
208            buffer_index,
209            expected: length,
210            actual,
211        });
212    }
213    drive_chunks::<P, _, _>(
214        BufferArray { buffers },
215        chunk_size,
216        "moirai global executor: for_each_chunk_buffers_mut_enumerated_with",
217        f,
218    );
219    Ok(())
220}
221
222/// Apply `f(index, a_chunk, b_chunk)` to paired `chunk_size`-element mutable
223/// chunks of two **distinct** buffers in parallel, scheduled by policy `P`.
224///
225/// Synchronous equivalent of
226/// `a.par_chunks_mut(n).zip(b.par_chunks_mut(n)).enumerate().for_each(f)`. The
227/// number of chunks is derived from `a`; `b` is chunked identically, so callers
228/// must ensure `b.len() >= a.len()` (typically equal). The two buffers must not
229/// alias.
230pub fn for_each_chunk_pair_mut_enumerated_with<P, A, B, F>(
231    a: &mut [A],
232    b: &mut [B],
233    chunk_size: usize,
234    f: F,
235) where
236    P: ExecutionPolicy,
237    A: Send,
238    B: Send,
239    F: Fn(usize, &mut [A], &mut [B]) + Send + Sync,
240{
241    // The paired pass stops at the shorter buffer, matching the sequential
242    // `zip` path, so a `b` shorter than `a` is processed up to `b`'s extent
243    // rather than reading out of bounds; `drive_chunks` takes the minimum
244    // partition count for exactly that reason.
245    drive_chunks::<P, _, _>(
246        PairShards { a, b },
247        chunk_size,
248        "moirai global executor: for_each_chunk_pair_mut_enumerated_with",
249        |chunk_index, (run_a, run_b)| f(chunk_index, run_a, run_b),
250    );
251}
252
253/// Apply `f(index, a_chunk, b_chunk, c_chunk, d_chunk)` to four **distinct**
254/// mutable buffers chunked identically, scheduled by policy `P`.
255///
256/// This is the four-buffer counterpart to
257/// [`for_each_chunk_pair_mut_enumerated_with`]. It is intended for fused
258/// statistics and stencil bookkeeping kernels where one authoritative pass
259/// updates several output arrays without allocating intermediate tuples.
260pub fn for_each_chunk_quad_mut_enumerated_with<P, A, B, C, D, F>(
261    a: &mut [A],
262    b: &mut [B],
263    c: &mut [C],
264    d: &mut [D],
265    chunk_size: usize,
266    f: F,
267) where
268    P: ExecutionPolicy,
269    A: Send,
270    B: Send,
271    C: Send,
272    D: Send,
273    F: Fn(usize, &mut [A], &mut [B], &mut [C], &mut [D]) + Send + Sync,
274{
275    assert_eq!(
276        a.len(),
277        b.len(),
278        "quad chunk buffers must have equal lengths"
279    );
280    assert_eq!(
281        a.len(),
282        c.len(),
283        "quad chunk buffers must have equal lengths"
284    );
285    assert_eq!(
286        a.len(),
287        d.len(),
288        "quad chunk buffers must have equal lengths"
289    );
290    drive_chunks::<P, _, _>(
291        QuadShards { a, b, c, d },
292        chunk_size,
293        "moirai global executor: for_each_chunk_quad_mut_enumerated_with",
294        |chunk_index, (run_a, run_b, run_c, run_d)| f(chunk_index, run_a, run_b, run_c, run_d),
295    );
296}
297
298/// Apply `f(index, a_chunk, b_chunk, c_chunk)` to three **distinct** mutable
299/// buffers chunked identically, scheduled by policy `P`.
300///
301/// This is the three-buffer counterpart to
302/// [`for_each_chunk_pair_mut_enumerated_with`].
303pub fn for_each_chunk_triple_mut_enumerated_with<P, A, B, C, F>(
304    a: &mut [A],
305    b: &mut [B],
306    c: &mut [C],
307    chunk_size: usize,
308    f: F,
309) where
310    P: ExecutionPolicy,
311    A: Send,
312    B: Send,
313    C: Send,
314    F: Fn(usize, &mut [A], &mut [B], &mut [C]) + Send + Sync,
315{
316    assert_eq!(
317        a.len(),
318        b.len(),
319        "triple chunk buffers must have equal lengths"
320    );
321    assert_eq!(
322        a.len(),
323        c.len(),
324        "triple chunk buffers must have equal lengths"
325    );
326    drive_chunks::<P, _, _>(
327        TripleShards { a, b, c },
328        chunk_size,
329        "moirai global executor: for_each_chunk_triple_mut_enumerated_with",
330        |chunk_index, (run_a, run_b, run_c)| f(chunk_index, run_a, run_b, run_c),
331    );
332}
333
334/// Like [`for_each_chunk_mut_with`] but also passes the zero-based chunk index to
335/// `f` (synchronous equivalent of
336/// `data.par_chunks_mut(chunk_size).enumerate().for_each(f)`).
337pub fn for_each_chunk_mut_enumerated_with<P, T, F>(data: &mut [T], chunk_size: usize, f: F)
338where
339    P: ExecutionPolicy,
340    T: Send,
341    F: Fn(usize, &mut [T]) + Send + Sync,
342{
343    drive_chunks::<P, _, _>(
344        SingleShards { data },
345        chunk_size,
346        "moirai global executor: for_each_chunk_mut_enumerated_with",
347        |chunk_index, (chunk,)| f(chunk_index, chunk),
348    );
349}