use super::*;
#[test]
fn test_worker_new_and_accessors() {
let (worker, _ch) = make_test_worker();
assert_eq!(worker.rank(), 0);
assert_eq!(worker.local_step(), 0);
assert_eq!(worker.current_version(), 0);
assert_eq!(worker.param_vars.len(), 2); }
#[test]
fn test_worker_snapshot_params() {
let (mut worker, _ch) = make_test_worker();
let snap = worker.snapshot_params();
assert_eq!(snap.rank, 0);
assert_eq!(snap.params.len(), 2); assert_eq!(snap.buffers.len(), 0); assert_eq!(snap.batch_count, 0);
assert_eq!(snap.params[0].shape(), &[2, 4]); assert_eq!(snap.params[1].shape(), &[2]); }
#[test]
fn test_worker_snapshot_is_send() {
let (mut worker, _ch) = make_test_worker();
let snap = worker.snapshot_params();
let (tx, rx) = mpsc::channel::<ParamSnapshot>();
tx.send(snap).unwrap();
let received = rx.recv().unwrap();
assert_eq!(received.rank, 0);
assert_eq!(received.params.len(), 2);
}
#[test]
fn test_worker_load_averaged() {
let (mut worker, _ch) = make_test_worker();
let cpu = TensorOptions { dtype: DType::Float32, device: Device::CPU };
let new_weight = Tensor::ones(&[2, 4], cpu).unwrap();
let new_bias = Tensor::ones(&[2], cpu).unwrap();
let update = AveragedParams {
params: vec![new_weight, new_bias],
buffers: vec![],
version: 42,
};
worker.load_averaged(&update).unwrap();
let dev = test_device();
if let Device::CUDA(idx) = dev {
crate::tensor::cuda_synchronize(idx);
}
assert_eq!(worker.current_version(), 42);
let snap = worker.snapshot_params();
let w_sum: f64 = snap.params[0].sum().unwrap().item().unwrap();
assert!((w_sum - 8.0).abs() < 1e-5, "weight should be all ones (sum=8), got {w_sum}");
let b_sum: f64 = snap.params[1].sum().unwrap().item().unwrap();
assert!((b_sum - 2.0).abs() < 1e-5, "bias should be all ones (sum=2), got {b_sum}");
}
#[test]
fn test_worker_load_averaged_easgd_blends() {
let alpha = 0.25;
let (mut worker, _ch) = make_test_worker_customized(0, 1, 4, |c| {
c.policy = ApplyPolicy::Async;
c.easgd_alpha = Some(alpha);
});
let pre_w = worker.param_vars[0].data().to_f32_vec().unwrap();
let pre_b = worker.param_vars[1].data().to_f32_vec().unwrap();
let cpu = TensorOptions { dtype: DType::Float32, device: Device::CPU };
let avg_w_val = 3.0;
let avg_b_val = -1.0;
let avg_w = Tensor::full(&[2, 4], avg_w_val, cpu).unwrap();
let avg_b = Tensor::full(&[2], avg_b_val, cpu).unwrap();
let update = AveragedParams {
params: vec![avg_w, avg_b],
buffers: vec![],
version: 7,
};
worker.load_averaged(&update).unwrap();
if let Device::CUDA(idx) = test_device() {
crate::tensor::cuda_synchronize(idx);
}
let post_w = worker.param_vars[0].data().to_f32_vec().unwrap();
let post_b = worker.param_vars[1].data().to_f32_vec().unwrap();
for (i, (pre, post)) in pre_w.iter().zip(&post_w).enumerate() {
let want = (1.0 - alpha) as f32 * pre + alpha as f32 * avg_w_val as f32;
assert!(
(post - want).abs() < 1e-5,
"weight[{i}]: want {want}, got {post} (pre {pre})"
);
}
for (i, (pre, post)) in pre_b.iter().zip(&post_b).enumerate() {
let want = (1.0 - alpha) as f32 * pre + alpha as f32 * avg_b_val as f32;
assert!(
(post - want).abs() < 1e-5,
"bias[{i}]: want {want}, got {post} (pre {pre})"
);
}
}
#[test]
fn test_update_subtracts_snapshot_steps_not_zeroes() {
let (mut worker, _ch) = make_test_worker();
worker.set_steps_since_avg(5);
worker.dispatch_control(ControlMsg::RequestParams).unwrap();
worker.set_steps_since_avg(worker.steps_since_avg() + 3);
let cpu = TensorOptions { dtype: DType::Float32, device: Device::CPU };
let update = AveragedParams {
params: vec![
Tensor::ones(&[2, 4], cpu).unwrap(),
Tensor::ones(&[2], cpu).unwrap(),
],
buffers: vec![],
version: 1,
};
worker.dispatch_control(ControlMsg::Update(update)).unwrap();
assert_eq!(worker.steps_since_avg(), 3, "overshoot steps must keep their mass credit");
let update2 = AveragedParams {
params: vec![
Tensor::ones(&[2, 4], cpu).unwrap(),
Tensor::ones(&[2], cpu).unwrap(),
],
buffers: vec![],
version: 2,
};
worker.dispatch_control(ControlMsg::Update(update2)).unwrap();
assert_eq!(worker.steps_since_avg(), 3, "second Update without a snapshot subtracts nothing");
}
#[test]
fn test_snapshot_never_aliases_live_params() {
let (mut worker, _ch) = make_test_worker();
let snap = worker.snapshot_params();
let before: f64 = snap.params[0].sum().unwrap().item().unwrap();
let cpu = TensorOptions { dtype: DType::Float32, device: Device::CPU };
let ones = Tensor::ones(&[2, 4], cpu).unwrap();
crate::autograd::no_grad(|| -> crate::tensor::Result<()> {
let live = worker.param_vars[0].data();
let src = if live.device() == Device::CPU {
ones
} else {
ones.to_device(live.device())?
};
live.copy_(&src, false)?;
Ok(())
})
.unwrap();
if let Device::CUDA(idx) = test_device() {
crate::tensor::cuda_synchronize(idx);
}
let after: f64 = snap.params[0].sum().unwrap().item().unwrap();
assert!(
(after - before).abs() < 1e-6,
"snapshot changed after live-param mutation (aliased storage): before={before} after={after}"
);
}
#[test]
fn test_worker_load_averaged_wrong_count() {
let (mut worker, _ch) = make_test_worker();
let update = AveragedParams {
params: vec![], buffers: vec![],
version: 1,
};
assert!(worker.load_averaged(&update).is_err());
}
#[test]
fn test_worker_train_step() {
let (mut worker, ch) = make_test_worker();
let opts = test_opts();
let batch = vec![
Tensor::randn(&[4, 4], opts).unwrap(),
Tensor::randn(&[4, 2], opts).unwrap(),
];
let (loss, ms) = worker.train_step(&batch, &mse_train).unwrap();
assert!(ms > 0.0);
assert!(loss > 0.0);
assert_eq!(worker.local_step(), 1);
assert!(ch.timing_rx.try_recv().is_err());
}
#[test]
fn test_worker_report_timing() {
let (worker, ch) = make_test_worker();
worker.report_timing(12.5, 2.0, None, 0.5, None).unwrap();
let msg = ch.timing_rx.recv().unwrap();
match msg {
TimingMsg::Batch { rank, batch_ms, step_count, .. } => {
assert_eq!(rank, 0);
assert!((batch_ms - 12.5).abs() < 1e-10);
assert_eq!(step_count, 0);
}
_ => panic!("expected Batch"),
}
}
#[test]
fn test_worker_report_epoch() {
let (worker, ch) = make_test_worker();
worker.report_epoch(0.5, 100, 5000.0, 5000.0, 0.0, 0.0).unwrap();
let msg = ch.metrics_rx.recv().unwrap();
assert_eq!(msg.rank, 0);
assert_eq!(msg.epoch, 0);
assert!((msg.avg_loss - 0.5).abs() < 1e-10);
assert_eq!(msg.batches_processed, 100);
}
#[test]
fn test_worker_handle_control_request_params() {
let (mut worker, ch) = make_test_worker();
ch.control_tx.send(ControlMsg::RequestParams).unwrap();
let shutdown = worker.handle_control().unwrap();
assert!(!shutdown);
let snap = ch.param_rx.recv().unwrap();
assert_eq!(snap.rank, 0);
assert_eq!(snap.params.len(), 2);
}
#[test]
fn test_worker_handle_control_update() {
let (mut worker, ch) = make_test_worker();
let dev = test_device();
let opts = TensorOptions { dtype: DType::Float32, device: dev };
let update = AveragedParams {
params: vec![
Tensor::zeros(&[2, 4], opts).unwrap(),
Tensor::zeros(&[2], opts).unwrap(),
],
buffers: vec![],
version: 7,
};
ch.control_tx.send(ControlMsg::Update(update)).unwrap();
let shutdown = worker.handle_control().unwrap();
assert!(!shutdown);
assert_eq!(worker.current_version(), 7);
}
#[test]
fn test_worker_handle_control_start_epoch() {
let (mut worker, ch) = make_test_worker();
assert!(worker.pending_plan.is_none());
ch.control_tx.send(ControlMsg::StartEpoch(EpochPlan {
epoch: 1, partition_offset: 0, partition_size: 750,
})).unwrap();
worker.handle_control().unwrap();
let plan = worker.pending_plan.take();
assert!(plan.is_some());
assert_eq!(plan.unwrap().partition_size, 750);
assert!(worker.pending_plan.is_none()); }
#[test]
fn test_worker_handle_control_shutdown() {
let (mut worker, ch) = make_test_worker();
ch.control_tx.send(ControlMsg::Shutdown).unwrap();
let shutdown = worker.handle_control().unwrap();
assert!(shutdown);
}
#[test]
fn test_worker_handle_control_sync_now_noop() {
let (mut worker, ch) = make_test_worker();
ch.control_tx.send(ControlMsg::SyncNow).unwrap();
let shutdown = worker.handle_control().unwrap();
assert!(!shutdown);
}
#[test]
fn test_worker_full_roundtrip() {
let (mut worker, ch) = make_test_worker();
let opts = test_opts();
let batch = vec![
Tensor::randn(&[4, 4], opts).unwrap(),
Tensor::randn(&[4, 2], opts).unwrap(),
];
worker.train_step(&batch, &mse_train).unwrap();
assert_eq!(worker.local_step(), 1);
ch.control_tx.send(ControlMsg::RequestParams).unwrap();
worker.handle_control().unwrap();
let snap = ch.param_rx.recv().unwrap();
assert_eq!(snap.batch_count, 1);
let update = AveragedParams {
params: snap.params,
buffers: snap.buffers,
version: 1,
};
ch.control_tx.send(ControlMsg::Update(update)).unwrap();
worker.handle_control().unwrap();
assert_eq!(worker.current_version(), 1);
let batch2 = vec![
Tensor::randn(&[4, 4], opts).unwrap(),
Tensor::randn(&[4, 2], opts).unwrap(),
];
worker.train_step(&batch2, &mse_train).unwrap();
assert_eq!(worker.local_step(), 2);
}
#[test]
fn test_worker_epoch_from_plan() {
let (mut worker, _ch) = make_test_worker();
assert_eq!(worker.current_epoch, 0);
worker.current_epoch = 3;
assert_eq!(worker.current_epoch, 3);
}
#[test]
fn test_worker_channels_create() {
let ((timing_tx, metrics_tx, param_tx, _final_param_tx, _control_rx), ch) =
GpuWorker::<Linear>::channels();
timing_tx.send(TimingMsg::Batch { rank: 0, batch_ms: 1.0, data_ms: 0.0, step_count: 0, param_norm: None, batch_loss: 0.1, sync_divergence: None }).unwrap();
let msg = ch.timing_rx.recv().unwrap();
assert!(matches!(msg, TimingMsg::Batch { rank: 0, .. }));
metrics_tx.send(MetricsMsg {
rank: 0, epoch: 0, avg_loss: 0.5, batches_processed: 10, epoch_ms: 100.0, share_complete_ms: 100.0, compute_only_ms: 100.0, data_starve_ms: 0.0,
samples_processed: 320, scalars: HashMap::new(),
}).unwrap();
let msg = ch.metrics_rx.recv().unwrap();
assert_eq!(msg.batches_processed, 10);
param_tx.send(ParamSnapshot {
rank: 0, params: vec![], buffers: vec![], batch_count: 0,
}).unwrap();
let snap = ch.param_rx.recv().unwrap();
assert_eq!(snap.rank, 0);
ch.control_tx.send(ControlMsg::Shutdown).unwrap();
}
fn epoch_metrics_fixture(epoch: usize, avg_loss: f64) -> EpochMetrics {
EpochMetrics {
epoch,
scalars: HashMap::new(),
per_rank: vec![],
avg_loss,
per_rank_loss: vec![],
per_rank_samples: vec![],
epoch_ms: 0.0,
per_rank_throughput: vec![],
per_rank_batch_share: vec![],
per_rank_share_complete_ms: vec![],
per_rank_compute_only_ms: vec![],
per_rank_data_starve_ms: vec![],
device_indices: vec![],
}
}
#[test]
fn test_worker_metrics_stream_forwards_every_epoch() {
let (mut worker, ch) = make_test_worker();
let rx = worker.enable_metrics_stream();
ch.control_tx
.send(ControlMsg::EpochAggregated(Box::new(epoch_metrics_fixture(0, 0.9))))
.unwrap();
ch.control_tx
.send(ControlMsg::EpochAggregated(Box::new(epoch_metrics_fixture(1, 0.5))))
.unwrap();
let shutdown = worker.handle_control().unwrap();
assert!(!shutdown);
let drained: Vec<EpochMetrics> = std::iter::from_fn(|| rx.try_recv().ok()).collect();
assert_eq!(drained.len(), 2, "both aggregated epochs must be streamed");
assert_eq!(drained[0].epoch, 0);
assert_eq!(drained[1].epoch, 1);
assert!((drained[1].avg_loss - 0.5).abs() < 1e-9);
let latest = worker.aggregated_metrics().lock().unwrap().clone();
assert_eq!(latest.unwrap().epoch, 1);
}
#[test]
fn test_worker_metrics_slot_without_stream() {
let (mut worker, ch) = make_test_worker();
ch.control_tx
.send(ControlMsg::EpochAggregated(Box::new(epoch_metrics_fixture(3, 0.1))))
.unwrap();
worker.handle_control().unwrap();
let latest = worker.aggregated_metrics().lock().unwrap().clone();
assert_eq!(latest.unwrap().epoch, 3);
}
#[test]
fn test_worker_eval_stream_forwards_each_broadcast() {
let (mut worker, ch) = make_test_worker();
let rx = worker.enable_eval_stream();
ch.control_tx
.send(ControlMsg::EvalBroadcast { epoch: 1, metric: 0.80 })
.unwrap();
ch.control_tx
.send(ControlMsg::EvalBroadcast { epoch: 3, metric: 0.91 })
.unwrap();
assert!(!worker.handle_control().unwrap());
let drained: Vec<(usize, f64)> = std::iter::from_fn(|| rx.try_recv().ok()).collect();
assert_eq!(drained.len(), 2, "both eval broadcasts must be streamed");
assert_eq!(drained[0].0, 1);
assert_eq!(drained[1].0, 3);
assert!((drained[1].1 - 0.91).abs() < 1e-9);
}
#[test]
fn test_worker_eval_broadcast_without_stream_is_noop() {
let (mut worker, ch) = make_test_worker();
ch.control_tx
.send(ControlMsg::EvalBroadcast { epoch: 2, metric: 0.5 })
.unwrap();
assert!(!worker.handle_control().unwrap());
}