use std::path::PathBuf;
use std::time::Duration;
use super::error::DistError;
#[non_exhaustive]
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Devices {
Single(usize),
All,
Count(usize),
List(Vec<usize>),
}
impl Devices {
pub fn parse(s: &str) -> Result<Self, DistError> {
let t = s.trim();
if t.eq_ignore_ascii_case("all") {
return Ok(Devices::All);
}
if t.contains(',') {
let mut list = Vec::new();
for part in t.split(',') {
let ord: usize = part.trim().parse().map_err(|_| {
DistError::Config(format!("bad device ordinal {part:?} in {s:?}"))
})?;
if list.contains(&ord) {
return Err(DistError::Config(format!(
"device ordinal {ord} repeats in {s:?}"
)));
}
list.push(ord);
}
return Ok(Devices::List(list));
}
let n: usize = t
.parse()
.map_err(|_| DistError::Config(format!("bad device spec {s:?}")))?;
if n == 0 {
return Err(DistError::Config("device count must be positive".into()));
}
Ok(if n == 1 {
Devices::Single(0)
} else {
Devices::Count(n)
})
}
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum ReduceContract {
#[default]
FixedOrder,
NcclSum,
}
#[non_exhaustive]
#[derive(Clone, Debug)]
pub enum Rendezvous {
File { dir: PathBuf, job_id: String },
}
impl Default for Rendezvous {
fn default() -> Self {
Rendezvous::File {
dir: std::env::temp_dir().join("mamba-rs-rendezvous"),
job_id: String::new(),
}
}
}
#[non_exhaustive]
#[derive(Clone, Debug)]
pub struct DistConfig {
pub devices: Devices,
pub logical_world: Option<usize>,
pub rendezvous: Rendezvous,
pub reduce: ReduceContract,
pub seed: u64,
pub init_timeout: Duration,
pub collective_timeout: Duration,
}
impl Default for DistConfig {
fn default() -> Self {
Self {
devices: Devices::Single(0),
logical_world: None,
rendezvous: Rendezvous::default(),
reduce: ReduceContract::default(),
seed: 0,
init_timeout: Duration::from_secs(120),
collective_timeout: Duration::from_secs(300),
}
}
}
impl DistConfig {
#[must_use]
pub fn with_devices(mut self, d: Devices) -> Self {
self.devices = d;
self
}
#[must_use]
pub fn with_logical_world(mut self, w: usize) -> Self {
self.logical_world = Some(w);
self
}
#[must_use]
pub fn with_rendezvous(mut self, r: Rendezvous) -> Self {
self.rendezvous = r;
self
}
#[must_use]
pub fn with_reduce(mut self, r: ReduceContract) -> Self {
self.reduce = r;
self
}
#[must_use]
pub fn with_seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
pub fn validate(&self) -> Result<(), DistError> {
if let Devices::List(l) = &self.devices
&& l.is_empty()
{
return Err(DistError::Config("empty device list".into()));
}
if let Devices::Count(0) = self.devices {
return Err(DistError::Config("device count must be positive".into()));
}
if let Some(0) = self.logical_world {
return Err(DistError::Config("logical world must be positive".into()));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn device_grammar() {
assert_eq!(Devices::parse("1").unwrap(), Devices::Single(0));
assert_eq!(Devices::parse("4").unwrap(), Devices::Count(4));
assert_eq!(Devices::parse("all").unwrap(), Devices::All);
assert_eq!(
Devices::parse("0,2,3").unwrap(),
Devices::List(vec![0, 2, 3])
);
assert!(Devices::parse("0,0").is_err());
assert!(Devices::parse("").is_err());
assert!(Devices::parse("0x2").is_err());
}
}