pub fn report_interval(total_samples: usize, batch_size: usize, reports_per_epoch: usize) -> f64 {
if batch_size == 0 || reports_per_epoch == 0 {
return f64::INFINITY;
}
let steps_per_epoch = (total_samples / batch_size) as f64;
steps_per_epoch / reports_per_epoch as f64
}
#[derive(Debug, Clone)]
pub struct ReportScheduler {
reports_per_epoch: usize,
epoch_work: f64,
fired: usize,
}
impl ReportScheduler {
pub fn new(reports_per_epoch: usize, epoch_work: f64) -> Self {
Self {
reports_per_epoch,
epoch_work,
fired: 0,
}
}
fn interval(&self) -> f64 {
self.epoch_work / self.reports_per_epoch as f64
}
pub fn on_sync(&mut self, in_epoch_work: f64) -> bool {
if self.reports_per_epoch == 0 || in_epoch_work <= 0.0 {
return false;
}
let interval = self.interval();
if !interval.is_finite() || interval <= 0.0 {
return false;
}
let reached = (in_epoch_work / interval).floor() as usize;
let due = reached.min(self.reports_per_epoch);
if due > self.fired {
self.fired = due;
true
} else {
false
}
}
pub fn reset_epoch(&mut self) {
self.fired = 0;
}
pub fn fired_this_epoch(&self) -> usize {
self.fired
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn interval_helper_matches_epoch_work_over_x() {
assert_eq!(report_interval(1000, 10, 4), 25.0);
assert_eq!(report_interval(1000, 10, 0), f64::INFINITY);
assert_eq!(report_interval(1000, 0, 4), f64::INFINITY);
}
#[test]
fn fires_on_each_threshold_crossing() {
let mut s = ReportScheduler::new(4, 100.0);
assert!(!s.on_sync(10.0)); assert!(!s.on_sync(24.0));
assert!(s.on_sync(25.0)); assert!(!s.on_sync(40.0)); assert!(s.on_sync(55.0)); assert!(s.on_sync(80.0)); assert!(s.on_sync(100.0)); assert!(!s.on_sync(120.0)); assert_eq!(s.fired_this_epoch(), 4);
}
#[test]
fn report_only_at_sync_boundaries() {
let mut s = ReportScheduler::new(4, 100.0);
let mut fired_at = Vec::new();
let mut work = 0.0;
for _ in 0..9 {
work += 3.0; if s.on_sync(work) {
fired_at.push(work);
}
}
assert_eq!(fired_at, vec![27.0]); }
#[test]
fn big_window_crossing_many_thresholds_fires_once() {
let mut s = ReportScheduler::new(10, 100.0); assert!(s.on_sync(35.0)); assert_eq!(s.fired_this_epoch(), 3); assert!(!s.on_sync(38.0)); assert!(s.on_sync(42.0)); }
#[test]
fn epoch_reset_restarts_thresholds() {
let mut s = ReportScheduler::new(2, 100.0); assert!(s.on_sync(60.0));
assert!(s.on_sync(100.0));
assert_eq!(s.fired_this_epoch(), 2);
s.reset_epoch();
assert_eq!(s.fired_this_epoch(), 0);
assert!(s.on_sync(50.0)); }
#[test]
fn single_epoch_degenerate_spreads_x_over_the_run() {
let mut s = ReportScheduler::new(5, 10_000.0); let mut fires = 0;
for step in (0..=10_000).step_by(500) {
if s.on_sync(step as f64) {
fires += 1;
}
}
assert_eq!(fires, 5);
}
#[test]
fn step_to_sample_proxy_is_scale_invariant() {
let bs = 32.0;
let mut in_steps = ReportScheduler::new(4, 100.0);
let mut in_samples = ReportScheduler::new(4, 100.0 * bs);
for step in (0..=100).step_by(7) {
let a = in_steps.on_sync(step as f64);
let b = in_samples.on_sync(step as f64 * bs);
assert_eq!(a, b, "at step {step}");
}
assert_eq!(in_steps.fired_this_epoch(), in_samples.fired_this_epoch());
}
#[test]
fn disabled_and_degenerate_never_fire() {
assert!(!ReportScheduler::new(0, 100.0).on_sync(50.0)); assert!(!ReportScheduler::new(4, 0.0).on_sync(50.0)); assert!(!ReportScheduler::new(4, 100.0).on_sync(0.0)); }
}