Skip to main content

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;