use crate::autograd::Variable;
use crate::graph::Graph;
use crate::nn::{Buffer, Module, Optimizer, Parameter};
use super::nccl::{NcclRankComm, ReduceOp};
use super::rendezvous::TcpRendezvous;
use super::config::TrainerConfig;
use super::ddp_run::{DdpBuilder, DdpHandle};
pub use super::el_che::ElChe;
use crate::tensor::{Device, Result, Tensor, TensorError};
#[cfg(test)]
pub(crate) static NCCL_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
pub struct Ddp {
comms: NcclRankComm,
device: Device,
params: Vec<Variable>,
buffers: Vec<Buffer>,
}
impl Ddp {
pub fn wrap(
model: &dyn Module,
device: Device,
global_rank: usize,
rdv: &TcpRendezvous,
) -> Result<Self> {
let world_size = rdv.world_size();
if global_rank >= world_size {
return Err(TensorError::new(&format!(
"Ddp::wrap: global_rank {global_rank} >= world_size {world_size}"
)));
}
if let Device::CUDA(idx) = device {
crate::tensor::set_current_cuda_device(idx);
}
let comms = NcclRankComm::init_rank(global_rank, world_size, rdv.unique_id())?;
let params: Vec<Variable> = model
.parameters()
.into_iter()
.map(|p| p.variable)
.collect();
crate::distributed::ddp_run::ensure_trainable_params(params.len(), "Ddp::wrap")?;
let buffers: Vec<Buffer> = model.buffers();
Ok(Ddp { comms, device, params, buffers })
}
pub fn from_comm(
comms: NcclRankComm,
model: &dyn Module,
device: Device,
) -> Result<Self> {
let params: Vec<Variable> = model
.parameters()
.into_iter()
.map(|p| p.variable)
.collect();
crate::distributed::ddp_run::ensure_trainable_params(params.len(), "Ddp::from_comm")?;
let buffers: Vec<Buffer> = model.buffers();
Ok(Ddp { comms, device, params, buffers })
}
pub fn average_params(&self) -> Result<()> {
let tensors: Vec<Tensor> = self.params.iter().map(|v| v.data()).collect();
if tensors.is_empty() {
return Ok(());
}
let refs: Vec<&Tensor> = tensors.iter().collect();
self.comms.all_reduce(&refs, ReduceOp::Avg)?;
Ok(())
}
pub fn make_divergence_scratch(&self) -> Result<Vec<Tensor>> {
self.params
.iter()
.map(|v| Tensor::zeros_like(&v.data()))
.collect()
}
pub fn average_params_with_divergence(
&self,
scratch: &[Tensor],
) -> Result<(f64, Option<f64>, Option<f64>)> {
let param_tensors: Vec<Tensor> = self.params.iter().map(|v| v.data()).collect();
if param_tensors.is_empty() {
return Ok((0.0, None, None));
}
if scratch.len() != param_tensors.len() {
return Err(TensorError::new(&format!(
"average_params_with_divergence: scratch.len() ({}) must equal \
number of parameters ({})",
scratch.len(),
param_tensors.len(),
)));
}
for (dst, src) in scratch.iter().zip(¶m_tensors) {
dst.copy_(src, false)?;
}
let refs: Vec<&Tensor> = param_tensors.iter().collect();
self.comms.all_reduce(&refs, ReduceOp::Avg)?;
crate::distributed::divergence::divergence_triple(scratch, ¶m_tensors)
}
pub fn sync_params(&self) -> Result<()> {
let p_tensors: Vec<Tensor> = self.params.iter().map(|v| v.data()).collect();
if !p_tensors.is_empty() {
let refs: Vec<&Tensor> = p_tensors.iter().collect();
self.comms.broadcast(&refs, 0)?;
}
let b_tensors: Vec<Tensor> = self.buffers.iter().map(|b| b.get()).collect();
if !b_tensors.is_empty() {
let refs: Vec<&Tensor> = b_tensors.iter().collect();
self.comms.broadcast(&refs, 0)?;
}
Ok(())
}
pub fn all_reduce_gradients(&self) -> Result<()> {
let grads: Vec<Tensor> = self.params.iter().filter_map(|v| v.grad()).collect();
if grads.is_empty() {
return Ok(());
}
let refs: Vec<&Tensor> = grads.iter().collect();
self.comms.all_reduce(&refs, ReduceOp::Avg)?;
Ok(())
}
pub fn sync_buffers(&self) -> Result<()> {
let tensors: Vec<Tensor> = self.buffers.iter().map(|b| b.get()).collect();
if tensors.is_empty() {
return Ok(());
}
let refs: Vec<&Tensor> = tensors.iter().collect();
self.comms.broadcast(&refs, 0)?;
Ok(())
}
pub fn weighted_all_reduce_gradients(&self, batch_counts: &[usize]) -> Result<()> {
if batch_counts.len() != self.comms.world_size() {
return Err(TensorError::new(&format!(
"weighted_all_reduce: batch_counts len ({}) != world_size ({})",
batch_counts.len(),
self.comms.world_size(),
)));
}
let total: usize = batch_counts.iter().sum();
if total == 0 {
return Err(TensorError::new(
"weighted_all_reduce: total batch count is 0",
));
}
let my_rank = self.comms.rank();
let weight = batch_counts[my_rank] as f64 / total as f64;
let grads: Vec<Tensor> = self.params
.iter()
.filter_map(|v| {
v.grad().inspect(|g| {
g.mul_scalar_(weight).ok();
})
})
.collect();
if grads.is_empty() {
return Ok(());
}
let refs: Vec<&Tensor> = grads.iter().collect();
self.comms.all_reduce(&refs, ReduceOp::Sum)?;
Ok(())
}
pub fn world_size(&self) -> usize {
self.comms.world_size()
}
pub fn rank(&self) -> usize {
self.comms.rank()
}
pub fn all_reduce_per_rank_f64(&self, local: &mut [f64]) -> Result<()> {
let world_size = self.comms.world_size();
if local.len() != world_size {
return Err(TensorError::new(&format!(
"all_reduce_per_rank_f64: vector len ({}) must equal world_size ({})",
local.len(),
world_size,
)));
}
let t = Tensor::from_f64(local, &[world_size as i64], self.device)?;
self.comms.all_reduce(&[&t], ReduceOp::Sum)?;
let out = t.to_f64_vec()?;
local.copy_from_slice(&out);
Ok(())
}
pub fn device(&self) -> Device {
self.device
}
}
pub struct Trainer;
impl Trainer {
pub fn builder<F, M, G, O, T>(
model_factory: F,
optim_factory: G,
train_fn: T,
) -> DdpBuilder<F, M, G, O, T>
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,
{
DdpHandle::new_builder(model_factory, optim_factory, train_fn)
}
pub fn run<F, M, G, O, T>(
model_factory: F,
optim_factory: G,
train_fn: T,
cfg: TrainerConfig<M>,
) -> Result<DdpHandle>
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 mut b = DdpHandle::new_builder(model_factory, optim_factory, train_fn)
.dataset(cfg.dataset)
.batch_size(cfg.batch_size)
.num_epochs(cfg.num_epochs)
.elche(cfg.elche)
.vram_pool(cfg.vram_pool)
.vram_max_usage(cfg.vram_max_usage)
.ram_max_usage(cfg.ram_max_usage)
.sample_cache(cfg.sample_cache)
.disk_stage(cfg.disk_stage_gb)
.augment(cfg.augment);
if let Some(dir) = cfg.disk_stage_dir {
b = b.disk_stage_dir(dir);
}
if let Some(f) = cfg.transform {
b = b.transform_fn(f);
}
if let Some(n) = cfg.max_grad_norm {
b = b.max_grad_norm(n);
}
if let Some(t) = cfg.max_failure {
b = b.max_failure(t);
}
if let Some(n) = cfg.checkpoint_every {
b = b.checkpoint_every(n);
}
if let Some(p) = cfg.save_path {
b = b.save_path(p);
}
if let Some(p) = cfg.resume_from {
b = b.resume_from(p);
}
if let Some(e) = cfg.checkpoint_at_epoch {
b = b.checkpoint_at_epoch(e);
}
if let Some(f) = cfg.outer_optimizer {
b = b.outer_optimizer_arc(f);
}
if let Some(f) = cfg.checkpoint_fn {
b = b.checkpoint_fn_arc(f);
}
if let Some(f) = cfg.epoch_fn {
b = b.epoch_fn_arc(f);
}
if let Some(f) = cfg.metrics_fn {
b = b.metrics_fn_arc(f);
}
if let Some(f) = cfg.scheduler_fn {
b = b.scheduler_fn_boxed(f);
}
let has_eval_fn = cfg.eval_fn.is_some();
if let Some(f) = cfg.eval_fn {
b = b.eval_fn_arc(f);
}
match (cfg.eval_every, has_eval_fn) {
(Some(n), _) => b = b.eval_every(crate::distributed::ddp_run::EvalCadence::Epochs(n)),
(None, true) => b = b.eval_every(crate::distributed::ddp_run::EvalCadence::Epochs(1)),
(None, false) => {}
}
if let Some(ds) = cfg.eval_dataset {
b = b.eval_dataset(ds);
}
if let Some(f) = cfg.eval_result_fn {
b = b.eval_result_fn_arc(f);
}
if let Some(t) = cfg.timeline {
b = b.timeline(t);
}
b = b.epoch_callback_policy(cfg.epoch_callback_policy);
if let Some(c) = cfg.cluster {
b = b.cluster(c);
}
b.run()
}
}
pub trait HasGraph {
fn graph(&self) -> &Graph;
}
impl HasGraph for Graph {
fn graph(&self) -> &Graph { self }
}
#[cfg(test)]
#[path = "ddp_tests.rs"]
mod tests;