molgfx_math/parallel.rs
1//! Threshold-guarded data parallelism with a thread-count-independent shape.
2//!
3//! Two rules make every helper here safe to drop into a load path that has to
4//! stay reproducible.
5//!
6//! **The partition is a pure function of the input length.** Work is split into
7//! fixed-size blocks, never into "one piece per thread", so the same input
8//! produces the same blocks on a laptop and on a build machine. A reduction
9//! then combines block results in index order. Floating-point addition is not
10//! associative, so a thread-count-dependent partition would give a
11//! thread-count-dependent answer; a fixed one cannot.
12//!
13//! **Small inputs stay on the calling thread.** Every entry point takes a
14//! minimum length and runs serially below it. A ligand must not pay for a
15//! thread pool, and a ribosome must not be denied one.
16//!
17//! Cost is `O(n / threads)` above the threshold and `O(n)` below it, with one
18//! `Vec` of block results — `blocks = n / BLOCK` entries — for the reducing
19//! forms and no allocation at all for the in-place forms.
20//!
21//! Every parallel branch schedules onto Rayon's process-wide registry. No
22//! scene, renderer or language binding constructs a private pool, so molgfx,
23//! molframe and independent Python `Engine` objects share the same Rust workers
24//! instead of multiplying thread counts and scratch memory.
25
26use rayon::prelude::*;
27
28/// Elements per block.
29///
30/// Sized so one block of the widest element this crate reduces over stays
31/// inside a typical L2 slice while remaining large enough that per-block
32/// scheduling overhead disappears against the work.
33pub const BLOCK: usize = 8192;
34
35/// Applies `f` to each `(offset, block)` of `slice`, in parallel above `min_len`.
36///
37/// Blocks are visited in an unspecified order, so `f` must not depend on
38/// ordering; it receives the block's start offset when it needs to address the
39/// source.
40pub fn for_each_block<T, F>(slice: &[T], min_len: usize, f: F)
41where
42 T: Sync,
43 F: Fn(usize, &[T]) + Send + Sync,
44{
45 if slice.len() < min_len {
46 for (index, block) in slice.chunks(BLOCK).enumerate() {
47 f(index * BLOCK, block);
48 }
49 return;
50 }
51 slice
52 .par_chunks(BLOCK)
53 .enumerate()
54 .for_each(|(index, block)| f(index * BLOCK, block));
55}
56
57/// Writes `dst` from `src` block-wise, in parallel above `min_len`.
58///
59/// Each block owns a disjoint output range, so no synchronisation is needed and
60/// the result is identical to the serial form byte for byte. `dst` is truncated
61/// to `src`'s length if it is longer.
62pub fn map_blocks_into<T, U, F>(src: &[T], dst: &mut [U], min_len: usize, f: F)
63where
64 T: Sync,
65 U: Send,
66 F: Fn(&[T], &mut [U]) + Send + Sync,
67{
68 let len = src.len().min(dst.len());
69 let (src, dst) = (&src[..len], &mut dst[..len]);
70 if len < min_len {
71 for (source, target) in src.chunks(BLOCK).zip(dst.chunks_mut(BLOCK)) {
72 f(source, target);
73 }
74 return;
75 }
76 src.par_chunks(BLOCK)
77 .zip(dst.par_chunks_mut(BLOCK))
78 .for_each(|(source, target)| f(source, target));
79}
80
81/// Writes `dst` from two aligned sources block-wise, in parallel above
82/// `min_len`.
83///
84/// The two sources are visited in index-aligned blocks, so `f` sees the same
85/// element triplets as the serial form and the result is identical to it byte
86/// for byte.
87pub fn map_zip_blocks_into<T, U, V, F>(a: &[T], b: &[U], dst: &mut [V], min_len: usize, f: F)
88where
89 T: Sync,
90 U: Sync,
91 V: Send,
92 F: Fn(&[T], &[U], &mut [V]) + Send + Sync,
93{
94 let len = a.len().min(b.len()).min(dst.len());
95 let (a, b, dst) = (&a[..len], &b[..len], &mut dst[..len]);
96 if len < min_len {
97 for ((source, other), target) in a
98 .chunks(BLOCK)
99 .zip(b.chunks(BLOCK))
100 .zip(dst.chunks_mut(BLOCK))
101 {
102 f(source, other, target);
103 }
104 return;
105 }
106 a.par_chunks(BLOCK)
107 .zip(b.par_chunks(BLOCK))
108 .zip(dst.par_chunks_mut(BLOCK))
109 .for_each(|((source, other), target)| f(source, other, target));
110}
111
112/// Reduces `slice` with a fixed block partition and an in-order combine.
113///
114/// `fold` collapses one block to a partial result; `combine` merges two
115/// partials. Because the partition depends only on `slice.len()` and the merge
116/// walks blocks in index order, the result does not depend on how many threads
117/// ran — which is what lets a float reduction keep a stable answer.
118pub fn reduce_blocks<T, A, Fold, Combine>(
119 slice: &[T],
120 min_len: usize,
121 identity: A,
122 fold: Fold,
123 combine: Combine,
124) -> A
125where
126 T: Sync,
127 A: Send + Clone,
128 Fold: Fn(&[T]) -> A + Send + Sync,
129 Combine: Fn(A, A) -> A,
130{
131 if slice.len() < min_len {
132 return slice.chunks(BLOCK).map(&fold).fold(identity, &combine);
133 }
134 let partials: Vec<A> = slice.par_chunks(BLOCK).map(&fold).collect();
135 partials.into_iter().fold(identity, &combine)
136}
137
138/// Whether `f` holds for every element, short-circuiting per block.
139///
140/// A block that fails stops that block immediately; other blocks already in
141/// flight run to completion, which costs nothing on the passing path and keeps
142/// the answer independent of scheduling.
143pub fn all_blocks<T, F>(slice: &[T], min_len: usize, f: F) -> bool
144where
145 T: Sync,
146 F: Fn(&[T]) -> bool + Send + Sync,
147{
148 if slice.len() < min_len {
149 return slice.chunks(BLOCK).all(&f);
150 }
151 slice.par_chunks(BLOCK).all(&f)
152}
153
154#[cfg(test)]
155#[path = "parallel_tests.rs"]
156mod tests;