use ruda_tensor::Backend;
use crate::local::tensor_map::{CollectiveTensorMap, get_peer_devices};
use crate::local::{broadcast_centralized, reduce_sum_centralized};
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(tensors))
)]
pub(crate) fn all_reduce_sum_centralized<B: Backend>(
tensors: CollectiveTensorMap<B>,
) -> CollectiveTensorMap<B> {
let peer_devices = get_peer_devices::<B>(&tensors);
let central_device = *tensors.keys().next().unwrap();
let central_tensor = reduce_sum_centralized::<B>(tensors, ¢ral_device);
broadcast_centralized::<B>(peer_devices, central_device, central_tensor)
}