use super::tree::all_reduce_sum_tree;
use crate::local::tensor_map::CollectiveTensorMap;
use crate::{PeerId, local::tensor_map};
use ruda_tensor::{Backend, Shape, Slice, TensorMetadata, tensor::FloatTensor};
use std::{collections::HashMap, ops::Range};
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(tensors))
)]
pub(crate) fn all_reduce_sum_ring<B: Backend>(
tensors: CollectiveTensorMap<B>,
) -> CollectiveTensorMap<B> {
let shape = tensor_map::get_common_shape::<B>(&tensors)
.expect("Cannot aggregate tensors with different sizes");
let slice_dim = get_slice_dim(&shape);
let slice_dim_size = shape[slice_dim];
let tensor_count = tensors.len();
if slice_dim_size < tensor_count {
return all_reduce_sum_tree::<B>(tensors, 2);
}
let mut sliced_tensors = slice_tensors::<B>(tensors, shape, slice_dim);
ring_cycles::<B>(&mut sliced_tensors, true);
ring_cycles::<B>(&mut sliced_tensors, false);
sliced_tensors
.into_iter()
.map(|(id, slices)| (id, B::float_cat(slices, slice_dim)))
.collect()
}
pub(crate) fn get_slice_dim(shape: &Shape) -> usize {
shape
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.cmp(b))
.map(|(index, _)| index)
.unwrap()
}
fn ring_cycles<B: Backend>(
sliced_tensors: &mut [(PeerId, Vec<B::FloatTensorPrimitive>)],
is_phase_one: bool,
) {
let tensor_count = sliced_tensors.len();
for cycle in 0..(tensor_count - 1) {
for i in 0..tensor_count {
let src_tensor_idx = i;
let dest_tensor_idx = (i + 1) % tensor_count;
let slice_idx = if is_phase_one {
(i + (tensor_count - 1) * cycle) % tensor_count
} else {
(i + 1 + (tensor_count - 1) * cycle) % tensor_count
};
let src_slice = sliced_tensors[src_tensor_idx].1.remove(slice_idx);
let mut dest_slice = sliced_tensors[dest_tensor_idx].1.remove(slice_idx);
let dest_device = B::float_device(&dest_slice);
let src_slice_on_dest = B::float_to_device(src_slice.clone(), &dest_device);
if is_phase_one {
dest_slice = B::float_add(dest_slice, src_slice_on_dest);
} else {
let slices: Vec<Slice> = dest_slice
.shape()
.iter()
.map(|&d| Slice::new(0, Some(d as isize), 1))
.collect();
dest_slice =
B::float_slice_assign(dest_slice, slices.as_slice(), src_slice_on_dest);
}
sliced_tensors[src_tensor_idx]
.1
.insert(slice_idx, src_slice);
sliced_tensors[dest_tensor_idx]
.1
.insert(slice_idx, dest_slice);
}
}
}
fn slice_tensors<B: Backend>(
mut tensors: HashMap<PeerId, FloatTensor<B>>,
shape: Shape,
slice_dim: usize,
) -> Vec<(PeerId, Vec<FloatTensor<B>>)> {
let ranges = get_ring_reduce_slice_ranges(shape[slice_dim], tensors.len());
let mut sliced_tensors = vec![];
for (id, tensor) in tensors.drain() {
let mut slices = vec![];
for range in &ranges {
let full_range = shape
.iter()
.enumerate()
.map(|(dim_idx, dim)| {
if dim_idx == slice_dim {
Slice::from(range.clone())
} else {
Slice::from(0..*dim)
}
})
.collect::<Vec<_>>();
let slice = B::float_slice(tensor.clone(), &full_range);
slices.push(slice);
}
sliced_tensors.push((id, slices));
}
sliced_tensors
}
pub(crate) fn get_ring_reduce_slice_ranges(
slice_dim_size: usize,
slice_count: usize,
) -> Vec<Range<usize>> {
let mut ranges: Vec<Range<usize>> = vec![];
let slice_size = slice_dim_size.div_ceil(slice_count);
for i in 0..slice_count {
let start = i * slice_size;
let end = start + slice_size;
ranges.push(Range { start, end });
}
ranges.last_mut().unwrap().end = slice_dim_size;
ranges
}