use super::*;
use crate::distributed::cluster_coordinator::{
ClusterCoordinator, ClusterCoordinatorConfig,
};
use crate::distributed::ddp::ElChe;
use crate::distributed::ddp_run::{ApplyPolicy, AverageBackend};
use crate::distributed::wire::SESSION_SALT_BYTES;
use std::net::Ipv4Addr;
use std::time::Instant;
const TEST_SALT: SessionSalt = [0x42u8; SESSION_SALT_BYTES];
fn coord_config_sync_nccl(world_size: usize) -> ClusterCoordinatorConfig {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Nccl,
world_size,
ElChe::new(world_size, 1),
)
.no_divergence_guard()
}
fn spawn_coord<D>(
world_size: usize,
drive: D,
) -> (u16, thread::JoinHandle<Result<()>>)
where
D: Send + 'static + FnOnce(&mut ClusterCoordinator) -> Result<()>,
{
let (listener, port) = ClusterCoordinator::bind(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
)
.expect("bind succeeds");
let h = thread::spawn(move || -> Result<()> {
let mut coord = ClusterCoordinator::start_from_listener(
listener,
TEST_SALT,
coord_config_sync_nccl(world_size),
)?;
let r = drive(&mut coord);
let _ = coord.shutdown();
r
});
(port, h)
}
#[test]
fn handshake_with_real_coordinator() {
let world_size = 1;
let world_size = world_size.max(2);
let (port, coord_handle) = spawn_coord(world_size, |coord| {
let _ = coord.tick();
Ok(())
});
let coord_real_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), port);
let (addr, _crelay_rx) =
spawn_relay(ChannelKind::Control, coord_real_addr, world_size, TEST_SALT);
fn raw_rank_handshake(addr: SocketAddr, rank: u32, ws: u32) {
let mut stream =
TcpStream::connect_timeout(&addr, Duration::from_secs(5)).unwrap();
stream
.set_read_timeout(Some(Duration::from_secs(5)))
.unwrap();
write_handshake_rank(&mut stream, rank, ws, &TEST_SALT).unwrap();
read_handshake_ack(&mut stream, &TEST_SALT).unwrap();
thread::sleep(Duration::from_millis(50));
}
let r0 = thread::spawn(move || raw_rank_handshake(addr, 0, world_size as u32));
let r1 = thread::spawn(move || raw_rank_handshake(addr, 1, world_size as u32));
r0.join().unwrap();
r1.join().unwrap();
coord_handle.join().unwrap().expect("coord drives clean");
}
#[test]
fn handshake_rejects_wrong_salt_on_worker_side() {
let world_size = 2;
let bad_salt: SessionSalt = [0u8; SESSION_SALT_BYTES];
let dummy_upstream = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 1);
let (relay_addr, relay_rx) =
spawn_relay(ChannelKind::Control, dummy_upstream, world_size, TEST_SALT);
let rank = thread::spawn(move || {
let mut s = TcpStream::connect_timeout(&relay_addr, Duration::from_secs(5)).unwrap();
let _ = write_handshake_rank(&mut s, 0, world_size as u32, &bad_salt);
let _ = read_handshake_ack(&mut s, &bad_salt);
});
let err = match relay_rx.recv().unwrap() {
Ok(_) => panic!("expected relay to reject wrong-salt handshake"),
Err(e) => e,
};
assert!(
err.to_string().contains("HMAC verification failed"),
"expected HMAC failure, got: {err}"
);
let _ = rank.join();
}
#[test]
#[ignore = "requires CUDA + NCCL — run via fdl cuda-test-nccl"]
fn end_to_end_sync_nccl_smoke() {
}
#[test]
#[ignore = "requires CUDA + NCCL + 2+ GPUs — run via fdl @cluster-test cuda-test-nccl"]
fn end_to_end_sync_nccl_via_coord_smoke() {
use crate::distributed::testing::discover_test_cluster;
use crate::distributed::nccl::NcclComms;
let cluster = match discover_test_cluster() {
Some(c) => c,
None => {
eprintln!(
"end_to_end_sync_nccl_via_coord_smoke: no cluster topology \
available (set FLODL_TESTING_CLUSTER_JSON via \
`fdl @cluster-test` or run on a CUDA host)"
);
return;
}
};
let total_ranks: usize = cluster.workers.iter().map(|h| h.ranks.len()).sum();
if total_ranks < 2 {
eprintln!(
"end_to_end_sync_nccl_via_coord_smoke: NCCL needs 2+ ranks \
(have {total_ranks}); skipping"
);
return;
}
let world_size = total_ranks;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let (coord_listener, coord_port) = CCoord::bind(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
)
.expect("coord bind succeeds");
let coord_real_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), coord_port);
let (coord_addr, _ctrl_relay_rx) =
spawn_relay(ChannelKind::Control, coord_real_addr, world_size, TEST_SALT);
let dead_for_coord = Arc::clone(&dead_ranks);
let total_samples = 16usize;
let batch_size = 4usize;
let config_for_coord = move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Nccl,
world_size,
crate::distributed::ddp::ElChe::new(world_size, 1),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.total_samples(total_samples)
.batch_size(batch_size)
.num_epochs(1)
};
let coord_thread = thread::spawn(move || -> Result<CCoord> {
CCoord::start_from_listener(
coord_listener,
[0u8; crate::distributed::wire::SESSION_SALT_BYTES],
config_for_coord(),
)
});
let devices: Vec<Device> = (0..world_size as u8)
.map(Device::CUDA)
.collect();
let group = NcclComms::new(&devices).expect("NcclComms::new succeeds");
let rank_comms = group.split().expect("split succeeds");
let ref_model = Linear::on_device(4, 2, Device::CPU).unwrap();
let initial_params: Vec<Tensor> = ref_model
.parameters()
.iter()
.map(|p| p.variable.data())
.collect();
let initial_buffers: Vec<Tensor> = ref_model
.buffers()
.iter()
.map(|b| b.get())
.collect();
drop(ref_model);
let salt = [0u8; crate::distributed::wire::SESSION_SALT_BYTES];
let mut worker_handles: Vec<thread::JoinHandle<Result<()>>> = Vec::new();
for (rank_id, comm) in rank_comms.into_iter().enumerate() {
let initial_params = initial_params.clone();
let initial_buffers = initial_buffers.clone();
let device = Device::CUDA(rank_id as u8);
worker_handles.push(thread::spawn(move || -> Result<()> {
let config = WorkerConfig {
rank: rank_id,
world_size,
device,
initial_params,
initial_buffers,
total_samples,
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,
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 dataset: Arc<dyn crate::data::BatchDataSet> =
Arc::new(TestDataset { n: total_samples });
let worker = ClusterWorker::connect_and_build(
coord_addr,
None, rank_id as u32,
salt,
config,
move |d| Linear::on_device(4, 2, d),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
dataset,
Some(comm),
RankCallbacks::default(),
)?;
worker.run_until_shutdown(mse_train).map(|_| ())
}));
}
let mut coord = coord_thread
.join()
.expect("coord thread join")
.expect("start_from_listener succeeds");
coord.dispatch_epoch(0).expect("dispatch_epoch(0) succeeds");
let start = Instant::now();
while coord.avg_count() == 0 {
if start.elapsed() > Duration::from_secs(30) {
panic!(
"end_to_end_sync_nccl_via_coord_smoke: avg_count never \
advanced (no NCCL AllReduce observed within 30s)"
);
}
coord.tick().expect("tick");
thread::sleep(Duration::from_millis(20));
}
assert!(coord.avg_count() >= 1, "at least one NCCL averaging cycle");
coord.shutdown_workers().expect("shutdown_workers");
coord.shutdown().expect("coord shutdown");
for h in worker_handles {
h.join().expect("worker thread join").expect("worker exits clean");
}
}
#[test]
#[ignore = "requires CUDA + NCCL + 2+ GPUs — run via fdl @cluster-test cuda-test-nccl"]
fn end_to_end_cadence_nccl_via_coord_smoke() {
use crate::distributed::testing::discover_test_cluster;
use crate::distributed::nccl::NcclComms;
let cluster = match discover_test_cluster() {
Some(c) => c,
None => {
eprintln!(
"end_to_end_cadence_nccl_via_coord_smoke: no cluster topology \
available (set FLODL_TESTING_CLUSTER_JSON via \
`fdl @cluster-test` or run on a CUDA host)"
);
return;
}
};
let total_ranks: usize = cluster.workers.iter().map(|h| h.ranks.len()).sum();
if total_ranks < 2 {
eprintln!(
"end_to_end_cadence_nccl_via_coord_smoke: NCCL needs 2+ ranks \
(have {total_ranks}); skipping"
);
return;
}
let world_size = total_ranks;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let (coord_listener, coord_port) = CCoord::bind(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
)
.expect("coord bind succeeds");
let coord_real_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), coord_port);
let (coord_addr, _ctrl_relay_rx) =
spawn_relay(ChannelKind::Control, coord_real_addr, world_size, TEST_SALT);
let dead_for_coord = Arc::clone(&dead_ranks);
let total_samples = 32usize;
let batch_size = 4usize;
let elche_anchor = 2usize;
let num_epochs = 4usize;
let config_for_coord = move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Cadence,
AverageBackend::Nccl,
world_size,
crate::distributed::ddp::ElChe::new(world_size, elche_anchor),
)
.no_divergence_guard()
.dead_ranks(dead_for_coord)
.total_samples(total_samples)
.batch_size(batch_size)
.num_epochs(num_epochs)
};
let coord_thread = thread::spawn(move || -> Result<CCoord> {
CCoord::start_from_listener(
coord_listener,
[0u8; crate::distributed::wire::SESSION_SALT_BYTES],
config_for_coord(),
)
});
let devices: Vec<Device> = (0..world_size as u8)
.map(Device::CUDA)
.collect();
let group = NcclComms::new(&devices).expect("NcclComms::new succeeds");
let rank_comms = group.split().expect("split succeeds");
let ref_model = Linear::on_device(4, 2, Device::CPU).unwrap();
let initial_params: Vec<Tensor> = ref_model
.parameters()
.iter()
.map(|p| p.variable.data())
.collect();
let initial_buffers: Vec<Tensor> = ref_model
.buffers()
.iter()
.map(|b| b.get())
.collect();
drop(ref_model);
let salt = [0u8; crate::distributed::wire::SESSION_SALT_BYTES];
let mut worker_handles: Vec<thread::JoinHandle<Result<()>>> = Vec::new();
for (rank_id, comm) in rank_comms.into_iter().enumerate() {
let initial_params = initial_params.clone();
let initial_buffers = initial_buffers.clone();
let device = Device::CUDA(rank_id as u8);
worker_handles.push(thread::spawn(move || -> Result<()> {
let config = WorkerConfig {
rank: rank_id,
world_size,
device,
initial_params,
initial_buffers,
total_samples,
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,
seed: 42,
max_grad_norm: None,
vram_pool: false,
easgd_alpha: None,
gamma: 1.0,
bf16_wire: false,
timeline: None,
policy: ApplyPolicy::Cadence,
save_path: None,
coord_liveness_timeout_secs:
crate::distributed::ddp_run::DEFAULT_COORD_LIVENESS_TIMEOUT_SECS,
};
let dataset: Arc<dyn crate::data::BatchDataSet> =
Arc::new(TestDataset { n: total_samples });
let worker = ClusterWorker::connect_and_build(
coord_addr,
None,
rank_id as u32,
salt,
config,
move |d| Linear::on_device(4, 2, d),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
dataset,
Some(comm),
RankCallbacks::default(),
)?;
worker.run_until_shutdown(mse_train).map(|_| ())
}));
}
let mut coord = coord_thread
.join()
.expect("coord thread join")
.expect("start_from_listener succeeds");
coord.dispatch_epoch(0).expect("dispatch_epoch(0) succeeds");
let start = Instant::now();
while coord.avg_count() < 2 {
if start.elapsed() > Duration::from_secs(60) {
panic!(
"end_to_end_cadence_nccl_via_coord_smoke: avg_count={} \
never reached 2 within 60s",
coord.avg_count(),
);
}
coord.tick().expect("tick");
thread::sleep(Duration::from_millis(20));
}
assert!(
coord.avg_count() >= 2,
"Cadence drives multiple AllReduce cycles: avg_count={}",
coord.avg_count(),
);
let coord_anchor = coord.el_che().anchor();
assert_eq!(
coord_anchor, elche_anchor,
"NoGuard: anchor stable at initial value ({elche_anchor}), got {coord_anchor}",
);
coord.shutdown_workers().expect("shutdown_workers");
coord.shutdown().expect("coord shutdown");
for h in worker_handles {
h.join().expect("worker thread join").expect("worker exits clean");
}
}
#[test]
#[ignore = "heavy in-process integration smoke; timing-flaky under parallel load — run explicitly"]
fn end_to_end_cadence_cpu_via_coord_smoke() {
let world_size = 2;
let total_samples = 32usize;
let batch_size = 4usize;
let elche_anchor = 1usize;
let num_epochs = 16usize;
let target_cycles = 12u64;
let controller = ClusterController::start(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
world_size,
TEST_SALT,
)
.expect("controller starts");
let controller_addr =
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), controller.port());
let (reduce_addr, _relay_rx) =
spawn_relay(ChannelKind::Data, controller_addr, world_size, TEST_SALT);
let (coord_listener, coord_port) = CCoord::bind(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
)
.expect("coord bind succeeds");
let coord_real_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), coord_port);
let (coord_addr, _ctrl_relay_rx) =
spawn_relay(ChannelKind::Control, coord_real_addr, world_size, TEST_SALT);
let config_for_coord = move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Cadence,
AverageBackend::Cpu,
world_size,
ElChe::new(world_size, elche_anchor).with_max_anchor(2),
)
.total_samples(total_samples)
.batch_size(batch_size)
.num_epochs(num_epochs)
};
let coord_thread = thread::spawn(move || -> Result<CCoord> {
CCoord::start_from_listener(coord_listener, TEST_SALT, config_for_coord())
});
let ref_model = Linear::on_device(4, 2, Device::CPU).unwrap();
let initial_params: Vec<Tensor> = ref_model
.parameters()
.iter()
.map(|p| p.variable.data())
.collect();
let initial_buffers: Vec<Tensor> = ref_model
.buffers()
.iter()
.map(|b| b.get())
.collect();
drop(ref_model);
let mut worker_handles: Vec<thread::JoinHandle<Result<()>>> = Vec::new();
for rank_id in 0..world_size {
let initial_params = initial_params.clone();
let initial_buffers = initial_buffers.clone();
worker_handles.push(thread::spawn(move || -> Result<()> {
let cpu_client = crate::distributed::cpu_reduce::CpuReduceClient::connect(
reduce_addr,
rank_id as u32,
world_size as u32,
TEST_SALT,
)?;
let config = WorkerConfig {
rank: rank_id,
world_size,
device: Device::CPU,
initial_params,
initial_buffers,
total_samples,
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,
seed: 42,
max_grad_norm: None,
vram_pool: false,
easgd_alpha: None,
gamma: 1.0,
bf16_wire: false,
timeline: None,
policy: ApplyPolicy::Cadence,
save_path: None,
coord_liveness_timeout_secs:
crate::distributed::ddp_run::DEFAULT_COORD_LIVENESS_TIMEOUT_SECS,
};
let dataset: Arc<dyn crate::data::BatchDataSet> =
Arc::new(TestDataset { n: total_samples });
let worker = ClusterWorker::connect_and_build(
coord_addr,
Some(cpu_client),
rank_id as u32,
TEST_SALT,
config,
move |d| Linear::on_device(4, 2, d),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
dataset,
None, RankCallbacks::default(),
)?;
worker.run_until_shutdown(mse_train).map(|_| ())
}));
}
let mut coord = coord_thread
.join()
.expect("coord thread join")
.expect("start_from_listener succeeds");
coord.dispatch_epoch(0).expect("dispatch_epoch(0) succeeds");
let start = Instant::now();
while coord.avg_count() < target_cycles {
if start.elapsed() > Duration::from_secs(30) {
panic!(
"end_to_end_cadence_cpu_via_coord_smoke: avg_count={} \
never reached {target_cycles} within 30s — CPU averaging \
stalled (re-arm wedge or window-blowup regression?)",
coord.avg_count(),
);
}
coord.tick().expect("tick");
thread::sleep(Duration::from_millis(5));
}
assert!(
coord.avg_count() >= target_cycles,
"CPU Cadence re-arms every window: avg_count={} (>= {target_cycles})",
coord.avg_count(),
);
let final_anchor = coord.el_che().anchor();
assert!(
final_anchor <= 2,
"window honored max_anchor cap: anchor={final_anchor}",
);
coord.shutdown_workers().expect("shutdown_workers");
coord.shutdown().expect("coord shutdown");
controller.shutdown().expect("controller shutdown");
for h in worker_handles {
h.join().expect("worker thread join").expect("worker exits clean");
}
}
use crate::distributed::cluster_coordinator::ClusterCoordinator as CCoord;
use crate::distributed::controller::ClusterController;
use crate::distributed::relay::agent::{ChannelKind, RelayChannel};
use crate::nn::Linear;
fn spawn_relay(
kind: ChannelKind,
upstream_addr: SocketAddr,
world_size: usize,
salt: SessionSalt,
) -> (SocketAddr, std::sync::mpsc::Receiver<Result<RelayChannel>>) {
let (listener, relay_port) =
RelayChannel::bind(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0)).unwrap();
let loopback = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), relay_port);
let ranks: Vec<u32> = (0..world_size as u32).collect();
let (tx, rx) = std::sync::mpsc::channel();
thread::spawn(move || {
let started = RelayChannel::start(
listener,
kind,
upstream_addr,
"test-host".into(),
ranks,
world_size,
salt,
);
let _ = tx.send(started);
});
(loopback, rx)
}
struct TestDataset {
n: usize,
}
impl crate::data::BatchDataSet for TestDataset {
fn len(&self) -> usize {
self.n
}
fn get_batch(&self, indices: &[usize]) -> Result<Vec<Tensor>> {
let n = indices.len() as i64;
let mut x_vals: Vec<f32> = Vec::with_capacity(indices.len() * 4);
let mut y_vals: Vec<f32> = Vec::with_capacity(indices.len() * 2);
for &idx in indices {
let f = idx as f32;
x_vals.extend_from_slice(&[
f / 10.0,
(f + 1.0) / 10.0,
(f + 2.0) / 10.0,
(f + 3.0) / 10.0,
]);
y_vals.extend_from_slice(&[f * 0.1, f * 0.2]);
}
Ok(vec![
Tensor::from_f32(&x_vals, &[n, 4], Device::CPU)?,
Tensor::from_f32(&y_vals, &[n, 2], Device::CPU)?,
])
}
}
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()
}
struct RecordingGuard {
captured: Arc<std::sync::Mutex<Vec<Vec<f64>>>>,
}
impl crate::distributed::ddp_run::convergence::ConvergenceGuard for RecordingGuard {
fn clone_box(
&self,
) -> Box<dyn crate::distributed::ddp_run::convergence::ConvergenceGuard> {
Box::new(RecordingGuard {
captured: self.captured.clone(),
})
}
fn report(
&mut self,
report: &crate::distributed::ddp_run::convergence::DivergenceReport,
_k_used: usize,
_k_max: usize,
) -> crate::distributed::ddp_run::convergence::ConvergenceAction {
self.captured.lock().unwrap().push(report.deltas.clone());
crate::distributed::ddp_run::convergence::ConvergenceAction::Stable
}
}
#[test]
fn end_to_end_sync_cpu_smoke() {
let world_size = 2usize;
let total_samples = 8usize;
let batch_size = 4usize;
let dead_ranks = crate::distributed::controller::DeadRanks::new(world_size);
let controller = ClusterController::start_with_dead_ranks(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
world_size,
TEST_SALT,
Arc::clone(&dead_ranks),
None,
None,
)
.expect("ClusterController::start_with_dead_ranks succeeds");
let controller_addr =
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), controller.port());
let (data_addr, _relay_rx) =
spawn_relay(ChannelKind::Data, controller_addr, world_size, TEST_SALT);
let (coord_listener, coord_port) = CCoord::bind(
SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0),
)
.expect("coord bind succeeds");
let coord_real_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), coord_port);
let (coord_addr, _ctrl_relay_rx) =
spawn_relay(ChannelKind::Control, coord_real_addr, world_size, TEST_SALT);
let captured_deltas: Arc<std::sync::Mutex<Vec<Vec<f64>>>> =
Arc::new(std::sync::Mutex::new(Vec::new()));
let captured_for_coord = Arc::clone(&captured_deltas);
let dead_ranks_for_coord = Arc::clone(&dead_ranks);
let config_for_coord = move || {
ClusterCoordinatorConfig::new(
ApplyPolicy::Sync,
AverageBackend::Cpu,
world_size,
crate::distributed::ddp::ElChe::new(world_size, 1),
)
.with_convergence_guard(Box::new(RecordingGuard {
captured: captured_for_coord,
}))
.dead_ranks(dead_ranks_for_coord)
.total_samples(total_samples)
.batch_size(batch_size)
.num_epochs(1)
};
let coord_thread = thread::spawn(move || -> Result<CCoord> {
CCoord::start_from_listener(coord_listener, TEST_SALT, config_for_coord())
});
let ref_model = Linear::on_device(4, 2, Device::CPU).unwrap();
let initial_params: Vec<Tensor> = ref_model
.parameters()
.iter()
.map(|p| p.variable.data())
.collect();
let initial_buffers: Vec<Tensor> = ref_model
.buffers()
.iter()
.map(|b| b.get())
.collect();
drop(ref_model);
let salt = TEST_SALT;
let mut worker_handles: Vec<thread::JoinHandle<Result<()>>> = Vec::new();
for rank_id in 0..world_size {
let initial_params = initial_params.clone();
let initial_buffers = initial_buffers.clone();
worker_handles.push(thread::spawn(move || -> Result<()> {
let config = WorkerConfig {
rank: rank_id,
world_size,
device: Device::CPU,
initial_params,
initial_buffers,
total_samples,
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,
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 dataset: Arc<dyn crate::data::BatchDataSet> =
Arc::new(TestDataset { n: total_samples });
let cpu_client = crate::distributed::cpu_reduce::CpuReduceClient::connect(
data_addr,
rank_id as u32,
world_size as u32,
salt,
)?;
let worker = ClusterWorker::connect_and_build(
coord_addr,
Some(cpu_client),
rank_id as u32,
salt,
config,
|d| Linear::on_device(4, 2, d),
|params| crate::nn::SGD::new(params, 0.01, 0.0),
dataset,
None, RankCallbacks::default(),
)?;
worker.run_until_shutdown(mse_train).map(|_| ())
}));
}
let mut coord = coord_thread
.join()
.expect("coord thread join")
.expect("start_from_listener succeeds");
coord.dispatch_epoch(0).expect("dispatch_epoch(0) succeeds");
let start = Instant::now();
let timed_out = loop {
if coord.avg_count() > 0 {
break false;
}
if start.elapsed() > Duration::from_secs(60) {
break true;
}
coord.tick().expect("tick");
thread::sleep(Duration::from_millis(10));
};
if timed_out {
coord.shutdown_workers().ok();
coord.shutdown().ok();
for h in worker_handles {
let _ = h.join();
}
controller.shutdown().ok();
panic!(
"end_to_end_sync_cpu_smoke: avg_count never advanced \
(no averaging cycle observed within 60s — likely \
parallel-load CPU starvation, see test comment)"
);
}
assert!(coord.avg_count() >= 1, "at least one averaging cycle");
let cycles = captured_deltas.lock().unwrap().clone();
let no_cycles = cycles.is_empty();
let (first_len, has_positive, first_dump) = if no_cycles {
(0, false, Vec::new())
} else {
let f = &cycles[0];
(
f.len(),
f.iter().any(|d| d.is_finite() && *d > 0.0),
f.clone(),
)
};
let len_ok = first_len == world_size;
let div_check_passed = !no_cycles && len_ok && has_positive;
coord.shutdown_workers().ok();
coord.shutdown().ok();
let worker_results: Vec<(usize, std::thread::Result<Result<()>>)> =
worker_handles
.into_iter()
.enumerate()
.map(|(rank_id, h)| (rank_id, h.join()))
.collect();
controller.shutdown().ok();
assert!(
div_check_passed,
"smoke divergence check failed: cycles_seen={} first_len={} \
(expected {}) any_positive={} first_deltas={:?}",
cycles.len(),
first_len,
world_size,
has_positive,
first_dump,
);
for (rank_id, r) in worker_results {
let r = r.expect("worker thread join");
r.unwrap_or_else(|e| {
panic!("worker rank {rank_id} run_until_shutdown: {e}");
});
}
}
fn inbound_test_rig(
world_size: usize,
coord_liveness_timeout_secs: u64,
) -> (
std::net::TcpStream, // coord side
std::thread::JoinHandle<()>, // inbound
std::sync::mpsc::Receiver<ControlMsg>, // control_rx
Arc<crate::distributed::controller::DeadRanks>, // ledger
) {
use std::net::{TcpListener, TcpStream};
let listener =
TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).expect("bind");
let addr = listener.local_addr().expect("addr");
let coord_side = TcpStream::connect(addr).expect("connect");
let (mut worker_side, _) = listener.accept().expect("accept");
worker_side
.set_read_timeout(Some(std::time::Duration::from_millis(100)))
.expect("set_read_timeout");
let shutdown = Arc::new(std::sync::atomic::AtomicBool::new(false));
let (control_tx, control_rx) = std::sync::mpsc::channel();
let (timing_tx, _timing_rx) = std::sync::mpsc::channel();
let dead = crate::distributed::controller::DeadRanks::new(world_size);
let mailbox = Arc::new(std::sync::Mutex::new(None));
let dead_for_loop = Arc::clone(&dead);
let handle = std::thread::spawn(move || {
inbound_loop(
0,
&mut worker_side,
&TEST_SALT,
&shutdown,
&control_tx,
&dead_for_loop,
&mailbox,
&timing_tx,
coord_liveness_timeout_secs,
);
});
(coord_side, handle, control_rx, dead)
}
#[test]
fn inbound_eof_without_shutdown_poisons_peer_ledger() {
let (coord_side, handle, control_rx, dead) = inbound_test_rig(3, 30);
drop(coord_side);
handle.join().expect("inbound join");
assert!(!dead.is_dead(0), "own rank must never be poisoned");
assert!(dead.is_dead(1) && dead.is_dead(2),
"all peers must be declared dead so the NCCL watchdog can \
abort an in-flight collective");
assert!(matches!(control_rx.recv(), Ok(ControlMsg::Shutdown)));
}
#[test]
fn inbound_eof_after_clean_shutdown_leaves_ledger_alone() {
let (mut coord_side, handle, control_rx, dead) = inbound_test_rig(3, 30);
let frame = crate::distributed::wire::ControlFrame::encode(
&TEST_SALT,
crate::distributed::wire::MsgKind::Control,
&crate::distributed::wire::ControlMsgWire::Shutdown,
)
.expect("encode shutdown");
let mut bytes = Vec::new();
frame.write_to(&mut bytes).expect("serialize frame");
crate::distributed::relay::mux::write_len_framed(&mut coord_side, &bytes)
.expect("write len-framed");
drop(coord_side);
handle.join().expect("inbound join");
assert_eq!(
dead.dead_count(),
0,
"clean Shutdown-then-EOF is the normal teardown — poisoning \
here would abort a final coherent reduce mid-flight"
);
assert!(matches!(control_rx.recv(), Ok(ControlMsg::Shutdown)));
}
#[test]
fn inbound_wedged_open_coord_trips_liveness_deadline() {
let (coord_side, handle, control_rx, dead) = inbound_test_rig(3, 1);
handle.join().expect("inbound join");
assert!(!dead.is_dead(0), "own rank must never be poisoned");
assert!(
dead.is_dead(1) && dead.is_dead(2),
"a coord silent past the liveness deadline must poison peers so \
the NCCL watchdog can abort an in-flight collective"
);
assert!(
matches!(control_rx.recv(), Ok(ControlMsg::Shutdown)),
"the recv-parked inner must be woken with Shutdown"
);
drop(coord_side);
}