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
12pub trait FileRecorder<B: Backend>:
14 Recorder<B, RecordArgs = PathBuf, RecordOutput = (), LoadArgs = PathBuf>
15{
16 fn file_extension() -> &'static str;
18}
19
20pub type DefaultFileRecorder<S> = NamedMpkFileRecorder<S>;
22
23#[derive(new, Debug, Default, Clone)]
25pub struct BinFileRecorder<S: PrecisionSettings> {
26 _settings: PhantomData<S>,
27}
28
29#[derive(new, Debug, Default, Clone)]
31pub struct BinGzFileRecorder<S: PrecisionSettings> {
32 _settings: PhantomData<S>,
33}
34
35#[derive(new, Debug, Default, Clone)]
37pub struct JsonGzFileRecorder<S: PrecisionSettings> {
38 _settings: PhantomData<S>,
39}
40
41#[derive(new, Debug, Default, Clone)]
43pub struct PrettyJsonFileRecorder<S: PrecisionSettings> {
44 _settings: PhantomData<S>,
45}
46
47#[derive(new, Debug, Default, Clone)]
49pub struct NamedMpkGzFileRecorder<S: PrecisionSettings> {
50 _settings: PhantomData<S>,
51}
52
53#[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 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 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 #[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}