Skip to main content

only_brain/
io.rs

1//! Saving and loading networks in the model file format.
2//!
3//! # Format
4//!
5//! Every file starts with an 8-byte magic, `ONLYBRN\0`, and a format version. Version
6//! 1, the one written today, continues as follows. Every number is little-endian.
7//!
8//! | Field | Type | Meaning |
9//! |---|---|---|
10//! | magic | 8 bytes | `ONLYBRN\0` |
11//! | version | `u16` | `1` |
12//! | inputs | `u32` | number of input neurons |
13//! | layers | `u32` | number of layers after the input layer |
14//!
15//! Then, for each layer after the input layer:
16//!
17//! | Field | Type | Meaning |
18//! |---|---|---|
19//! | neurons | `u32` | number of neurons in this layer |
20//! | activation | `u8` | 0 sigmoid, 1 tanh, 2 ReLU, 3 binary step, 4 identity |
21//! | weights | `f64` × neurons × previous layer's neurons | one row per neuron |
22//! | biases | `f64` × neurons | one per neuron |
23//!
24//! The weights and biases, read in file order, are exactly the network's [flat
25//! parameter view](crate::NeuralNetwork#flat-parameter-view).
26//!
27//! Files written before 0.3 have no magic: they are the bincode 1 encoding of the
28//! network, with one activation function for the whole network. They still load, and
29//! every layer gets that activation function.
30
31use std::fs::File;
32use std::io::{BufReader, BufWriter, Read, Write};
33use std::path::Path;
34
35use crate::neural_network::{LayerData, NetworkData};
36use crate::{ActivationFunction, NeuralNetwork};
37
38mod legacy;
39
40/// The bytes every model file written since 0.3 starts with.
41const MAGIC: [u8; 8] = *b"ONLYBRN\0";
42
43/// The model format version [`dump_model`] and [`write_model`] write. [`load_model`] and
44/// [`read_model`] read this version and every earlier one.
45pub const MODEL_FORMAT_VERSION: u16 = 1;
46
47/// Something that went wrong while saving or loading a model.
48#[derive(Debug, thiserror::Error)]
49#[non_exhaustive]
50pub enum ModelError {
51    /// The file could not be read or written.
52    #[error("could not access the model file")]
53    Io(#[from] std::io::Error),
54
55    /// The data is not a model written by this library.
56    #[error("data is not a model produced by this library")]
57    UnrecognizedFormat,
58
59    /// The model was written by a newer version of this library, in a format this
60    /// version does not know.
61    #[error("model format version {found} is not supported, the newest known is {newest}")]
62    UnsupportedVersion {
63        /// The version stored in the file.
64        found: u16,
65        /// The newest version this library reads, [`MODEL_FORMAT_VERSION`].
66        newest: u16,
67    },
68
69    /// The data starts like a model but is truncated or corrupted.
70    #[error("model data is corrupted: {reason}")]
71    Malformed {
72        /// What was wrong.
73        reason: &'static str,
74    },
75
76    /// The stored network has a different shape than the type asked for.
77    #[error("model has {found} {end} neurons but {expected} were expected")]
78    DimensionMismatch {
79        /// Which end of the network mismatched.
80        end: &'static str,
81        /// The width required by the requested type.
82        expected: usize,
83        /// The width actually found in the file.
84        found: usize,
85    },
86
87    /// The file decoded to a network with no layers at all.
88    #[error("model file contains a network with no layers")]
89    EmptyNetwork,
90
91    /// A stored layer does not fit with itself or with the layer before it, which means
92    /// the file was corrupted or edited by hand.
93    #[error("layer {layer} of the model is inconsistent: {reason}")]
94    InconsistentLayer {
95        /// The layer number, counting the input layer as 0.
96        layer: usize,
97        /// What disagreed.
98        reason: &'static str,
99    },
100}
101
102/// Writes a model to `path`, in the newest format version.
103///
104/// # Errors
105///
106/// Returns [`ModelError::Io`] if the file cannot be written.
107///
108/// # Example
109///
110/// ```no_run
111/// # use only_brain::{dump_model, NeuralNetwork};
112/// # fn main() -> Result<(), Box<dyn std::error::Error>> {
113/// let nn = NeuralNetwork::<2, 1>::new(&[2]);
114/// dump_model(&nn, "model.bin")?;
115/// # Ok(())
116/// # }
117/// ```
118pub fn dump_model<const IN: usize, const OUT: usize>(
119    model: &NeuralNetwork<IN, OUT>,
120    path: impl AsRef<Path>,
121) -> Result<(), ModelError> {
122    let mut file = BufWriter::new(File::create(path)?);
123    write_model(model, &mut file)?;
124    file.flush()?;
125
126    Ok(())
127}
128
129/// Reads a model from `path`, checking that its shape matches `IN` and `OUT`.
130///
131/// Files written by every earlier version of this library load too. Every layer is
132/// checked against its neighbours, so a corrupted file is reported here rather than
133/// panicking later in [`NeuralNetwork::feed_forward`], and the file must end where the
134/// model does.
135///
136/// # Errors
137///
138/// Returns [`ModelError::Io`] if the file cannot be read,
139/// [`ModelError::UnrecognizedFormat`] if it is not a model,
140/// [`ModelError::UnsupportedVersion`] if a newer version of this library wrote it,
141/// [`ModelError::Malformed`] if it is truncated or corrupted,
142/// [`ModelError::DimensionMismatch`] if the stored network does not fit
143/// `NeuralNetwork<IN, OUT>`, and [`ModelError::EmptyNetwork`] or
144/// [`ModelError::InconsistentLayer`] if its layers do not form a network at all.
145///
146/// # Example
147///
148/// ```no_run
149/// # use only_brain::{load_model, NeuralNetwork};
150/// # fn main() -> Result<(), Box<dyn std::error::Error>> {
151/// let nn: NeuralNetwork<2, 1> = load_model("model.bin")?;
152/// # Ok(())
153/// # }
154/// ```
155pub fn load_model<const IN: usize, const OUT: usize>(
156    path: impl AsRef<Path>,
157) -> Result<NeuralNetwork<IN, OUT>, ModelError> {
158    let mut file = BufReader::new(File::open(path)?);
159    let model = read_model(&mut file)?;
160
161    if file.read(&mut [0])? != 0 {
162        return Err(ModelError::Malformed {
163            reason: "the file continues after the model",
164        });
165    }
166
167    Ok(model)
168}
169
170/// Writes a model to any writer, in the newest format version.
171///
172/// The model is encoded in memory and written with a single `write_all`, so the writer
173/// does not need to be buffered. Use it with a `Vec<u8>` to keep a model in memory,
174/// for example to store it in a database.
175///
176/// # Errors
177///
178/// Returns [`ModelError::Io`] if writing fails.
179///
180/// # Example
181///
182/// ```
183/// # use only_brain::{read_model, write_model, NeuralNetwork};
184/// # fn main() -> Result<(), Box<dyn std::error::Error>> {
185/// let nn = NeuralNetwork::<2, 1>::new(&[2]);
186///
187/// let mut bytes = Vec::new();
188/// write_model(&nn, &mut bytes)?;
189///
190/// let copy: NeuralNetwork<2, 1> = read_model(bytes.as_slice())?;
191/// assert_eq!(copy, nn);
192/// # Ok(())
193/// # }
194/// ```
195pub fn write_model<const IN: usize, const OUT: usize>(
196    model: &NeuralNetwork<IN, OUT>,
197    mut writer: impl Write,
198) -> Result<(), ModelError> {
199    let layers = model.layers();
200    let mut bytes = Vec::with_capacity(MAGIC.len() + 10 + 8 * model.parameter_count() + 5 * layers.len());
201
202    bytes.extend_from_slice(&MAGIC);
203    bytes.extend_from_slice(&MODEL_FORMAT_VERSION.to_le_bytes());
204    bytes.extend_from_slice(&format_count(IN).to_le_bytes());
205    bytes.extend_from_slice(&format_count(layers.len()).to_le_bytes());
206
207    for layer in layers {
208        bytes.extend_from_slice(&format_count(layer.size()).to_le_bytes());
209        bytes.push(layer.activation().code());
210        // The weights row by row, then the biases: the flat parameter view.
211        for row in layer.weights().row_iter() {
212            for value in &row {
213                bytes.extend_from_slice(&value.to_le_bytes());
214            }
215        }
216        for value in layer.biases() {
217            bytes.extend_from_slice(&value.to_le_bytes());
218        }
219    }
220
221    writer.write_all(&bytes)?;
222    Ok(())
223}
224
225/// Reads one model from any reader, checking that its shape matches `IN` and `OUT`.
226///
227/// Reading stops where the model ends, so several models can be read one after the
228/// other from the same stream. Wrap slow readers such as files in a
229/// [`std::io::BufReader`].
230///
231/// # Errors
232///
233/// The same as [`load_model`], except that data after the model is left unread rather
234/// than reported.
235pub fn read_model<const IN: usize, const OUT: usize>(
236    reader: impl Read,
237) -> Result<NeuralNetwork<IN, OUT>, ModelError> {
238    let mut decoder = Decoder::new(reader);
239
240    let start: [u8; 8] = decoder.bytes().map_err(|_| ModelError::UnrecognizedFormat)?;
241    let data = if start == MAGIC {
242        match decoder.u16()? {
243            1 => read_v1(&mut decoder)?,
244            found => {
245                return Err(ModelError::UnsupportedVersion {
246                    found,
247                    newest: MODEL_FORMAT_VERSION,
248                })
249            }
250        }
251    } else {
252        legacy::read(u64::from_le_bytes(start), &mut decoder)?
253    };
254
255    NeuralNetwork::try_from(data)
256}
257
258/// Reads the rest of a version 1 model, after the magic and the version.
259fn read_v1<R: Read>(decoder: &mut Decoder<R>) -> Result<NetworkData, ModelError> {
260    let mut inputs = decoder.u32()? as usize;
261    let layer_count = decoder.u32()?;
262
263    let mut layers = Vec::new();
264    for _ in 0..layer_count {
265        let neurons = decoder.u32()? as usize;
266        // Every later read then consumes data, so a corrupted count runs out of input
267        // instead of looping without end.
268        if neurons == 0 || inputs == 0 {
269            return Err(ModelError::Malformed {
270                reason: "a layer has no neurons",
271            });
272        }
273        let activation = ActivationFunction::from_code(decoder.u8()?).ok_or(ModelError::Malformed {
274            reason: "unknown activation function code",
275        })?;
276
277        let mut weights = Vec::new();
278        for _ in 0..neurons {
279            weights.push(decoder.f64s(inputs)?);
280        }
281        let biases = decoder.f64s(neurons)?;
282
283        layers.push(LayerData {
284            activation,
285            weights,
286            biases,
287        });
288        inputs = neurons;
289    }
290
291    Ok(NetworkData { layers })
292}
293
294/// A count as the format stores it.
295fn format_count(count: usize) -> u32 {
296    u32::try_from(count).expect("a layer cannot have more than u32::MAX neurons")
297}
298
299/// Reads exactly enough bytes to fill `bytes`, reporting a short read as a malformed
300/// model.
301fn fill(reader: &mut impl Read, bytes: &mut [u8]) -> Result<(), ModelError> {
302    reader.read_exact(bytes).map_err(|error| match error.kind() {
303        std::io::ErrorKind::UnexpectedEof => ModelError::Malformed {
304            reason: "the data ends before the model does",
305        },
306        _ => ModelError::Io(error),
307    })
308}
309
310/// How many floats [`Decoder::f64s`] reads per chunk.
311const DECODE_CHUNK: usize = 512;
312
313/// Reads little-endian values, reporting a short read as a malformed model.
314struct Decoder<R> {
315    reader: R,
316    /// Scratch for bulk reads, so each read does not initialise its own buffer.
317    buffer: Box<[u8; 8 * DECODE_CHUNK]>,
318}
319
320impl<R: Read> Decoder<R> {
321    fn new(reader: R) -> Self {
322        Self {
323            reader,
324            buffer: Box::new([0; 8 * DECODE_CHUNK]),
325        }
326    }
327
328    fn bytes<const N: usize>(&mut self) -> Result<[u8; N], ModelError> {
329        let mut bytes = [0; N];
330        fill(&mut self.reader, &mut bytes)?;
331        Ok(bytes)
332    }
333
334    fn u8(&mut self) -> Result<u8, ModelError> {
335        Ok(u8::from_le_bytes(self.bytes()?))
336    }
337
338    fn u16(&mut self) -> Result<u16, ModelError> {
339        Ok(u16::from_le_bytes(self.bytes()?))
340    }
341
342    fn u32(&mut self) -> Result<u32, ModelError> {
343        Ok(u32::from_le_bytes(self.bytes()?))
344    }
345
346    fn u64(&mut self) -> Result<u64, ModelError> {
347        Ok(u64::from_le_bytes(self.bytes()?))
348    }
349
350    /// Reads `count` floats, in bulk but one bounded chunk at a time. The count comes
351    /// from the data, so memory is reserved as values arrive rather than up front, and a
352    /// corrupted count fails at the end of the data instead of exhausting memory.
353    fn f64s(&mut self, count: usize) -> Result<Vec<f64>, ModelError> {
354        let mut values = Vec::with_capacity(count.min(8 * DECODE_CHUNK));
355        let mut remaining = count;
356        while remaining > 0 {
357            let chunk = remaining.min(DECODE_CHUNK);
358            let bytes = &mut self.buffer[..8 * chunk];
359            fill(&mut self.reader, bytes)?;
360            // The slice holds 8 * chunk bytes, so there is no remainder.
361            let (chunks, _) = bytes.as_chunks::<8>();
362            values.extend(chunks.iter().map(|value| f64::from_le_bytes(*value)));
363            remaining -= chunk;
364        }
365        Ok(values)
366    }
367}