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