use std::net::{SocketAddr, TcpStream};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, mpsc};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use crate::autograd::Variable;
use crate::data::BatchDataSet;
use crate::distributed::wire::{read_handshake_ack, write_handshake_rank};
use crate::distributed::ddp_run::{
ControlMsg, EpochFn, EpochPlan, GpuWorker, RankCallbacks, TimingMsg, WorkerConfig,
};
use crate::distributed::nccl::NcclRankComm;
use crate::distributed::relay::mux::{try_read_len_framed, write_len_framed, LenFramedRead};
use crate::distributed::wire::{
ControlFrame, ControlMsgWire, MsgKind, SessionSalt,
};
use crate::distributed::wire_convert::{
control_wire_to_msg, metrics_msg_to_wire, timing_msg_to_wire,
};
use crate::distributed::nccl_session::PendingNcclSession;
use crate::nn::{Module, Optimizer, Parameter};
use crate::tensor::{Device, Result, Tensor, TensorError};
pub(crate) const RESOURCE_SAMPLE_INTERVAL_MS: u64 = 500;
pub struct ClusterWorker<M: Module> {
inner: Option<GpuWorker<M>>,
bridges: Vec<JoinHandle<()>>,
shutdown_flag: Arc<AtomicBool>,
#[allow(dead_code)]
local_dead_ranks: Arc<crate::distributed::controller::DeadRanks>,
#[allow(dead_code)]
nccl_session_mailbox: Arc<std::sync::Mutex<Option<PendingNcclSession>>>,
epoch_fn: Option<EpochFn<M>>,
last_epoch_fired: usize,
final_param_rx: Option<mpsc::Receiver<crate::distributed::ddp_run::ParamSnapshot>>,
}
impl<M: Module + 'static> ClusterWorker<M> {
#[allow(clippy::too_many_arguments)]
pub(crate) fn connect_and_build<F, G, O>(
coord_addr: SocketAddr,
cpu_client: Option<crate::distributed::cpu_reduce::CpuReduceClient>,
rank_id: u32,
salt: SessionSalt,
config: WorkerConfig,
model_factory: F,
optim_factory: G,
dataset: Arc<dyn BatchDataSet>,
nccl_comm: Option<NcclRankComm>,
rank_callbacks: RankCallbacks<M>,
) -> Result<Self>
where
F: FnOnce(Device) -> Result<M>,
G: FnOnce(&[Parameter]) -> O,
O: Optimizer + 'static,
{
let RankCallbacks {
checkpoint_fn,
epoch_fn,
eval_fn,
eval_dataset,
outer_optimizer_factory,
} = rank_callbacks;
if rank_id as usize >= config.world_size {
return Err(TensorError::new(&format!(
"cluster_worker: rank_id {rank_id} >= world_size {}",
config.world_size,
)));
}
let stream = crate::distributed::wire::connect_with_retry(coord_addr, "cluster_worker coord")?;
stream
.set_read_timeout(Some(Duration::from_secs(10)))
.map_err(|e| {
TensorError::new(&format!("cluster_worker: set_read_timeout: {e}"))
})?;
stream
.set_write_timeout(Some(crate::distributed::wire::write_stall_timeout()))
.map_err(|e| {
TensorError::new(&format!("cluster_worker: set_write_timeout: {e}"))
})?;
let mut handshake_stream = stream;
write_handshake_rank(
&mut handshake_stream,
rank_id,
config.world_size as u32,
&salt,
)?;
read_handshake_ack(&mut handshake_stream, &salt)?;
handshake_stream
.set_read_timeout(Some(Duration::from_millis(250)))
.map_err(|e| {
TensorError::new(&format!("cluster_worker: set_read_timeout: {e}"))
})?;
let read_stream = handshake_stream;
let mut write_stream = read_stream.try_clone().map_err(|e| {
TensorError::new(&format!(
"cluster_worker: stream try_clone for bridge split: {e}"
))
})?;
let (timing_tx, timing_rx) = mpsc::channel::<TimingMsg>();
let timing_tx_for_param_bridge = timing_tx.clone();
let timing_tx_for_heartbeat = timing_tx.clone();
let timing_tx_for_inbound = timing_tx.clone();
let (metrics_tx, metrics_rx) = mpsc::channel::<crate::distributed::ddp_run::MetricsMsg>();
let (param_tx, param_rx) =
mpsc::channel::<crate::distributed::ddp_run::ParamSnapshot>();
let (final_param_tx, final_param_rx) =
mpsc::channel::<crate::distributed::ddp_run::ParamSnapshot>();
let (control_tx, control_rx) = mpsc::channel::<ControlMsg>();
let control_tx_for_param_bridge = control_tx.clone();
let local_dead_ranks =
crate::distributed::controller::DeadRanks::new(config.world_size);
let nccl_session_mailbox: Arc<std::sync::Mutex<Option<PendingNcclSession>>> =
Arc::new(std::sync::Mutex::new(None));
let outer_optimizer = outer_optimizer_factory.as_ref().map(|f| f());
let mut inner = GpuWorker::<M>::new(
&config,
model_factory,
optim_factory,
dataset,
nccl_comm,
checkpoint_fn,
eval_fn,
eval_dataset,
timing_tx,
metrics_tx,
param_tx,
final_param_tx,
control_rx,
outer_optimizer,
)?;
inner.attach_nccl_session_mailbox(Arc::clone(&nccl_session_mailbox));
inner.attach_local_dead_ranks(Arc::clone(&local_dead_ranks));
let nccl_abort_slot: crate::distributed::ddp_run::NcclAbortSlot =
Arc::new(std::sync::Mutex::new(inner.nccl_abort_handle()));
inner.attach_nccl_abort_slot(Arc::clone(&nccl_abort_slot));
let shutdown_flag = Arc::new(AtomicBool::new(false));
let mut bridges: Vec<JoinHandle<()>> = Vec::new();
let salt_in = salt;
let shutdown_in = Arc::clone(&shutdown_flag);
let rank_in = config.rank;
let coord_liveness_timeout_in = config.coord_liveness_timeout_secs;
let mut read_stream_for_bridge = read_stream;
let dead_for_inbound = Arc::clone(&local_dead_ranks);
let mailbox_for_inbound = Arc::clone(&nccl_session_mailbox);
bridges.push(
thread::Builder::new()
.name(format!("flodl-worker-inbound:r{rank_in}"))
.spawn(move || {
inbound_loop(
rank_in,
&mut read_stream_for_bridge,
&salt_in,
&shutdown_in,
&control_tx,
&dead_for_inbound,
&mailbox_for_inbound,
&timing_tx_for_inbound,
coord_liveness_timeout_in,
);
})
.map_err(|e| {
TensorError::new(&format!(
"cluster_worker: spawn inbound bridge for rank {rank_in}: {e}"
))
})?,
);
let salt_out = salt;
let shutdown_out = Arc::clone(&shutdown_flag);
let rank_out = config.rank;
bridges.push(
thread::Builder::new()
.name(format!("flodl-worker-outbound:r{rank_out}"))
.spawn(move || {
outbound_loop(
rank_out,
&mut write_stream,
&salt_out,
&shutdown_out,
timing_rx,
metrics_rx,
);
})
.map_err(|e| {
TensorError::new(&format!(
"cluster_worker: spawn outbound bridge for rank {rank_out}: {e}"
))
})?,
);
let rank_for_bridge = rank_id as u64;
let gamma_for_bridge = config.gamma;
bridges.push(
thread::Builder::new()
.name(format!("flodl-worker-param-bridge:r{rank_out}"))
.spawn(move || {
param_bridge_loop(
rank_for_bridge,
param_rx,
cpu_client,
control_tx_for_param_bridge,
timing_tx_for_param_bridge,
gamma_for_bridge,
);
})
.map_err(|e| {
TensorError::new(&format!(
"cluster_worker: spawn param bridge: {e}"
))
})?,
);
let final_param_rx_for_handle = Some(final_param_rx);
let shutdown_for_hb = Arc::clone(&shutdown_flag);
let rank_for_hb = rank_id as usize;
bridges.push(
thread::Builder::new()
.name(format!("flodl-worker-heartbeat:r{rank_out}"))
.spawn(move || {
heartbeat_loop(
rank_for_hb,
timing_tx_for_heartbeat,
shutdown_for_hb,
);
})
.map_err(|e| {
TensorError::new(&format!(
"cluster_worker: spawn heartbeat thread: {e}"
))
})?,
);
let spawn_watchdog = nccl_abort_slot
.lock()
.expect("nccl abort slot poisoned")
.is_some();
if spawn_watchdog {
let shutdown_for_wd = Arc::clone(&shutdown_flag);
let dead_for_wd = Arc::clone(&local_dead_ranks);
let slot_for_wd = Arc::clone(&nccl_abort_slot);
let rank_for_wd = rank_id as usize;
bridges.push(
thread::Builder::new()
.name(format!("flodl-worker-nccl-watchdog:r{rank_out}"))
.spawn(move || {
nccl_watchdog_loop(
rank_for_wd,
slot_for_wd,
dead_for_wd,
shutdown_for_wd,
);
})
.map_err(|e| {
TensorError::new(&format!(
"cluster_worker: spawn NCCL watchdog: {e}"
))
})?,
);
}
Ok(ClusterWorker {
inner: Some(inner),
bridges,
shutdown_flag,
local_dead_ranks,
nccl_session_mailbox,
epoch_fn,
last_epoch_fired: usize::MAX,
final_param_rx: final_param_rx_for_handle,
})
}
pub fn inner(&self) -> &GpuWorker<M> {
self.inner
.as_ref()
.expect("inner GpuWorker present until run_until_shutdown drops it")
}
pub fn inner_mut(&mut self) -> &mut GpuWorker<M> {
self.inner
.as_mut()
.expect("inner GpuWorker present until run_until_shutdown drops it")
}
pub fn run_until_shutdown<T>(
mut self,
train_fn: T,
) -> Result<Option<crate::distributed::ddp_run::ParamSnapshot>>
where
T: Fn(&M, &[Tensor]) -> Result<Variable>,
{
let prof = self.inner().prof_enabled();
let mut wait_ns: u128 = 0;
let mut run_ns: u128 = 0;
let mut wait_pre_update_ns: u128 = 0;
let mut wait_post_update_ns: u128 = 0;
let exit_clean = (|| -> Result<bool> {
loop {
let prev_update = if prof { self.inner().last_update_at() } else { None };
let w0 = std::time::Instant::now();
let plan = self.inner_mut().wait_for_epoch_plan()?;
if prof {
let w_elapsed = w0.elapsed().as_nanos();
wait_ns += w_elapsed;
match self.inner().last_update_at() {
Some(u) if Some(u) != prev_update && u >= w0 => {
let pre = u.duration_since(w0).as_nanos();
wait_pre_update_ns += pre;
wait_post_update_ns += w_elapsed.saturating_sub(pre);
}
_ => wait_pre_update_ns += w_elapsed,
}
}
match plan {
Some(plan) => {
self.fire_epoch_callback(plan.epoch);
let r0 = std::time::Instant::now();
let shutdown = self.inner_mut().run_epoch_plan(&plan, &train_fn)?;
if prof {
run_ns += r0.elapsed().as_nanos();
}
if shutdown {
return Ok(true);
}
}
None => return Ok(true),
}
}
})();
if prof {
let inner = self.inner();
let run_s = run_ns as f64 / 1e9;
let wait_s = wait_ns as f64 / 1e9;
let pre_s = wait_pre_update_ns as f64 / 1e9;
let post_s = wait_post_update_ns as f64 / 1e9;
let h2d_s = inner.h2d_wait_ms_total() / 1e3;
let snap_s = inner.snapshot_readout_ms_total() / 1e3;
let snap_n = inner.snapshot_readout_count();
let snap_per = if snap_n > 0 { snap_s / snap_n as f64 } else { 0.0 };
let compute_s = inner.compute_ms_run_total() / 1e3;
let data_s = inner.data_ms_run_total() / 1e3;
let other_s = (run_s - compute_s - data_s).max(0.0);
let ctrl_msgs = inner.ctrl_msgs_handled();
let tot = (run_s + wait_s).max(1e-9);
eprintln!(
"[worker-prof] rank={} run_epoch={:.1}s ({:.0}%) wait={:.1}s ({:.0}%) \
| run-split: compute={:.1}s data={:.1}s other(ctrl/sync)={:.1}s \
| wait-split: pre_update(reduce+slowrank)={:.1}s post_update(dispatch)={:.2}s \
| h2d_writeback={:.2}s | snapshot_readout={:.1}s ({} calls, {:.0}ms/call) \
| ctrl_msgs={}",
inner.rank(),
run_s,
100.0 * run_s / tot,
wait_s,
100.0 * wait_s / tot,
compute_s,
data_s,
other_s,
pre_s,
post_s,
h2d_s,
snap_s,
snap_n,
snap_per * 1e3,
ctrl_msgs,
);
if let Some((rss, pss, shared, private)) = smaps_rollup_kb() {
let gb = |kb: u64| kb as f64 / (1024.0 * 1024.0);
eprintln!(
"[worker-mem] rank={} rss={:.2}GB pss={:.2}GB \
private={:.2}GB shared={:.2}GB",
inner.rank(), gb(rss), gb(pss), gb(private), gb(shared),
);
}
}
let final_snapshot = self.teardown(exit_clean.is_ok());
exit_clean.map(|_| final_snapshot)
}
pub(crate) fn fire_epoch_callback(&mut self, epoch: usize) {
if epoch == self.last_epoch_fired {
return;
}
self.last_epoch_fired = epoch;
let inner = self
.inner
.as_mut()
.expect("inner GpuWorker present until teardown");
if inner.epoch_callback_role() != Some(inner.rank()) {
return;
}
if let Some(ref f) = self.epoch_fn {
let start = std::time::Instant::now();
f(epoch, inner);
let elapsed_ms = start.elapsed().as_secs_f64() * 1000.0;
inner.report_epoch_fn_elapsed(epoch, elapsed_ms);
}
}
pub(crate) fn teardown(
&mut self,
clean: bool,
) -> Option<crate::distributed::ddp_run::ParamSnapshot> {
let mut inner = self.inner.take()?;
let _ = inner.drain_pending_shutdown();
inner.abort_nccl();
inner.send_final_snapshot();
if clean {
inner.report_exiting();
}
let final_snapshot = self.final_param_rx.take().and_then(|rx| {
match rx.try_recv() {
Ok(snap) => Some(snap),
Err(mpsc::TryRecvError::Empty) => rx
.recv_timeout(std::time::Duration::from_millis(500))
.ok(),
Err(mpsc::TryRecvError::Disconnected) => None,
}
});
drop(inner);
self.shutdown_flag.store(true, Ordering::SeqCst);
for handle in self.bridges.drain(..) {
let _ = handle.join();
}
final_snapshot
}
}
impl<M: Module> Drop for ClusterWorker<M> {
fn drop(&mut self) {
self.shutdown_flag.store(true, Ordering::SeqCst);
self.inner.take();
for handle in self.bridges.drain(..) {
let _ = handle.join();
}
}
}
fn smaps_rollup_kb() -> Option<(u64, u64, u64, u64)> {
let text = std::fs::read_to_string("/proc/self/smaps_rollup").ok()?;
let (mut rss, mut pss) = (None, None);
let (mut shared, mut private) = (0u64, 0u64);
for line in text.lines() {
let mut it = line.split_whitespace();
let (Some(key), Some(val)) = (it.next(), it.next()) else { continue };
let Ok(kb) = val.parse::<u64>() else { continue };
match key {
"Rss:" => rss = Some(kb),
"Pss:" => pss = Some(kb),
"Shared_Clean:" | "Shared_Dirty:" => shared += kb,
"Private_Clean:" | "Private_Dirty:" => private += kb,
_ => {}
}
}
Some((rss?, pss?, shared, private))
}
mod bridges;
mod param_bridge;
#[cfg(test)]
pub(crate) use param_bridge::sumcount_reduce;
use bridges::*;
use param_bridge::*;
#[cfg(test)]
#[path = "../cluster_worker_tests.rs"]
mod tests;