Skip to main content

moirai_core/communication/
collective.rs

1#![expect(
2    clippy::unwrap_used,
3    reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
4)]
5
6/// CSR-shaped chunked buffer: one contiguous flat allocation plus a
7/// chunk-offset table.
8///
9/// Replaces the jagged `Vec<Vec<T>>` layout previously used by the collective
10/// operations. `offsets[i]..offsets[i+1]` is chunk `i`, so element traversal
11/// is a single contiguous pass over `flat` instead of a per-chunk pointer
12/// chase, and [`ChunkedVec::into_flat`] hands the storage back with no
13/// re-flatten pass.
14#[derive(Debug, Clone, Default, PartialEq, Eq)]
15pub struct ChunkedVec<T> {
16    flat: Vec<T>,
17    offsets: Vec<usize>,
18}
19
20impl<T> ChunkedVec<T> {
21    /// Number of chunks in the buffer.
22    #[must_use]
23    pub fn num_chunks(&self) -> usize {
24        self.offsets.len().saturating_sub(1)
25    }
26
27    /// Total number of elements across all chunks.
28    #[must_use]
29    pub fn len(&self) -> usize {
30        self.flat.len()
31    }
32
33    /// True when the buffer holds no elements.
34    #[must_use]
35    pub fn is_empty(&self) -> bool {
36        self.flat.is_empty()
37    }
38
39    /// Iterate the chunk slices in order.
40    pub fn chunks(&self) -> impl Iterator<Item = &[T]> + '_ {
41        self.offsets
42            .windows(2)
43            .map(move |window| &self.flat[window[0]..window[1]])
44    }
45
46    /// Consume the buffer, returning the contiguous element storage.
47    #[must_use]
48    pub fn into_flat(self) -> Vec<T> {
49        self.flat
50    }
51}
52
53/// Efficient collective operations for group communication
54pub struct CollectiveOps;
55
56impl CollectiveOps {
57    /// All-reduce operation: combine values from all participants
58    pub fn all_reduce<T, F>(values: Vec<T>, op: F) -> Vec<T>
59    where
60        T: Clone + Send,
61        F: Fn(T, T) -> T + Sync,
62    {
63        if values.is_empty() {
64            return vec![];
65        }
66
67        let result_len = values.len();
68
69        // Tree reduction for efficiency
70        let mut current = values;
71        while current.len() > 1 {
72            let mut next = Vec::with_capacity(current.len().div_ceil(2));
73
74            for chunk in current.chunks(2) {
75                if chunk.len() == 2 {
76                    next.push(op(chunk[0].clone(), chunk[1].clone()));
77                } else {
78                    next.push(chunk[0].clone());
79                }
80            }
81
82            current = next;
83        }
84
85        // Broadcast result to all
86        vec![current[0].clone(); result_len]
87    }
88
89    /// Scatter operation: distribute data into one CSR-shaped chunked buffer
90    /// with one chunk per participant.
91    ///
92    /// The result allocates once (`flat`) plus the chunk-offset table, instead
93    /// of one `Vec` per participant. Empty input and a zero participant count
94    /// produce an empty buffer rather than the historical `chunks(0)` panic.
95    pub fn scatter<T: Clone>(data: Vec<T>, num_participants: usize) -> ChunkedVec<T> {
96        let num_participants = num_participants.max(1);
97        let chunk_size = data.len().div_ceil(num_participants).max(1);
98        let mut offsets = Vec::with_capacity(num_participants + 1);
99        offsets.push(0);
100        let mut flat = Vec::with_capacity(data.len());
101        for chunk in data.chunks(chunk_size) {
102            flat.extend_from_slice(chunk);
103            offsets.push(flat.len());
104        }
105        ChunkedVec { flat, offsets }
106    }
107
108    /// Gather operation: collect the chunked buffer back into one contiguous
109    /// `Vec`. This is O(1): the flat buffer is returned directly, with no
110    /// re-flatten pass over per-chunk allocations.
111    pub fn gather<T>(chunks: ChunkedVec<T>) -> Vec<T> {
112        chunks.into_flat()
113    }
114
115    /// All-to-all communication pattern: transpose the chunked buffer so
116    /// result chunk `j` holds element `j` of every input chunk. Columns at or
117    /// beyond the chunk count are dropped, matching the historical contract.
118    pub fn all_to_all<T: Clone>(data: ChunkedVec<T>) -> ChunkedVec<T> {
119        let columns = data.num_chunks();
120        if columns == 0 {
121            return ChunkedVec {
122                flat: Vec::new(),
123                offsets: vec![0],
124            };
125        }
126        let mut offsets = Vec::with_capacity(columns + 1);
127        offsets.push(0);
128        for column in 0..columns {
129            let count = data.chunks().filter(|chunk| chunk.len() > column).count();
130            offsets.push(offsets.last().unwrap() + count);
131        }
132        let mut flat = Vec::with_capacity(*offsets.last().unwrap());
133        for column in 0..columns {
134            for chunk in data.chunks() {
135                if let Some(item) = chunk.get(column) {
136                    flat.push(item.clone());
137                }
138            }
139        }
140        ChunkedVec { flat, offsets }
141    }
142
143    /// Zero-copy scatter operation using slices
144    #[deprecated(
145        since = "0.5.0",
146        note = "use [`CollectiveOps::scatter`] which returns a CSR-shaped `ChunkedVec`\
147                (flat buffer + offset table) with the same chunk tiling; this slice-array\
148                form is superseded by the flat layout."
149    )]
150    pub fn scatter_zero_copy<T>(data: &[T], num_chunks: usize) -> Vec<&[T]> {
151        let chunk_size = data.len() / num_chunks;
152        let mut chunks = Vec::with_capacity(num_chunks);
153
154        for i in 0..num_chunks {
155            let start = i * chunk_size;
156            let end = if i == num_chunks - 1 {
157                data.len()
158            } else {
159                (i + 1) * chunk_size
160            };
161            chunks.push(&data[start..end]);
162        }
163
164        chunks
165    }
166
167    /// Zero-copy gather operation using iterators
168    #[deprecated(
169        since = "0.5.0",
170        note = "use [`CollectiveOps::gather`] which hands back the contiguous `ChunkedVec`\
171                storage in O(1); iteration over the flat buffer replaces this slice iterator."
172    )]
173    pub fn gather_zero_copy<'a, T, I>(chunks: I) -> impl Iterator<Item = &'a T>
174    where
175        I: IntoIterator<Item = &'a [T]>,
176        T: 'a,
177    {
178        chunks.into_iter().flat_map(|chunk| chunk.iter())
179    }
180
181    /// Zero-copy all-reduce operation
182    pub fn all_reduce_zero_copy<T, F>(data: &[T], op: F) -> T
183    where
184        T: Clone,
185        F: Fn(&T, &T) -> T,
186    {
187        data.iter()
188            .skip(1)
189            .fold(data[0].clone(), |acc, item| op(&acc, item))
190    }
191}