Skip to main content

fpm_rs/callbacks/
early_stop.rs

1use crate::{Result, diagnostics::DiagnosticRequest};
2
3use super::{Callback, CallbackAction, CallbackHook, StepContext};
4
5/// Stops after a configured number of iterations without sufficient objective improvement.
6pub struct StopOnPlateau {
7    patience: usize,
8    minimum_improvement: f64,
9    best: f64,
10    stale_iterations: usize,
11}
12
13impl StopOnPlateau {
14    /// Creates a stopper with at least one iteration of patience and a non-negative
15    /// absolute `minimum_improvement`.
16    pub fn new(patience: usize, minimum_improvement: f64) -> Self {
17        Self {
18            patience: patience.max(1),
19            minimum_improvement: minimum_improvement.max(0.0),
20            best: f64::INFINITY,
21            stale_iterations: 0,
22        }
23    }
24}
25
26impl Callback for StopOnPlateau {
27    fn requires(&self) -> Vec<DiagnosticRequest> {
28        vec![DiagnosticRequest::Objective]
29    }
30
31    fn requires_for(&self, hook: CallbackHook, _iteration: usize) -> Vec<DiagnosticRequest> {
32        if hook == CallbackHook::IterationEnd {
33            self.requires()
34        } else {
35            Vec::new()
36        }
37    }
38
39    fn on_start(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
40        self.best = f64::INFINITY;
41        self.stale_iterations = 0;
42        for record in &context.trace.iterations {
43            if self.best - record.objective > self.minimum_improvement {
44                self.best = record.objective;
45                self.stale_iterations = 0;
46            } else {
47                self.stale_iterations += 1;
48            }
49        }
50        Ok(CallbackAction::Continue)
51    }
52
53    fn on_iteration_end(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
54        if let Some(objective) = context.diagnostics.objective {
55            if self.best - objective > self.minimum_improvement {
56                self.best = objective;
57                self.stale_iterations = 0;
58            } else {
59                self.stale_iterations += 1;
60            }
61        }
62        Ok(if self.stale_iterations >= self.patience {
63            CallbackAction::Stop
64        } else {
65            CallbackAction::Continue
66        })
67    }
68}