use std::io::Read;
use super::Decoder;
use crate::neural_network::{LayerData, NetworkData};
use crate::{ActivationFunction, ModelError};
pub(super) fn read<R: Read>(layer_count: u64, decoder: &mut Decoder<R>) -> Result<NetworkData, ModelError> {
read_network(layer_count, decoder).map_err(|error| match error {
ModelError::Malformed { .. } => ModelError::UnrecognizedFormat,
error => error,
})
}
fn read_network<R: Read>(layer_count: u64, decoder: &mut Decoder<R>) -> Result<NetworkData, ModelError> {
let mut layers = Vec::new();
for index in 0..layer_count {
layers.push(read_layer(decoder, index as usize + 1)?);
}
let activation = match decoder.u8()? {
0 => ActivationFunction::default(),
1 => u8::try_from(decoder.u32()?)
.ok()
.and_then(ActivationFunction::from_code)
.ok_or(ModelError::UnrecognizedFormat)?,
_ => return Err(ModelError::UnrecognizedFormat),
};
for layer in &mut layers {
layer.activation = activation;
}
Ok(NetworkData { layers })
}
fn read_layer<R: Read>(decoder: &mut Decoder<R>, layer: usize) -> Result<LayerData, ModelError> {
let size = decoder.u64()?;
let columns_first = read_floats(decoder)?;
let rows = decoder.u64()?;
let columns = decoder.u64()?;
if rows.checked_mul(columns) != Some(columns_first.len() as u64) || (rows > 0 && columns == 0) {
return Err(ModelError::UnrecognizedFormat);
}
let biases = read_floats(decoder)?;
if decoder.u64()? != biases.len() as u64 {
return Err(ModelError::UnrecognizedFormat);
}
if size != rows {
return Err(ModelError::InconsistentLayer {
layer,
reason: "its size does not match its number of weight rows",
});
}
let (rows, columns) = (rows as usize, columns as usize);
let weights = (0..rows)
.map(|row| (0..columns).map(|column| columns_first[column * rows + row]).collect())
.collect();
Ok(LayerData {
activation: ActivationFunction::default(),
weights,
biases,
})
}
fn read_floats<R: Read>(decoder: &mut Decoder<R>) -> Result<Vec<f64>, ModelError> {
let length = decoder.u64()?;
let length = usize::try_from(length).map_err(|_| ModelError::UnrecognizedFormat)?;
decoder.f64s(length)
}