use std::time::Instant;
use crate::distributed::ddp_run::AverageBackend;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum CpuAvgPhase {
Idle,
Pending,
}
#[derive(Debug)]
pub(super) enum CycleMachine {
NcclInline,
Cpu {
phase: CpuAvgPhase,
pending_since: Option<Instant>,
throttled: Vec<bool>,
},
}
#[derive(Debug)]
pub(super) struct AvgCycleState {
pub(super) started_at: Option<Instant>,
pub(super) step_snapshot: Vec<usize>,
pub(super) acked: Vec<bool>,
pub(super) divergence: Vec<Option<f64>>,
pub(super) pre_norm: Vec<Option<f64>>,
pub(super) post_norm: Option<f64>,
pub(super) last_sync_ms: f64,
pub(super) sync_lag_ms: Vec<Option<f64>>,
pub(super) upload_ms: Vec<Option<f64>>,
pub(super) machine: CycleMachine,
}
impl AvgCycleState {
pub(super) fn new(backend: AverageBackend, world_size: usize) -> Self {
AvgCycleState {
started_at: None,
step_snapshot: vec![0; world_size],
acked: vec![true; world_size],
divergence: vec![None; world_size],
pre_norm: vec![None; world_size],
post_norm: None,
last_sync_ms: 0.0,
sync_lag_ms: vec![None; world_size],
upload_ms: vec![None; world_size],
machine: match backend {
AverageBackend::Nccl => CycleMachine::NcclInline,
AverageBackend::Cpu => CycleMachine::Cpu {
phase: CpuAvgPhase::Idle,
pending_since: None,
throttled: vec![false; world_size],
},
},
}
}
pub(super) fn arm(&mut self, now: Instant, last_step_count: &[usize]) {
self.started_at = Some(now);
self.step_snapshot.copy_from_slice(last_step_count);
self.acked.fill(false);
}
pub(super) fn sync_ack_step_meaningful(&self) -> bool {
matches!(self.machine, CycleMachine::NcclInline)
}
pub(super) fn note_batch_ack(
&mut self,
rank: usize,
step_count: usize,
sync_divergence: Option<f64>,
) {
if let Some(div) = sync_divergence {
self.divergence[rank] = Some(div);
}
if rank < self.acked.len()
&& !self.acked[rank]
&& step_count > self.step_snapshot[rank]
{
self.acked[rank] = true;
self.capture_sync_elapsed_if_complete();
}
}
pub(super) fn note_sync_ack(
&mut self,
rank: usize,
step_count: usize,
divergence: Option<f64>,
pre_norm: Option<f64>,
post_norm: Option<f64>,
) {
if let Some(div) = divergence {
self.divergence[rank] = Some(div);
}
if let Some(p) = pre_norm {
self.pre_norm[rank] = Some(p);
}
if let Some(p) = post_norm {
match self.post_norm {
None => self.post_norm = Some(p),
Some(prev) => debug_assert!(
(prev - p).abs() <= 1e-6 * prev.abs().max(1.0),
"post_norm rank-disagreement: prev={prev} new={p} (rank {rank})"
),
}
}
if rank < self.acked.len()
&& !self.acked[rank]
&& step_count > self.step_snapshot[rank]
{
self.acked[rank] = true;
if let Some(start) = self.started_at {
if rank < self.sync_lag_ms.len() {
self.sync_lag_ms[rank] =
Some(start.elapsed().as_secs_f64() * 1000.0);
}
}
self.capture_sync_elapsed_if_complete();
}
}
pub(super) fn note_snapshot_ready(&mut self, rank: usize) {
if rank < self.upload_ms.len() {
if let Some(start) = self.started_at {
self.upload_ms[rank] =
Some(start.elapsed().as_secs_f64() * 1000.0);
}
}
}
pub(super) fn capture_sync_elapsed_if_complete(&mut self) {
if self.acked.iter().all(|&a| a) {
if let Some(start) = self.started_at.take() {
self.last_sync_ms = start.elapsed().as_secs_f64() * 1000.0;
}
}
}
pub(super) fn all_alive_acked(
&self,
mut is_dead: impl FnMut(usize) -> bool,
) -> bool {
(0..self.acked.len()).all(|r| is_dead(r) || self.acked[r])
}
pub(super) fn all_alive_diverged(
&self,
mut is_dead: impl FnMut(usize) -> bool,
) -> bool {
(0..self.divergence.len())
.all(|r| is_dead(r) || self.divergence[r].is_some())
}
pub(super) fn take_last_sync_ms(&mut self) -> f64 {
std::mem::take(&mut self.last_sync_ms)
}
pub(super) fn divergence_report(
&self,
) -> crate::distributed::ddp_run::convergence::DivergenceReport {
let pre_norms: Option<Vec<f64>> =
if self.pre_norm.iter().all(|p| p.is_some()) {
Some(self.pre_norm.iter().map(|p| p.unwrap()).collect())
} else {
None
};
crate::distributed::ddp_run::convergence::DivergenceReport {
deltas: self
.divergence
.iter()
.map(|d| d.unwrap_or(0.0))
.collect(),
pre_norms,
post_norm: self.post_norm,
}
}
pub(super) fn reset_divergence_signals(&mut self) {
crate::distributed::ddp_run::convergence::reset_divergence_signals(
&mut self.divergence,
&mut self.pre_norm,
&mut self.post_norm,
);
}
pub(super) fn reset_upload_markers(&mut self) {
for slot in &mut self.upload_ms {
*slot = None;
}
}
pub(super) fn begin_cpu_pending(&mut self, now: Instant) {
if let CycleMachine::Cpu { phase, pending_since, .. } = &mut self.machine {
*phase = CpuAvgPhase::Pending;
*pending_since = Some(now);
}
}
pub(super) fn cpu_pending(&self) -> bool {
matches!(
self.machine,
CycleMachine::Cpu { phase: CpuAvgPhase::Pending, .. }
)
}
pub(super) fn abort_cpu_pending(&mut self) {
if let CycleMachine::Cpu { phase, .. } = &mut self.machine {
*phase = CpuAvgPhase::Idle;
}
}
pub(super) fn finish_cpu_pending(&mut self) -> Option<Instant> {
if let CycleMachine::Cpu { phase, pending_since, .. } = &mut self.machine {
*phase = CpuAvgPhase::Idle;
pending_since.take()
} else {
None
}
}
pub(super) fn cpu_pending_since(&self) -> Option<Instant> {
match &self.machine {
CycleMachine::Cpu { pending_since, .. } => *pending_since,
CycleMachine::NcclInline => None,
}
}
pub(super) fn is_throttled(&self, rank: usize) -> bool {
match &self.machine {
CycleMachine::Cpu { throttled, .. } => throttled[rank],
CycleMachine::NcclInline => false,
}
}
pub(super) fn set_throttled(&mut self, rank: usize) {
if let CycleMachine::Cpu { throttled, .. } = &mut self.machine {
throttled[rank] = true;
}
}
pub(super) fn throttle_all(&mut self) {
if let CycleMachine::Cpu { throttled, .. } = &mut self.machine {
for t in throttled.iter_mut() {
*t = true;
}
}
}
pub(super) fn clear_throttled(&mut self) {
if let CycleMachine::Cpu { throttled, .. } = &mut self.machine {
for t in throttled.iter_mut() {
*t = false;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::{Duration, Instant};
fn cpu_state(n: usize) -> AvgCycleState {
AvgCycleState::new(AverageBackend::Cpu, n)
}
fn nccl_state(n: usize) -> AvgCycleState {
AvgCycleState::new(AverageBackend::Nccl, n)
}
#[test]
fn arm_snapshots_steps_and_clears_acks() {
let mut c = nccl_state(2);
assert!(c.acked.iter().all(|&a| a), "at rest = settled");
c.arm(Instant::now(), &[5, 9]);
assert_eq!(c.step_snapshot, vec![5, 9]);
assert!(c.acked.iter().all(|&a| !a));
assert!(c.started_at.is_some());
}
#[test]
fn batch_ack_requires_step_past_snapshot() {
let mut c = nccl_state(2);
c.arm(Instant::now(), &[5, 9]);
c.note_batch_ack(0, 5, None); assert!(!c.acked[0]);
c.note_batch_ack(0, 6, None);
assert!(c.acked[0]);
assert!(c.started_at.is_some());
c.note_batch_ack(1, 10, None);
assert!(c.started_at.is_none(), "all-acked takes started_at");
assert!(c.last_sync_ms >= 0.0);
}
#[test]
fn sync_ack_step_meaningful_only_on_nccl() {
assert!(nccl_state(2).sync_ack_step_meaningful());
assert!(!cpu_state(2).sync_ack_step_meaningful());
}
#[test]
fn sync_ack_populates_evidence_and_lag() {
let mut c = cpu_state(2);
c.arm(
Instant::now()
.checked_sub(Duration::from_millis(10))
.unwrap(),
&[3, 3],
);
c.note_sync_ack(0, 4, Some(0.5), Some(1.0), Some(2.0));
assert_eq!(c.divergence[0], Some(0.5));
assert_eq!(c.pre_norm[0], Some(1.0));
assert_eq!(c.post_norm, Some(2.0));
assert!(c.acked[0]);
assert!(c.sync_lag_ms[0].is_some_and(|ms| ms > 0.0));
assert!(!c.all_alive_diverged(|_| false));
c.note_sync_ack(1, 4, Some(0.7), Some(1.0), None);
assert!(c.all_alive_diverged(|_| false));
assert!(c.all_alive_acked(|_| false));
}
#[test]
fn dead_ranks_count_as_acked_and_diverged() {
let mut c = cpu_state(2);
c.arm(Instant::now(), &[0, 0]);
c.note_sync_ack(0, 1, Some(0.1), None, None);
assert!(c.all_alive_acked(|r| r == 1));
assert!(c.all_alive_diverged(|r| r == 1));
}
#[test]
fn snapshot_ready_dropped_after_finalize() {
let mut c = cpu_state(1);
c.arm(Instant::now(), &[0]);
c.note_snapshot_ready(0);
assert!(c.upload_ms[0].is_some());
c.reset_upload_markers();
assert!(c.upload_ms[0].is_none());
c.note_batch_ack(0, 1, None);
assert!(c.started_at.is_none());
c.note_snapshot_ready(0);
assert!(c.upload_ms[0].is_none(), "late straggler dropped");
}
#[test]
fn cpu_pending_window_lifecycle() {
let mut c = cpu_state(2);
assert!(!c.cpu_pending());
c.begin_cpu_pending(Instant::now());
assert!(c.cpu_pending());
assert!(c.cpu_pending_since().is_some());
let start = c.finish_cpu_pending();
assert!(start.is_some());
assert!(!c.cpu_pending());
assert!(c.cpu_pending_since().is_none());
c.begin_cpu_pending(Instant::now());
c.abort_cpu_pending();
assert!(!c.cpu_pending());
assert!(c.cpu_pending_since().is_some());
}
#[test]
fn cpu_machine_hooks_are_noops_on_nccl() {
let mut c = nccl_state(2);
c.begin_cpu_pending(Instant::now());
assert!(!c.cpu_pending());
assert!(c.cpu_pending_since().is_none());
assert!(c.finish_cpu_pending().is_none());
c.set_throttled(0);
c.throttle_all();
assert!(!c.is_throttled(0));
}
#[test]
fn throttle_ledger_on_cpu() {
let mut c = cpu_state(3);
c.set_throttled(1);
assert!(c.is_throttled(1) && !c.is_throttled(0));
c.throttle_all();
assert!((0..3).all(|r| c.is_throttled(r)));
c.clear_throttled();
assert!((0..3).all(|r| !c.is_throttled(r)));
}
#[test]
fn divergence_report_all_or_none_pre_norms() {
let mut c = cpu_state(2);
c.arm(Instant::now(), &[0, 0]);
c.note_sync_ack(0, 1, Some(0.5), Some(1.0), Some(2.0));
c.note_sync_ack(1, 1, Some(0.3), None, None);
let report = c.divergence_report();
assert_eq!(report.deltas, vec![0.5, 0.3]);
assert!(report.pre_norms.is_none(), "partial pre-norms omitted");
assert_eq!(report.post_norm, Some(2.0));
c.reset_divergence_signals();
assert!(c.divergence.iter().all(|d| d.is_none()));
assert!(c.post_norm.is_none());
}
}