use std::sync::Arc;
use std::sync::Mutex;
use std::sync::mpsc;
use crate::autograd::Variable;
use crate::data::BatchDataSet;
use crate::nn::buffer::Buffer;
use crate::tensor::cuda_event::CudaEvent;
use crate::tensor::cuda_stream::CudaStream;
use crate::distributed::nccl::{NcclAbortHandle, NcclRankComm};
use crate::nn::{Module, Optimizer};
use crate::tensor::{Device, Tensor};
use super::{
CheckpointFn, EpochMetrics, EvalFn, TimingMsg, MetricsMsg,
ParamSnapshot, ControlMsg, EpochPlan,
};
mod constructor;
mod control;
mod epoch_plan;
mod reporting;
mod stager;
mod sync;
pub(crate) use epoch_plan::EpochState;
#[cfg(test)]
pub(crate) use sync::weighted_allreduce_nccl;
pub(crate) type NcclAbortSlot =
Arc<Mutex<Option<Arc<NcclAbortHandle>>>>;
pub struct GpuWorker<M: Module> {
model: M,
optimizer: Box<dyn Optimizer>,
pub(super) param_vars: Vec<Variable>,
buffer_list: Vec<Buffer>,
rank: usize,
world_size: usize,
device: Device,
compute_stream: Option<CudaStream>,
comm_stream: Option<CudaStream>,
copy_done: Option<CudaEvent>,
pending_param_h2d: bool,
last_h2d_wait_ms: f64,
last_update_at: Option<std::time::Instant>,
h2d_wait_ms_total: f64,
prof_enabled: bool,
snapshot_ns_total: u128,
snapshot_count: u64,
compute_ms_run_total: f64,
data_ms_run_total: f64,
ctrl_msgs_handled: u64,
nccl_comm: Option<NcclRankComm>,
nccl_abort_handle: Option<Arc<NcclAbortHandle>>,
nccl_abort_slot: Option<NcclAbortSlot>,
nccl_session_mailbox: Option<
Arc<Mutex<Option<crate::distributed::nccl_session::PendingNcclSession>>>,
>,
local_dead_ranks: Option<Arc<crate::distributed::controller::DeadRanks>>,
timing_tx: mpsc::Sender<TimingMsg>,
metrics_tx: mpsc::Sender<MetricsMsg>,
param_tx: mpsc::Sender<ParamSnapshot>,
final_param_tx: mpsc::Sender<ParamSnapshot>,
control_rx: mpsc::Receiver<ControlMsg>,
dataset: Arc<dyn BatchDataSet>,
pub(super) partition: Vec<usize>,
batch_size: usize,
base_seed: u64,
augment: usize,
transform: Option<crate::data::TransformFn>,
local_step: usize,
nccl_sync_seq: usize,
steps_since_avg: usize,
steps_at_snapshot: usize,
gamma: f64,
current_version: u64,
pub(super) current_epoch: usize,
pub(super) pending_plan: Option<EpochPlan>,
global_step: usize,
scheduler: Option<Arc<dyn crate::nn::Scheduler>>,
lr_scale: f64,
aggregated_metrics: Arc<Mutex<Option<EpochMetrics>>>,
metrics_stream_tx: Option<mpsc::Sender<EpochMetrics>>,
eval_stream_tx: Option<mpsc::Sender<(usize, f64)>>,
pub(super) checkpoint_fn: Option<CheckpointFn<M>>,
pub(super) epoch_callback_role: Option<usize>,
pub(super) eval_fn: Option<EvalFn<M>>,
pub(super) eval_dataset: Option<Arc<dyn BatchDataSet>>,
save_path: Option<String>,
prefetch: Option<crate::data::prefetch::PrefetchWorker>,
stager: Option<stager::StagerHandle>,
per_sample_bytes: usize,
vram_max_usage: f64,
activation_peak_bytes: usize,
vram_pool_budget_sent: bool,
max_grad_norm: Option<f64>,
easgd_alpha: Option<f64>,
timeline: Option<std::sync::Arc<crate::monitor::Timeline>>,
pre_sync_scratch: Option<Vec<Tensor>>,
pre_sync_buffer_scratch: Option<Vec<Tensor>>,
outer_optimizer: Option<Box<dyn crate::distributed::OuterOptimizer>>,
outer_prev_global: Option<Vec<Tensor>>,
snapshot_pinned_params: Vec<Tensor>,
snapshot_pinned_buffers: Vec<Tensor>,
pinned_fallback_logged: bool,
_grad_accumulators: Vec<crate::tensor::GradAccumulatorHandle>,
}
#[allow(dead_code)]
pub(crate) struct WorkerChannels {
pub timing_rx: mpsc::Receiver<TimingMsg>,
pub metrics_rx: mpsc::Receiver<MetricsMsg>,
pub param_rx: mpsc::Receiver<ParamSnapshot>,
pub final_param_rx: mpsc::Receiver<ParamSnapshot>,
pub control_tx: mpsc::Sender<ControlMsg>,
}
#[allow(clippy::type_complexity)]
pub(crate) type WorkerEndpoints = (
mpsc::Sender<TimingMsg>,
mpsc::Sender<MetricsMsg>,
mpsc::Sender<ParamSnapshot>,
mpsc::Sender<ParamSnapshot>, mpsc::Receiver<ControlMsg>,
);
impl<M: Module> GpuWorker<M> {
pub fn rank(&self) -> usize {
self.rank
}
pub fn last_update_at(&self) -> Option<std::time::Instant> {
self.last_update_at
}
pub fn h2d_wait_ms_total(&self) -> f64 {
self.h2d_wait_ms_total
}
pub fn snapshot_readout_ms_total(&self) -> f64 {
self.snapshot_ns_total as f64 / 1e6
}
pub fn snapshot_readout_count(&self) -> u64 {
self.snapshot_count
}
pub fn prof_enabled(&self) -> bool {
self.prof_enabled
}
pub fn compute_ms_run_total(&self) -> f64 {
self.compute_ms_run_total
}
pub fn data_ms_run_total(&self) -> f64 {
self.data_ms_run_total
}
pub fn ctrl_msgs_handled(&self) -> u64 {
self.ctrl_msgs_handled
}
pub fn epoch_callback_role(&self) -> Option<usize> {
self.epoch_callback_role
}
pub fn device(&self) -> Device {
self.device
}
pub fn local_step(&self) -> usize {
self.local_step
}
#[cfg(test)]
pub(crate) fn steps_since_avg(&self) -> usize {
self.steps_since_avg
}
#[cfg(test)]
pub(crate) fn set_steps_since_avg(&mut self, n: usize) {
self.steps_since_avg = n;
}
pub fn current_version(&self) -> u64 {
self.current_version
}
pub fn current_epoch(&self) -> usize {
self.current_epoch
}
pub fn nccl_abort_handle(&self) -> Option<Arc<NcclAbortHandle>> {
self.nccl_abort_handle.clone()
}
pub(crate) fn attach_nccl_session_mailbox(
&mut self,
mailbox: Arc<Mutex<Option<crate::distributed::nccl_session::PendingNcclSession>>>,
) {
self.nccl_session_mailbox = Some(mailbox);
}
pub(crate) fn attach_local_dead_ranks(
&mut self,
dead_ranks: Arc<crate::distributed::controller::DeadRanks>,
) {
self.local_dead_ranks = Some(dead_ranks);
}
pub(crate) fn attach_nccl_abort_slot(&mut self, slot: NcclAbortSlot) {
self.nccl_abort_slot = Some(slot);
}
pub fn replace_nccl_comm(&mut self, new_comm: NcclRankComm) {
let handle = new_comm.abort_handle();
if let Some(slot) = &self.nccl_abort_slot {
*slot.lock().expect("nccl abort slot poisoned") = Some(Arc::clone(&handle));
}
self.nccl_abort_handle = Some(handle);
self.nccl_comm = Some(new_comm);
}
pub fn set_lr(&mut self, lr: f64) {
self.optimizer.set_lr(lr);
}
pub fn current_lr(&self) -> f64 {
self.optimizer.lr()
}
pub fn scale_lr(&mut self, factor: f64) {
self.optimizer.scale_lr(factor);
}
pub fn set_lr_scale(&mut self, scale: f64) {
self.lr_scale = scale;
}
pub fn set_scheduler(&mut self, sched: Arc<dyn crate::nn::Scheduler>) {
self.scheduler = Some(sched);
}
pub fn model(&self) -> &M {
&self.model
}
}