use ruda_tensor::{Backend, DeviceOps};
use std::collections::HashMap;
use crate::{
PeerId,
local::tensor_map::{CollectiveTensorMap, PeerDeviceMap},
};
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(devices, tensor))
)]
pub(crate) fn broadcast_tree<B: Backend>(
mut devices: PeerDeviceMap<B>,
root: PeerId,
tensor: B::FloatTensorPrimitive,
arity: u32,
) -> CollectiveTensorMap<B> {
let mut devices_vec = vec![];
let root_device = devices.remove(&root).unwrap();
for (id, tensor) in devices.drain() {
devices_vec.push((id, tensor));
}
devices_vec.sort_by(|a, b| {
let dev_a = &a.1;
let dev_b = &b.1;
dev_a.id().cmp(&dev_b.id())
});
devices_vec.insert(0, (root, root_device));
let out = broadcast_tree_inner::<B>(tensor, devices_vec, arity);
let mut tensors = HashMap::new();
for (id, tensor) in out {
tensors.insert(id, tensor);
}
tensors
}
fn broadcast_tree_inner<B: Backend>(
tensor: B::FloatTensorPrimitive,
mut all_devices: Vec<(PeerId, B::Device)>,
arity: u32,
) -> Vec<(PeerId, B::FloatTensorPrimitive)> {
let mut parents = vec![];
let mut children_groups = vec![];
while !all_devices.is_empty() {
let mut children = vec![];
let parent = all_devices.remove(0);
for _ in 0..arity {
if all_devices.is_empty() {
break;
}
children.push(all_devices.remove(0));
}
parents.push(parent);
children_groups.push(children);
}
let mut parents = if parents.len() > 1 {
broadcast_tree_inner::<B>(tensor, parents, arity)
} else {
let root = parents.first().unwrap();
vec![(root.0, tensor)]
};
let mut tensors = vec![];
for children in children_groups {
let parent = parents.remove(0);
for (child_id, child_device) in children {
let child_tensor = B::float_to_device(parent.1.clone(), &child_device);
tensors.push((child_id, child_tensor));
}
tensors.push(parent);
}
tensors
}