fpm_rs/callbacks/
checkpoint.rs1use 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}