Skip to main content

fpm_rs/callbacks/
checkpoint.rs

1use std::{fs, path::PathBuf};
2
3use crate::{Result, reconstruction::ReconstructionCheckpoint};
4
5use super::{Callback, CallbackAction, StepContext};
6
7/// Writes resumable JSON checkpoints at a fixed iteration frequency.
8pub struct CheckpointEvery {
9    frequency: usize,
10    directory: PathBuf,
11}
12
13impl CheckpointEvery {
14    /// Creates a checkpoint callback; zero `frequency` is normalized to one.
15    pub fn new(frequency: usize, directory: impl Into<PathBuf>) -> Self {
16        Self {
17            frequency: frequency.max(1),
18            directory: directory.into(),
19        }
20    }
21}
22
23impl Callback for CheckpointEvery {
24    fn on_start(&mut self, _context: &StepContext<'_>) -> Result<CallbackAction> {
25        fs::create_dir_all(&self.directory)?;
26        Ok(CallbackAction::Continue)
27    }
28
29    fn on_iteration_end(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
30        if context.iteration.is_multiple_of(self.frequency) {
31            ReconstructionCheckpoint::capture(context.iteration, context.state, context.trace)
32                .save(
33                    self.directory
34                        .join(format!("checkpoint_{:05}.json", context.iteration)),
35                )?;
36        }
37        Ok(CallbackAction::Continue)
38    }
39}