use crate::{
CollectiveConfig, CollectiveError, PeerId, ReduceOperation,
global::node::base::Node,
local::{
AllReduceOp, AllReduceResult, BroadcastOp, BroadcastResult, ReduceOp, ReduceResult,
client::LocalCollectiveClient,
},
};
use ruda_tensor::{Backend, TensorMetadata, tensor::FloatTensor};
use ruda_communication::websocket::{WebSocket, WsServer};
use std::sync::{MutexGuard, OnceLock};
use std::{
any::{Any, TypeId},
collections::HashMap,
fmt::Debug,
sync::{
Arc, Mutex,
mpsc::{Receiver, SyncSender},
},
};
use tokio::runtime::{Builder, Runtime};
type Network = WebSocket;
pub(crate) type RegisterResult = Result<(), CollectiveError>;
pub(crate) type FinishResult = Result<(), CollectiveError>;
pub(crate) struct LocalCollectiveServer<B: Backend> {
message_rec: Receiver<Message<B>>,
config: Option<CollectiveConfig>,
peers: Vec<PeerId>,
callbacks_register: Vec<SyncSender<RegisterResult>>,
devices: HashMap<PeerId, B::Device>,
all_reduce_op: Option<AllReduceOp<B>>,
reduce_op: Option<ReduceOp<B>>,
broadcast_op: Option<BroadcastOp<B>>,
global_client: Option<Node<B, Network>>,
}
#[derive(Debug)]
pub(crate) enum Message<B: Backend> {
Register {
device_id: PeerId,
device: B::Device,
config: CollectiveConfig,
callback: SyncSender<RegisterResult>,
},
AllReduce {
device_id: PeerId,
tensor: B::FloatTensorPrimitive,
op: ReduceOperation,
callback: SyncSender<AllReduceResult<B::FloatTensorPrimitive>>,
},
Reduce {
device_id: PeerId,
tensor: B::FloatTensorPrimitive,
op: ReduceOperation,
root: PeerId,
callback: SyncSender<ReduceResult<B::FloatTensorPrimitive>>,
},
Broadcast {
device_id: PeerId,
tensor: Option<B::FloatTensorPrimitive>,
callback: SyncSender<BroadcastResult<B::FloatTensorPrimitive>>,
},
Reset,
Finish {
id: PeerId,
callback: SyncSender<FinishResult>,
},
}
type LocalClientBox = Box<dyn Any + Send + Sync>;
static BACKEND_CLIENT_MAP: OnceLock<Mutex<HashMap<TypeId, LocalClientBox>>> = OnceLock::new();
pub(crate) fn get_backend_client_map() -> MutexGuard<'static, HashMap<TypeId, LocalClientBox>> {
BACKEND_CLIENT_MAP
.get_or_init(Default::default)
.lock()
.unwrap()
}
pub(crate) fn get_collective_client<B: Backend>() -> LocalCollectiveClient<B> {
let typeid = TypeId::of::<B>();
let mut state_map = get_backend_client_map();
match state_map.get(&typeid) {
Some(val) => val.downcast_ref().cloned().unwrap(),
None => {
let client = LocalCollectiveServer::<B>::setup(LocalCollectiveClientConfig::default());
state_map.insert(typeid, Box::new(client.clone()));
client
}
}
}
static SERVER_RUNTIME: OnceLock<Arc<Runtime>> = OnceLock::new();
pub(crate) fn get_collective_server_runtime() -> Arc<Runtime> {
SERVER_RUNTIME
.get_or_init(|| {
Builder::new_multi_thread()
.enable_all()
.build()
.expect("Unable to initialize runtime")
.into()
})
.clone()
}
pub struct LocalCollectiveClientConfig {
pub channel_capacity: usize,
}
impl Default for LocalCollectiveClientConfig {
fn default() -> Self {
Self {
channel_capacity: 50,
}
}
}
impl From<usize> for LocalCollectiveClientConfig {
fn from(capacity: usize) -> Self {
Self {
channel_capacity: capacity,
}
}
}
impl<B: Backend> LocalCollectiveServer<B> {
fn new(rec: Receiver<Message<B>>) -> Self {
Self {
message_rec: rec,
config: None,
peers: vec![],
devices: HashMap::new(),
all_reduce_op: None,
reduce_op: None,
broadcast_op: None,
callbacks_register: vec![],
global_client: None,
}
}
pub(crate) fn setup<C>(cfg: C) -> LocalCollectiveClient<B>
where
C: Into<LocalCollectiveClientConfig>,
{
let cfg = cfg.into();
let (tx, rx) = std::sync::mpsc::sync_channel(cfg.channel_capacity);
get_collective_server_runtime().spawn(async {
let typeid = TypeId::of::<B>();
log::info!("Starting server for backend: {typeid:?}");
let mut server = LocalCollectiveServer::new(rx);
loop {
match server.message_rec.recv() {
Ok(message) => server.process_message(message).await,
Err(err) => {
log::error!(
"Error receiving message from local collective server: {err:?}"
);
break;
}
}
}
});
LocalCollectiveClient { channel: tx }
}
async fn process_message(&mut self, message: Message<B>) {
match message {
Message::Register {
device_id,
device,
config,
callback,
} => {
self.process_register_message(device_id, device, config, &callback)
.await
}
Message::AllReduce {
device_id,
tensor,
op,
callback,
} => {
self.process_all_reduce_message(device_id, tensor, op, callback)
.await
}
Message::Reduce {
device_id,
tensor,
op,
root,
callback,
} => {
self.process_reduce_message(device_id, tensor, op, root, callback)
.await
}
Message::Broadcast {
device_id,
tensor,
callback,
} => {
self.process_broadcast_message(device_id, tensor, callback)
.await
}
Message::Reset => self.reset(),
Message::Finish { id, callback } => self.process_finish_message(id, callback).await,
}
}
async fn process_register_message(
&mut self,
device_id: PeerId,
device: B::Device,
config: CollectiveConfig,
callback: &SyncSender<RegisterResult>,
) {
if !config.is_valid() {
callback.send(Err(CollectiveError::InvalidConfig)).unwrap();
return;
}
if self.peers.contains(&device_id) {
callback
.send(Err(CollectiveError::MultipleRegister))
.unwrap();
return;
}
if self.peers.is_empty() || self.config.is_none() {
self.config = Some(config);
} else if let Some(cfg) = &self.config
&& *cfg != config
{
callback
.send(Err(CollectiveError::RegisterParamsMismatch))
.unwrap();
return;
}
self.peers.push(device_id);
self.callbacks_register.push(callback.clone());
self.devices.insert(device_id, device);
let config = self.config.as_ref().unwrap();
let global_params = config.global_register_params();
if let Some(global_params) = &global_params
&& self.global_client.is_none()
{
let server = WsServer::new(global_params.data_service_port);
let client = Node::new(&global_params.global_address, server);
self.global_client = Some(client)
}
if self.peers.len() == config.num_devices {
let mut register_result = Ok(());
if let Some(global_params) = global_params {
let client = self
.global_client
.as_mut()
.expect("Global client should be initialized");
register_result = client
.register(self.peers.clone(), global_params)
.await
.map_err(CollectiveError::Global);
};
self.callbacks_register
.drain(..)
.for_each(|tx| tx.send(register_result.clone()).unwrap());
}
}
async fn process_all_reduce_message(
&mut self,
peer_id: PeerId,
tensor: FloatTensor<B>,
op: ReduceOperation,
callback: SyncSender<AllReduceResult<FloatTensor<B>>>,
) {
if !self.peers.contains(&peer_id) {
callback
.send(Err(CollectiveError::RegisterNotFirstOperation))
.unwrap();
return;
}
if self.all_reduce_op.is_none() {
self.all_reduce_op = Some(AllReduceOp::new(tensor.shape(), op));
}
let mut all_reduce_op = self.all_reduce_op.take().unwrap();
let res =
all_reduce_op.register_call(peer_id, tensor, callback.clone(), op, self.peers.len());
match res {
Ok(is_ready) => {
if is_ready {
all_reduce_op
.execute(self.config.as_ref().unwrap(), &mut self.global_client)
.await;
} else {
self.all_reduce_op = Some(all_reduce_op)
}
}
Err(err) => all_reduce_op.fail(err),
}
}
async fn process_reduce_message(
&mut self,
peer_id: PeerId,
tensor: FloatTensor<B>,
op: ReduceOperation,
root: PeerId,
callback: SyncSender<ReduceResult<B::FloatTensorPrimitive>>,
) {
if !self.peers.contains(&root) {
callback
.send(Err(CollectiveError::RegisterNotFirstOperation))
.unwrap();
return;
}
if self.reduce_op.is_none() {
self.reduce_op = Some(ReduceOp::new(tensor.shape(), op, root));
}
let mut reduce_op = self.reduce_op.take().unwrap();
let res = reduce_op.register_call(
peer_id,
tensor,
callback.clone(),
op,
root,
self.peers.len(),
);
match res {
Ok(is_ready) => {
if is_ready {
reduce_op
.execute(root, self.config.as_ref().unwrap(), &mut self.global_client)
.await;
} else {
self.reduce_op = Some(reduce_op)
}
}
Err(err) => reduce_op.fail(err),
}
}
async fn process_broadcast_message(
&mut self,
caller: PeerId,
tensor: Option<FloatTensor<B>>,
callback: SyncSender<BroadcastResult<B::FloatTensorPrimitive>>,
) {
if self.config.is_none() {
callback
.send(Err(CollectiveError::RegisterNotFirstOperation))
.unwrap();
return;
}
if !self.peers.contains(&caller) {
callback
.send(Err(CollectiveError::RegisterNotFirstOperation))
.unwrap();
return;
}
if self.broadcast_op.is_none() {
self.broadcast_op = Some(BroadcastOp::new());
}
let device = self.devices.get(&caller).unwrap().clone();
let mut broadcast_op = self.broadcast_op.take().unwrap();
let res =
broadcast_op.register_call(caller, tensor, callback.clone(), device, self.peers.len());
match res {
Ok(is_ready) => {
if is_ready {
broadcast_op
.execute(self.config.as_ref().unwrap(), &mut self.global_client)
.await;
} else {
self.broadcast_op = Some(broadcast_op)
}
}
Err(err) => broadcast_op.fail(err),
}
}
fn reset(&mut self) {
self.peers.clear();
self.all_reduce_op = None;
self.reduce_op = None;
self.broadcast_op = None;
}
async fn process_finish_message(&mut self, id: PeerId, callback: SyncSender<RegisterResult>) {
if self.config.is_none() {
callback
.send(Err(CollectiveError::RegisterNotFirstOperation))
.unwrap();
return;
}
if !self.peers.contains(&id) {
callback
.send(Err(CollectiveError::MultipleUnregister))
.unwrap();
return;
}
self.peers.retain(|x| *x != id);
if self.peers.is_empty()
&& let Some(mut global_client) = self.global_client.take()
{
global_client.finish().await;
}
callback.send(Ok(())).unwrap();
}
}