use crate::error::{OptimError, Result};
use scirs2_core::ndarray::{Array1, ScalarOperand};
use scirs2_core::numeric::Float;
use std::fmt::Debug;
type RankSegments<A> = Vec<Vec<A>>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReduceOp {
Sum,
Mean,
Max,
Min,
Product,
}
impl ReduceOp {
pub fn apply<A: Float>(self, a: A, b: A) -> A {
match self {
ReduceOp::Sum | ReduceOp::Mean => a + b,
ReduceOp::Product => a * b,
ReduceOp::Max => a.max(b),
ReduceOp::Min => a.min(b),
}
}
pub fn identity<A: Float>(self) -> A {
match self {
ReduceOp::Sum | ReduceOp::Mean => A::zero(),
ReduceOp::Product => A::one(),
ReduceOp::Max => A::neg_infinity(),
ReduceOp::Min => A::infinity(),
}
}
pub fn needs_mean_finalize(self) -> bool {
matches!(self, ReduceOp::Mean)
}
}
pub trait CollectiveTransport<A> {
fn world_size(&self) -> usize;
fn send_to_successor(&mut self, src_rank: usize, segment: Vec<A>) -> Result<()>;
fn recv_from_predecessor(&mut self, dst_rank: usize) -> Result<Vec<A>>;
}
#[derive(Debug)]
pub struct LocalTransport<A> {
world_size: usize,
mailboxes: Vec<Option<Vec<A>>>,
}
impl<A> LocalTransport<A> {
pub fn new(world_size: usize) -> Result<Self> {
if world_size == 0 {
return Err(OptimError::InvalidConfig(
"world_size must be at least 1".to_string(),
));
}
let mut mailboxes = Vec::with_capacity(world_size);
for _ in 0..world_size {
mailboxes.push(None);
}
Ok(Self {
world_size,
mailboxes,
})
}
}
impl<A> CollectiveTransport<A> for LocalTransport<A> {
fn world_size(&self) -> usize {
self.world_size
}
fn send_to_successor(&mut self, src_rank: usize, segment: Vec<A>) -> Result<()> {
if src_rank >= self.world_size {
return Err(OptimError::InvalidConfig(format!(
"src_rank {src_rank} out of range for world_size {}",
self.world_size
)));
}
let dst = (src_rank + 1) % self.world_size;
if self.mailboxes[dst].is_some() {
return Err(OptimError::InvalidState(format!(
"mailbox for rank {dst} already holds an uncollected message"
)));
}
self.mailboxes[dst] = Some(segment);
Ok(())
}
fn recv_from_predecessor(&mut self, dst_rank: usize) -> Result<Vec<A>> {
if dst_rank >= self.world_size {
return Err(OptimError::InvalidConfig(format!(
"dst_rank {dst_rank} out of range for world_size {}",
self.world_size
)));
}
self.mailboxes[dst_rank].take().ok_or_else(|| {
OptimError::InvalidState(format!("no message waiting in mailbox for rank {dst_rank}"))
})
}
}
#[inline]
fn modn(value: isize, n: usize) -> usize {
let m = n as isize;
(((value % m) + m) % m) as usize
}
fn compute_segment_offsets(l: usize, n: usize) -> Vec<usize> {
let base = l / n;
let rem = l % n;
let mut offsets = Vec::with_capacity(n + 1);
let mut acc = 0usize;
offsets.push(acc);
for k in 0..n {
let len = if k < rem { base + 1 } else { base };
acc += len;
offsets.push(acc);
}
offsets
}
fn ring_step<A, T>(
transport: &mut T,
states: &mut [RankSegments<A>],
step: usize,
send_offset: isize,
recv_offset: isize,
op: Option<ReduceOp>,
) -> Result<()>
where
A: Float,
T: CollectiveTransport<A>,
{
let n = states.len();
let step_i = step as isize;
for (r, segments) in states.iter().enumerate() {
let send_chunk = modn(r as isize - step_i + send_offset, n);
let segment = segments[send_chunk].clone();
transport.send_to_successor(r, segment)?;
}
for (r, segments) in states.iter_mut().enumerate() {
let recv_chunk = modn(r as isize - step_i + recv_offset, n);
let incoming = transport.recv_from_predecessor(r)?;
match op {
Some(reduce_op) => {
let slot = &mut segments[recv_chunk];
if slot.len() != incoming.len() {
return Err(OptimError::DimensionMismatch(format!(
"segment length mismatch at rank {r}, chunk {recv_chunk}: \
local {} vs received {}",
slot.len(),
incoming.len()
)));
}
for (dst, src) in slot.iter_mut().zip(incoming.iter()) {
*dst = reduce_op.apply(*dst, *src);
}
}
None => {
segments[recv_chunk] = incoming;
}
}
}
Ok(())
}
fn reduce_scatter<A, T>(
transport: &mut T,
states: &mut [RankSegments<A>],
op: ReduceOp,
) -> Result<()>
where
A: Float,
T: CollectiveTransport<A>,
{
let n = states.len();
for step in 0..(n - 1) {
ring_step(transport, states, step, 0, -1, Some(op))?;
}
Ok(())
}
fn all_gather_reduced<A, T>(transport: &mut T, states: &mut [RankSegments<A>]) -> Result<()>
where
A: Float,
T: CollectiveTransport<A>,
{
let n = states.len();
for step in 0..(n - 1) {
ring_step(transport, states, step, 1, 0, None)?;
}
Ok(())
}
fn all_gather_slots<A, T>(transport: &mut T, states: &mut [RankSegments<A>]) -> Result<()>
where
A: Float,
T: CollectiveTransport<A>,
{
let n = states.len();
for step in 0..(n - 1) {
ring_step(transport, states, step, 0, -1, None)?;
}
Ok(())
}
fn flatten_segments<A: Float>(segments: RankSegments<A>, capacity: usize) -> Array1<A> {
let mut flat = Vec::with_capacity(capacity);
for segment in segments {
flat.extend(segment);
}
Array1::from_vec(flat)
}
#[derive(Debug, Clone, Copy)]
pub struct RingAllReduce {
world_size: usize,
}
impl RingAllReduce {
pub fn new(world_size: usize) -> Result<Self> {
if world_size == 0 {
return Err(OptimError::InvalidConfig(
"world_size must be at least 1".to_string(),
));
}
Ok(Self { world_size })
}
pub fn world_size(&self) -> usize {
self.world_size
}
fn validate_equal_length(&self, inputs: &[Array1<impl Float>]) -> Result<usize> {
if inputs.is_empty() {
return Err(OptimError::InvalidConfig(
"inputs must not be empty".to_string(),
));
}
if inputs.len() != self.world_size {
return Err(OptimError::DimensionMismatch(format!(
"expected {} inputs (one per rank), got {}",
self.world_size,
inputs.len()
)));
}
let len = inputs[0].len();
if len == 0 {
return Err(OptimError::InvalidConfig(
"input vectors must have non-zero length".to_string(),
));
}
for (rank, vector) in inputs.iter().enumerate() {
if vector.len() != len {
return Err(OptimError::DimensionMismatch(format!(
"rank {rank} length {} does not match rank 0 length {len}",
vector.len()
)));
}
}
Ok(len)
}
pub fn all_reduce_all<A>(&self, inputs: &[Array1<A>], op: ReduceOp) -> Result<Vec<Array1<A>>>
where
A: Float + ScalarOperand + Debug,
{
let len = self.validate_equal_length(inputs)?;
let n = self.world_size;
if n == 1 {
return Ok(vec![inputs[0].clone()]);
}
let offsets = compute_segment_offsets(len, n);
let mut states: Vec<RankSegments<A>> = Vec::with_capacity(n);
for vector in inputs {
let slice = vector.as_slice().ok_or_else(|| {
OptimError::InvalidConfig("input vector must be contiguous".to_string())
})?;
let segments: RankSegments<A> = (0..n)
.map(|k| slice[offsets[k]..offsets[k + 1]].to_vec())
.collect();
states.push(segments);
}
let mut transport = LocalTransport::<A>::new(n)?;
reduce_scatter(&mut transport, &mut states, op)?;
all_gather_reduced(&mut transport, &mut states)?;
if op.needs_mean_finalize() {
let denom = A::from(n).ok_or_else(|| {
OptimError::InvalidConfig("cannot represent world_size as scalar".to_string())
})?;
for segments in states.iter_mut() {
for segment in segments.iter_mut() {
for value in segment.iter_mut() {
*value = *value / denom;
}
}
}
}
let results = states
.into_iter()
.map(|segments| flatten_segments(segments, len))
.collect();
Ok(results)
}
pub fn all_gather<A>(&self, inputs: &[Array1<A>]) -> Result<Vec<Array1<A>>>
where
A: Float + ScalarOperand + Debug,
{
if inputs.is_empty() {
return Err(OptimError::InvalidConfig(
"inputs must not be empty".to_string(),
));
}
if inputs.len() != self.world_size {
return Err(OptimError::DimensionMismatch(format!(
"expected {} inputs (one per rank), got {}",
self.world_size,
inputs.len()
)));
}
for (rank, vector) in inputs.iter().enumerate() {
if vector.is_empty() {
return Err(OptimError::InvalidConfig(format!(
"rank {rank} contribution must have non-zero length"
)));
}
}
let total: usize = inputs.iter().map(|vector| vector.len()).sum();
let n = self.world_size;
if n == 1 {
return Ok(vec![inputs[0].clone()]);
}
let mut states: Vec<RankSegments<A>> = Vec::with_capacity(n);
for (rank, vector) in inputs.iter().enumerate() {
let slice = vector.as_slice().ok_or_else(|| {
OptimError::InvalidConfig("input vector must be contiguous".to_string())
})?;
let mut segments: RankSegments<A> = vec![Vec::new(); n];
segments[rank] = slice.to_vec();
states.push(segments);
}
let mut transport = LocalTransport::<A>::new(n)?;
all_gather_slots(&mut transport, &mut states)?;
let results = states
.into_iter()
.map(|segments| flatten_segments(segments, total))
.collect();
Ok(results)
}
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::Array1;
fn naive_reduce(inputs: &[Array1<f64>], op: ReduceOp) -> Array1<f64> {
let len = inputs[0].len();
let mut out = vec![op.identity::<f64>(); len];
for vector in inputs {
for (acc, &value) in out.iter_mut().zip(vector.iter()) {
*acc = op.apply(*acc, value);
}
}
if op.needs_mean_finalize() {
let denom = inputs.len() as f64;
for acc in out.iter_mut() {
*acc /= denom;
}
}
Array1::from_vec(out)
}
fn assert_all_ranks_eq(results: &[Array1<f64>], expected: &Array1<f64>) {
for (rank, result) in results.iter().enumerate() {
assert_eq!(result.len(), expected.len(), "rank {rank} length mismatch");
for (index, (&got, &want)) in result.iter().zip(expected.iter()).enumerate() {
assert!(
(got - want).abs() < 1e-9,
"rank {rank} coordinate {index}: got {got}, want {want}"
);
}
}
}
#[test]
fn test_new_rejects_zero_world_size() {
assert!(RingAllReduce::new(0).is_err());
assert!(RingAllReduce::new(1).is_ok());
assert!(RingAllReduce::new(8).is_ok());
}
#[test]
fn test_ring_all_reduce_sum_matches_naive() {
let n = 4;
let l = 10;
let inputs: Vec<Array1<f64>> = (0..n)
.map(|r| Array1::from_vec((0..l).map(|i| (r * 10 + i) as f64).collect()))
.collect();
let ring = RingAllReduce::new(n).unwrap();
let results = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
let expected = Array1::from_vec((0..l).map(|i| 60.0 + 4.0 * i as f64).collect());
assert_eq!(results.len(), n);
assert_all_ranks_eq(&results, &expected);
assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Sum));
}
#[test]
fn test_ring_all_reduce_mean() {
let n = 3;
let l = 9;
let inputs: Vec<Array1<f64>> = (0..n)
.map(|r| Array1::from_vec((0..l).map(|i| (r + i) as f64).collect()))
.collect();
let ring = RingAllReduce::new(n).unwrap();
let results = ring.all_reduce_all(&inputs, ReduceOp::Mean).unwrap();
let expected = Array1::from_vec((0..l).map(|i| (i + 1) as f64).collect());
assert_all_ranks_eq(&results, &expected);
assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Mean));
}
#[test]
fn test_ring_all_reduce_max() {
let n = 4;
let l = 7;
let inputs: Vec<Array1<f64>> = (0..n)
.map(|r| Array1::from_vec((0..l).map(|i| (i * 4 + r) as f64).collect()))
.collect();
let ring = RingAllReduce::new(n).unwrap();
let results = ring.all_reduce_all(&inputs, ReduceOp::Max).unwrap();
let expected = Array1::from_vec((0..l).map(|i| (i * 4 + 3) as f64).collect());
assert_all_ranks_eq(&results, &expected);
assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Max));
}
#[test]
fn test_ring_all_reduce_min() {
let n = 4;
let l = 7;
let inputs: Vec<Array1<f64>> = (0..n)
.map(|r| Array1::from_vec((0..l).map(|i| (i * 4 + r) as f64).collect()))
.collect();
let ring = RingAllReduce::new(n).unwrap();
let results = ring.all_reduce_all(&inputs, ReduceOp::Min).unwrap();
let expected = Array1::from_vec((0..l).map(|i| (i * 4) as f64).collect());
assert_all_ranks_eq(&results, &expected);
assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Min));
}
#[test]
fn test_ring_all_reduce_product() {
let n = 3;
let l = 5;
let values = [2.0f64, 3.0, 0.5];
let inputs: Vec<Array1<f64>> = (0..n)
.map(|r| Array1::from_vec(vec![values[r]; l]))
.collect();
let ring = RingAllReduce::new(n).unwrap();
let results = ring.all_reduce_all(&inputs, ReduceOp::Product).unwrap();
let expected = Array1::from_vec(vec![3.0; l]);
assert_all_ranks_eq(&results, &expected);
assert_all_ranks_eq(&results, &naive_reduce(&inputs, ReduceOp::Product));
}
#[test]
fn test_world_size_one_identity() {
let ring = RingAllReduce::new(1).unwrap();
let inputs = vec![Array1::from_vec(vec![1.0f64, 2.0, 3.0])];
let sum = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
assert_eq!(sum.len(), 1);
assert_all_ranks_eq(&sum, &inputs[0]);
let mean = ring.all_reduce_all(&inputs, ReduceOp::Mean).unwrap();
assert_all_ranks_eq(&mean, &inputs[0]);
let gathered = ring.all_gather(&inputs).unwrap();
assert_all_ranks_eq(&gathered, &inputs[0]);
}
#[test]
fn test_length_not_divisible_by_world_size() {
let n = 3;
let l = 7;
let inputs: Vec<Array1<f64>> = (0..n)
.map(|r| Array1::from_vec((0..l).map(|i| (r * 100 + i) as f64).collect()))
.collect();
let ring = RingAllReduce::new(n).unwrap();
let results = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
let expected = Array1::from_vec((0..l).map(|i| 300.0 + 3.0 * i as f64).collect());
assert_all_ranks_eq(&results, &expected);
}
#[test]
fn test_length_smaller_than_world_size() {
let n = 4;
let inputs: Vec<Array1<f64>> = (0..n)
.map(|r| Array1::from_vec(vec![r as f64, (r + 1) as f64]))
.collect();
let ring = RingAllReduce::new(n).unwrap();
let results = ring.all_reduce_all(&inputs, ReduceOp::Sum).unwrap();
let expected = Array1::from_vec(vec![6.0, 10.0]);
assert_all_ranks_eq(&results, &expected);
}
#[test]
fn test_all_gather_round_trip() {
let inputs = vec![
Array1::from_vec(vec![1.0f64, 2.0]),
Array1::from_vec(vec![3.0, 4.0, 5.0]),
Array1::from_vec(vec![6.0]),
Array1::from_vec(vec![7.0, 8.0, 9.0, 10.0]),
];
let ring = RingAllReduce::new(4).unwrap();
let results = ring.all_gather(&inputs).unwrap();
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]);
assert_eq!(results.len(), 4);
assert_all_ranks_eq(&results, &expected);
}
#[test]
fn test_invalid_inputs_are_rejected() {
let ring = RingAllReduce::new(3).unwrap();
let empty: Vec<Array1<f64>> = vec![];
assert!(ring.all_reduce_all(&empty, ReduceOp::Sum).is_err());
assert!(ring.all_gather(&empty).is_err());
let wrong_count = vec![Array1::from_vec(vec![1.0f64]), Array1::from_vec(vec![2.0])];
assert!(ring.all_reduce_all(&wrong_count, ReduceOp::Sum).is_err());
assert!(ring.all_gather(&wrong_count).is_err());
let inconsistent = vec![
Array1::from_vec(vec![1.0f64, 2.0]),
Array1::from_vec(vec![3.0, 4.0]),
Array1::from_vec(vec![5.0]),
];
assert!(ring.all_reduce_all(&inconsistent, ReduceOp::Sum).is_err());
let zero_len = vec![
Array1::<f64>::from_vec(vec![]),
Array1::from_vec(vec![]),
Array1::from_vec(vec![]),
];
assert!(ring.all_reduce_all(&zero_len, ReduceOp::Sum).is_err());
assert!(ring.all_gather(&zero_len).is_err());
}
#[test]
fn test_compute_segment_offsets() {
assert_eq!(compute_segment_offsets(7, 3), vec![0, 3, 5, 7]);
assert_eq!(compute_segment_offsets(10, 4), vec![0, 3, 6, 8, 10]);
assert_eq!(compute_segment_offsets(6, 3), vec![0, 2, 4, 6]);
assert_eq!(compute_segment_offsets(2, 4), vec![0, 1, 2, 2, 2]);
let offsets = compute_segment_offsets(10, 4);
let lengths: Vec<usize> = offsets.windows(2).map(|w| w[1] - w[0]).collect();
let max_len = *lengths.iter().max().unwrap();
let min_len = *lengths.iter().min().unwrap();
assert!(max_len - min_len <= 1);
assert_eq!(lengths.iter().sum::<usize>(), 10);
}
#[test]
fn test_local_transport_mailbox_semantics() {
let mut transport = LocalTransport::<f64>::new(3).unwrap();
assert_eq!(transport.world_size(), 3);
transport.send_to_successor(0, vec![1.0, 2.0]).unwrap();
assert!(transport.recv_from_predecessor(0).is_err());
assert_eq!(transport.recv_from_predecessor(1).unwrap(), vec![1.0, 2.0]);
assert!(transport.recv_from_predecessor(1).is_err());
transport.send_to_successor(0, vec![3.0]).unwrap();
assert!(transport.send_to_successor(0, vec![4.0]).is_err());
assert!(transport.send_to_successor(3, vec![0.0]).is_err());
assert!(transport.recv_from_predecessor(3).is_err());
}
#[test]
fn test_reduce_op_identity_and_apply() {
assert_eq!(ReduceOp::Sum.identity::<f64>(), 0.0);
assert_eq!(ReduceOp::Product.identity::<f64>(), 1.0);
assert_eq!(ReduceOp::Max.identity::<f64>(), f64::NEG_INFINITY);
assert_eq!(ReduceOp::Min.identity::<f64>(), f64::INFINITY);
assert_eq!(ReduceOp::Sum.apply(2.0, 3.0), 5.0);
assert_eq!(ReduceOp::Product.apply(2.0, 3.0), 6.0);
assert_eq!(ReduceOp::Max.apply(2.0, 3.0), 3.0);
assert_eq!(ReduceOp::Min.apply(2.0, 3.0), 2.0);
assert_eq!(ReduceOp::Mean.apply(2.0, 3.0), 5.0);
assert!(ReduceOp::Mean.needs_mean_finalize());
assert!(!ReduceOp::Sum.needs_mean_finalize());
}
}