Skip to main content

fpm_rs/callbacks/
csv_logger.rs

1use std::{fs::File, path::PathBuf};
2
3use crate::{Result, diagnostics::DiagnosticRequest};
4
5use super::{Callback, CallbackAction, CallbackHook, StepContext};
6
7pub struct CsvLogger {
8    path: PathBuf,
9    writer: Option<csv::Writer<File>>,
10}
11
12impl CsvLogger {
13    pub fn new(path: impl Into<PathBuf>) -> Self {
14        Self {
15            path: path.into(),
16            writer: None,
17        }
18    }
19}
20
21impl Callback for CsvLogger {
22    fn requires(&self) -> Vec<DiagnosticRequest> {
23        vec![DiagnosticRequest::Loss]
24    }
25
26    fn requires_for(&self, hook: CallbackHook, _iteration: usize) -> Vec<DiagnosticRequest> {
27        if hook == CallbackHook::IterationEnd {
28            self.requires()
29        } else {
30            Vec::new()
31        }
32    }
33
34    fn on_start(&mut self, _context: &StepContext<'_>) -> Result<CallbackAction> {
35        let mut writer = csv::Writer::from_path(&self.path)?;
36        writer.write_record(["iteration", "loss"])?;
37        self.writer = Some(writer);
38        Ok(CallbackAction::Continue)
39    }
40
41    fn on_iteration_end(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
42        if let (Some(writer), Some(loss)) = (&mut self.writer, context.diagnostics.loss) {
43            writer.serialize((context.iteration, loss))?;
44            writer.flush()?;
45        }
46        Ok(CallbackAction::Continue)
47    }
48}