use std::fmt::Display;
use ruda_communication::Address;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct CollectiveConfig {
pub(crate) num_devices: usize,
pub(crate) local_all_reduce_strategy: AllReduceStrategy,
pub(crate) local_reduce_strategy: ReduceStrategy,
pub(crate) local_broadcast_strategy: BroadcastStrategy,
pub(crate) num_nodes: Option<u32>,
pub(crate) global_address: Option<Address>,
pub(crate) node_address: Option<Address>,
pub(crate) data_service_port: Option<u16>,
pub(crate) global_all_reduce_strategy: Option<AllReduceStrategy>,
pub(crate) global_reduce_strategy: Option<ReduceStrategy>,
pub(crate) global_broadcast_strategy: Option<BroadcastStrategy>,
}
impl Default for CollectiveConfig {
fn default() -> Self {
Self::new()
}
}
impl Display for CollectiveConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let num_devices = self.num_devices;
let local_all_reduce_strategy = self.local_all_reduce_strategy;
let local_reduce_strategy = self.local_reduce_strategy;
let local_broadcast_strategy = self.local_broadcast_strategy;
let num_nodes = self.num_nodes;
let global_address = &self.global_address;
let node_address = &self.node_address;
let data_service_port = self.data_service_port;
let global_all_reduce_strategy = self.global_all_reduce_strategy;
let global_reduce_strategy = self.global_reduce_strategy;
let global_broadcast_strategy = self.global_broadcast_strategy;
write!(
f,
r#"
CollectiveConfig {{
num_devices: {num_devices:?},
local_all_reduce_strategy: {local_all_reduce_strategy:?},
local_reduce_strategy: {local_reduce_strategy:?},
local_broadcast_strategy: {local_broadcast_strategy:?},
num_nodes: {num_nodes:?},
global_address: {global_address:?},
node_address: {node_address:?},
data_service_port: {data_service_port:?},
global_all_reduce_strategy: {global_all_reduce_strategy:?},
global_reduce_strategy: {global_reduce_strategy:?},
global_broadcast_strategy: {global_broadcast_strategy:?},
}}
"#
)
}
}
impl CollectiveConfig {
fn new() -> Self {
Self {
num_devices: 1,
local_all_reduce_strategy: AllReduceStrategy::Tree(2),
local_reduce_strategy: ReduceStrategy::Tree(2),
local_broadcast_strategy: BroadcastStrategy::Tree(2),
num_nodes: None,
global_address: None,
node_address: None,
data_service_port: None,
global_all_reduce_strategy: Some(AllReduceStrategy::Ring),
global_reduce_strategy: Some(ReduceStrategy::Tree(2)),
global_broadcast_strategy: Some(BroadcastStrategy::Tree(2)),
}
}
pub fn with_num_devices(mut self, num: usize) -> Self {
self.num_devices = num;
self
}
pub fn with_local_all_reduce_strategy(mut self, strategy: AllReduceStrategy) -> Self {
self.local_all_reduce_strategy = strategy;
self
}
pub fn with_local_reduce_strategy(mut self, strategy: ReduceStrategy) -> Self {
self.local_reduce_strategy = strategy;
self
}
pub fn with_local_broadcast_strategy(mut self, strategy: BroadcastStrategy) -> Self {
self.local_broadcast_strategy = strategy;
self
}
pub fn with_num_nodes(mut self, n: u32) -> Self {
self.num_nodes = Some(n);
self
}
pub fn with_global_address(mut self, addr: Address) -> Self {
self.global_address = Some(addr);
self
}
pub fn with_node_address(mut self, addr: Address) -> Self {
self.node_address = Some(addr);
self
}
pub fn with_data_service_port(mut self, port: u16) -> Self {
self.data_service_port = Some(port);
self
}
pub fn with_global_all_reduce_strategy(mut self, strategy: AllReduceStrategy) -> Self {
self.global_all_reduce_strategy = Some(strategy);
self
}
pub fn with_global_reduce_strategy(mut self, strategy: ReduceStrategy) -> Self {
self.global_reduce_strategy = Some(strategy);
self
}
pub fn with_global_broadcast_strategy(mut self, strategy: BroadcastStrategy) -> Self {
self.global_broadcast_strategy = Some(strategy);
self
}
pub fn is_valid(&self) -> bool {
match (
self.num_nodes,
&self.global_address,
&self.node_address,
self.data_service_port,
) {
(None, None, None, None) => true,
(Some(_), Some(_), Some(_), Some(_)) => true,
_ => false,
}
}
pub(crate) fn global_register_params(&self) -> Option<GlobalRegisterParams> {
match (
self.num_nodes,
&self.global_address,
&self.node_address,
self.data_service_port,
) {
(None, None, None, None) => None,
(Some(num_nodes), Some(global_addr), Some(node_addr), Some(data_service_port)) => {
Some(GlobalRegisterParams {
num_nodes,
global_address: global_addr.clone(),
node_address: node_addr.clone(),
data_service_port,
})
}
_ => None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GlobalRegisterParams {
pub global_address: Address,
pub node_address: Address,
pub data_service_port: u16,
pub num_nodes: u32,
}
#[derive(Debug, PartialEq, Clone, Serialize, Deserialize)]
pub struct SharedAllReduceParams {
pub op: ReduceOperation,
pub local_strategy: AllReduceStrategy,
pub global_strategy: Option<AllReduceStrategy>,
}
#[derive(Debug, PartialEq, Clone, Serialize, Deserialize)]
pub struct SharedReduceParams {}
#[derive(Debug, PartialEq, Clone, Serialize, Deserialize)]
pub struct SharedBroadcastParams {
pub op: ReduceOperation,
pub local_strategy: BroadcastStrategy,
pub global_strategy: Option<BroadcastStrategy>,
}
#[derive(Debug, PartialEq, Clone, Copy, Serialize, Deserialize)]
pub enum ReduceOperation {
Sum,
Mean,
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, Serialize, Deserialize)]
pub enum AllReduceStrategy {
Centralized,
Tree(u32),
Ring,
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, Serialize, Deserialize)]
pub enum ReduceStrategy {
Centralized,
Tree(u32),
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, Serialize, Deserialize)]
pub enum BroadcastStrategy {
Centralized,
Tree(u32),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct PeerId(pub u32);
impl core::fmt::Display for PeerId {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "PeerId({})", self.0)
}
}
impl From<u32> for PeerId {
fn from(value: u32) -> Self {
Self(value)
}
}
impl From<i32> for PeerId {
fn from(value: i32) -> Self {
Self(value as u32)
}
}
impl From<usize> for PeerId {
fn from(value: usize) -> Self {
Self(value as u32)
}
}