use crate::autograd::Variable;
use crate::distributed::cluster_worker::ClusterWorker;
use crate::nn::Module;
use crate::tensor::cuda_stream::StreamGuard;
use crate::tensor::{Device, Result, Tensor, TensorError};
use super::worker::{EpochState, GpuWorker};
use super::{EpochMetrics, EpochPlan, TrainedState};
#[derive(Debug, Clone, Copy)]
pub struct StepOutcome {
pub shutdown: bool,
}
pub struct Worker<M: Module> {
inner: WorkerInner<M>,
epoch_state: Option<EpochState>,
pending_plan: Option<EpochPlan>,
compute_start: Option<std::time::Instant>,
active_guard: Option<StreamGuard>,
shutdown: bool,
finished: bool,
forensics: Option<RankForensics>,
metrics_rx: Option<std::sync::mpsc::Receiver<EpochMetrics>>,
eval_rx: Option<std::sync::mpsc::Receiver<(usize, f64)>>,
}
struct RankForensics {
save_path: Option<String>,
global_rank: usize,
world_size: usize,
}
enum WorkerInner<M: Module> {
Single {
worker: GpuWorker<M>,
num_epochs: usize,
total_samples: usize,
next_epoch: usize,
},
Cluster(ClusterWorker<M>),
}
impl<M: Module + 'static> Worker<M> {
pub(crate) fn single(
worker: GpuWorker<M>,
num_epochs: usize,
total_samples: usize,
) -> Self {
Worker {
inner: WorkerInner::Single {
worker,
num_epochs,
total_samples,
next_epoch: 0,
},
epoch_state: None,
pending_plan: None,
compute_start: None,
active_guard: None,
shutdown: false,
finished: false,
forensics: None, metrics_rx: None, eval_rx: None, }
}
pub(crate) fn cluster(
mut cluster: ClusterWorker<M>,
save_path: Option<String>,
global_rank: usize,
world_size: usize,
) -> Self {
let metrics_rx = cluster.inner_mut().enable_metrics_stream();
let eval_rx = cluster.inner_mut().enable_eval_stream();
Worker {
inner: WorkerInner::Cluster(cluster),
epoch_state: None,
pending_plan: None,
compute_start: None,
active_guard: None,
shutdown: false,
finished: false,
forensics: Some(RankForensics {
save_path,
global_rank,
world_size,
}),
metrics_rx: Some(metrics_rx),
eval_rx: Some(eval_rx),
}
}
fn worker_mut(&mut self) -> &mut GpuWorker<M> {
match &mut self.inner {
WorkerInner::Single { worker, .. } => worker,
WorkerInner::Cluster(cw) => cw.inner_mut(),
}
}
fn worker_ref(&self) -> &GpuWorker<M> {
match &self.inner {
WorkerInner::Single { worker, .. } => worker,
WorkerInner::Cluster(cw) => cw.inner(),
}
}
pub fn model(&self) -> &M {
self.worker_ref().model()
}
pub fn epoch_metrics(&self) -> Option<EpochMetrics> {
self.worker_ref()
.aggregated_metrics()
.lock()
.ok()
.and_then(|g| (*g).clone())
}
pub fn poll_metrics(&self) -> Vec<EpochMetrics> {
match &self.metrics_rx {
Some(rx) => {
let mut out = Vec::new();
while let Ok(m) = rx.try_recv() {
out.push(m);
}
out
}
None => Vec::new(),
}
}
pub fn poll_eval(&self) -> Vec<(usize, f64)> {
match &self.eval_rx {
Some(rx) => {
let mut out = Vec::new();
while let Ok(e) = rx.try_recv() {
out.push(e);
}
out
}
None => Vec::new(),
}
}
pub fn request_eval(&self) {
self.worker_ref()
.report_intent(crate::distributed::wire::IntentKind::EvalNow);
}
pub fn request_checkpoint(&self) {
self.worker_ref()
.report_intent(crate::distributed::wire::IntentKind::CheckpointNow);
}
pub fn next_plan(&mut self) -> Result<Option<EpochPlan>> {
if self.shutdown {
return Ok(None);
}
self.active_guard = None;
if let Some(mut st) = self.epoch_state.take() {
if !st.shutdown() {
self.worker_mut().end_epoch(&mut st)?;
}
}
self.pending_plan = None;
self.compute_start = None;
let plan = match &mut self.inner {
WorkerInner::Single {
worker,
num_epochs,
total_samples,
next_epoch,
} => {
if *next_epoch >= *num_epochs {
None
} else {
let epoch = *next_epoch;
*next_epoch += 1;
let _ = worker;
Some(EpochPlan {
epoch,
partition_offset: 0,
partition_size: *total_samples,
})
}
}
WorkerInner::Cluster(cw) => match cw.inner_mut().wait_for_epoch_plan()? {
Some(plan) => {
cw.fire_epoch_callback(plan.epoch);
Some(plan)
}
None => None,
},
};
if plan.is_none() {
self.shutdown = matches!(self.inner, WorkerInner::Cluster(_));
}
self.pending_plan = plan.clone();
Ok(plan)
}
pub fn next_batch(&mut self) -> Result<Option<Vec<Tensor>>> {
self.active_guard = None;
if self.epoch_state.is_none() {
let plan = match self.pending_plan.take() {
Some(plan) => plan,
None => return Ok(None),
};
let st = self.worker_mut().begin_epoch(&plan)?;
self.epoch_state = Some(st);
}
let mut st = self
.epoch_state
.take()
.expect("epoch_state present after lazy setup");
match self.worker_mut().next_batch_inner(&mut st) {
Ok(Some(batch)) => {
self.worker_mut().sync_before_forward()?;
self.active_guard = self.worker_ref().compute_stream_guard();
self.compute_start = Some(std::time::Instant::now());
self.epoch_state = Some(st);
Ok(Some(batch))
}
Ok(None) => {
if st.shutdown() {
self.shutdown = true; } else {
self.worker_mut().end_epoch(&mut st)?;
}
Ok(None)
}
Err(e) => {
self.epoch_state = Some(st);
Err(e)
}
}
}
pub fn step(&mut self, loss: &Variable) -> Result<StepOutcome> {
let mut st = self.epoch_state.take().ok_or_else(|| {
TensorError::new(
"Worker::step called with no batch in flight; call next_batch() first",
)
})?;
let loss_val: f64 = loss.data().item()?;
self.worker_mut().optimizer_step_and_bookkeep()?;
let ms = self
.compute_start
.take()
.map(|t| t.elapsed().as_secs_f64() * 1000.0)
.unwrap_or(0.0);
self.worker_mut().after_step(&mut st, loss_val, ms)?;
let shutdown = st.shutdown();
if shutdown {
self.shutdown = true;
}
self.epoch_state = Some(st);
self.active_guard = None;
Ok(StepOutcome { shutdown })
}
pub fn finish(mut self) -> Result<TrainedState> {
self.active_guard = None;
if let Some(mut st) = self.epoch_state.take() {
if !st.shutdown() {
self.worker_mut().end_epoch(&mut st)?;
}
}
let state = match &mut self.inner {
WorkerInner::Single { worker, .. } => {
let snap = worker.snapshot_params();
TrainedState {
params: snap
.params
.iter()
.map(|t| t.to_device(Device::CPU))
.collect::<Result<Vec<_>>>()?,
buffers: snap
.buffers
.iter()
.map(|t| t.to_device(Device::CPU))
.collect::<Result<Vec<_>>>()?,
}
}
WorkerInner::Cluster(cluster) => {
let final_snapshot = cluster.teardown(true);
final_snapshot
.map(|snap| TrainedState {
params: snap.params,
buffers: snap.buffers,
})
.unwrap_or(TrainedState {
params: Vec::new(),
buffers: Vec::new(),
})
}
};
self.finished = true;
Ok(state)
}
}
impl<M: Module> Drop for Worker<M> {
fn drop(&mut self) {
if self.finished {
return;
}
let Some(f) = self.forensics.take() else {
return; };
let panicking = std::thread::panicking();
let reason = if panicking {
"cooperative Worker dropped during a panic (user loop panicked)".to_string()
} else {
"cooperative Worker dropped without finish() (user loop returned Err)".to_string()
};
if let Some(stem) = &f.save_path {
let record = crate::distributed::RankDeathRecord::new(
f.global_rank,
f.world_size,
reason.clone(),
);
let path =
crate::distributed::CheckpointBundle::rank_death_path(stem, f.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()
),
}
}
if panicking {
eprintln!("flodl cluster rank: {reason}");
} else {
eprintln!("flodl cluster rank: {reason}; exiting to unblock peers");
crate::distributed::ddp_run::clean_process_exit(1);
}
}
}