use std::time::{Duration, Instant};
use crate::autograd::{NoGradGuard, Variable};
use crate::tensor::cuda_stream::StreamGuard;
use crate::distributed::nccl::ReduceOp;
use crate::nn::Module;
use crate::tensor::{Device, Result, Tensor, TensorError, TensorOptions};
use super::super::{
AveragedParams, ParamSnapshot,
};
use super::GpuWorker;
fn pinned_like(t: &Tensor, as_bf16: bool) -> Result<Tensor> {
let dtype = if as_bf16 && t.dtype() == crate::tensor::DType::Float32 {
crate::tensor::DType::BFloat16
} else {
t.dtype()
};
let opts = TensorOptions { dtype, device: Device::CPU };
Tensor::empty(&t.shape(), opts)?.pin_memory()
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn weighted_allreduce_nccl(
comm: &crate::distributed::nccl::NcclRankComm,
stream: Option<&crate::tensor::cuda_stream::CudaStream>,
param_refs: &[&Tensor],
buffer_refs: &[&Tensor],
n_i: f64,
gamma: f64,
device: Device,
rank: usize,
seq: usize,
) -> Result<()> {
crate::debug!(
" ddp-areduce: rank {rank} seq={seq} COUNT enter (n_i={n_i}, gamma={gamma})"
);
let n_eff = crate::distributed::realized_work::gamma_mass(n_i, gamma);
let _stream_guard = stream.map(StreamGuard::new);
let mover = crate::distributed::realized_work::mover_mass(n_i);
let count = Tensor::from_f32(&[n_eff as f32, mover as f32], &[2], device)?;
let totals = match stream {
Some(s) => {
comm.all_reduce_on_stream(&[&count], ReduceOp::Sum, s)?;
s.synchronize()?;
count.to_f64_vec()?
}
None => {
comm.all_reduce(&[&count], ReduceOp::Sum)?;
count.to_f64_vec()?
}
};
let (total_n, total_movers) = (totals[0], totals[1]);
crate::debug!(
" ddp-areduce: rank {rank} seq={seq} COUNT exit (total_n={total_n}, total_movers={total_movers})"
);
if !crate::distributed::realized_work::is_realized(total_n) {
crate::debug!(
" ddp-areduce: rank {rank} seq={seq} SKIP param-reduce (total_n={total_n}) \
-- 1 collective this seq (peers doing more would deadlock)"
);
return Ok(());
}
let factor = (n_eff / total_n) as f32;
crate::debug!(
" ddp-areduce: rank {rank} seq={seq} PARAM enter (nparams={}, factor={factor})",
param_refs.len()
);
match stream {
Some(s) => {
comm.all_reduce_premul_sum(param_refs, factor, Some(s))?;
}
None => {
comm.all_reduce_premul_sum(param_refs, factor, None)?;
}
}
crate::debug!(" ddp-areduce: rank {rank} seq={seq} PARAM exit");
if !buffer_refs.is_empty() {
let buf_factor = (mover / total_movers) as f32;
crate::debug!(
" ddp-areduce: rank {rank} seq={seq} BUFFER enter (nbuffers={}, factor={buf_factor})",
buffer_refs.len()
);
comm.all_reduce_premul_sum(buffer_refs, buf_factor, stream)?;
crate::debug!(" ddp-areduce: rank {rank} seq={seq} BUFFER exit");
}
if let Some(s) = stream {
s.synchronize()?;
}
Ok(())
}
impl<M: Module> GpuWorker<M> {
pub fn snapshot_params(&mut self) -> ParamSnapshot {
if let Some(stream) = &self.compute_stream {
let _ = stream.synchronize();
}
if let Some(stream) = &self.comm_stream {
let _ = stream.synchronize();
}
let (params, buffers) = if self.comm_stream.is_some() {
match self.read_params_pinned() {
Ok(out) => out,
Err(e) => {
if !self.pinned_fallback_logged {
self.pinned_fallback_logged = true;
eprintln!(
"flodl ddp: rank {} pinned snapshot readout failed ({e}); \
falling back to per-param synchronous D2H (slower)",
self.rank
);
}
self.read_params_passthrough()
}
}
} else {
self.read_params_passthrough()
};
ParamSnapshot {
rank: self.rank,
params,
buffers,
batch_count: self.steps_since_avg,
}
}
fn read_params_pinned(&mut self) -> Result<(Vec<Tensor>, Vec<Tensor>)> {
if self.snapshot_pinned_params.is_empty() && !self.param_vars.is_empty() {
let mut bufs = Vec::with_capacity(self.param_vars.len());
for v in &self.param_vars {
bufs.push(pinned_like(&v.data(), self.bf16_wire)?);
}
self.snapshot_pinned_params = bufs;
}
if self.snapshot_pinned_buffers.is_empty() && !self.buffer_list.is_empty() {
let mut bufs = Vec::with_capacity(self.buffer_list.len());
for b in &self.buffer_list {
bufs.push(pinned_like(&b.get(), false)?);
}
self.snapshot_pinned_buffers = bufs;
}
let stream = self.comm_stream.as_ref().ok_or_else(|| {
TensorError::new("read_params_pinned: comm_stream absent")
})?;
{
let _guard = StreamGuard::new(stream);
for (dst, v) in self.snapshot_pinned_params.iter().zip(&self.param_vars) {
dst.copy_(&v.data(), true)?;
}
for (dst, b) in self.snapshot_pinned_buffers.iter().zip(&self.buffer_list) {
dst.copy_(&b.get(), true)?;
}
}
stream.synchronize()?;
Ok((
self.snapshot_pinned_params.clone(),
self.snapshot_pinned_buffers.clone(),
))
}
pub fn snapshot_params_exact(&mut self) -> ParamSnapshot {
if let Some(stream) = &self.compute_stream {
let _ = stream.synchronize();
}
if let Some(stream) = &self.comm_stream {
let _ = stream.synchronize();
}
let (params, buffers) = self.read_params_passthrough();
ParamSnapshot {
rank: self.rank,
params,
buffers,
batch_count: self.steps_since_avg,
}
}
fn read_params_passthrough(&self) -> (Vec<Tensor>, Vec<Tensor>) {
fn cpu_detached(t: Tensor) -> Tensor {
if t.device() != Device::CPU {
return t.to_device(Device::CPU).unwrap_or(t);
}
Tensor::zeros_like(&t)
.and_then(|c| c.copy_(&t, false).map(|()| c))
.unwrap_or(t)
}
let params = self.param_vars.iter()
.map(|v| cpu_detached(v.data()))
.collect();
let buffers = self.buffer_list.iter()
.map(|b| cpu_detached(b.get()))
.collect();
(params, buffers)
}
pub fn load_averaged(&mut self, update: &AveragedParams) -> Result<()> {
if update.params.len() != self.param_vars.len() {
return Err(TensorError::new(&format!(
"load_averaged: expected {} params, got {}",
self.param_vars.len(), update.params.len()
)));
}
if let Some(stream) = &self.compute_stream {
stream.synchronize()?;
}
let non_blocking = self.comm_stream.is_some();
let _guard = self.comm_stream.as_ref().map(StreamGuard::new);
let resets_inner = self
.outer_optimizer
.as_ref()
.is_some_and(|o| o.resets_inner());
{
let _no_grad = NoGradGuard::new();
match self.easgd_alpha {
None => {
for (var, src) in self.param_vars.iter().zip(&update.params) {
var.data().copy_(src, non_blocking)?;
}
}
Some(alpha) => {
let max_numel = update
.params
.iter()
.map(|t| t.numel())
.max()
.unwrap_or(0);
let staging = Tensor::empty(
&[max_numel],
TensorOptions {
dtype: crate::tensor::DType::Float32,
device: self.device,
},
)?;
for (var, src) in self.param_vars.iter().zip(&update.params) {
let dst = var.data();
let staged = staging
.narrow(0, 0, src.numel())?
.reshape(&dst.shape())?;
staged.copy_(src, non_blocking)?;
crate::tensor::Tensor::foreach_lerp_scalar_(
std::slice::from_ref(&dst),
std::slice::from_ref(&staged),
alpha,
)?;
}
}
}
}
for (buf, src) in self.buffer_list.iter().zip(&update.buffers) {
buf.get().copy_(src, non_blocking)?;
}
if resets_inner {
self.optimizer.reset_state();
}
if let (Some(ev), Some(stream)) = (&self.copy_done, &self.comm_stream) {
ev.record_on(stream)?;
}
self.pending_param_h2d = true;
self.current_version = update.version;
if self.prof_enabled {
self.last_update_at = Some(std::time::Instant::now());
}
Ok(())
}
pub(super) fn sync_now_nccl(&mut self) -> Result<(Option<f64>, Option<f64>, Option<f64>)> {
const MAX_REBUILD_ATTEMPTS: usize = 32;
let _diag_start = Instant::now();
if self.nccl_comm.is_none() {
return Ok((None, None, None));
}
if let Some(ref cs) = self.compute_stream {
cs.synchronize()?;
}
let param_tensors: Vec<_> = self.param_vars.iter().map(|v| v.data()).collect();
let buffer_tensors: Vec<Tensor> = self
.buffer_list
.iter()
.map(|b| b.get())
.filter(|t| t.dtype() == crate::tensor::DType::Float32)
.collect();
let mut nccl_ms_total = 0.0_f64;
let seq = self.nccl_sync_seq;
self.nccl_sync_seq += 1;
let rank = self.rank;
for attempt in 0..MAX_REBUILD_ATTEMPTS {
if let Some(ref scratch) = self.pre_sync_scratch {
let _guard = self.comm_stream.as_ref().map(StreamGuard::new);
let _no_grad = NoGradGuard::new();
if attempt == 0 {
for (dst, src) in scratch.iter().zip(¶m_tensors) {
dst.copy_(src, true)?; }
} else {
for (dst, src) in param_tensors.iter().zip(scratch.iter()) {
dst.copy_(src, true)?; }
}
if let Some(ref buf_scratch) = self.pre_sync_buffer_scratch {
if attempt == 0 {
for (dst, src) in buf_scratch.iter().zip(&buffer_tensors) {
dst.copy_(src, true)?; }
} else {
for (dst, src) in buffer_tensors.iter().zip(buf_scratch.iter()) {
dst.copy_(src, true)?; }
}
}
} else if attempt > 0 {
return Err(TensorError::new(
"sync_now_nccl: NCCL aborted but pre_sync_scratch is None; \
cannot restore params for retry. Cluster NCCL mode must \
allocate scratch unconditionally.",
));
}
let param_refs: Vec<&Tensor> = param_tensors.iter().collect();
let buffer_refs: Vec<&Tensor> = buffer_tensors.iter().collect();
let n_i = self.steps_since_avg as f64;
let device = self.device;
let comm = self.nccl_comm.as_ref().expect("nccl_comm present");
let nccl_start = Instant::now();
let attempt_result: Result<()> = weighted_allreduce_nccl(
comm,
self.comm_stream.as_ref(),
¶m_refs,
&buffer_refs,
n_i,
self.gamma,
device,
rank,
seq,
);
nccl_ms_total += nccl_start.elapsed().as_secs_f64() * 1000.0;
match attempt_result {
Ok(()) => {
let divg_start = Instant::now();
let divergence = if let Some(ref scratch) = self.pre_sync_scratch {
let pre_norm_tensors = Tensor::foreach_norm(scratch, 2.0)?;
let mut pre_sq = 0.0f64;
for n in &pre_norm_tensors {
let v: f64 = n.item()?;
pre_sq += v * v;
}
let pre_norm = pre_sq.sqrt();
Tensor::foreach_add_list_(scratch, ¶m_tensors, -1.0)?;
let diff_norms = Tensor::foreach_norm(scratch, 2.0)?;
let post_norms = Tensor::foreach_norm(¶m_tensors, 2.0)?;
let mut diff_sq = 0.0f64;
for n in &diff_norms {
let v: f64 = n.item()?;
diff_sq += v * v;
}
let mut post_sq = 0.0f64;
for n in &post_norms {
let v: f64 = n.item()?;
post_sq += v * v;
}
let post_norm = post_sq.sqrt();
let div = if post_norm > 1e-10 {
diff_sq.sqrt() / post_norm
} else {
0.0
};
crate::verbose!(
" ddp-worker: rank {} sync divergence={:.6} \
(||delta||={:.4}, ||pre||={:.4}, ||post||={:.4})",
self.rank, div, diff_sq.sqrt(), pre_norm, post_norm,
);
(Some(div), Some(post_norm), Some(pre_norm))
} else {
(None, None, None)
};
if let Some(mut outer) = self.outer_optimizer.take() {
let _guard = self.comm_stream.as_ref().map(StreamGuard::new);
let new_global = {
let prev: &[Tensor] =
self.outer_prev_global.as_deref().unwrap_or(¶m_tensors);
outer.outer_step(prev, ¶m_tensors)?
};
{
let _no_grad = NoGradGuard::new();
for (p, ng) in param_tensors.iter().zip(&new_global) {
p.copy_(ng, true)?;
}
}
self.outer_prev_global = Some(new_global);
let resets_inner = outer.resets_inner();
self.outer_optimizer = Some(outer);
if resets_inner {
self.optimizer.reset_state();
}
}
if let (Some(ev), Some(stream)) =
(&self.copy_done, &self.comm_stream)
{
ev.record_on(stream)?;
}
let divg_ms = divg_start.elapsed().as_secs_f64() * 1000.0;
let total_ms = _diag_start.elapsed().as_secs_f64() * 1000.0;
crate::verbose!(
" ddp-sync-diag: rank {} sync_total={:.1}ms (nccl={:.1}ms divg={:.1}ms attempts={})",
self.rank, total_ms, nccl_ms_total, divg_ms, attempt + 1,
);
return Ok(divergence);
}
Err(e) => {
let aborted = self
.nccl_abort_handle
.as_ref()
.is_some_and(|h| h.is_aborted());
if !aborted {
return Err(e);
}
if attempt + 1 == MAX_REBUILD_ATTEMPTS {
return Err(TensorError::new(&format!(
"sync_now_nccl: rank {} hit max NCCL rebuild \
attempts ({}) without successful AllReduce",
self.rank, MAX_REBUILD_ATTEMPTS,
)));
}
crate::verbose!(
" ddp-worker: rank {} NCCL collective aborted on \
attempt {} (err: {}), waiting for new comm and \
retrying",
self.rank,
attempt + 1,
e,
);
let pending = self.wait_for_nccl_session()?;
let uid_bytes: [u8; crate::distributed::NCCL_UNIQUE_ID_BYTES] =
pending.uid_bytes.as_slice().try_into().map_err(|_| {
TensorError::new(
"sync_now_nccl: NewNcclSession uid_bytes \
wrong length (expected NCCL_UNIQUE_ID_BYTES)",
)
})?;
let uid =
crate::distributed::nccl::NcclUniqueId::from_bytes(uid_bytes);
let new_comm = crate::distributed::nccl::NcclRankComm::init_rank(
pending.new_rank,
pending.new_world_size,
&uid,
)?;
self.replace_nccl_comm(new_comm);
}
}
}
Err(TensorError::new(&format!(
"sync_now_nccl: rank {} unexpected exit from retry loop",
self.rank,
)))
}
pub(super) fn wait_for_nccl_session(
&self,
) -> Result<crate::distributed::nccl_session::PendingNcclSession> {
let mailbox = self.nccl_session_mailbox.as_ref().ok_or_else(|| {
TensorError::new(
"sync_now_nccl: NCCL aborted but no session mailbox attached; \
cluster_worker must call attach_nccl_session_mailbox before \
run_until_shutdown.",
)
})?;
let start = Instant::now();
let max_wait = Duration::from_secs(60);
loop {
if let Ok(mut g) = mailbox.lock() {
if let Some(p) = g.take() {
return Ok(p);
}
}
if let Some(ref dead_ranks) = self.local_dead_ranks {
let dead = dead_ranks.dead_count();
if dead >= self.world_size.saturating_sub(1) {
return Err(TensorError::new(&format!(
"sync_now_nccl: rank {} is lone NCCL survivor \
({} of {} ranks dead); no rendezvous possible",
self.rank, dead, self.world_size,
)));
}
}
if start.elapsed() > max_wait {
return Err(TensorError::new(&format!(
"sync_now_nccl: rank {} timed out waiting for new \
NCCL session after {:?}",
self.rank, max_wait,
)));
}
std::thread::sleep(Duration::from_millis(50));
}
}
pub(crate) fn compute_stream_guard(&self) -> Option<StreamGuard> {
self.compute_stream.as_ref().map(StreamGuard::new)
}
pub(crate) fn sync_before_forward(&mut self) -> Result<()> {
if self.pending_param_h2d
&& let Some(stream) = &self.comm_stream
{
let t = Instant::now();
stream.synchronize()?;
self.last_h2d_wait_ms = t.elapsed().as_secs_f64() * 1000.0;
if self.prof_enabled {
self.h2d_wait_ms_total += self.last_h2d_wait_ms;
}
self.pending_param_h2d = false;
}
Ok(())
}
pub fn train_step(
&mut self,
batch: &[Tensor],
train_fn: &impl Fn(&M, &[Tensor]) -> Result<Variable>,
) -> Result<(f64, f64)> {
self.sync_before_forward()?;
let _stream_guard = self.compute_stream.as_ref().map(StreamGuard::new);
let start = Instant::now();
let loss = train_fn(&self.model, batch)?;
let loss_val: f64 = loss.data().item()?;
loss.backward()?;
self.optimizer_step_and_bookkeep()?;
let elapsed_ms = start.elapsed().as_secs_f64() * 1000.0;
Ok((loss_val, elapsed_ms))
}
pub(crate) fn optimizer_step_and_bookkeep(&mut self) -> Result<()> {
let _stream_guard = self.compute_stream.as_ref().map(StreamGuard::new);
if let Some(max_norm) = self.max_grad_norm {
let params: Vec<Tensor> = self.model.parameters()
.iter()
.filter(|p| p.variable.grad().is_some())
.map(|p| p.variable.data())
.collect();
if !params.is_empty() {
Tensor::clip_grad_norm_fused(¶ms, max_norm)?;
}
}
if let Some(ref sched) = self.scheduler {
let base = sched.lr(self.global_step + self.steps_since_avg);
self.optimizer.set_lr(base * self.lr_scale);
}
self.optimizer.step()?;
self.optimizer.zero_grad();
self.local_step += 1;
self.steps_since_avg += 1;
Ok(())
}
}