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
7pub struct CheckpointEvery {
8    frequency: usize,
9    directory: PathBuf,
10}
11
12impl CheckpointEvery {
13    pub fn new(frequency: usize, directory: impl Into<PathBuf>) -> Self {
14        Self {
15            frequency: frequency.max(1),
16            directory: directory.into(),
17        }
18    }
19}
20
21impl Callback for CheckpointEvery {
22    fn on_start(&mut self, _context: &StepContext<'_>) -> Result<CallbackAction> {
23        fs::create_dir_all(&self.directory)?;
24        Ok(CallbackAction::Continue)
25    }
26
27    fn on_iteration_end(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
28        if context.iteration.is_multiple_of(self.frequency) {
29            ReconstructionCheckpoint::capture(context.iteration, context.state, context.history)
30                .save(
31                    self.directory
32                        .join(format!("checkpoint_{:05}.json", context.iteration)),
33                )?;
34        }
35        Ok(CallbackAction::Continue)
36    }
37}