#[derive(Debug)]
pub(crate) struct WindowLedger {
steps: Vec<usize>,
wall_ms: Vec<f64>,
delivered_ms: Vec<f64>,
delivered_batches: Vec<usize>,
first_batch_ms: Vec<f64>,
loss_sum: Vec<f64>,
loss_count: Vec<usize>,
}
impl WindowLedger {
pub(crate) fn new(world_size: usize) -> Self {
WindowLedger {
steps: vec![0; world_size],
wall_ms: vec![0.0; world_size],
delivered_ms: vec![0.0; world_size],
delivered_batches: vec![0; world_size],
first_batch_ms: vec![0.0; world_size],
loss_sum: vec![0.0; world_size],
loss_count: vec![0; world_size],
}
}
pub(crate) fn record_batch(&mut self, rank: usize, batch_ms: f64, data_ms: f64) {
if rank >= self.steps.len() {
return;
}
self.steps[rank] = self.steps[rank].saturating_add(1);
self.wall_ms[rank] += batch_ms;
if self.steps[rank] > 1 {
self.delivered_ms[rank] += batch_ms + data_ms;
self.delivered_batches[rank] += 1;
} else {
self.first_batch_ms[rank] = batch_ms + data_ms;
}
}
pub(crate) fn record_batch_loss(&mut self, rank: usize, loss: f64) {
if rank >= self.loss_sum.len() || !loss.is_finite() {
return;
}
self.loss_sum[rank] += loss;
self.loss_count[rank] += 1;
}
pub(crate) fn mean_loss(&self, rank: usize) -> Option<f64> {
let n = *self.loss_count.get(rank)?;
if n == 0 {
return None;
}
Some(self.loss_sum[rank] / n as f64)
}
pub(crate) fn absorb_callback_cost(&mut self, rank: usize, elapsed_ms: f64) {
if let Some(w) = self.wall_ms.get_mut(rank) {
*w = (*w - elapsed_ms).max(0.0);
}
if let Some(d) = self.delivered_ms.get_mut(rank) {
*d = (*d - elapsed_ms).max(0.0);
}
}
pub(crate) fn steps(&self, rank: usize) -> usize {
self.steps[rank]
}
pub(crate) fn steps_all(&self) -> &[usize] {
&self.steps
}
pub(crate) fn min_steps(&self) -> usize {
self.steps.iter().copied().min().unwrap_or(0)
}
pub(crate) fn max_steps(&self) -> usize {
self.steps.iter().copied().max().unwrap_or(0)
}
pub(crate) fn total_steps(&self) -> usize {
self.steps.iter().sum()
}
pub(crate) fn wall_ms(&self, rank: usize) -> f64 {
self.wall_ms[rank]
}
pub(crate) fn wall_ms_all(&self) -> &[f64] {
&self.wall_ms
}
pub(crate) fn per_batch_wall_ms(&self, rank: usize) -> f64 {
let steps = self.steps.get(rank).copied().unwrap_or(0);
if steps == 0 {
return f64::INFINITY;
}
let wall = self.wall_ms.get(rank).copied().unwrap_or(0.0);
wall / steps as f64
}
pub(crate) fn has_delivered_sample(&self, rank: usize) -> bool {
self.delivered_batches[rank] > 0 && self.delivered_ms[rank] > 0.0
}
pub(crate) fn delivered_ms(&self, rank: usize) -> f64 {
self.delivered_ms[rank]
}
pub(crate) fn delivered_ms_all(&self) -> &[f64] {
&self.delivered_ms
}
pub(crate) fn delivered_batches(&self, rank: usize) -> usize {
self.delivered_batches[rank]
}
pub(crate) fn delivered_batches_all(&self) -> &[usize] {
&self.delivered_batches
}
pub(crate) fn fill_excess_ms(&self, rank: usize) -> f64 {
if !self.has_delivered_sample(rank) {
return 0.0;
}
let marginal = self.delivered_ms[rank] / self.delivered_batches[rank] as f64;
(self.first_batch_ms[rank] - marginal).max(0.0)
}
pub(crate) fn reset_timing(&mut self) {
for a in &mut self.wall_ms {
*a = 0.0;
}
for a in &mut self.delivered_ms {
*a = 0.0;
}
for n in &mut self.delivered_batches {
*n = 0;
}
for f in &mut self.first_batch_ms {
*f = 0.0;
}
for l in &mut self.loss_sum {
*l = 0.0;
}
for n in &mut self.loss_count {
*n = 0;
}
}
pub(crate) fn reset_steps(&mut self) {
for s in &mut self.steps {
*s = 0;
}
}
#[cfg(test)]
pub(crate) fn set_steps_for_test(&mut self, rank: usize, n: usize) {
self.steps[rank] = n;
}
#[cfg(test)]
pub(crate) fn set_wall_ms_for_test(&mut self, rank: usize, ms: f64) {
self.wall_ms[rank] = ms;
}
#[cfg(test)]
pub(crate) fn set_delivered_for_test(&mut self, rank: usize, ms: f64, batches: usize) {
self.delivered_ms[rank] = ms;
self.delivered_batches[rank] = batches;
}
}
#[cfg(test)]
mod tests {
use super::WindowLedger;
#[test]
fn first_batch_fills_then_marginal_accumulates() {
let mut l = WindowLedger::new(2);
l.record_batch(0, 10.0, 5.0);
assert_eq!(l.steps(0), 1);
assert_eq!(l.wall_ms(0), 10.0);
assert!(!l.has_delivered_sample(0));
assert_eq!(l.fill_excess_ms(0), 0.0); l.record_batch(0, 8.0, 2.0);
l.record_batch(0, 8.0, 2.0);
assert_eq!(l.steps(0), 3);
assert!(l.has_delivered_sample(0));
assert_eq!(l.delivered_ms(0), 20.0);
assert_eq!(l.delivered_batches(0), 2);
assert!((l.fill_excess_ms(0) - 5.0).abs() < 1e-12);
assert_eq!(l.steps(1), 0);
}
#[test]
fn fill_excess_clamps_at_zero() {
let mut l = WindowLedger::new(1);
l.record_batch(0, 1.0, 0.0); l.record_batch(0, 8.0, 2.0); assert_eq!(l.fill_excess_ms(0), 0.0);
}
#[test]
fn absorb_callback_cost_clamps_and_ignores_oob() {
let mut l = WindowLedger::new(1);
l.record_batch(0, 10.0, 0.0);
l.record_batch(0, 10.0, 5.0);
l.absorb_callback_cost(0, 12.0);
assert_eq!(l.wall_ms(0), 8.0);
assert_eq!(l.delivered_ms(0), 3.0);
l.absorb_callback_cost(0, 100.0); assert_eq!(l.wall_ms(0), 0.0);
assert_eq!(l.delivered_ms(0), 0.0);
l.absorb_callback_cost(7, 1.0); }
#[test]
fn record_batch_ignores_oob_rank() {
let mut l = WindowLedger::new(1);
l.record_batch(3, 1.0, 1.0);
assert_eq!(l.total_steps(), 0);
}
#[test]
fn resets_are_split_steps_vs_timing() {
let mut l = WindowLedger::new(2);
l.record_batch(0, 10.0, 5.0);
l.record_batch(0, 10.0, 5.0);
l.record_batch(1, 4.0, 1.0);
l.reset_timing();
assert_eq!(l.wall_ms(0), 0.0);
assert!(!l.has_delivered_sample(0));
assert_eq!(l.fill_excess_ms(0), 0.0);
assert_eq!(l.steps(0), 2);
l.reset_steps();
assert_eq!(l.steps(0), 0);
assert_eq!(l.steps(1), 0);
}
#[test]
fn cohort_step_stats() {
let mut l = WindowLedger::new(3);
l.record_batch(0, 1.0, 0.0);
l.record_batch(0, 1.0, 0.0);
l.record_batch(2, 1.0, 0.0);
assert_eq!(l.min_steps(), 0);
assert_eq!(l.max_steps(), 2);
assert_eq!(l.total_steps(), 3);
}
#[test]
fn mean_loss_is_absent_until_a_batch_reports() {
let mut l = WindowLedger::new(2);
assert_eq!(l.mean_loss(0), None);
l.record_batch_loss(0, 0.4);
l.record_batch_loss(0, 0.6);
assert!((l.mean_loss(0).unwrap() - 0.5).abs() < 1e-12);
assert_eq!(l.mean_loss(1), None);
l.record_batch_loss(9, 1.0);
assert_eq!(l.mean_loss(9), None);
}
#[test]
fn non_finite_loss_never_poisons_the_window_mean() {
let mut l = WindowLedger::new(1);
l.record_batch_loss(0, 0.5);
l.record_batch_loss(0, f64::NAN);
l.record_batch_loss(0, f64::INFINITY);
assert_eq!(l.mean_loss(0), Some(0.5));
}
#[test]
fn loss_resets_with_the_window_timing() {
let mut l = WindowLedger::new(1);
l.record_batch_loss(0, 2.0);
assert_eq!(l.mean_loss(0), Some(2.0));
l.reset_timing();
assert_eq!(l.mean_loss(0), None);
l.record_batch_loss(0, 3.0);
assert_eq!(l.mean_loss(0), Some(3.0));
}
#[test]
fn loss_count_is_independent_of_step_resets() {
let mut l = WindowLedger::new(1);
l.record_batch(0, 1.0, 0.0);
l.record_batch_loss(0, 4.0);
l.record_batch(0, 1.0, 0.0);
l.record_batch_loss(0, 6.0);
l.reset_steps(); assert_eq!(l.steps(0), 0);
assert_eq!(l.mean_loss(0), Some(5.0)); }
#[test]
fn per_batch_wall_is_infinite_cold() {
let mut l = WindowLedger::new(1);
assert_eq!(l.per_batch_wall_ms(0), f64::INFINITY);
assert_eq!(l.per_batch_wall_ms(9), f64::INFINITY); l.record_batch(0, 6.0, 0.0);
l.record_batch(0, 4.0, 0.0);
assert!((l.per_batch_wall_ms(0) - 5.0).abs() < 1e-12);
}
}