use std::collections::HashMap;
use crate::{
PeerId,
local::tensor_map::{CollectiveTensorMap, PeerDeviceMap},
};
use ruda_tensor::Backend;
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(devices, tensor))
)]
pub(crate) fn broadcast_centralized<B: Backend>(
mut devices: PeerDeviceMap<B>,
central: PeerId,
tensor: B::FloatTensorPrimitive,
) -> CollectiveTensorMap<B> {
let mut output = HashMap::new();
devices
.remove(¢ral)
.expect("Central device id is in `devices`");
for (dest, dest_device) in devices {
let tensor = B::float_to_device(tensor.clone(), &dest_device);
output.insert(dest, tensor);
}
output.insert(central, tensor);
output
}