1use crate::error::{OptimError, Result};
36use scirs2_core::ndarray::{Array1, ScalarOperand};
37use scirs2_core::numeric::Float;
38use std::fmt::Debug;
39
40type RankSegments<A> = Vec<Vec<A>>;
42
43#[derive(Debug, Clone, Copy, PartialEq, Eq)]
50pub enum ReduceOp {
51 Sum,
53 Mean,
55 Max,
57 Min,
59 Product,
61}
62
63impl ReduceOp {
64 pub fn apply<A: Float>(self, a: A, b: A) -> A {
70 match self {
71 ReduceOp::Sum | ReduceOp::Mean => a + b,
72 ReduceOp::Product => a * b,
73 ReduceOp::Max => a.max(b),
74 ReduceOp::Min => a.min(b),
75 }
76 }
77
78 pub fn identity<A: Float>(self) -> A {
82 match self {
83 ReduceOp::Sum | ReduceOp::Mean => A::zero(),
84 ReduceOp::Product => A::one(),
85 ReduceOp::Max => A::neg_infinity(),
86 ReduceOp::Min => A::infinity(),
87 }
88 }
89
90 pub fn needs_mean_finalize(self) -> bool {
94 matches!(self, ReduceOp::Mean)
95 }
96}
97
98pub trait CollectiveTransport<A> {
110 fn world_size(&self) -> usize;
112
113 fn send_to_successor(&mut self, src_rank: usize, segment: Vec<A>) -> Result<()>;
115
116 fn recv_from_predecessor(&mut self, dst_rank: usize) -> Result<Vec<A>>;
118}
119
120#[derive(Debug)]
126pub struct LocalTransport<A> {
127 world_size: usize,
128 mailboxes: Vec<Option<Vec<A>>>,
129}
130
131impl<A> LocalTransport<A> {
132 pub fn new(world_size: usize) -> Result<Self> {
136 if world_size == 0 {
137 return Err(OptimError::InvalidConfig(
138 "world_size must be at least 1".to_string(),
139 ));
140 }
141 let mut mailboxes = Vec::with_capacity(world_size);
142 for _ in 0..world_size {
143 mailboxes.push(None);
144 }
145 Ok(Self {
146 world_size,
147 mailboxes,
148 })
149 }
150}
151
152impl<A> CollectiveTransport<A> for LocalTransport<A> {
153 fn world_size(&self) -> usize {
154 self.world_size
155 }
156
157 fn send_to_successor(&mut self, src_rank: usize, segment: Vec<A>) -> Result<()> {
158 if src_rank >= self.world_size {
159 return Err(OptimError::InvalidConfig(format!(
160 "src_rank {src_rank} out of range for world_size {}",
161 self.world_size
162 )));
163 }
164 let dst = (src_rank + 1) % self.world_size;
165 if self.mailboxes[dst].is_some() {
166 return Err(OptimError::InvalidState(format!(
167 "mailbox for rank {dst} already holds an uncollected message"
168 )));
169 }
170 self.mailboxes[dst] = Some(segment);
171 Ok(())
172 }
173
174 fn recv_from_predecessor(&mut self, dst_rank: usize) -> Result<Vec<A>> {
175 if dst_rank >= self.world_size {
176 return Err(OptimError::InvalidConfig(format!(
177 "dst_rank {dst_rank} out of range for world_size {}",
178 self.world_size
179 )));
180 }
181 self.mailboxes[dst_rank].take().ok_or_else(|| {
182 OptimError::InvalidState(format!("no message waiting in mailbox for rank {dst_rank}"))
183 })
184 }
185}
186
187#[inline]
189fn modn(value: isize, n: usize) -> usize {
190 let m = n as isize;
191 (((value % m) + m) % m) as usize
192}
193
194fn compute_segment_offsets(l: usize, n: usize) -> Vec<usize> {
198 let base = l / n;
199 let rem = l % n;
200 let mut offsets = Vec::with_capacity(n + 1);
201 let mut acc = 0usize;
202 offsets.push(acc);
203 for k in 0..n {
204 let len = if k < rem { base + 1 } else { base };
205 acc += len;
206 offsets.push(acc);
207 }
208 offsets
209}
210
211fn ring_step<A, T>(
219 transport: &mut T,
220 states: &mut [RankSegments<A>],
221 step: usize,
222 send_offset: isize,
223 recv_offset: isize,
224 op: Option<ReduceOp>,
225) -> Result<()>
226where
227 A: Float,
228 T: CollectiveTransport<A>,
229{
230 let n = states.len();
231 let step_i = step as isize;
232
233 for (r, segments) in states.iter().enumerate() {
236 let send_chunk = modn(r as isize - step_i + send_offset, n);
237 let segment = segments[send_chunk].clone();
238 transport.send_to_successor(r, segment)?;
239 }
240
241 for (r, segments) in states.iter_mut().enumerate() {
243 let recv_chunk = modn(r as isize - step_i + recv_offset, n);
244 let incoming = transport.recv_from_predecessor(r)?;
245 match op {
246 Some(reduce_op) => {
247 let slot = &mut segments[recv_chunk];
248 if slot.len() != incoming.len() {
249 return Err(OptimError::DimensionMismatch(format!(
250 "segment length mismatch at rank {r}, chunk {recv_chunk}: \
251 local {} vs received {}",
252 slot.len(),
253 incoming.len()
254 )));
255 }
256 for (dst, src) in slot.iter_mut().zip(incoming.iter()) {
257 *dst = reduce_op.apply(*dst, *src);
258 }
259 }
260 None => {
261 segments[recv_chunk] = incoming;
262 }
263 }
264 }
265
266 Ok(())
267}
268
269fn reduce_scatter<A, T>(
273 transport: &mut T,
274 states: &mut [RankSegments<A>],
275 op: ReduceOp,
276) -> Result<()>
277where
278 A: Float,
279 T: CollectiveTransport<A>,
280{
281 let n = states.len();
282 for step in 0..(n - 1) {
283 ring_step(transport, states, step, 0, -1, Some(op))?;
284 }
285 Ok(())
286}
287
288fn all_gather_reduced<A, T>(transport: &mut T, states: &mut [RankSegments<A>]) -> Result<()>
291where
292 A: Float,
293 T: CollectiveTransport<A>,
294{
295 let n = states.len();
296 for step in 0..(n - 1) {
297 ring_step(transport, states, step, 1, 0, None)?;
298 }
299 Ok(())
300}
301
302fn all_gather_slots<A, T>(transport: &mut T, states: &mut [RankSegments<A>]) -> Result<()>
305where
306 A: Float,
307 T: CollectiveTransport<A>,
308{
309 let n = states.len();
310 for step in 0..(n - 1) {
311 ring_step(transport, states, step, 0, -1, None)?;
312 }
313 Ok(())
314}
315
316fn flatten_segments<A: Float>(segments: RankSegments<A>, capacity: usize) -> Array1<A> {
319 let mut flat = Vec::with_capacity(capacity);
320 for segment in segments {
321 flat.extend(segment);
322 }
323 Array1::from_vec(flat)
324}
325
326#[derive(Debug, Clone, Copy)]
333pub struct RingAllReduce {
334 world_size: usize,
335}
336
337impl RingAllReduce {
338 pub fn new(world_size: usize) -> Result<Self> {
342 if world_size == 0 {
343 return Err(OptimError::InvalidConfig(
344 "world_size must be at least 1".to_string(),
345 ));
346 }
347 Ok(Self { world_size })
348 }
349
350 pub fn world_size(&self) -> usize {
352 self.world_size
353 }
354
355 fn validate_equal_length(&self, inputs: &[Array1<impl Float>]) -> Result<usize> {
358 if inputs.is_empty() {
359 return Err(OptimError::InvalidConfig(
360 "inputs must not be empty".to_string(),
361 ));
362 }
363 if inputs.len() != self.world_size {
364 return Err(OptimError::DimensionMismatch(format!(
365 "expected {} inputs (one per rank), got {}",
366 self.world_size,
367 inputs.len()
368 )));
369 }
370 let len = inputs[0].len();
371 if len == 0 {
372 return Err(OptimError::InvalidConfig(
373 "input vectors must have non-zero length".to_string(),
374 ));
375 }
376 for (rank, vector) in inputs.iter().enumerate() {
377 if vector.len() != len {
378 return Err(OptimError::DimensionMismatch(format!(
379 "rank {rank} length {} does not match rank 0 length {len}",
380 vector.len()
381 )));
382 }
383 }
384 Ok(len)
385 }
386
387 pub fn all_reduce_all<A>(&self, inputs: &[Array1<A>], op: ReduceOp) -> Result<Vec<Array1<A>>>
393 where
394 A: Float + ScalarOperand + Debug,
395 {
396 let len = self.validate_equal_length(inputs)?;
397 let n = self.world_size;
398
399 if n == 1 {
402 return Ok(vec![inputs[0].clone()]);
403 }
404
405 let offsets = compute_segment_offsets(len, n);
407 let mut states: Vec<RankSegments<A>> = Vec::with_capacity(n);
408 for vector in inputs {
409 let slice = vector.as_slice().ok_or_else(|| {
410 OptimError::InvalidConfig("input vector must be contiguous".to_string())
411 })?;
412 let segments: RankSegments<A> = (0..n)
413 .map(|k| slice[offsets[k]..offsets[k + 1]].to_vec())
414 .collect();
415 states.push(segments);
416 }
417
418 let mut transport = LocalTransport::<A>::new(n)?;
419 reduce_scatter(&mut transport, &mut states, op)?;
420 all_gather_reduced(&mut transport, &mut states)?;
421
422 if op.needs_mean_finalize() {
424 let denom = A::from(n).ok_or_else(|| {
425 OptimError::InvalidConfig("cannot represent world_size as scalar".to_string())
426 })?;
427 for segments in states.iter_mut() {
428 for segment in segments.iter_mut() {
429 for value in segment.iter_mut() {
430 *value = *value / denom;
431 }
432 }
433 }
434 }
435
436 let results = states
437 .into_iter()
438 .map(|segments| flatten_segments(segments, len))
439 .collect();
440 Ok(results)
441 }
442
443 pub fn all_gather<A>(&self, inputs: &[Array1<A>]) -> Result<Vec<Array1<A>>>
450 where
451 A: Float + ScalarOperand + Debug,
452 {
453 if inputs.is_empty() {
454 return Err(OptimError::InvalidConfig(
455 "inputs must not be empty".to_string(),
456 ));
457 }
458 if inputs.len() != self.world_size {
459 return Err(OptimError::DimensionMismatch(format!(
460 "expected {} inputs (one per rank), got {}",
461 self.world_size,
462 inputs.len()
463 )));
464 }
465 for (rank, vector) in inputs.iter().enumerate() {
466 if vector.is_empty() {
467 return Err(OptimError::InvalidConfig(format!(
468 "rank {rank} contribution must have non-zero length"
469 )));
470 }
471 }
472
473 let total: usize = inputs.iter().map(|vector| vector.len()).sum();
474 let n = self.world_size;
475
476 if n == 1 {
478 return Ok(vec![inputs[0].clone()]);
479 }
480
481 let mut states: Vec<RankSegments<A>> = Vec::with_capacity(n);
484 for (rank, vector) in inputs.iter().enumerate() {
485 let slice = vector.as_slice().ok_or_else(|| {
486 OptimError::InvalidConfig("input vector must be contiguous".to_string())
487 })?;
488 let mut segments: RankSegments<A> = vec![Vec::new(); n];
489 segments[rank] = slice.to_vec();
490 states.push(segments);
491 }
492
493 let mut transport = LocalTransport::<A>::new(n)?;
494 all_gather_slots(&mut transport, &mut states)?;
495
496 let results = states
497 .into_iter()
498 .map(|segments| flatten_segments(segments, total))
499 .collect();
500 Ok(results)
501 }
502}
503
504#[cfg(test)]
505mod tests {
506 use super::*;
507 use scirs2_core::ndarray::Array1;
508
509 fn naive_reduce(inputs: &[Array1<f64>], op: ReduceOp) -> Array1<f64> {
511 let len = inputs[0].len();
512 let mut out = vec![op.identity::<f64>(); len];
513 for vector in inputs {
514 for (acc, &value) in out.iter_mut().zip(vector.iter()) {
515 *acc = op.apply(*acc, value);
516 }
517 }
518 if op.needs_mean_finalize() {
519 let denom = inputs.len() as f64;
520 for acc in out.iter_mut() {
521 *acc /= denom;
522 }
523 }
524 Array1::from_vec(out)
525 }
526
527 fn assert_all_ranks_eq(results: &[Array1<f64>], expected: &Array1<f64>) {
528 for (rank, result) in results.iter().enumerate() {
529 assert_eq!(result.len(), expected.len(), "rank {rank} length mismatch");
530 for (index, (&got, &want)) in result.iter().zip(expected.iter()).enumerate() {
531 assert!(
532 (got - want).abs() < 1e-9,
533 "rank {rank} coordinate {index}: got {got}, want {want}"
534 );
535 }
536 }
537 }
538
539 #[test]
540 fn test_new_rejects_zero_world_size() {
541 assert!(RingAllReduce::new(0).is_err());
542 assert!(RingAllReduce::new(1).is_ok());
543 assert!(RingAllReduce::new(8).is_ok());
544 }
545
546 #[test]
547 fn test_ring_all_reduce_sum_matches_naive() {
548 let n = 4;
550 let l = 10;
551 let inputs: Vec<Array1<f64>> = (0..n)
552 .map(|r| Array1::from_vec((0..l).map(|i| (r * 10 + i) as f64).collect()))
553 .collect();
554
555 let ring = RingAllReduce::new(n).unwrap();
556 let results = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
557
558 let expected = Array1::from_vec((0..l).map(|i| 60.0 + 4.0 * i as f64).collect());
560 assert_eq!(results.len(), n);
561 assert_all_ranks_eq(&results, &expected);
562 assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Sum));
563 }
564
565 #[test]
566 fn test_ring_all_reduce_mean() {
567 let n = 3;
569 let l = 9;
570 let inputs: Vec<Array1<f64>> = (0..n)
571 .map(|r| Array1::from_vec((0..l).map(|i| (r + i) as f64).collect()))
572 .collect();
573
574 let ring = RingAllReduce::new(n).unwrap();
575 let results = ring.all_reduce_all(&inputs, ReduceOp::Mean).unwrap();
576
577 let expected = Array1::from_vec((0..l).map(|i| (i + 1) as f64).collect());
579 assert_all_ranks_eq(&results, &expected);
580 assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Mean));
581 }
582
583 #[test]
584 fn test_ring_all_reduce_max() {
585 let n = 4;
587 let l = 7;
588 let inputs: Vec<Array1<f64>> = (0..n)
589 .map(|r| Array1::from_vec((0..l).map(|i| (i * 4 + r) as f64).collect()))
590 .collect();
591
592 let ring = RingAllReduce::new(n).unwrap();
593 let results = ring.all_reduce_all(&inputs, ReduceOp::Max).unwrap();
594
595 let expected = Array1::from_vec((0..l).map(|i| (i * 4 + 3) as f64).collect());
597 assert_all_ranks_eq(&results, &expected);
598 assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Max));
599 }
600
601 #[test]
602 fn test_ring_all_reduce_min() {
603 let n = 4;
604 let l = 7;
605 let inputs: Vec<Array1<f64>> = (0..n)
606 .map(|r| Array1::from_vec((0..l).map(|i| (i * 4 + r) as f64).collect()))
607 .collect();
608
609 let ring = RingAllReduce::new(n).unwrap();
610 let results = ring.all_reduce_all(&inputs, ReduceOp::Min).unwrap();
611
612 let expected = Array1::from_vec((0..l).map(|i| (i * 4) as f64).collect());
614 assert_all_ranks_eq(&results, &expected);
615 assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Min));
616 }
617
618 #[test]
619 fn test_ring_all_reduce_product() {
620 let n = 3;
622 let l = 5;
623 let values = [2.0f64, 3.0, 0.5];
624 let inputs: Vec<Array1<f64>> = (0..n)
625 .map(|r| Array1::from_vec(vec![values[r]; l]))
626 .collect();
627
628 let ring = RingAllReduce::new(n).unwrap();
629 let results = ring.all_reduce_all(&inputs, ReduceOp::Product).unwrap();
630
631 let expected = Array1::from_vec(vec![3.0; l]);
632 assert_all_ranks_eq(&results, &expected);
633 assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Product));
634 }
635
636 #[test]
637 fn test_world_size_one_identity() {
638 let ring = RingAllReduce::new(1).unwrap();
639 let inputs = vec![Array1::from_vec(vec![1.0f64, 2.0, 3.0])];
640
641 let sum = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
642 assert_eq!(sum.len(), 1);
643 assert_all_ranks_eq(&sum, &inputs[0]);
644
645 let mean = ring.all_reduce_all(&inputs, ReduceOp::Mean).unwrap();
646 assert_all_ranks_eq(&mean, &inputs[0]);
647
648 let gathered = ring.all_gather(&inputs).unwrap();
649 assert_all_ranks_eq(&gathered, &inputs[0]);
650 }
651
652 #[test]
653 fn test_length_not_divisible_by_world_size() {
654 let n = 3;
656 let l = 7;
657 let inputs: Vec<Array1<f64>> = (0..n)
658 .map(|r| Array1::from_vec((0..l).map(|i| (r * 100 + i) as f64).collect()))
659 .collect();
660
661 let ring = RingAllReduce::new(n).unwrap();
662 let results = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
663
664 let expected = Array1::from_vec((0..l).map(|i| 300.0 + 3.0 * i as f64).collect());
666 assert_all_ranks_eq(&results, &expected);
667 }
668
669 #[test]
670 fn test_length_smaller_than_world_size() {
671 let n = 4;
673 let inputs: Vec<Array1<f64>> = (0..n)
674 .map(|r| Array1::from_vec(vec![r as f64, (r + 1) as f64]))
675 .collect();
676
677 let ring = RingAllReduce::new(n).unwrap();
678 let results = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
679
680 let expected = Array1::from_vec(vec![6.0, 10.0]);
682 assert_all_ranks_eq(&results, &expected);
683 }
684
685 #[test]
686 fn test_all_gather_round_trip() {
687 let inputs = vec![
689 Array1::from_vec(vec![1.0f64, 2.0]),
690 Array1::from_vec(vec![3.0, 4.0, 5.0]),
691 Array1::from_vec(vec![6.0]),
692 Array1::from_vec(vec![7.0, 8.0, 9.0, 10.0]),
693 ];
694
695 let ring = RingAllReduce::new(4).unwrap();
696 let results = ring.all_gather(&inputs).unwrap();
697
698 let expected = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0]);
699 assert_eq!(results.len(), 4);
700 assert_all_ranks_eq(&results, &expected);
701 }
702
703 #[test]
704 fn test_invalid_inputs_are_rejected() {
705 let ring = RingAllReduce::new(3).unwrap();
706
707 let empty: Vec<Array1<f64>> = vec![];
709 assert!(ring.all_reduce_all(&empty, ReduceOp::Sum).is_err());
710 assert!(ring.all_gather(&empty).is_err());
711
712 let wrong_count = vec![Array1::from_vec(vec![1.0f64]), Array1::from_vec(vec![2.0])];
714 assert!(ring.all_reduce_all(&wrong_count, ReduceOp::Sum).is_err());
715 assert!(ring.all_gather(&wrong_count).is_err());
716
717 let inconsistent = vec![
719 Array1::from_vec(vec![1.0f64, 2.0]),
720 Array1::from_vec(vec![3.0, 4.0]),
721 Array1::from_vec(vec![5.0]),
722 ];
723 assert!(ring.all_reduce_all(&inconsistent, ReduceOp::Sum).is_err());
724
725 let zero_len = vec![
727 Array1::<f64>::from_vec(vec![]),
728 Array1::from_vec(vec![]),
729 Array1::from_vec(vec![]),
730 ];
731 assert!(ring.all_reduce_all(&zero_len, ReduceOp::Sum).is_err());
732 assert!(ring.all_gather(&zero_len).is_err());
733 }
734
735 #[test]
736 fn test_compute_segment_offsets() {
737 assert_eq!(compute_segment_offsets(7, 3), vec![0, 3, 5, 7]);
738 assert_eq!(compute_segment_offsets(10, 4), vec![0, 3, 6, 8, 10]);
739 assert_eq!(compute_segment_offsets(6, 3), vec![0, 2, 4, 6]);
740 assert_eq!(compute_segment_offsets(2, 4), vec![0, 1, 2, 2, 2]);
742
743 let offsets = compute_segment_offsets(10, 4);
745 let lengths: Vec<usize> = offsets.windows(2).map(|w| w[1] - w[0]).collect();
746 let max_len = *lengths.iter().max().unwrap();
747 let min_len = *lengths.iter().min().unwrap();
748 assert!(max_len - min_len <= 1);
749 assert_eq!(lengths.iter().sum::<usize>(), 10);
750 }
751
752 #[test]
753 fn test_local_transport_mailbox_semantics() {
754 let mut transport = LocalTransport::<f64>::new(3).unwrap();
755 assert_eq!(transport.world_size(), 3);
756
757 transport.send_to_successor(0, vec![1.0, 2.0]).unwrap();
759 assert!(transport.recv_from_predecessor(0).is_err());
761 assert_eq!(transport.recv_from_predecessor(1).unwrap(), vec![1.0, 2.0]);
763 assert!(transport.recv_from_predecessor(1).is_err());
764
765 transport.send_to_successor(0, vec![3.0]).unwrap();
767 assert!(transport.send_to_successor(0, vec![4.0]).is_err());
768
769 assert!(transport.send_to_successor(3, vec![0.0]).is_err());
771 assert!(transport.recv_from_predecessor(3).is_err());
772 }
773
774 #[test]
775 fn test_reduce_op_identity_and_apply() {
776 assert_eq!(ReduceOp::Sum.identity::<f64>(), 0.0);
777 assert_eq!(ReduceOp::Product.identity::<f64>(), 1.0);
778 assert_eq!(ReduceOp::Max.identity::<f64>(), f64::NEG_INFINITY);
779 assert_eq!(ReduceOp::Min.identity::<f64>(), f64::INFINITY);
780
781 assert_eq!(ReduceOp::Sum.apply(2.0, 3.0), 5.0);
782 assert_eq!(ReduceOp::Product.apply(2.0, 3.0), 6.0);
783 assert_eq!(ReduceOp::Max.apply(2.0, 3.0), 3.0);
784 assert_eq!(ReduceOp::Min.apply(2.0, 3.0), 2.0);
785 assert_eq!(ReduceOp::Mean.apply(2.0, 3.0), 5.0);
787 assert!(ReduceOp::Mean.needs_mean_finalize());
788 assert!(!ReduceOp::Sum.needs_mean_finalize());
789 }
790}