moirai_core/communication/
collective.rs1#![expect(
2 clippy::unwrap_used,
3 reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
4)]
5
6#[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 #[must_use]
23 pub fn num_chunks(&self) -> usize {
24 self.offsets.len().saturating_sub(1)
25 }
26
27 #[must_use]
29 pub fn len(&self) -> usize {
30 self.flat.len()
31 }
32
33 #[must_use]
35 pub fn is_empty(&self) -> bool {
36 self.flat.is_empty()
37 }
38
39 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 #[must_use]
48 pub fn into_flat(self) -> Vec<T> {
49 self.flat
50 }
51}
52
53pub struct CollectiveOps;
55
56impl CollectiveOps {
57 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 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 vec![current[0].clone(); result_len]
87 }
88
89 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 pub fn gather<T>(chunks: ChunkedVec<T>) -> Vec<T> {
112 chunks.into_flat()
113 }
114
115 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 #[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 #[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 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}