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 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}