use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct DdpConfig {
pub mode: Option<String>,
pub policy: Option<String>,
pub backend: Option<String>,
pub anchor: Option<serde_json::Value>,
pub max_anchor: Option<u32>,
pub overhead_target: Option<f64>,
pub divergence_threshold: Option<f64>,
pub max_batch_diff: Option<serde_json::Value>,
pub speed_hint: Option<SpeedHint>,
pub partition_ratios: Option<Vec<f64>>,
pub progressive: Option<serde_json::Value>,
pub max_grad_norm: Option<f64>,
pub lr_scale_ratio: Option<f64>,
pub snapshot_timeout: Option<u32>,
pub checkpoint_every: Option<u32>,
pub timeline: Option<bool>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct SpeedHint {
pub slow_rank: usize,
pub ratio: f64,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct TrainingConfig {
pub epochs: Option<u32>,
pub batch_size: Option<u32>,
pub batches_per_epoch: Option<u32>,
pub lr: Option<f64>,
pub seed: Option<u64>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct OutputConfig {
pub dir: Option<String>,
pub timeline: Option<bool>,
pub monitor: Option<u16>,
}
pub const DEFAULT_CONTROLLER_PORT: u16 = 1337;
fn default_controller_port() -> u16 {
DEFAULT_CONTROLLER_PORT
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ClusterConfig {
pub controller: ClusterController,
pub workers: Vec<ClusterWorker>,
#[serde(default, skip_serializing_if = "std::collections::BTreeMap::is_empty")]
pub env: std::collections::BTreeMap<String, String>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ClusterController {
pub host: String,
#[serde(default = "default_controller_port")]
pub port: u16,
pub path: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub docker: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub arch: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub data_path: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub join: Option<ClusterJoin>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ClusterJoin {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub min_rank_start: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub join_timeout: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub target_ranks: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_join_timeout: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub open_admission: Option<bool>,
}
impl ClusterController {
pub fn effective_data_path(&self) -> &str {
self.data_path.as_deref().unwrap_or(DEFAULT_DATA_PATH)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum LocalDevices {
All,
Explicit(Vec<u8>),
}
impl LocalDevices {
pub fn is_all(&self) -> bool {
matches!(self, LocalDevices::All)
}
pub fn as_explicit(&self) -> Option<&[u8]> {
match self {
LocalDevices::All => None,
LocalDevices::Explicit(v) => Some(v.as_slice()),
}
}
}
impl Serialize for LocalDevices {
fn serialize<S: serde::Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
match self {
LocalDevices::All => s.serialize_str("all"),
LocalDevices::Explicit(v) => v.serialize(s),
}
}
}
impl<'de> Deserialize<'de> for LocalDevices {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
use serde::de::Error;
let v = Value::deserialize(d)?;
match v {
Value::String(s) if s == "all" => Ok(LocalDevices::All),
Value::String(s) => Err(D::Error::custom(format!(
"local_devices: expected \"all\" or array of device indices, got string {s:?}"
))),
Value::Array(arr) => {
let mut out = Vec::with_capacity(arr.len());
for (i, item) in arr.iter().enumerate() {
let n = item.as_u64().ok_or_else(|| {
D::Error::custom(format!(
"local_devices[{i}]: expected integer device index, got {item}"
))
})?;
let d = u8::try_from(n).map_err(|_| {
D::Error::custom(format!(
"local_devices[{i}]: value {n} does not fit in u8"
))
})?;
out.push(d);
}
Ok(LocalDevices::Explicit(out))
}
_ => Err(D::Error::custom(
"local_devices: expected \"all\" or array of device indices",
)),
}
}
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct SshConfig {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub target: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub port: Option<u16>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub identity_file: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub options: Vec<String>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ClusterWorker {
pub host: String,
#[serde(default)]
pub ranks: Vec<usize>,
pub local_devices: LocalDevices,
pub nccl_socket_ifname: String,
pub path: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub ssh: Option<SshConfig>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub tunnel: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub arch: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub data_path: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub docker: Option<String>,
#[serde(default, skip_serializing_if = "std::collections::BTreeMap::is_empty")]
pub env: std::collections::BTreeMap<String, String>,
}
pub const DEFAULT_DATA_PATH: &str = "/flodl/data";
impl ClusterWorker {
pub fn effective_data_path(&self) -> &str {
self.data_path.as_deref().unwrap_or(DEFAULT_DATA_PATH)
}
}
impl ClusterConfig {
pub fn world_size(&self) -> usize {
self.workers.iter().map(|w| w.ranks.len()).sum()
}
pub fn spans_multiple_hosts(&self) -> bool {
self.workers.len() > 1
}
pub fn validate(&self) -> Result<(), String> {
if self.controller.host.trim().is_empty() {
return Err("cluster.controller.host must be non-empty".into());
}
if self.controller.path.trim().is_empty() {
return Err("cluster.controller.path must be non-empty".into());
}
if self.workers.is_empty() {
return Err("cluster.workers must be non-empty".into());
}
for k in self.env.keys() {
if crate::cluster::is_reserved_cluster_env_key(k) {
return Err(format!(
"cluster.env: key {k:?} is reserved (launcher-owned) and \
cannot be set via env — it would override the launcher's \
per-rank value"
));
}
}
for (i, w) in self.workers.iter().enumerate() {
for k in w.env.keys() {
if crate::cluster::is_reserved_cluster_env_key(k) {
return Err(format!(
"cluster.workers[{i}] ({:?}): env key {k:?} is reserved \
(launcher-owned) and cannot be set via env — it would \
override the launcher's per-rank value",
w.host,
));
}
}
}
let multi_host = self.spans_multiple_hosts();
for (i, w) in self.workers.iter().enumerate() {
if w.host.trim().is_empty() {
return Err(format!("cluster.workers[{i}].host must be non-empty"));
}
if multi_host && w.nccl_socket_ifname.trim().is_empty() {
return Err(format!(
"cluster.workers[{i}] ({:?}): nccl_socket_ifname must be \
non-empty when the cluster spans multiple workers",
w.host
));
}
if w.path.trim().is_empty() {
return Err(format!(
"cluster.workers[{i}] ({:?}): path (project checkout dir) \
must be non-empty",
w.host
));
}
}
if self.workers.iter().all(|w| !w.ranks.is_empty()) {
for (i, w) in self.workers.iter().enumerate() {
if let Some(devs) = w.local_devices.as_explicit() {
if w.ranks.len() != devs.len() {
return Err(format!(
"cluster.workers[{i}] ({:?}): ranks ({}) and local_devices ({}) length mismatch",
w.host,
w.ranks.len(),
devs.len()
));
}
}
}
let mut all: Vec<usize> = self
.workers
.iter()
.flat_map(|w| w.ranks.iter().copied())
.collect();
let ws = all.len();
all.sort_unstable();
let expected: Vec<usize> = (0..ws).collect();
if all != expected {
return Err(format!(
"cluster: ranks across workers must be exactly 0..{ws} with no \
duplicates or gaps, got sorted-unique sequence {all:?}"
));
}
}
Ok(())
}
pub fn populate_ranks(&mut self, device_counts: &[usize]) -> Result<(), String> {
if device_counts.len() != self.workers.len() {
return Err(format!(
"populate_ranks: device_counts len {} != workers len {}",
device_counts.len(),
self.workers.len(),
));
}
let mut next_rank = 0usize;
for (i, w) in self.workers.iter_mut().enumerate() {
let count = device_counts[i];
if count == 0 {
return Err(format!(
"populate_ranks: worker[{i}] ({:?}) reported 0 devices",
w.host,
));
}
w.ranks = (next_rank..next_rank + count).collect();
next_rank += count;
}
Ok(())
}
pub fn canonical_json(&self) -> Result<String, String> {
serde_json::to_string_pretty(self)
.map_err(|e| format!("cluster: JSON serialization failed: {e}"))
}
pub fn local_envelope_for(&self, worker: &ClusterWorker) -> Value {
let mut worker_obj = serde_json::Map::new();
worker_obj.insert("host".into(), Value::String(worker.host.clone()));
worker_obj.insert(
"ranks".into(),
Value::Array(worker.ranks.iter().map(|r| Value::from(*r)).collect()),
);
worker_obj.insert(
"local_devices".into(),
match &worker.local_devices {
LocalDevices::All => Value::String("all".into()),
LocalDevices::Explicit(v) => {
Value::Array(v.iter().map(|d| Value::from(*d)).collect())
}
},
);
worker_obj.insert(
"nccl_socket_ifname".into(),
Value::String(worker.nccl_socket_ifname.clone()),
);
worker_obj.insert("path".into(), Value::String(worker.path.clone()));
if let Some(a) = &worker.arch {
worker_obj.insert("arch".into(), Value::String(a.clone()));
}
worker_obj.insert(
"data_path".into(),
Value::String(worker.effective_data_path().into()),
);
let mut controller_obj = serde_json::Map::new();
controller_obj.insert("host".into(), Value::String(self.controller.host.clone()));
controller_obj.insert("port".into(), Value::from(self.controller.port));
let mut envelope = serde_json::Map::new();
envelope.insert("controller".into(), Value::Object(controller_obj));
envelope.insert("world_size".into(), Value::from(self.world_size()));
envelope.insert("num_workers".into(), Value::from(self.workers.len()));
envelope.insert("worker".into(), Value::Object(worker_obj));
Value::Object(envelope)
}
pub fn ssh_target<'a>(&'a self, worker: &'a ClusterWorker) -> &'a str {
worker
.ssh
.as_ref()
.and_then(|s| s.target.as_deref())
.unwrap_or(&worker.host)
}
}