Skip to main content

tboard/
lib.rs

1use pyo3::prelude::*;
2use pyo3::types::{PyBytes, PyDict, PyList};
3
4use ::tboard as tb;
5use tb::Error;
6
7#[allow(unused)]
8fn w_py(err: PyErr) -> Error {
9    Error::msg(err)
10}
11
12#[allow(unused)]
13fn w<E: std::error::Error>(err: E) -> PyErr {
14    pyo3::exceptions::PyValueError::new_err(err.to_string())
15}
16
17#[macro_export]
18macro_rules! py_bail {
19    ($msg:literal $(,)?) => {
20        return Err(pyo3::exceptions::PyValueError::new_err(format!($msg)))
21    };
22    ($err:expr $(,)?) => {
23        return Err(pyo3::exceptions::PyValueError::new_err(format!($err)))
24    };
25    ($fmt:expr, $($arg:tt)*) => {
26        return Err(pyo3::exceptions::PyValueError::new_err(format!($fmt, $($arg)*)))
27    };
28}
29
30#[derive(Copy, Clone, PartialEq, Eq)]
31enum OnError {
32    Log,
33    Raise,
34}
35
36#[pyclass]
37struct EventIter {
38    reader: tb::SummaryReader<std::fs::File>,
39}
40
41#[pymethods]
42impl EventIter {
43    fn __iter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
44        slf
45    }
46
47    fn __next__(mut slf: PyRefMut<'_, Self>, py: Python) -> PyResult<Option<PyObject>> {
48        use ::tboard::tensorboard::event::What;
49        use ::tboard::tensorboard::summary::value::Value::SimpleValue;
50        use std::ops::DerefMut;
51
52        let slf = slf.deref_mut();
53        match slf.reader.next() {
54            None => Ok(None),
55            Some(ok_or_err) => {
56                let event = ok_or_err.map_err(w)?;
57                let dict = PyDict::new(py);
58                dict.set_item("wall_time", event.wall_time)?;
59                dict.set_item("step", event.step)?;
60                dict.set_item("source_metadata", event.source_metadata.map(|v| v.writer))?;
61                let mut values = vec![];
62                if let Some(what) = event.what {
63                    match what {
64                        What::Summary(summary) => {
65                            dict.set_item("kind", "summary")?;
66                            for value in summary.value.iter() {
67                                let v = PyDict::new(py);
68                                v.set_item("tag", &value.tag)?;
69                                v.set_item("node_name", &value.node_name)?;
70                                match value.value {
71                                    Some(SimpleValue(s)) => v.set_item("value", s)?,
72                                    Some(_) => {}
73                                    None => v.set_item("value", None::<usize>)?,
74                                }
75                                // v.set_item("metadata", value.metadata);
76                                values.push(v)
77                            }
78                        }
79                        What::MetaGraphDef(def) => {
80                            dict.set_item("kind", "meta_graph_def")?;
81                            dict.set_item("meta_graph_def", PyBytes::new(py, &def))?;
82                        }
83                        What::TaggedRunMetadata(trm) => {
84                            dict.set_item("kind", "tagged_run_metadata")?;
85                            dict.set_item("tag", trm.tag)?;
86                            dict.set_item("run_metadata", trm.run_metadata)?;
87                        }
88                        What::FileVersion(version) => {
89                            dict.set_item("kind", "file_version")?;
90                            dict.set_item("file_version", version)?;
91                        }
92                        What::SessionLog(sl) => {
93                            dict.set_item("kind", "session_log")?;
94                            dict.set_item("status", sl.status)?;
95                            dict.set_item("msg", sl.msg)?;
96                            dict.set_item("checkpoint_path", sl.checkpoint_path)?;
97                        }
98                        What::LogMessage(lm) => {
99                            dict.set_item("kind", "log_message")?;
100                            dict.set_item("level", lm.level)?;
101                            dict.set_item("message", lm.message)?;
102                        }
103                        What::GraphDef(gd) => {
104                            dict.set_item("kind", "graph_def")?;
105                            dict.set_item("graph_def", PyBytes::new(py, &gd))?;
106                        }
107                    }
108                }
109                let what = PyList::new(py, values.iter());
110                dict.set_item("what", what)?;
111                Ok(Some(dict.into()))
112            }
113        }
114    }
115}
116
117#[pyclass]
118struct EventReader {
119    filename: std::path::PathBuf,
120}
121
122#[pymethods]
123impl EventReader {
124    #[new]
125    #[pyo3(signature = (filename,))]
126    fn new(filename: &str) -> PyResult<Self> {
127        let filename = std::path::PathBuf::from(filename);
128        if !filename.is_file() {
129            py_bail!("{filename:?} is not a file")
130        }
131        Ok(Self { filename })
132    }
133
134    fn __iter__(slf: PyRef<'_, Self>) -> PyResult<EventIter> {
135        let reader = std::fs::File::open(&slf.filename)?;
136        let reader = tb::SummaryReader::new(reader);
137        Ok(EventIter { reader })
138    }
139}
140
141#[pyclass]
142struct EventWriter {
143    inner: tb::EventWriter<std::fs::File>,
144    on_error: OnError,
145    logdir: String,
146}
147
148impl EventWriter {
149    fn handle_err(&self, r: tb::Result<()>) -> PyResult<()> {
150        match self.on_error {
151            OnError::Raise => r.map_err(w),
152            OnError::Log => {
153                if let Err(err) = r {
154                    eprintln!("error logging to {:?}: {err:?}", self.inner.filename());
155                }
156                Ok(())
157            }
158        }
159    }
160}
161
162#[pymethods]
163impl EventWriter {
164    #[new]
165    #[pyo3(signature = (logdir, on_error="raise"))]
166    fn new(logdir: String, on_error: &str) -> PyResult<Self> {
167        let inner = tb::EventWriter::create(&logdir).map_err(w)?;
168        let on_error = match on_error {
169            "raise" => OnError::Raise,
170            "log" => OnError::Log,
171            on_error => py_bail!("on_error can only be 'raise' or 'log', got '{on_error}'"),
172        };
173        Ok(Self { inner, logdir, on_error })
174    }
175
176    #[pyo3(signature = (tag, scalar_value, global_step=0))]
177    fn add_scalar(&mut self, tag: &str, scalar_value: f32, global_step: i64) -> PyResult<()> {
178        let res = self.inner.write_scalar(global_step, tag, scalar_value);
179        self.handle_err(res)?;
180        self.flush()
181    }
182
183    fn flush(&mut self) -> PyResult<()> {
184        let res = self.inner.flush();
185        self.handle_err(res)
186    }
187
188    #[getter]
189    fn logdir(&self) -> &str {
190        &self.logdir
191    }
192
193    #[getter]
194    fn filename(&self) -> Option<&str> {
195        self.inner.filename().and_then(|v| v.to_str())
196    }
197}
198
199#[pymodule]
200fn tboard(_py: Python, m: &PyModule) -> PyResult<()> {
201    m.add_class::<EventReader>()?;
202    m.add_class::<EventWriter>()?;
203    Ok(())
204}