Skip to main content

ruda_model/record/
file.rs

1use super::{PrecisionSettings, Recorder, RecorderError, bin_config};
2use ruda_tensor::api::backend::Backend;
3use core::marker::PhantomData;
4use flate2::{Compression, read::GzDecoder, write::GzEncoder};
5use serde::{Serialize, de::DeserializeOwned};
6use std::io::BufReader;
7use std::{fs::File, path::PathBuf};
8
9mod writer;
10use writer::RecordWriter;
11
12/// Recorder trait specialized to save and load data to and from files.
13pub trait FileRecorder<B: Backend>:
14    Recorder<B, RecordArgs = PathBuf, RecordOutput = (), LoadArgs = PathBuf>
15{
16    /// File extension of the format used by the recorder.
17    fn file_extension() -> &'static str;
18}
19
20/// Default [file recorder](FileRecorder).
21pub type DefaultFileRecorder<S> = NamedMpkFileRecorder<S>;
22
23/// File recorder using the [bincode format](bincode).
24#[derive(new, Debug, Default, Clone)]
25pub struct BinFileRecorder<S: PrecisionSettings> {
26    _settings: PhantomData<S>,
27}
28
29/// File recorder using the [bincode format](bincode) compressed with gzip.
30#[derive(new, Debug, Default, Clone)]
31pub struct BinGzFileRecorder<S: PrecisionSettings> {
32    _settings: PhantomData<S>,
33}
34
35/// File recorder using the [json format](serde_json) compressed with gzip.
36#[derive(new, Debug, Default, Clone)]
37pub struct JsonGzFileRecorder<S: PrecisionSettings> {
38    _settings: PhantomData<S>,
39}
40
41/// File recorder using [pretty json format](serde_json) for easy readability.
42#[derive(new, Debug, Default, Clone)]
43pub struct PrettyJsonFileRecorder<S: PrecisionSettings> {
44    _settings: PhantomData<S>,
45}
46
47/// File recorder using the [named msgpack](rmp_serde) format compressed with gzip.
48#[derive(new, Debug, Default, Clone)]
49pub struct NamedMpkGzFileRecorder<S: PrecisionSettings> {
50    _settings: PhantomData<S>,
51}
52
53/// File recorder using the [named msgpack](rmp_serde) format.
54#[derive(new, Debug, Default, Clone)]
55pub struct NamedMpkFileRecorder<S: PrecisionSettings> {
56    _settings: PhantomData<S>,
57}
58
59impl<S: PrecisionSettings, B: Backend> FileRecorder<B> for BinGzFileRecorder<S> {
60    fn file_extension() -> &'static str {
61        "bin.gz"
62    }
63}
64impl<S: PrecisionSettings, B: Backend> FileRecorder<B> for BinFileRecorder<S> {
65    fn file_extension() -> &'static str {
66        "bin"
67    }
68}
69impl<S: PrecisionSettings, B: Backend> FileRecorder<B> for JsonGzFileRecorder<S> {
70    fn file_extension() -> &'static str {
71        "json.gz"
72    }
73}
74impl<S: PrecisionSettings, B: Backend> FileRecorder<B> for PrettyJsonFileRecorder<S> {
75    fn file_extension() -> &'static str {
76        "json"
77    }
78}
79
80impl<S: PrecisionSettings, B: Backend> FileRecorder<B> for NamedMpkGzFileRecorder<S> {
81    fn file_extension() -> &'static str {
82        "mpk.gz"
83    }
84}
85
86impl<S: PrecisionSettings, B: Backend> FileRecorder<B> for NamedMpkFileRecorder<S> {
87    fn file_extension() -> &'static str {
88        "mpk"
89    }
90}
91
92macro_rules! str2reader {
93    (
94        $file:expr
95    ) => {{
96        $file.set_extension(<Self as FileRecorder<B>>::file_extension());
97        let path = $file.as_path();
98
99        File::open(path)
100            .map_err(|err| match err.kind() {
101                std::io::ErrorKind::NotFound => RecorderError::FileNotFound(err.to_string()),
102                _ => RecorderError::Unknown(err.to_string()),
103            })
104            .map(|file| BufReader::new(file))
105    }};
106}
107
108macro_rules! str2writer {
109    (
110        $file:expr
111    ) => {{
112        $file.set_extension(<Self as FileRecorder<B>>::file_extension());
113        let path = $file.as_path();
114
115        log::debug!("Writing to file: {:?}", path);
116
117        RecordWriter::new(path)
118    }};
119}
120
121impl<S: PrecisionSettings, B: Backend> Recorder<B> for BinGzFileRecorder<S> {
122    type Settings = S;
123    type RecordArgs = PathBuf;
124    type RecordOutput = ();
125    type LoadArgs = PathBuf;
126
127    fn save_item<I: Serialize>(
128        &self,
129        item: I,
130        mut file: Self::RecordArgs,
131    ) -> Result<(), RecorderError> {
132        let config = bin_config();
133        let writer = str2writer!(file)?;
134        let mut writer = GzEncoder::new(writer, Compression::default());
135
136        bincode::serde::encode_into_std_write(&item, &mut writer, config)
137            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
138
139        writer.finish()
140            .map_err(|err| RecorderError::Unknown(err.to_string()))?
141            .commit()
142    }
143
144    fn load_item<I: DeserializeOwned>(
145        &self,
146        file: &mut Self::LoadArgs,
147    ) -> Result<I, RecorderError> {
148        let reader = str2reader!(file)?;
149        let mut reader = GzDecoder::new(reader);
150        let state = bincode::serde::decode_from_std_read(&mut reader, bin_config())
151            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
152
153        Ok(state)
154    }
155}
156
157impl<S: PrecisionSettings, B: Backend> Recorder<B> for BinFileRecorder<S> {
158    type Settings = S;
159    type RecordArgs = PathBuf;
160    type RecordOutput = ();
161    type LoadArgs = PathBuf;
162
163    fn save_item<I: Serialize>(
164        &self,
165        item: I,
166        mut file: Self::RecordArgs,
167    ) -> Result<(), RecorderError> {
168        let config = bin_config();
169        let mut writer = str2writer!(file)?;
170        bincode::serde::encode_into_std_write(&item, &mut writer, config)
171            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
172        writer.commit()
173    }
174
175    fn load_item<I: DeserializeOwned>(
176        &self,
177        file: &mut Self::LoadArgs,
178    ) -> Result<I, RecorderError> {
179        let mut reader = str2reader!(file)?;
180        let state = bincode::serde::decode_from_std_read(&mut reader, bin_config())
181            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
182        Ok(state)
183    }
184}
185
186impl<S: PrecisionSettings, B: Backend> Recorder<B> for JsonGzFileRecorder<S> {
187    type Settings = S;
188    type RecordArgs = PathBuf;
189    type RecordOutput = ();
190    type LoadArgs = PathBuf;
191
192    fn save_item<I: Serialize>(
193        &self,
194        item: I,
195        mut file: Self::RecordArgs,
196    ) -> Result<(), RecorderError> {
197        let writer = str2writer!(file)?;
198        let mut writer = GzEncoder::new(writer, Compression::default());
199        serde_json::to_writer(&mut writer, &item)
200            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
201
202        writer.finish()
203            .map_err(|err| RecorderError::Unknown(err.to_string()))?
204            .commit()
205    }
206
207    fn load_item<I: DeserializeOwned>(
208        &self,
209        file: &mut Self::LoadArgs,
210    ) -> Result<I, RecorderError> {
211        let reader = str2reader!(file)?;
212        let reader = GzDecoder::new(reader);
213        let state = serde_json::from_reader(reader)
214            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
215
216        Ok(state)
217    }
218}
219
220impl<S: PrecisionSettings, B: Backend> Recorder<B> for PrettyJsonFileRecorder<S> {
221    type Settings = S;
222    type RecordArgs = PathBuf;
223    type RecordOutput = ();
224    type LoadArgs = PathBuf;
225
226    fn save_item<I: Serialize>(
227        &self,
228        item: I,
229        mut file: Self::RecordArgs,
230    ) -> Result<(), RecorderError> {
231        let mut writer = str2writer!(file)?;
232        serde_json::to_writer_pretty(&mut writer, &item)
233            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
234        writer.commit()
235    }
236
237    fn load_item<I: DeserializeOwned>(
238        &self,
239        file: &mut Self::LoadArgs,
240    ) -> Result<I, RecorderError> {
241        let reader = str2reader!(file)?;
242        let state = serde_json::from_reader(reader)
243            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
244
245        Ok(state)
246    }
247}
248
249impl<S: PrecisionSettings, B: Backend> Recorder<B> for NamedMpkGzFileRecorder<S> {
250    type Settings = S;
251    type RecordArgs = PathBuf;
252    type RecordOutput = ();
253    type LoadArgs = PathBuf;
254
255    fn save_item<I: Serialize>(
256        &self,
257        item: I,
258        mut file: Self::RecordArgs,
259    ) -> Result<(), RecorderError> {
260        let writer = str2writer!(file)?;
261        let mut writer = GzEncoder::new(writer, Compression::default());
262        rmp_serde::encode::write_named(&mut writer, &item)
263            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
264
265        writer.finish()
266            .map_err(|err| RecorderError::Unknown(err.to_string()))?
267            .commit()
268    }
269
270    fn load_item<I: DeserializeOwned>(
271        &self,
272        file: &mut Self::LoadArgs,
273    ) -> Result<I, RecorderError> {
274        let reader = str2reader!(file)?;
275        let reader = GzDecoder::new(reader);
276        let state = rmp_serde::decode::from_read(reader)
277            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
278
279        Ok(state)
280    }
281}
282
283impl<S: PrecisionSettings, B: Backend> Recorder<B> for NamedMpkFileRecorder<S> {
284    type Settings = S;
285    type RecordArgs = PathBuf;
286    type RecordOutput = ();
287    type LoadArgs = PathBuf;
288
289    fn save_item<I: Serialize>(
290        &self,
291        item: I,
292        mut file: Self::RecordArgs,
293    ) -> Result<(), RecorderError> {
294        let mut writer = str2writer!(file)?;
295
296        rmp_serde::encode::write_named(&mut writer, &item)
297            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
298
299        writer.commit()
300    }
301
302    fn load_item<I: DeserializeOwned>(
303        &self,
304        file: &mut Self::LoadArgs,
305    ) -> Result<I, RecorderError> {
306        let reader = str2reader!(file)?;
307        let state = rmp_serde::decode::from_read(reader)
308            .map_err(|err| RecorderError::Unknown(err.to_string()))?;
309
310        Ok(state)
311    }
312}
313
314#[allow(deprecated)]
315#[cfg(test)]
316mod tests {
317    use super::*;
318    use crate::config::Config;
319    use crate::module::Ignored;
320    use crate::test_utils::SimpleLinear;
321    use crate::{
322        TestBackend,
323        module::Module,
324        record::{BinBytesRecorder, FullPrecisionSettings},
325    };
326    use ruda_tensor::api::backend::Backend;
327    use ruda_tensor::api::{Device, Tensor};
328
329    #[inline(always)]
330    fn file_path(file: &str) -> PathBuf {
331        std::env::temp_dir().as_path().join(file)
332    }
333
334    #[test]
335    fn test_can_save_and_load_jsongz_format() {
336        test_can_save_and_load(JsonGzFileRecorder::<FullPrecisionSettings>::default())
337    }
338
339    #[test]
340    fn test_can_save_and_load_bin_format() {
341        test_can_save_and_load(BinFileRecorder::<FullPrecisionSettings>::default())
342    }
343
344    #[test]
345    fn test_can_save_and_load_bingz_format() {
346        test_can_save_and_load(BinGzFileRecorder::<FullPrecisionSettings>::default())
347    }
348
349    #[test]
350    fn test_can_save_and_load_pretty_json_format() {
351        test_can_save_and_load(PrettyJsonFileRecorder::<FullPrecisionSettings>::default())
352    }
353
354    #[test]
355    fn test_can_save_and_load_mpkgz_format() {
356        test_can_save_and_load(NamedMpkGzFileRecorder::<FullPrecisionSettings>::default())
357    }
358
359    #[test]
360    fn test_can_save_and_load_mpk_format() {
361        test_can_save_and_load(NamedMpkFileRecorder::<FullPrecisionSettings>::default())
362    }
363
364    fn test_can_save_and_load<Recorder>(recorder: Recorder)
365    where
366        Recorder: FileRecorder<TestBackend>,
367    {
368        let filename = "ruda_test_file_recorder";
369
370        let device = Default::default();
371        let mut model_before = create_model(&device);
372
373        // NOTE: Non-module fields currently act like `#[module(skip)]`, meaning their state
374        // is not persistent. These fields hold `EmptyRecord`s.
375        // So `model_bytes_after == model_bytes_before` because the changes do not persist in the record.
376        model_before.tensor = Tensor::full([4], 2., &device);
377        model_before.arr = [3, 3];
378        model_before.int = 1;
379        model_before.ignore = Ignored(PaddingConfig2d::Valid);
380
381        recorder
382            .record(model_before.clone().into_record(), file_path(filename))
383            .unwrap();
384
385        let model_after =
386            create_model(&device).load_record(recorder.load(file_path(filename), &device).unwrap());
387
388        // State is not persisted for empty record fields
389        assert_eq!(model_after.arr, [2, 2]);
390        assert_eq!(model_after.int, 0);
391        assert_eq!(model_after.ignore.0, PaddingConfig2d::Same);
392
393        let byte_recorder = BinBytesRecorder::<FullPrecisionSettings>::default();
394        let model_bytes_before = byte_recorder
395            .record(model_before.into_record(), ())
396            .unwrap();
397        let model_bytes_after = byte_recorder.record(model_after.into_record(), ()).unwrap();
398
399        assert_eq!(model_bytes_after, model_bytes_before);
400    }
401
402    #[derive(Config, Debug, PartialEq, Eq)]
403    pub enum PaddingConfig2d {
404        Same,
405        Valid,
406        Explicit(usize, usize),
407    }
408
409    // Dummy model with different record types
410    #[derive(Module, Debug)]
411    pub struct Model<B: Backend> {
412        linear1: SimpleLinear<B>,
413        phantom: PhantomData<B>,
414        tensor: Tensor<B, 1>,
415        arr: [usize; 2],
416        int: usize,
417        ignore: Ignored<PaddingConfig2d>,
418    }
419
420    pub fn create_model(device: &Device<TestBackend>) -> Model<TestBackend> {
421        let linear1 = SimpleLinear::new(32, 32, device);
422
423        Model {
424            linear1,
425            phantom: PhantomData,
426            tensor: Tensor::zeros([2], device),
427            arr: [2, 2],
428            int: 0,
429            ignore: Ignored(PaddingConfig2d::Same),
430        }
431    }
432}