pub(crate) use super::*;
pub(crate) use crate::autograd::Variable;
pub(crate) use crate::nn::{Linear, Module};
pub(crate) use crate::tensor::{DType, Tensor, TensorError, TensorOptions, test_device, test_opts};
pub(crate) use std::sync::mpsc;
mod builder_and_resume;
mod checkpoint;
mod cooperative;
mod worker;
#[test]
fn test_apply_policy_variants() {
let policies = [ApplyPolicy::Sync, ApplyPolicy::Cadence, ApplyPolicy::Async];
assert_eq!(policies.len(), 3);
assert_eq!(ApplyPolicy::Sync, ApplyPolicy::Sync);
assert_ne!(ApplyPolicy::Sync, ApplyPolicy::Async);
}
#[test]
fn test_average_backend_variants() {
let backends = [AverageBackend::Nccl, AverageBackend::Cpu];
assert_eq!(backends.len(), 2);
assert_eq!(AverageBackend::Nccl, AverageBackend::Nccl);
assert_ne!(AverageBackend::Nccl, AverageBackend::Cpu);
}
#[test]
fn test_control_msg_variants() {
let _req = ControlMsg::RequestParams;
let _sync = ControlMsg::SyncNow;
let _throttle = ControlMsg::Throttle;
let _start = ControlMsg::StartEpoch(EpochPlan {
epoch: 0, partition_offset: 0, partition_size: 1000,
});
let _ckpt = ControlMsg::Checkpoint { version: 42, target_rank: 0 };
let _shutdown = ControlMsg::Shutdown;
let _update = ControlMsg::Update(AveragedParams {
params: vec![],
buffers: vec![],
version: 0,
});
}
#[test]
fn test_timing_msg_send() {
fn assert_send<T: Send>() {}
assert_send::<TimingMsg>();
}
#[test]
fn test_metrics_msg_send() {
fn assert_send<T: Send>() {}
assert_send::<MetricsMsg>();
}
#[test]
fn test_param_snapshot_send() {
fn assert_send<T: Send>() {}
assert_send::<ParamSnapshot>();
}
#[test]
fn test_averaged_params_send() {
fn assert_send<T: Send>() {}
assert_send::<AveragedParams>();
}
#[test]
fn test_control_msg_send() {
fn assert_send<T: Send>() {}
assert_send::<ControlMsg>();
}
#[test]
fn test_worker_config_send() {
fn assert_send<T: Send>() {}
assert_send::<WorkerConfig>();
}
#[test]
fn test_worker_config_clone() {
let cfg = WorkerConfig {
rank: 0,
world_size: 2,
device: Device::CPU,
initial_params: vec![],
initial_buffers: vec![],
total_samples: 10000,
augment: 1,
transform: None,
vram_max_usage: 0.90,
ram_max_usage: 0.50,
sample_cache: true,
disk_stage_gb: 0,
disk_stage_dir: None,
batch_size: 32,
seed: 42,
max_grad_norm: None,
vram_pool: false,
easgd_alpha: None,
gamma: 1.0,
bf16_wire: false,
timeline: None,
policy: ApplyPolicy::Sync,
save_path: None,
coord_liveness_timeout_secs:
crate::distributed::ddp_run::DEFAULT_COORD_LIVENESS_TIMEOUT_SECS,
};
let cfg2 = cfg.clone();
assert_eq!(cfg2.rank, 0);
assert_eq!(cfg2.world_size, 2);
assert_eq!(cfg2.total_samples, 10000);
}
pub(super) struct TestDataset {
n: usize,
}
impl crate::data::BatchDataSet for TestDataset {
fn len(&self) -> usize { self.n }
fn get_batch(&self, indices: &[usize]) -> crate::tensor::Result<Vec<Tensor>> {
let n = indices.len() as i64;
let opts = TensorOptions { dtype: DType::Float32, device: Device::CPU };
Ok(vec![
Tensor::randn(&[n, 4], opts)?,
Tensor::randn(&[n, 2], opts)?,
])
}
}
pub(super) fn mse_train(model: &Linear, batch: &[Tensor]) -> Result<Variable> {
let input = Variable::new(batch[0].clone(), false);
let target = Variable::new(batch[1].clone(), false);
let output = model.forward(&input)?;
let diff = output.sub(&target)?;
diff.mul(&diff)?.mean()
}
pub(super) fn make_test_worker() -> (GpuWorker<Linear>, WorkerChannels) {
make_test_worker_with(0, 1, 4)
}
pub(super) fn make_test_worker_with(
rank: usize,
world_size: usize,
dataset_size: usize,
) -> (GpuWorker<Linear>, WorkerChannels) {
make_test_worker_customized(rank, world_size, dataset_size, |_| {})
}
pub(super) fn make_test_worker_customized(
rank: usize,
world_size: usize,
dataset_size: usize,
tweak: impl FnOnce(&mut WorkerConfig),
) -> (GpuWorker<Linear>, WorkerChannels) {
let dev = test_device();
let tmp_model = Linear::on_device(4, 2, dev).unwrap();
let tmp_params: Vec<Tensor> = tmp_model.parameters().iter()
.map(|p| p.variable.data())
.collect();
let tmp_buffers: Vec<Tensor> = tmp_model.buffers().iter()
.map(|b| b.get())
.collect();
drop(tmp_model);
let mut config = WorkerConfig {
rank,
world_size,
device: dev,
initial_params: tmp_params,
initial_buffers: tmp_buffers,
total_samples: dataset_size,
augment: 1,
transform: None,
vram_max_usage: 0.90,
ram_max_usage: 0.50,
sample_cache: true,
disk_stage_gb: 0,
disk_stage_dir: None,
batch_size: 4,
seed: 42,
max_grad_norm: None,
vram_pool: false,
easgd_alpha: None,
gamma: 1.0,
bf16_wire: false,
timeline: None,
policy: ApplyPolicy::Sync,
save_path: None,
coord_liveness_timeout_secs:
crate::distributed::ddp_run::DEFAULT_COORD_LIVENESS_TIMEOUT_SECS,
};
tweak(&mut config);
let ((timing_tx, metrics_tx, param_tx, final_param_tx, control_rx), channels) =
GpuWorker::<Linear>::channels();
let dataset: Arc<dyn crate::data::BatchDataSet> =
Arc::new(TestDataset { n: dataset_size });
let worker = GpuWorker::new(
&config,
|d| Linear::on_device(4, 2, d),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
dataset,
None, None, None, None, timing_tx,
metrics_tx,
param_tx,
final_param_tx,
control_rx,
None, ).unwrap();
(worker, channels)
}
#[test]
fn test_zero_param_model_is_rejected_loudly() {
let err = super::ensure_trainable_params(0, "ddp: single device")
.expect_err("zero-parameter model must be rejected");
let msg = err.to_string();
assert!(msg.contains("zero trainable parameters"), "unexpected: {msg}");
assert!(msg.contains("Module::parameters()"), "unexpected: {msg}");
super::ensure_trainable_params(1, "ddp: single device").unwrap();
}