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}