ruCCL 0.21.26

Ruda collective communication algorithms and orchestration.
Documentation
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};

/// Must be synchronized between all nodes for collective operations to work
pub(crate) struct NodeState {
    pub node_id: NodeId,
    pub nodes: HashMap<NodeId, Address>,
    pub num_global_devices: u32,
}

/// A node talks to the global orchestrator as well as other nodes with a peer-to-peer service
pub(crate) struct Node<B, P>
where
    B: Backend,
    P: Protocol,
{
    // State is written during `register` and read during other operations,
    // sometimes by multiple threads (ex. syncing during an all-reduce)
    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(())
    }

    /// Performs an all-reduce
    ///
    /// Reads the NodeState
    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:?}");
        }
        // A disconnected data peer must not make explicit shutdown wait
        // forever after the control worker has already aborted this epoch.
        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");
        }
    }
}