use super::{EpochMetrics, MetricsMsg};
pub(crate) fn equal_sizes(world_size: usize, total: usize) -> Vec<usize> {
let base = total / world_size;
let remainder = total % world_size;
(0..world_size)
.map(|r| base + if r < remainder { 1 } else { 0 })
.collect()
}
pub(crate) fn throughput_sizes(
el_che: &crate::distributed::ddp::ElChe,
total: usize,
) -> Vec<usize> {
let ms = el_che.ms_per_batch();
let throughputs: Vec<f64> = ms.iter().map(|&m| 1.0 / m.max(0.001)).collect();
let total_tp: f64 = throughputs.iter().sum();
if total_tp <= 0.0 {
return equal_sizes(ms.len(), total);
}
let mut sizes: Vec<usize> = throughputs.iter()
.map(|t| ((t / total_tp) * total as f64).floor() as usize)
.collect();
let assigned: usize = sizes.iter().sum();
let mut remaining = total.saturating_sub(assigned);
if remaining > 0 {
let mut rank_order: Vec<usize> = (0..ms.len()).collect();
rank_order.sort_by(|&a, &b| throughputs[b].partial_cmp(&throughputs[a]).unwrap_or(std::cmp::Ordering::Equal));
for &rank in &rank_order {
if remaining == 0 { break; }
sizes[rank] += 1;
remaining -= 1;
}
}
sizes
}
pub(crate) fn ratio_to_sizes(ratios: &[f64], total: usize) -> Vec<usize> {
let sum: f64 = ratios.iter().sum();
let norm: Vec<f64> = if sum > 0.0 {
ratios.iter().map(|r| r / sum).collect()
} else {
vec![1.0 / ratios.len() as f64; ratios.len()]
};
let mut sizes: Vec<usize> = norm.iter()
.map(|r| (r * total as f64).floor() as usize)
.collect();
let assigned: usize = sizes.iter().sum();
let mut remaining = total.saturating_sub(assigned);
if remaining > 0 {
let mut rank_order: Vec<usize> = (0..ratios.len()).collect();
rank_order.sort_by(|&a, &b| norm[b].partial_cmp(&norm[a]).unwrap_or(std::cmp::Ordering::Equal));
for &rank in &rank_order {
if remaining == 0 { break; }
sizes[rank] += 1;
remaining -= 1;
}
}
sizes
}
pub(crate) fn aggregate_epoch_metrics(
epoch: usize,
msgs: &[MetricsMsg],
device_indices: &[u8],
bc_share: &[f64],
) -> EpochMetrics {
let world_size = device_indices.len();
let mut rank_batches: Vec<usize> = vec![0; world_size];
let mut rank_samples: Vec<usize> = vec![0; world_size];
let mut rank_loss_sum: Vec<f64> = vec![0.0; world_size];
let mut rank_time_ms: Vec<f64> = vec![0.0; world_size];
let mut rank_share_complete_ms: Vec<f64> = vec![0.0; world_size];
let mut rank_compute_only_ms: Vec<f64> = vec![0.0; world_size];
let mut rank_data_starve_ms: Vec<f64> = vec![0.0; world_size];
let mut rank_scalars: Vec<std::collections::HashMap<String, (f64, usize)>> =
(0..world_size).map(|_| std::collections::HashMap::new()).collect();
for m in msgs {
let r = m.rank.min(world_size - 1);
rank_batches[r] += m.batches_processed;
rank_samples[r] += m.samples_processed;
rank_loss_sum[r] += m.avg_loss * m.batches_processed as f64;
rank_time_ms[r] += m.epoch_ms;
rank_share_complete_ms[r] += m.share_complete_ms;
rank_compute_only_ms[r] += m.compute_only_ms;
rank_data_starve_ms[r] += m.data_starve_ms;
for (k, (sum, count)) in &m.scalars {
let entry = rank_scalars[r].entry(k.clone()).or_insert((0.0, 0));
entry.0 += sum;
entry.1 += count;
}
}
let total_batches: usize = rank_batches.iter().sum();
let avg_loss = if total_batches > 0 {
rank_loss_sum.iter().sum::<f64>() / total_batches as f64
} else {
0.0
};
let per_rank_loss: Vec<Option<f64>> = (0..world_size)
.map(|r| {
if rank_batches[r] > 0 {
Some(rank_loss_sum[r] / rank_batches[r] as f64)
} else {
None
}
})
.collect();
let epoch_ms = rank_time_ms.iter().copied().fold(0.0_f64, f64::max);
let per_rank: Vec<std::collections::HashMap<String, f64>> = rank_scalars
.iter()
.map(|scalars| {
scalars
.iter()
.map(|(k, (sum, count))| {
(k.clone(), if *count > 0 { sum / *count as f64 } else { 0.0 })
})
.collect()
})
.collect();
let mut scalars: std::collections::HashMap<String, f64> = std::collections::HashMap::new();
let mut weights: std::collections::HashMap<String, f64> = std::collections::HashMap::new();
for (r, rank_sc) in rank_scalars.iter().enumerate() {
let w = rank_batches[r] as f64;
for (k, (sum, count)) in rank_sc {
if *count > 0 {
let mean = sum / *count as f64;
*scalars.entry(k.clone()).or_default() += mean * w;
*weights.entry(k.clone()).or_default() += w;
}
}
}
for (k, v) in &mut scalars {
if let Some(w) = weights.get(k) {
if *w > 0.0 {
*v /= *w;
}
}
}
let per_rank_throughput: Vec<f64> = (0..world_size).map(|r| {
let denom = if rank_share_complete_ms[r] > 0.0 {
rank_share_complete_ms[r]
} else {
rank_time_ms[r]
};
if denom > 0.0 { rank_samples[r] as f64 / denom } else { 0.0 }
}).collect();
let per_rank_batch_share: Vec<f64> = if bc_share.len() == world_size {
let sum: f64 = bc_share.iter().sum();
if sum > 0.0 {
bc_share.to_vec()
} else {
vec![1.0 / world_size as f64; world_size]
}
} else {
vec![1.0 / world_size as f64; world_size]
};
EpochMetrics {
epoch, scalars, per_rank, avg_loss, epoch_ms,
per_rank_loss,
per_rank_samples: rank_samples,
per_rank_throughput, per_rank_batch_share,
per_rank_share_complete_ms: rank_share_complete_ms,
per_rank_compute_only_ms: rank_compute_only_ms,
per_rank_data_starve_ms: rank_data_starve_ms,
device_indices: device_indices.to_vec(),
}
}
#[cfg(test)]
mod tests {
use super::aggregate_epoch_metrics;
use super::super::MetricsMsg;
use std::collections::HashMap;
#[test]
fn test_aggregate_epoch_metrics() {
let mut scalars_r0 = HashMap::new();
scalars_r0.insert("loss".to_string(), (3.0, 3_usize)); scalars_r0.insert("acc".to_string(), (1.8, 3));
let mut scalars_r1 = HashMap::new();
scalars_r1.insert("loss".to_string(), (4.0, 2_usize)); scalars_r1.insert("acc".to_string(), (0.8, 2));
let msgs = vec![
MetricsMsg {
rank: 0, epoch: 0, avg_loss: 0.5, batches_processed: 60,
epoch_ms: 1000.0, share_complete_ms: 1000.0, compute_only_ms: 1000.0, data_starve_ms: 0.0, samples_processed: 1920, scalars: scalars_r0,
},
MetricsMsg {
rank: 1, epoch: 0, avg_loss: 0.7, batches_processed: 40,
epoch_ms: 1200.0, share_complete_ms: 1200.0, compute_only_ms: 1200.0, data_starve_ms: 0.0, samples_processed: 1280, scalars: scalars_r1,
},
];
let dev_indices = vec![0_u8, 1];
let bc_share = vec![0.6_f64, 0.4];
let m = aggregate_epoch_metrics(0, &msgs, &dev_indices, &bc_share);
assert_eq!(m.epoch, 0);
assert!((m.avg_loss - 0.58).abs() < 1e-9);
assert_eq!(m.epoch_ms, 1200.0);
assert!((m.scalars["loss"] - 1.4).abs() < 1e-9);
assert!((m.scalars["acc"] - 0.52).abs() < 1e-9);
assert_eq!(m.per_rank.len(), 2);
assert!((m.per_rank[0]["loss"] - 1.0).abs() < 1e-9);
assert!((m.per_rank[1]["loss"] - 2.0).abs() < 1e-9);
assert!((m.per_rank_loss[0].unwrap() - 0.5).abs() < 1e-9);
assert!((m.per_rank_loss[1].unwrap() - 0.7).abs() < 1e-9);
assert_eq!(m.per_rank_samples, vec![1920, 1280]);
assert!((m.per_rank_throughput[0] - 1.92).abs() < 1e-9);
assert!((m.per_rank_throughput[1] - 1280.0 / 1200.0).abs() < 1e-9);
assert!((m.per_rank_batch_share[0] - 0.6).abs() < 1e-9);
assert!((m.per_rank_batch_share[1] - 0.4).abs() < 1e-9);
assert_eq!(m.device_indices, vec![0, 1]);
}
#[test]
fn test_aggregate_epoch_metrics_progressive() {
let msgs = vec![
MetricsMsg {
rank: 0, epoch: 0, avg_loss: 0.5, batches_processed: 20,
epoch_ms: 300.0, share_complete_ms: 300.0, compute_only_ms: 300.0, data_starve_ms: 0.0, samples_processed: 640,
scalars: [("loss".to_string(), (2.0, 2_usize))].into(),
},
MetricsMsg {
rank: 0, epoch: 0, avg_loss: 0.4, batches_processed: 20,
epoch_ms: 300.0, share_complete_ms: 300.0, compute_only_ms: 300.0, data_starve_ms: 0.0, samples_processed: 640,
scalars: [("loss".to_string(), (1.6, 2_usize))].into(),
},
MetricsMsg {
rank: 0, epoch: 0, avg_loss: 0.6, batches_processed: 20,
epoch_ms: 300.0, share_complete_ms: 300.0, compute_only_ms: 300.0, data_starve_ms: 0.0, samples_processed: 640,
scalars: [("loss".to_string(), (1.8, 2_usize))].into(),
},
MetricsMsg {
rank: 1, epoch: 0, avg_loss: 0.7, batches_processed: 20,
epoch_ms: 500.0, share_complete_ms: 500.0, compute_only_ms: 500.0, data_starve_ms: 0.0, samples_processed: 640,
scalars: [("loss".to_string(), (2.8, 2_usize))].into(),
},
MetricsMsg {
rank: 1, epoch: 0, avg_loss: 0.8, batches_processed: 20,
epoch_ms: 500.0, share_complete_ms: 500.0, compute_only_ms: 500.0, data_starve_ms: 0.0, samples_processed: 640,
scalars: [("loss".to_string(), (3.2, 2_usize))].into(),
},
];
let dev_indices = vec![0_u8, 1];
let bc_share = vec![0.6_f64, 0.4];
let m = aggregate_epoch_metrics(0, &msgs, &dev_indices, &bc_share);
assert_eq!(m.per_rank_throughput.len(), 2, "should have world_size entries");
assert_eq!(m.per_rank_batch_share.len(), 2);
assert_eq!(m.per_rank.len(), 2);
assert_eq!(m.device_indices, vec![0, 1]);
assert!((m.per_rank_throughput[0] - 1920.0 / 900.0).abs() < 1e-6);
assert!((m.per_rank_throughput[1] - 1280.0 / 1000.0).abs() < 1e-6);
assert!((m.per_rank_batch_share[0] - 0.6).abs() < 1e-9);
assert!((m.per_rank_batch_share[1] - 0.4).abs() < 1e-9);
assert_eq!(m.epoch_ms, 1000.0);
assert!((m.per_rank[0]["loss"] - 0.9).abs() < 1e-9);
assert!((m.per_rank[1]["loss"] - 1.5).abs() < 1e-9);
assert!((m.scalars["loss"] - 1.14).abs() < 1e-9);
}
}