use std::fs::File;
use std::io::{BufReader, BufWriter, Read, Write};
use std::path::Path;
use crate::neural_network::{LayerData, NetworkData};
use crate::{ActivationFunction, NeuralNetwork};
mod legacy;
const MAGIC: [u8; 8] = *b"ONLYBRN\0";
pub const MODEL_FORMAT_VERSION: u16 = 1;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ModelError {
#[error("could not access the model file")]
Io(#[from] std::io::Error),
#[error("data is not a model produced by this library")]
UnrecognizedFormat,
#[error("model format version {found} is not supported, the newest known is {newest}")]
UnsupportedVersion {
found: u16,
newest: u16,
},
#[error("model data is corrupted: {reason}")]
Malformed {
reason: &'static str,
},
#[error("model has {found} {end} neurons but {expected} were expected")]
DimensionMismatch {
end: &'static str,
expected: usize,
found: usize,
},
#[error("model file contains a network with no layers")]
EmptyNetwork,
#[error("layer {layer} of the model is inconsistent: {reason}")]
InconsistentLayer {
layer: usize,
reason: &'static str,
},
}
pub fn dump_model<const IN: usize, const OUT: usize>(
model: &NeuralNetwork<IN, OUT>,
path: impl AsRef<Path>,
) -> Result<(), ModelError> {
let mut file = BufWriter::new(File::create(path)?);
write_model(model, &mut file)?;
file.flush()?;
Ok(())
}
pub fn load_model<const IN: usize, const OUT: usize>(
path: impl AsRef<Path>,
) -> Result<NeuralNetwork<IN, OUT>, ModelError> {
let mut file = BufReader::new(File::open(path)?);
let model = read_model(&mut file)?;
if file.read(&mut [0])? != 0 {
return Err(ModelError::Malformed {
reason: "the file continues after the model",
});
}
Ok(model)
}
pub fn write_model<const IN: usize, const OUT: usize>(
model: &NeuralNetwork<IN, OUT>,
mut writer: impl Write,
) -> Result<(), ModelError> {
let layers = model.layers();
let mut bytes = Vec::with_capacity(MAGIC.len() + 10 + 8 * model.parameter_count() + 5 * layers.len());
bytes.extend_from_slice(&MAGIC);
bytes.extend_from_slice(&MODEL_FORMAT_VERSION.to_le_bytes());
bytes.extend_from_slice(&format_count(IN).to_le_bytes());
bytes.extend_from_slice(&format_count(layers.len()).to_le_bytes());
for layer in layers {
bytes.extend_from_slice(&format_count(layer.size()).to_le_bytes());
bytes.push(layer.activation().code());
for row in layer.weights().row_iter() {
for value in &row {
bytes.extend_from_slice(&value.to_le_bytes());
}
}
for value in layer.biases() {
bytes.extend_from_slice(&value.to_le_bytes());
}
}
writer.write_all(&bytes)?;
Ok(())
}
pub fn read_model<const IN: usize, const OUT: usize>(
reader: impl Read,
) -> Result<NeuralNetwork<IN, OUT>, ModelError> {
let mut decoder = Decoder::new(reader);
let start: [u8; 8] = decoder.bytes().map_err(|_| ModelError::UnrecognizedFormat)?;
let data = if start == MAGIC {
match decoder.u16()? {
1 => read_v1(&mut decoder)?,
found => {
return Err(ModelError::UnsupportedVersion {
found,
newest: MODEL_FORMAT_VERSION,
})
}
}
} else {
legacy::read(u64::from_le_bytes(start), &mut decoder)?
};
NeuralNetwork::try_from(data)
}
fn read_v1<R: Read>(decoder: &mut Decoder<R>) -> Result<NetworkData, ModelError> {
let mut inputs = decoder.u32()? as usize;
let layer_count = decoder.u32()?;
let mut layers = Vec::new();
for _ in 0..layer_count {
let neurons = decoder.u32()? as usize;
if neurons == 0 || inputs == 0 {
return Err(ModelError::Malformed {
reason: "a layer has no neurons",
});
}
let activation = ActivationFunction::from_code(decoder.u8()?).ok_or(ModelError::Malformed {
reason: "unknown activation function code",
})?;
let mut weights = Vec::new();
for _ in 0..neurons {
weights.push(decoder.f64s(inputs)?);
}
let biases = decoder.f64s(neurons)?;
layers.push(LayerData {
activation,
weights,
biases,
});
inputs = neurons;
}
Ok(NetworkData { layers })
}
fn format_count(count: usize) -> u32 {
u32::try_from(count).expect("a layer cannot have more than u32::MAX neurons")
}
fn fill(reader: &mut impl Read, bytes: &mut [u8]) -> Result<(), ModelError> {
reader.read_exact(bytes).map_err(|error| match error.kind() {
std::io::ErrorKind::UnexpectedEof => ModelError::Malformed {
reason: "the data ends before the model does",
},
_ => ModelError::Io(error),
})
}
const DECODE_CHUNK: usize = 512;
struct Decoder<R> {
reader: R,
buffer: Box<[u8; 8 * DECODE_CHUNK]>,
}
impl<R: Read> Decoder<R> {
fn new(reader: R) -> Self {
Self {
reader,
buffer: Box::new([0; 8 * DECODE_CHUNK]),
}
}
fn bytes<const N: usize>(&mut self) -> Result<[u8; N], ModelError> {
let mut bytes = [0; N];
fill(&mut self.reader, &mut bytes)?;
Ok(bytes)
}
fn u8(&mut self) -> Result<u8, ModelError> {
Ok(u8::from_le_bytes(self.bytes()?))
}
fn u16(&mut self) -> Result<u16, ModelError> {
Ok(u16::from_le_bytes(self.bytes()?))
}
fn u32(&mut self) -> Result<u32, ModelError> {
Ok(u32::from_le_bytes(self.bytes()?))
}
fn u64(&mut self) -> Result<u64, ModelError> {
Ok(u64::from_le_bytes(self.bytes()?))
}
fn f64s(&mut self, count: usize) -> Result<Vec<f64>, ModelError> {
let mut values = Vec::with_capacity(count.min(8 * DECODE_CHUNK));
let mut remaining = count;
while remaining > 0 {
let chunk = remaining.min(DECODE_CHUNK);
let bytes = &mut self.buffer[..8 * chunk];
fill(&mut self.reader, bytes)?;
let (chunks, _) = bytes.as_chunks::<8>();
values.extend(chunks.iter().map(|value| f64::from_le_bytes(*value)));
remaining -= chunk;
}
Ok(values)
}
}