use super::*;
pub(super) fn param_bridge_loop(
rank: u64,
param_rx: mpsc::Receiver<crate::distributed::ddp_run::ParamSnapshot>,
cpu_client: Option<crate::distributed::cpu_reduce::CpuReduceClient>,
control_tx: mpsc::Sender<ControlMsg>,
timing_tx: mpsc::Sender<TimingMsg>,
gamma: f64,
) {
use crate::distributed::ddp_run::{AveragedParams, ParamSnapshot};
let Some(mut client) = cpu_client else {
while param_rx.recv().is_ok() {}
return;
};
let inject_shutdown = || {
let _ = control_tx.send(ControlMsg::Shutdown);
};
let mut version: u64 = 0;
let mut pre_scratch: Option<Vec<Tensor>> = None;
while let Ok(snapshot) = param_rx.recv() {
let ParamSnapshot {
rank: snap_rank,
params,
buffers,
batch_count: n_i,
} = snapshot;
debug_assert_eq!(
snap_rank as u64, rank,
"param bridge: snapshot.rank mismatch with bridge rank"
);
if pre_scratch.is_none() {
let allocated: Result<Vec<Tensor>> = params
.iter()
.map(|t| {
Tensor::zeros(
&t.shape(),
crate::tensor::TensorOptions {
dtype: crate::tensor::DType::Float32,
device: crate::tensor::Device::CPU,
},
)
})
.collect();
match allocated {
Ok(s) => pre_scratch = Some(s),
Err(e) => {
eprintln!(
"cluster_worker: param bridge r{rank} scratch alloc: {e}"
);
inject_shutdown();
return;
}
}
}
let scratch = pre_scratch.as_ref().expect("scratch just allocated");
let mut copy_failed = false;
for (dst, src) in scratch.iter().zip(params.iter()) {
if let Err(e) = dst.copy_(src, false) {
eprintln!(
"cluster_worker: param bridge r{rank} pre_scratch copy_: {e}"
);
copy_failed = true;
break;
}
}
if copy_failed {
inject_shutdown();
return;
}
let _ = timing_tx.send(TimingMsg::SnapshotReady {
rank: rank as usize,
});
let world = client.world_size() as usize;
let mut counts = vec![0.0f64; world];
if (rank as usize) < world {
counts[rank as usize] = n_i as f64;
}
if let Err(e) = client.all_reduce_per_rank_f64(&mut counts) {
eprintln!("cluster_worker: param bridge r{rank} count gather: {e}");
inject_shutdown();
return;
}
let total_n: f64 = counts.iter().sum();
let my_w = crate::distributed::realized_work::gamma_mass(n_i as f64, gamma);
let avg_params = if total_n == 0.0 {
params.clone()
} else {
match sumcount_reduce(
&mut client,
¶ms,
my_w,
crate::distributed::cpu_reduce::DECODE_SLOT_PARAMS,
) {
Ok(v) => v,
Err(e) => {
eprintln!(
"cluster_worker: param bridge r{rank} all_reduce params: {e}"
);
inject_shutdown();
return;
}
}
};
let (divergence, post_norm, pre_norm) =
match crate::distributed::divergence::divergence_triple(scratch, &avg_params) {
Ok(t) => t,
Err(e) => {
eprintln!(
"cluster_worker: param bridge r{rank} divergence: {e}"
);
inject_shutdown();
return;
}
};
let f32_buffer_idx: Vec<usize> = buffers
.iter()
.enumerate()
.filter(|(_, b)| b.dtype() == crate::tensor::DType::Float32)
.map(|(i, _)| i)
.collect();
let avg_buffers = if f32_buffer_idx.is_empty() || total_n == 0.0 {
buffers.clone()
} else {
let subset: Vec<Tensor> =
f32_buffer_idx.iter().map(|&i| buffers[i].clone()).collect();
let my_indicator = crate::distributed::realized_work::mover_mass(n_i as f64);
match sumcount_reduce(
&mut client,
&subset,
my_indicator,
crate::distributed::cpu_reduce::DECODE_SLOT_BUFFERS,
) {
Ok(reduced) => {
let mut merged = buffers.clone();
for (k, &i) in f32_buffer_idx.iter().enumerate() {
merged[i] = reduced[k].clone();
}
merged
}
Err(e) => {
eprintln!(
"cluster_worker: param bridge r{rank} all_reduce buffers: {e}"
);
inject_shutdown();
return;
}
}
};
version += 1;
let avg = AveragedParams {
params: avg_params,
buffers: avg_buffers,
version,
};
if control_tx.send(ControlMsg::Update(avg)).is_err() {
return;
}
let _ = timing_tx.send(TimingMsg::SyncAck {
rank: rank as usize,
step_count: 0,
divergence: Some(divergence),
post_norm,
pre_norm,
});
}
client.log_profile_summary();
}
pub(crate) fn sumcount_reduce(
client: &mut crate::distributed::cpu_reduce::CpuReduceClient,
tensors: &[Tensor],
my_weight: f64,
decode_slot: usize,
) -> Result<Vec<Tensor>> {
let refs: Vec<&Tensor> = tensors.iter().collect();
client.arm_pinned_decode(decode_slot);
let (consensus, realized) = client.all_reduce_scaled(
&refs,
my_weight,
crate::distributed::controller::RoundKind::Model,
my_weight,
)?;
if !crate::distributed::realized_work::is_realized(realized) {
crate::debug!(
"cluster_worker: reduce round realized no work (mass 0); keeping local state"
);
return Ok(tensors.to_vec());
}
Ok(consensus)
}