use super::*;
#[test]
fn test_builder_single_gpu_fallback() {
if crate::tensor::usable_cuda_devices().len() >= 2 {
return;
}
let ddp = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.dataset(Arc::new(TestDataset { n: 100 }))
.batch_size(4)
.num_epochs(2)
.policy(ApplyPolicy::Sync)
.backend(AverageBackend::Cpu) .run()
.unwrap();
assert!(ddp.world_size() >= 1);
let state = ddp.join().unwrap();
assert_eq!(state.params.len(), 2);
assert_eq!(state.buffers.len(), 0);
}
#[test]
fn test_ddp_handle_send_sync() {
fn assert_send<T: Send>() {}
assert_send::<DdpHandle>();
assert_send::<TrainedState>();
}
#[test]
fn test_builder_with_defaults() {
let ddp = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.dataset(Arc::new(TestDataset { n: 100 }))
.batch_size(4)
.num_epochs(2)
.backend(AverageBackend::Cpu)
.run()
.unwrap();
assert!(ddp.world_size() >= 1);
let state = ddp.join().unwrap();
assert_eq!(state.params.len(), 2);
}
#[test]
fn test_builder_with_all_options() {
let ddp = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.dataset(Arc::new(TestDataset { n: 8 }))
.batch_size(4)
.num_epochs(1)
.policy(ApplyPolicy::Sync)
.backend(AverageBackend::Cpu)
.overhead_target(0.15)
.max_anchor(100)
.anchor(5)
.divergence_threshold(0.1)
.max_batch_diff(10)
.run()
.unwrap();
let state = ddp.join().unwrap();
assert_eq!(state.params.len(), 2);
}
#[test]
#[should_panic(expected = "dataset is required")]
fn test_builder_missing_dataset_panics() {
let _ = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.batch_size(4)
.num_epochs(2)
.run();
}
#[test]
#[should_panic(expected = "batch_size is required")]
fn test_builder_missing_batch_size_panics() {
let _ = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.dataset(Arc::new(TestDataset { n: 100 }))
.num_epochs(2)
.run();
}
#[test]
#[should_panic(expected = "num_epochs is required")]
fn test_builder_missing_num_epochs_panics() {
let _ = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.dataset(Arc::new(TestDataset { n: 100 }))
.batch_size(4)
.run();
}
#[test]
fn resume_from_loads_meta_and_seeds_coord_config() {
use crate::distributed::{
CheckpointBundle, CheckpointMeta, ElCheState, SaveReason,
};
use crate::distributed::el_che::Phase;
use super::orchestrator::build_coord_config_from_builder;
let dir = std::env::temp_dir().join(format!(
"flodl_resume_e2e_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
let stem = dir.join("ckpt").to_string_lossy().into_owned();
let elche_state = ElCheState {
anchor: 14,
anchor_rank: Some(1),
smoothed_ms_per_batch: vec![3.5, 5.5],
phase: Phase::Stable,
calibration_count: 17,
trend_history: Some(vec![0.005, 0.01, 0.02, 0.025]),
};
let meta = CheckpointMeta::new(
4, 9_876, 33, 2, SaveReason::GracefulShutdown,
)
.with_elche_state(elche_state.clone());
let meta_path = CheckpointBundle::meta_path(&stem);
meta.write_to_file(&meta_path).unwrap();
let user_config = DdpRunConfig::new().with_resume_from(stem.clone());
let coord_config = build_coord_config_from_builder(
ApplyPolicy::Cadence,
AverageBackend::Cpu,
&user_config,
None,
None,
None,
2,
100,
4,
1,
)
.expect("resume meta loads cleanly");
assert_eq!(coord_config.start_epoch, 4);
assert_eq!(coord_config.start_global_step, 9_876);
assert_eq!(coord_config.start_avg_count, 33);
assert_eq!(
coord_config.start_elche_state.as_ref(),
Some(&elche_state),
"ElCheState carries through resume_from"
);
let history = coord_config
.convergence_guard
.trend_history()
.expect("TrendGuard surfaces a non-empty history after resume");
assert_eq!(history, vec![0.005, 0.01, 0.02, 0.025]);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn resume_from_missing_meta_errors() {
use super::orchestrator::build_coord_config_from_builder;
let user_config = DdpRunConfig::new()
.with_resume_from("/nonexistent/path/that/cannot/exist/ckpt");
let result = build_coord_config_from_builder(
ApplyPolicy::Cadence,
AverageBackend::Cpu,
&user_config,
None,
None,
None,
2,
100,
4,
1,
);
let err = match result {
Ok(_) => panic!("missing meta file must error, got Ok"),
Err(e) => e,
};
let msg = err.to_string();
assert!(
msg.contains("read") || msg.contains("CheckpointMeta"),
"expected read-error message, got: {msg}"
);
}
#[test]
fn test_worker_current_epoch_accessor() {
let (mut worker, _ch) = make_test_worker();
assert_eq!(worker.current_epoch(), 0);
worker.current_epoch = 1;
assert_eq!(worker.current_epoch(), 1);
}
#[test]
fn test_worker_set_lr() {
let (mut worker, _ch) = make_test_worker();
worker.set_lr(0.1);
let opts = test_opts();
let batch = vec![
Tensor::randn(&[4, 4], opts).unwrap(),
Tensor::randn(&[4, 2], opts).unwrap(),
];
let (loss, _) = worker.train_step(&batch, &mse_train).unwrap();
assert!(loss > 0.0);
}
#[test]
fn test_epoch_fn_called_per_epoch() {
use std::sync::atomic::{AtomicUsize, Ordering};
let counter = Arc::new(AtomicUsize::new(0));
let epochs_seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let counter_c = counter.clone();
let epochs_c = epochs_seen.clone();
let num_epochs = 3;
let ddp = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.dataset(Arc::new(TestDataset { n: 100 }))
.batch_size(4)
.num_epochs(num_epochs)
.backend(AverageBackend::Cpu)
.policy(ApplyPolicy::Sync)
.epoch_fn(move |epoch, worker| {
counter_c.fetch_add(1, Ordering::Relaxed);
epochs_c.lock().unwrap().push(epoch);
assert_eq!(worker.current_epoch(), epoch);
})
.run()
.unwrap();
let world = ddp.world_size();
let _state = ddp.join().unwrap();
let got_counter = counter.load(Ordering::Relaxed);
let expected_counter = num_epochs * world;
let mut seen = epochs_seen.lock().unwrap().clone();
seen.sort();
let mut expected_epochs: Vec<usize> = (0..num_epochs).cycle().take(num_epochs * world).collect();
expected_epochs.sort();
assert_eq!(
got_counter, expected_counter,
"epoch_fn fire count mismatch — got {got_counter}, expected {expected_counter}. \
world_size={world}, num_epochs={num_epochs}, epochs_seen={seen:?}.",
);
assert_eq!(
seen, expected_epochs,
"epoch_fn epoch-index set mismatch — got {seen:?}, expected {expected_epochs:?}. \
world_size={world}, num_epochs={num_epochs}, counter={got_counter}.",
);
}
#[test]
fn test_epoch_fn_set_lr() {
use std::sync::atomic::{AtomicUsize, Ordering};
let call_count = Arc::new(AtomicUsize::new(0));
let call_count_c = call_count.clone();
let ddp = crate::distributed::Trainer::builder(
|dev| Linear::on_device(4, 2, dev),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
mse_train,
)
.dataset(Arc::new(TestDataset { n: 100 }))
.batch_size(4)
.num_epochs(3)
.backend(AverageBackend::Cpu)
.policy(ApplyPolicy::Sync)
.epoch_fn(move |epoch, worker| {
let lr = 0.01 * (1.0 - epoch as f64 * 0.3);
worker.set_lr(lr);
call_count_c.fetch_add(1, Ordering::Relaxed);
})
.run()
.unwrap();
let world = ddp.world_size();
let _state = ddp.join().unwrap();
assert_eq!(call_count.load(Ordering::Relaxed), 3 * world);
}
#[test]
fn test_worker_send_final_snapshot() {
let (mut worker, ch) = make_test_worker();
worker.send_final_snapshot();
let snap = ch.final_param_rx.recv().unwrap();
assert_eq!(snap.params.len(), 2); assert_eq!(snap.rank, 0);
}
#[test]
fn max_overshoot_plumbs_into_coord_config_and_pins_it() {
use super::orchestrator::build_coord_config_from_builder;
let user_config = DdpRunConfig::new().with_max_overshoot(7);
let coord_config = build_coord_config_from_builder(
ApplyPolicy::Async,
AverageBackend::Cpu,
&user_config,
None, None, None,
2, 100, 4, 1,
)
.expect("build");
assert_eq!(coord_config.overshoot_initial, 7);
assert_eq!(coord_config.overshoot_ceiling, 7);
assert!(
!coord_config.overshoot_auto,
"a user-set max_overshoot pins the bound and disables auto-tune"
);
}
#[test]
fn unset_max_overshoot_leaves_auto_tune_defaults() {
use super::orchestrator::build_coord_config_from_builder;
let user_config = DdpRunConfig::new();
let coord_config = build_coord_config_from_builder(
ApplyPolicy::Async,
AverageBackend::Cpu,
&user_config,
None, None, None,
2, 100, 4, 1,
)
.expect("build");
assert!(
coord_config.overshoot_auto,
"unset max_overshoot keeps auto-tune on"
);
assert_eq!(coord_config.overshoot_ceiling, 15);
}