use std::sync::Arc;
use crate::autograd::Variable;
use crate::data::BatchDataSet;
use crate::distributed::ddp_run::{
ApplyPolicy, AverageBackend, DdpRunConfig, EpochCallbackPolicy,
RankCallbacks, SchedulerFn, TrainedState, Worker, WorkerConfig,
};
use crate::nn::{Module, Optimizer, Parameter};
use crate::tensor::{DType, Device, Result, Tensor, TensorError};
use super::DdpHandle;
fn parse_or_resolve_socket_addr(addr: &str) -> Result<std::net::SocketAddr> {
use std::net::ToSocketAddrs;
if let Ok(s) = addr.parse::<std::net::SocketAddr>() {
return Ok(s);
}
let mut iter = addr.to_socket_addrs().map_err(|e| {
TensorError::new(&format!("ddp: resolve addr '{addr}': {e}"))
})?;
iter.next().ok_or_else(|| {
TensorError::new(&format!("ddp: resolve addr '{addr}': no addresses returned"))
})
}
pub(crate) fn write_rank_death_record(
save_path: Option<&str>,
global_rank: usize,
world_size: usize,
reason: String,
) {
let Some(stem) = save_path else {
return;
};
let record =
crate::distributed::RankDeathRecord::new(global_rank, world_size, reason);
let path = crate::distributed::CheckpointBundle::rank_death_path(stem, global_rank);
match record.write_to_file(&path) {
Ok(()) => eprintln!(
"flodl cluster rank: wrote death record to {}",
path.display()
),
Err(werr) => eprintln!(
"flodl cluster rank: failed to write death record to {}: {werr}",
path.display()
),
}
}
fn rank_fires_callbacks(
policy: EpochCallbackPolicy,
_global_rank: usize,
world_size: usize,
) -> Result<bool> {
match policy {
EpochCallbackPolicy::Rank(n) => {
if n >= world_size {
return Err(crate::tensor::TensorError::new(&format!(
"EpochCallbackPolicy::Rank({n}) out of bounds (world_size={world_size}). \
Pick a rank in 0..{world_size}."
)));
}
Ok(true)
}
EpochCallbackPolicy::Fastest => Ok(true),
}
}
impl DdpHandle {
#[allow(clippy::too_many_arguments)]
pub(super) fn run_cluster_rank_via_coord<F, M, G, O, T>(
cluster: crate::distributed::cluster::LocalCluster,
policy: ApplyPolicy,
backend: AverageBackend,
model_factory: F,
optim_factory: G,
train_fn: T,
dataset: Arc<dyn BatchDataSet>,
batch_size: usize,
num_epochs: usize,
config: DdpRunConfig,
scheduler_fn: Option<SchedulerFn>,
rank_callbacks: RankCallbacks<M>,
) -> Result<Self>
where
F: Fn(Device) -> Result<M> + Send + Sync + 'static,
M: Module + 'static,
G: Fn(&[Parameter]) -> O + Send + Sync + 'static,
O: Optimizer + 'static,
T: Fn(&M, &[Tensor]) -> Result<Variable> + Send + Sync + 'static,
{
let save_path = config.save_path.clone();
let (global_rank, device) = cluster.my_rank()?;
let world_size = cluster.world_size();
let total_samples = dataset.len() * config.augment.max(1);
let policy_label = match policy {
ApplyPolicy::Sync => "Sync",
ApplyPolicy::Cadence => "Cadence",
ApplyPolicy::Async => "Async",
};
let backend_label = match backend {
AverageBackend::Nccl => "Nccl",
AverageBackend::Cpu => "Cpu",
};
crate::verbose!(
" ddp: cluster rank {global_rank}/{world_size} on {device:?} \
({policy_label}+{backend_label} via_coord, save_path={save_path:?})"
);
let coord_addr_str = format!(
"127.0.0.1:{}",
cluster.controller.port.saturating_add(
crate::distributed::relay::RELAY_CONTROL_LOOPBACK_OFFSET,
)
);
let controller_addr_str = match backend {
AverageBackend::Cpu => Some(format!(
"127.0.0.1:{}",
cluster.controller.port.saturating_add(
crate::distributed::relay::RELAY_DATA_LOOPBACK_OFFSET,
)
)),
AverageBackend::Nccl => None,
};
let mut meta = serde_json::json!({
"mode": format!("cluster-rank {policy_label}+{backend_label} via_coord"),
"global_rank": global_rank,
"world_size": world_size,
"device": format!("{device:?}"),
"batch_size": batch_size,
"num_epochs": num_epochs,
"total_samples": total_samples,
"coord_addr": coord_addr_str,
"save_path": save_path,
});
if let Some(ref addr) = controller_addr_str {
meta["controller_addr"] = serde_json::json!(addr);
}
let training_meta = Some(meta);
let worker_outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(
move || -> Result<TrainedState> {
let cluster_worker = Self::build_cluster_worker(
cluster,
global_rank,
device,
world_size,
policy,
backend,
model_factory,
optim_factory,
dataset,
batch_size,
config,
scheduler_fn,
rank_callbacks,
)?;
let final_snapshot = cluster_worker.run_until_shutdown(train_fn)?;
Ok(final_snapshot
.map(|snap| TrainedState {
params: snap.params,
buffers: snap.buffers,
})
.unwrap_or(TrainedState {
params: Vec::new(),
buffers: Vec::new(),
}))
},
));
let write_death_record =
|reason: String| write_rank_death_record(save_path.as_deref(), global_rank, world_size, reason);
let final_state = match worker_outcome {
Ok(Ok(state)) => state,
Ok(Err(e)) => {
eprintln!("flodl cluster rank: rank failed: {e}");
write_death_record(e.to_string());
super::clean_process_exit(1);
}
Err(panic) => {
let msg = panic
.downcast_ref::<&str>()
.map(|s| s.to_string())
.or_else(|| panic.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "unknown panic payload".to_string());
eprintln!("flodl cluster rank: rank panicked: {msg}");
write_death_record(format!("panic: {msg}"));
std::panic::resume_unwind(panic);
}
};
Ok(DdpHandle {
devices: vec![device],
final_state: Some(final_state),
metrics_rx: None,
launcher_driver: None,
launcher_abort: None,
architecture_svg: None,
graph_label: None,
graph_hash: None,
training_meta,
})
}
#[allow(clippy::too_many_arguments)]
pub(super) fn run_cluster_rank_worker<F, M, G, O>(
cluster: crate::distributed::cluster::LocalCluster,
policy: ApplyPolicy,
backend: AverageBackend,
model_factory: F,
optim_factory: G,
dataset: Arc<dyn BatchDataSet>,
batch_size: usize,
config: DdpRunConfig,
scheduler_fn: Option<SchedulerFn>,
rank_callbacks: RankCallbacks<M>,
) -> Result<Worker<M>>
where
F: Fn(Device) -> Result<M> + Send + Sync + 'static,
M: Module + 'static,
G: Fn(&[Parameter]) -> O + Send + Sync + 'static,
O: Optimizer + 'static,
{
let save_path = config.save_path.clone();
let (global_rank, device) = cluster.my_rank()?;
let world_size = cluster.world_size();
match Self::build_cluster_worker(
cluster,
global_rank,
device,
world_size,
policy,
backend,
model_factory,
optim_factory,
dataset,
batch_size,
config,
scheduler_fn,
rank_callbacks,
) {
Ok(cluster_worker) => Ok(Worker::cluster(
cluster_worker,
save_path,
global_rank,
world_size,
)),
Err(e) => {
write_rank_death_record(
save_path.as_deref(), global_rank, world_size, e.to_string(),
);
Err(e)
}
}
}
#[allow(clippy::too_many_arguments)]
fn build_cluster_worker<F, M, G, O>(
cluster: crate::distributed::cluster::LocalCluster,
global_rank: usize,
device: Device,
world_size: usize,
policy: ApplyPolicy,
backend: AverageBackend,
model_factory: F,
optim_factory: G,
dataset: Arc<dyn BatchDataSet>,
batch_size: usize,
config: DdpRunConfig,
scheduler_fn: Option<SchedulerFn>,
rank_callbacks: RankCallbacks<M>,
) -> Result<crate::distributed::cluster_worker::ClusterWorker<M>>
where
F: Fn(Device) -> Result<M> + Send + Sync + 'static,
M: Module + 'static,
G: Fn(&[Parameter]) -> O + Send + Sync + 'static,
O: Optimizer + 'static,
{
let RankCallbacks {
checkpoint_fn,
epoch_fn,
eval_fn,
eval_dataset,
outer_optimizer_factory,
} = rank_callbacks;
let save_path = config.save_path.clone();
let total_samples = dataset.len() * config.augment.max(1);
let fires_callbacks = rank_fires_callbacks(
config.epoch_callback_policy, global_rank, world_size,
)?;
let epoch_fn = if fires_callbacks { epoch_fn } else { None };
let checkpoint_fn = if fires_callbacks { checkpoint_fn } else { None };
let eval_fn = if fires_callbacks { eval_fn } else { None };
let eval_dataset = if fires_callbacks { eval_dataset } else { None };
let coord_port = cluster.controller.port.saturating_add(
crate::distributed::relay::RELAY_CONTROL_LOOPBACK_OFFSET,
);
let coord_addr =
parse_or_resolve_socket_addr(&format!("127.0.0.1:{coord_port}"))?;
let controller_addr_str = match backend {
AverageBackend::Cpu => {
let controller_port = cluster.controller.port.saturating_add(
crate::distributed::relay::RELAY_DATA_LOOPBACK_OFFSET,
);
Some(format!("127.0.0.1:{controller_port}"))
}
AverageBackend::Nccl => None,
};
let session_salt = cluster.salt;
let dataset_sig = [0u8; 32];
#[cfg(feature = "cuda")]
if matches!(backend, AverageBackend::Nccl) {
if let crate::tensor::Device::CUDA(idx) = device {
crate::tensor::set_current_cuda_device(idx);
}
}
let nccl_comm = match backend {
AverageBackend::Nccl => {
let rdv = cluster.rendezvous(dataset_sig)?;
let comm = crate::distributed::nccl::NcclRankComm::init_rank(
global_rank, world_size, rdv.unique_id(),
)?;
drop(rdv);
Some(comm)
}
AverageBackend::Cpu => None,
};
let mut cpu_client = match &controller_addr_str {
Some(addr) => {
let mut client = crate::distributed::cpu_reduce::CpuReduceClient::connect(
parse_or_resolve_socket_addr(addr)?,
global_rank as u32,
world_size as u32,
session_salt,
)?;
client.set_bf16_wire(config.elche.bf16_wire);
Some(client)
}
None => None,
};
let tmp_model = model_factory(device)?;
let initial_params_local: Vec<Tensor> = tmp_model
.parameters().iter().map(|p| p.variable.data()).collect();
crate::distributed::ddp_run::ensure_trainable_params(
initial_params_local.len(), "ddp: cluster rank",
)?;
let initial_buffers_local: Vec<Tensor> = tmp_model
.buffers().iter().map(|b| b.get()).collect();
let staging_bytes = {
use crate::distributed::wire;
let wire_bytes = wire::tensors_wire_bytes(&initial_params_local)
+ wire::tensors_wire_bytes(&initial_buffers_local);
wire::set_frame_ceiling(wire::derive_frame_ceiling(wire_bytes));
wire_bytes
};
if let Some(client) = &mut cpu_client
&& device.is_cuda()
&& policy.is_barrier_paced()
{
let local_ranks = cluster
.this_worker()
.map(|w| w.ranks.len())
.unwrap_or(1);
let affordable = crate::sys::mem_info().is_some_and(|m| {
crate::distributed::cpu_reduce::pinned_decode_affordable(
staging_bytes as u64,
local_ranks,
m.available_bytes,
)
});
client.set_pinned_decode(affordable);
if !affordable {
crate::verbose!(
"ddp: rank {global_rank} pinned consensus decode disabled \
(RAM headroom: staging {}MB x {local_ranks} local rank(s) \
needs more MemAvailable); using fresh-alloc decode",
staging_bytes / (1 << 20),
);
}
}
match (&nccl_comm, &mut cpu_client) {
(Some(comm), _) => {
if !initial_params_local.is_empty() {
let refs: Vec<&Tensor> = initial_params_local.iter().collect();
comm.broadcast(&refs, 0)?;
}
if !initial_buffers_local.is_empty() {
let refs: Vec<&Tensor> = initial_buffers_local.iter().collect();
comm.broadcast(&refs, 0)?;
}
}
(None, Some(client)) => {
if !initial_params_local.is_empty() {
let refs: Vec<&Tensor> = initial_params_local.iter().collect();
let broadcast = client.broadcast_from_root(&refs, 0)?;
crate::autograd::no_grad(|| -> crate::tensor::Result<()> {
for (dst, src) in initial_params_local.iter().zip(&broadcast) {
dst.copy_(src, false)?;
}
Ok(())
})?;
}
let f32_buffers: Vec<&Tensor> = initial_buffers_local
.iter().filter(|b| b.dtype() == DType::Float32).collect();
if !f32_buffers.is_empty() {
let broadcast = client.broadcast_from_root(&f32_buffers, 0)?;
crate::autograd::no_grad(|| -> crate::tensor::Result<()> {
for (dst, src) in f32_buffers.iter().zip(&broadcast) {
dst.copy_(src, false)?;
}
Ok(())
})?;
}
}
(None, None) => {
return Err(TensorError::new(
"build_cluster_worker: neither NCCL comm nor CPU reduce \
client was constructed (backend bootstrap bug)",
));
}
}
let initial_params: Vec<Tensor> = initial_params_local.iter()
.map(|t| t.to_device(Device::CPU).and_then(|t| t.pin_memory()))
.collect::<Result<Vec<_>>>()?;
let initial_buffers: Vec<Tensor> = initial_buffers_local.iter()
.map(|t| t.to_device(Device::CPU).and_then(|t| t.pin_memory()))
.collect::<Result<Vec<_>>>()?;
drop(tmp_model);
let worker_config = WorkerConfig {
rank: global_rank,
world_size,
device,
initial_params,
initial_buffers,
total_samples,
batch_size,
augment: config.augment.max(1),
transform: config.transform.clone(),
seed: crate::distributed::ddp_run::resolve_shuffle_seed(
config.resume_from.as_deref(),
)?,
max_grad_norm: config.max_grad_norm,
vram_pool: config.vram_pool,
vram_max_usage: config.vram_max_usage,
ram_max_usage: config.ram_max_usage,
sample_cache: config.sample_cache,
disk_stage_gb: config.disk_stage_gb,
disk_stage_dir: config.disk_stage_dir.clone(),
easgd_alpha: config.elche.easgd_alpha,
gamma: config.elche.gamma,
bf16_wire: config.elche.bf16_wire,
timeline: config.timeline.clone(),
policy,
save_path,
coord_liveness_timeout_secs: config.heartbeat_timeout_secs.unwrap_or_else(
|| crate::distributed::wire::scaled_deadline_secs(
crate::distributed::ddp_run::DEFAULT_COORD_LIVENESS_TIMEOUT_SECS,
),
),
};
let mut cluster_worker =
crate::distributed::cluster_worker::ClusterWorker::connect_and_build(
coord_addr,
cpu_client,
global_rank as u32,
session_salt,
worker_config,
model_factory,
optim_factory,
dataset,
nccl_comm,
RankCallbacks {
checkpoint_fn,
epoch_fn,
eval_fn,
eval_dataset,
outer_optimizer_factory,
},
)?;
if let Some(f) = scheduler_fn {
cluster_worker.inner_mut().set_scheduler(f(world_size));
}
if let Some(stem) = config.resume_from.as_ref() {
if let Err(e) = cluster_worker.inner_mut().resume_outer_momentum(stem) {
eprintln!(
"cluster_worker: rank {global_rank} outer-momentum resume \
failed ({e}); starting from zero momentum"
);
}
}
Ok(cluster_worker)
}
}