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}