1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
//! Reporting and observability hooks: aggregated_metrics, report_*, compute_param_norm, send_final_snapshot, abort_nccl.
use std::sync::mpsc;
use std::sync::{Arc, Mutex};
use crate::nn::Module;
use crate::tensor::{Result, Tensor, TensorError};
use super::super::{
ControlMsg, EpochMetrics,
MetricsMsg, TimingMsg,
};
use super::GpuWorker;
impl<M: Module> GpuWorker<M> {
/// Shared handle to the most recent aggregated [`EpochMetrics`]
/// broadcast from the coord. Clone this `Arc` so the user's
/// `Graph` (or any other reader) sees updates as they arrive
/// without coupling through the worker. Returns `None` inside the
/// mutex until the coord has aggregated at least one epoch.
pub fn aggregated_metrics(&self) -> Arc<Mutex<Option<EpochMetrics>>> {
Arc::clone(&self.aggregated_metrics)
}
/// Arm the cooperative-tier full metrics stream and return its receiver.
///
/// Called once, on the cooperative `Worker` cluster path, before the user
/// runs their first `next_plan` (so no `EpochAggregated` frame can be
/// dispatched before the sender is installed — dispatch only happens on
/// this worker thread). From here on `dispatch_control` forwards every
/// aggregated epoch to the returned receiver *as well as* the latest-only
/// `aggregated_metrics` slot, so the user drains the same per-epoch series
/// the managed `DdpHandle` gets from the coordinator's launcher sink.
///
/// Managed / setup-mode workers never call this, so `metrics_stream_tx`
/// stays `None` and nothing accumulates.
pub(crate) fn enable_metrics_stream(&mut self) -> mpsc::Receiver<EpochMetrics> {
let (tx, rx) = mpsc::channel();
self.metrics_stream_tx = Some(tx);
rx
}
/// Arm the cooperative-tier eval stream and return its receiver. Sibling
/// of [`Self::enable_metrics_stream`]: from here on `dispatch_control`
/// forwards every `EvalBroadcast` frame (`(epoch, metric)`) to the
/// returned receiver, so the user's
/// [`crate::distributed::Worker::poll_eval`] surfaces the controller-
/// elected eval. Managed / setup-mode workers never call this.
pub(crate) fn enable_eval_stream(&mut self) -> mpsc::Receiver<(usize, f64)> {
let (tx, rx) = mpsc::channel();
self.eval_stream_tx = Some(tx);
rx
}
/// Drain any queued `Shutdown` / `ShutdownWithSave` messages from
/// `control_rx` and process them. Called from `ClusterWorker`'s
/// teardown path so that — when the worker exits the main loop
/// with an error (e.g. lone NCCL survivor bailing out of
/// `wait_for_nccl_session`) — any pending coord-sent
/// `ShutdownWithSave` still gets handled and the rank-side
/// checkpoint bundle gets written. Non-shutdown messages in the
/// queue are dropped (the worker is on its way out).
///
/// Returns `true` if a shutdown frame was processed.
pub fn drain_pending_shutdown(&mut self) -> bool {
let mut handled = false;
while let Ok(msg) = self.control_rx.try_recv() {
match msg {
ControlMsg::ShutdownWithSave { reason } => {
if let Some(stem) = self.save_path.clone() {
self.write_checkpoint_bundle(&stem, reason);
} else {
crate::verbose!(
" ddp-worker: rank {} drain_pending_shutdown saw \
ShutdownWithSave but save_path is unset; \
exiting without saving",
self.rank,
);
}
handled = true;
}
ControlMsg::Shutdown => {
handled = true;
}
_ => {
// Drop — worker is exiting, other messages are stale.
}
}
}
handled
}
/// Report the wall-time the rank just spent inside `epoch_fn`. Called
/// by the cluster worker's main loop on the role rank after firing
/// the user closure, so the coord can time-exclude callback cost from
/// the coord's window ledger and update `last_epoch_fn_elapsed_ms_ewma`.
/// Fire-and-forget: a disconnected timing channel is non-fatal here
/// (the loop is exiting anyway).
pub fn report_epoch_fn_elapsed(&self, epoch: usize, elapsed_ms: f64) {
let _ = self.timing_tx.send(TimingMsg::EpochFnElapsed {
rank: self.rank,
epoch,
elapsed_ms,
});
}
/// Send a timing report to the coordinator.
///
/// Also emits a `TimingMsg::LrUpdate` piggyback message so the
/// coordinator's LR-aware meta-controller (when enabled) can track the
/// LR trajectory between averaging cycles. Cheap fire-and-forget; the
/// coordinator caches only the most recent value per rank.
pub fn report_timing(
&self,
batch_ms: f64,
data_ms: f64,
param_norm: Option<f64>,
batch_loss: f64,
sync_divergence: Option<f64>,
) -> Result<()> {
let res = self.timing_tx.send(TimingMsg::Batch {
rank: self.rank,
batch_ms,
data_ms,
step_count: self.local_step,
param_norm,
batch_loss,
sync_divergence,
}).map_err(|_| TensorError::new("timing channel disconnected"));
// Piggyback the current LR after the primary Batch. Failures are
// tolerated — the meta layer simply observes a stale value next cycle.
let _ = self.timing_tx.send(TimingMsg::LrUpdate {
rank: self.rank,
lr: self.optimizer.lr(),
});
res
}
/// Emit a cooperative-tier user intent (eval / checkpoint request) to the
/// controller on `timing_tx`. Fire-and-forget: a disconnected channel is
/// non-fatal (single-device has no controller listening, so the send just
/// drops — the request is a no-op there, as documented on
/// [`Worker::request_eval`](crate::distributed::ddp_run::Worker::request_eval)).
pub(crate) fn report_intent(&self, kind: crate::distributed::wire::IntentKind) {
let _ = self.timing_tx.send(TimingMsg::Intent {
rank: self.rank,
kind,
});
}
/// Compute the L2 norm of all model parameters.
///
/// Uses `Tensor::foreach_norm` for a single batched CUDA kernel instead
/// of per-parameter norm calls. Returns the global L2 norm (sqrt of sum
/// of squared per-tensor norms). Used for NCCL divergence detection.
pub(super) fn compute_param_norm(&self) -> Result<f64> {
let data: Vec<Tensor> = self.param_vars.iter().map(|v| v.data()).collect();
if data.is_empty() {
return Ok(0.0);
}
let norms = Tensor::foreach_norm(&data, 2.0)?;
let mut total_sq = 0.0f64;
for n in &norms {
let val: f64 = n.item()?;
total_sq += val * val;
}
Ok(total_sq.sqrt())
}
/// Send the final parameter snapshot on the dedicated channel before exiting.
///
/// This uses `final_param_tx` (not `param_tx`) to avoid racing with
/// CPU averaging snapshot collection on the same channel, and the
/// EXACT readout (never the pinned staging) so the trained weights
/// carry full f32 precision even when `bf16_wire` stages the
/// averaging-plane snapshots in bf16.
pub fn send_final_snapshot(&mut self) {
let snap = self.snapshot_params_exact();
let _ = self.final_param_tx.send(snap);
}
/// Abort the NCCL communicator, unblocking any stuck collective.
///
/// Must be called before [`Self::send_final_snapshot`] when the training loop
/// exits due to shutdown. A pending AllReduce on `comm_stream` (from a
/// SyncNow whose peer died) would block `to_device(CPU)` in snapshot_params
/// because the CUDA default stream synchronizes with all other streams.
pub fn abort_nccl(&mut self) {
if let Some(comm) = self.nccl_comm.take() {
let _ = comm.abort_handle().abort();
}
}
/// Notify the coordinator that this worker completed CLEANLY and is
/// about to exit.
///
/// Clean completion only: the coordinator's `Exiting` latch
/// suppresses both death detectors for this rank (heartbeat
/// staleness + launcher-reported child exits). Calling this on an
/// error exit masks the death — no ledger declare, no ElChe
/// recompute, no partition redistribution — and a cadence cohort
/// wedges on the dead rank's unfinished window. Error paths must
/// exit silently and let the detectors fire.
pub fn report_exiting(&self) {
let _ = self.timing_tx.send(TimingMsg::Exiting { rank: self.rank });
}
/// Send epoch-end metrics to the coordinator.
///
/// Drains the thread-local scalar accumulator populated by
/// [`record_scalar()`](super::super::record_scalar) calls during this epoch.
pub fn report_epoch(
&self,
avg_loss: f64,
batches: usize,
epoch_ms: f64,
share_complete_ms: f64,
compute_only_ms: f64,
data_starve_ms: f64,
) -> Result<()> {
let scalars = super::super::drain_scalars();
self.metrics_tx.send(MetricsMsg {
rank: self.rank,
epoch: self.current_epoch,
avg_loss,
batches_processed: batches,
epoch_ms,
samples_processed: batches * self.batch_size,
share_complete_ms,
compute_only_ms,
data_starve_ms,
scalars,
}).map_err(|_| TensorError::new("metrics channel disconnected"))
}
}