fpm_rs/callbacks/
early_stop.rs1use crate::{Result, diagnostics::DiagnosticRequest};
2
3use super::{Callback, CallbackAction, CallbackHook, StepContext};
4
5pub struct StopOnPlateau {
7 patience: usize,
8 minimum_improvement: f64,
9 best: f64,
10 stale_iterations: usize,
11}
12
13impl StopOnPlateau {
14 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}