use crate::local::tensor_map::{CollectiveTensorMap, PeerDeviceMap};
use crate::{
BroadcastStrategy, CollectiveConfig, CollectiveError, PeerId,
local::{broadcast_centralized, broadcast_tree},
node::base::Node,
};
use ruda_tensor::Backend;
#[allow(unused_imports)] use ruda_tensor::TensorMetadata;
use ruda_communication::Protocol;
use std::sync::mpsc::SyncSender;
pub struct BroadcastOp<B: Backend> {
calls: Vec<BroadcastOpCall<B>>,
tensor: Option<B::FloatTensorPrimitive>,
root: Option<PeerId>,
}
pub struct BroadcastOpCall<B: Backend> {
caller: PeerId,
device: B::Device,
result_sender: SyncSender<BroadcastResult<B::FloatTensorPrimitive>>,
}
pub(crate) type BroadcastResult<T> = Result<T, CollectiveError>;
impl<B: Backend> BroadcastOp<B> {
pub fn new() -> Self {
Self {
calls: vec![],
tensor: None,
root: None,
}
}
pub fn effective_root(&self) -> PeerId {
self.root.unwrap_or(self.calls.first().unwrap().caller)
}
#[allow(dead_code)]
pub fn peers(&self) -> Vec<PeerId> {
self.calls.iter().map(|c| c.caller).collect()
}
fn peer_devices(&self) -> PeerDeviceMap<B> {
self.calls
.iter()
.map(|call| (call.caller, call.device.clone()))
.collect()
}
pub fn register_call(
&mut self,
caller: PeerId,
input: Option<B::FloatTensorPrimitive>,
result_sender: SyncSender<BroadcastResult<B::FloatTensorPrimitive>>,
device: B::Device,
peer_count: usize,
) -> Result<bool, CollectiveError> {
if input.is_some() {
if self.tensor.is_some() {
return Err(CollectiveError::BroadcastMultipleTensors);
}
self.tensor = input;
}
self.calls.push(BroadcastOpCall {
caller,
device,
result_sender,
});
Ok(self.calls.len() == peer_count)
}
#[cfg_attr(feature = "tracing", tracing::instrument(
level="trace",
skip(self, config, global_client),
fields(
self.peers = ?self.peers(),
self.shape = ?self.tensor.as_ref().map(|t| t.shape()),
self.dtype = ?self.tensor.as_ref().map(|t| t.dtype()),
)
))]
pub async fn execute<P: Protocol>(
mut self,
config: &CollectiveConfig,
global_client: &mut Option<Node<B, P>>,
) {
match self.broadcast(config, global_client).await {
Ok(mut tensors) => {
self.calls.iter().for_each(|call| {
let result = tensors
.remove(&call.caller)
.expect("tensor/peer internal mismatch.");
call.result_sender.send(Ok(result)).unwrap();
});
assert_eq!(tensors.len(), 0, "tensor/peer internal mismatch.");
}
Err(err) => {
self.fail(err);
}
}
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(self, config, global_client))
)]
async fn broadcast<P: Protocol>(
&mut self,
config: &CollectiveConfig,
global_client: &mut Option<Node<B, P>>,
) -> Result<CollectiveTensorMap<B>, CollectiveError> {
if let Some(global_client) = &global_client {
let strategy = config
.global_broadcast_strategy
.expect("global_broadcast_strategy not defined");
self.tensor = Some(
global_client
.broadcast(self.tensor.clone(), strategy)
.await
.map_err(CollectiveError::Global)?,
)
}
let Some(tensor) = self.tensor.take() else {
return Err(CollectiveError::BroadcastNoTensor);
};
let root = self.effective_root();
let peer_devices = self.peer_devices();
Ok(match config.local_broadcast_strategy {
BroadcastStrategy::Tree(arity) => {
broadcast_tree::<B>(peer_devices, root, tensor, arity)
}
BroadcastStrategy::Centralized => {
broadcast_centralized::<B>(peer_devices, root, tensor)
}
})
}
pub fn fail(self, err: CollectiveError) {
self.calls.iter().for_each(|call| {
call.result_sender.send(Err(err.clone())).unwrap();
});
}
}