use ruda_tensor::{Backend, TensorMetadata};
use ruda_communication::Protocol;
use ruda_communication::data_service::TensorDataServer;
use ruda_communication::{Address, ProtocolServer, data_service::TensorDataService};
use std::collections::HashMap;
use std::{marker::PhantomData, sync::Arc};
use tokio::sync::RwLock;
use tokio_util::sync::CancellationToken;
use crate::node::sync::SyncService;
use crate::{
AllReduceStrategy, PeerId, ReduceOperation,
global::{
node::{
centralized::centralized_all_reduce_sum, ring::ring_all_reduce_sum,
tree::tree_all_reduce_sum, worker::GlobalClientWorker,
},
shared::{GlobalCollectiveError, RemoteRequest, RemoteResponse},
},
local::server::get_collective_server_runtime,
};
use crate::{BroadcastStrategy, GlobalRegisterParams, NodeId, ReduceStrategy};
pub(crate) struct NodeState {
pub node_id: NodeId,
pub nodes: HashMap<NodeId, Address>,
pub num_global_devices: u32,
}
pub(crate) struct Node<B, P>
where
B: Backend,
P: Protocol,
{
state: Arc<RwLock<Option<NodeState>>>,
data_service: Arc<TensorDataService<B, P>>,
sync_service: Arc<SyncService<P>>,
worker: GlobalClientWorker<P::Client>,
operation_lock: tokio::sync::Mutex<()>,
_n: PhantomData<P>,
}
impl<B, P> Node<B, P>
where
B: Backend,
P: Protocol,
{
pub fn new(global_address: &Address, comms_server: P::Server) -> Self {
let state = Arc::new(tokio::sync::RwLock::new(None));
let cancel_token = CancellationToken::new();
let data_service = Arc::new(TensorDataService::new(cancel_token.clone()));
let sync_service = Arc::new(SyncService::new(state.clone()));
let runtime = get_collective_server_runtime();
let server = comms_server
.route_tensor_data_service(data_service.clone())
.route("/sync", {
let sync_service = sync_service.clone();
async move |channel: <P::Server as ProtocolServer>::Channel| {
sync_service.handle_sync_connection(channel).await;
}
})
.serve({
let cancel_token = cancel_token.clone();
async move { cancel_token.cancelled().await }
});
runtime.spawn(server);
let worker = GlobalClientWorker::new(&runtime, cancel_token.clone(), global_address);
Self {
state,
data_service,
sync_service,
worker,
operation_lock: tokio::sync::Mutex::new(()),
_n: PhantomData,
}
}
pub async fn register(
&mut self,
peers: Vec<PeerId>,
global_params: GlobalRegisterParams,
) -> Result<(), GlobalCollectiveError> {
let req = RemoteRequest::Register {
node_addr: global_params.node_address,
num_nodes: global_params.num_nodes,
peers,
};
match self.worker.request(req).await {
RemoteResponse::Register {
node_id,
nodes,
num_global_devices,
} => {
let mut state = self.state.write().await;
*state = Some(NodeState {
node_id,
nodes,
num_global_devices,
});
}
RemoteResponse::Error(err) => {
self.worker.abort();
return Err(err);
}
resp => {
log::error!("Response to a register request should be an ack, not {resp:?}");
return Err(GlobalCollectiveError::WrongOrchestratorResponse);
}
}
Ok(())
}
async fn guarded<T>(&self, work: impl std::future::Future<Output=Result<T, GlobalCollectiveError>>)
-> Result<T, GlobalCollectiveError>
{
let mut guard = super::worker::CancelOnDrop { token: self.worker.token(), completed: false };
let queued = async {
let _serial = self.operation_lock.lock().await;
if self.worker.token().is_cancelled() { return Err(GlobalCollectiveError::CommunicatorAborted); }
work.await
};
let token = self.worker.token();
let result = tokio::select! {
_ = token.cancelled() => Err(GlobalCollectiveError::CommunicatorAborted),
value = tokio::time::timeout(self.worker.policy().collective_timeout, queued) =>
value.map_err(|_| GlobalCollectiveError::OperationTimeout).and_then(|result| result),
};
guard.completed = result.is_ok();
result
}
pub async fn all_reduce(&self, tensor: B::FloatTensorPrimitive, strategy: AllReduceStrategy,
op: ReduceOperation) -> Result<B::FloatTensorPrimitive, GlobalCollectiveError>
{ self.guarded(self.all_reduce_inner(tensor, strategy, op)).await }
pub async fn reduce(&self, tensor: B::FloatTensorPrimitive, strategy: ReduceStrategy,
root: PeerId, op: ReduceOperation) -> Result<Option<B::FloatTensorPrimitive>, GlobalCollectiveError>
{ self.guarded(self.reduce_inner(tensor, strategy, root, op)).await }
pub async fn broadcast(&self, tensor: Option<B::FloatTensorPrimitive>, strategy: BroadcastStrategy,
device: &B::Device) -> Result<B::FloatTensorPrimitive, GlobalCollectiveError>
{ self.guarded(self.broadcast_inner(tensor, strategy, device)).await }
async fn all_reduce_inner(
&self,
tensor: B::FloatTensorPrimitive,
strategy: AllReduceStrategy,
op: ReduceOperation,
) -> Result<B::FloatTensorPrimitive, GlobalCollectiveError> {
let state = self.state.read().await;
let Some(ref state) = *state else {
return Err(GlobalCollectiveError::AllReduceBeforeRegister);
};
self.begin(crate::global::shared::CollectiveSpec::AllReduce {
op, strategy, shape: tensor.shape().to_vec(), dtype: tensor.dtype(),
}).await?;
let node = state.node_id;
let nodes = &state.nodes;
let mut result = match strategy {
AllReduceStrategy::Centralized => {
centralized_all_reduce_sum(
node,
nodes,
&self.data_service,
self.sync_service.clone(),
tensor,
)
.await?
}
AllReduceStrategy::Tree(arity) => {
tree_all_reduce_sum(
node,
nodes,
self.data_service.clone(),
self.sync_service.clone(),
tensor,
arity,
)
.await?
}
AllReduceStrategy::Ring => {
ring_all_reduce_sum(
node,
nodes,
self.data_service.clone(),
self.sync_service.clone(),
tensor,
)
.await?
}
};
if op == ReduceOperation::Mean {
result = B::float_div_scalar(result, (state.num_global_devices as f32).into());
}
Ok(result)
}
async fn begin(&self, spec: crate::global::shared::CollectiveSpec)
-> Result<(NodeId, u64), GlobalCollectiveError>
{
match self.worker.request(RemoteRequest::Begin(spec)).await {
RemoteResponse::Begin { root_node, transfer_id } => Ok((root_node, transfer_id)),
RemoteResponse::Error(error) => Err(error),
_ => Err(GlobalCollectiveError::WrongOrchestratorResponse),
}
}
async fn reduce_inner(&self, tensor: B::FloatTensorPrimitive, strategy: ReduceStrategy,
root: PeerId, op: ReduceOperation)
-> Result<Option<B::FloatTensorPrimitive>, GlobalCollectiveError>
{
let guard = self.state.read().await;
let state = guard.as_ref().ok_or(GlobalCollectiveError::CollectiveBeforeRegister)?;
let (root_node, transfer) = self.begin(crate::global::shared::CollectiveSpec::Reduce {
root, op, strategy, shape: tensor.shape().to_vec(), dtype: tensor.dtype(),
}).await?;
let arity = match strategy { ReduceStrategy::Centralized => None, ReduceStrategy::Tree(k) => Some(k) };
let result = super::rooted::reduce_sum::<B, P>(state.node_id, root_node, &state.nodes,
&self.data_service, tensor, arity, transfer).await?;
let result = result.map(|tensor| if op == ReduceOperation::Mean {
B::float_div_scalar(tensor, (state.num_global_devices as f32).into())
} else { tensor });
self.sync_service.sync().await;
Ok(result)
}
async fn broadcast_inner(&self, tensor: Option<B::FloatTensorPrimitive>, strategy: BroadcastStrategy,
device: &B::Device) -> Result<B::FloatTensorPrimitive, GlobalCollectiveError>
{
let guard = self.state.read().await;
let state = guard.as_ref().ok_or(GlobalCollectiveError::CollectiveBeforeRegister)?;
let metadata = tensor.as_ref().map(|t| (t.shape().to_vec(), t.dtype()));
let (root_node, transfer) = self.begin(crate::global::shared::CollectiveSpec::Broadcast {
strategy, metadata,
}).await?;
let arity = match strategy { BroadcastStrategy::Centralized => None, BroadcastStrategy::Tree(k) => Some(k) };
let result = super::rooted::broadcast::<B, P>(state.node_id, root_node, &state.nodes,
&self.data_service, tensor, device, arity, transfer).await?;
self.sync_service.sync().await;
Ok(result)
}
pub async fn finish(&mut self) {
let res = self.worker.close_connection().await;
if let Err(err) = res {
log::error!("Global collective client error: {err:?}");
}
if tokio::time::timeout(self.worker.policy().request_timeout,
self.data_service.close()).await.is_err()
{
log::warn!("Global collective data-channel shutdown timed out");
}
}
}