use crate::{PeerId, local::tensor_map::CollectiveTensorMap};
use ruda_tensor::{Backend, DeviceOps};
use std::collections::HashMap;
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(tensors))
)]
pub(crate) fn all_reduce_sum_tree<B: Backend>(
tensors: CollectiveTensorMap<B>,
arity: u32,
) -> CollectiveTensorMap<B> {
let mut input = tensors.into_iter().collect::<Vec<_>>();
input.sort_by(|a, b| {
let dev_a = B::float_device(&a.1);
let dev_b = B::float_device(&b.1);
dev_a.id().cmp(&dev_b.id())
});
let out = all_reduce_sum_tree_inner::<B>(input, arity);
let mut tensors = HashMap::new();
for (id, tensor) in out {
tensors.insert(id, tensor);
}
tensors
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(tensors))
)]
fn all_reduce_sum_tree_inner<B: Backend>(
mut tensors: Vec<(PeerId, B::FloatTensorPrimitive)>,
arity: u32,
) -> Vec<(PeerId, B::FloatTensorPrimitive)> {
let mut parent_tensors = vec![];
let mut children_groups = vec![];
while !tensors.is_empty() {
let mut children = vec![];
let (parent, mut parent_tensor) = tensors.remove(0);
let parent_device = B::float_device(&parent_tensor);
for _ in 0..arity {
if tensors.is_empty() {
break;
}
let (child, mut child_tensor) = tensors.remove(0);
let child_device = B::float_device(&child_tensor);
children.push((child, child_device));
child_tensor = B::float_to_device(child_tensor, &parent_device);
parent_tensor = B::float_add(parent_tensor, child_tensor);
}
parent_tensors.push((parent, parent_tensor));
children_groups.push(children);
}
if parent_tensors.len() > 1 {
parent_tensors = all_reduce_sum_tree_inner::<B>(parent_tensors, arity);
}
for (parent, parent_tensor) in parent_tensors {
let children = children_groups.remove(0);
for (child, child_device) in children {
tensors.push((
child,
B::float_to_device(parent_tensor.clone(), &child_device),
));
}
tensors.push((parent, parent_tensor));
}
tensors
}