mod binary;
mod pure;
use std::{io, path::Path, sync::Arc};
use nalgebra::Const;
use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use serde::de::DeserializeOwned;
use crate::Residual;
use crate::ad::Gradient;
use super::ParametersAD;
pub use binary::*;
pub use pure::*;
struct DatasetData {
inputs: Array2<f64>,
target: Array1<f64>,
}
#[derive(Clone)]
struct DatasetStorage {
data: Arc<DatasetData>,
name: Option<String>,
}
impl DatasetStorage {
fn from_records<R: DatasetRecord>(records: Vec<R>) -> Self {
let n = records.len();
let inputs = Array2::from_shape_fn((n, R::N_INPUTS), |(i, j)| records[i].input(j));
let target = Array1::from_iter(records.iter().map(DatasetRecord::target));
Self {
data: Arc::new(DatasetData { inputs, target }),
name: None,
}
}
fn from_csv<R: DatasetRecord>(path: &Path) -> Result<Self, csv::Error> {
let records = csv::Reader::from_path(path)?
.deserialize()
.collect::<Result<Vec<R>, _>>()?;
Ok(Self::from_records(records))
}
fn from_reader<R: DatasetRecord>(reader: impl io::Read) -> Result<Self, csv::Error> {
let records = csv::Reader::from_reader(reader)
.deserialize()
.collect::<Result<Vec<R>, _>>()?;
Ok(Self::from_records(records))
}
fn inputs(&self) -> ArrayView2<'_, f64> {
self.data.inputs.view()
}
fn target(&self) -> ArrayView1<'_, f64> {
self.data.target.view()
}
fn name(&self) -> Option<&str> {
self.name.as_deref()
}
fn set_name(&mut self, name: String) {
self.name = Some(name);
}
}
pub trait DatasetRecord: DeserializeOwned {
const N_INPUTS: usize;
fn input(&self, column: usize) -> f64;
fn target(&self) -> f64;
}
pub trait Dataset {
fn inputs(&self) -> ArrayView2<'_, f64>;
fn target(&self) -> ArrayView1<'_, f64>;
fn name(&self) -> &str;
fn input_names(&self) -> &'static [&'static str];
fn target_name(&self) -> &'static str;
fn evaluate<E: Residual + Sync>(&self, eos: &E) -> (Array1<f64>, Array1<bool>);
}
macro_rules! define_dataset_ad {
($($p:literal),+ $(,)?) => {
pub const GRADIENT_SLOTS: &[usize] = &[$($p),+];
pub trait DatasetAD<const N: usize>: Dataset {
fn evaluate_ad_const<T: ParametersAD<Const<N>>, const P: usize>(
&self,
names: [String; P],
parameters: &[f64],
inputs: ArrayView2<f64>,
) -> (Array1<f64>, Array2<f64>, Array1<bool>)
where
T::Lifted<Gradient<P>>: Sync;
fn evaluate_ad<T: ParametersAD<Const<N>>>(
&self,
param_names: &[String],
parameters: &[f64],
) -> (Array1<f64>, Array2<f64>, Array1<bool>)
where
$(T::Lifted<Gradient<$p>>: Sync,)*
{
fn to_const<const P: usize>(names: &[String]) -> [String; P] {
names.to_vec().try_into().expect("parameter count mismatch")
}
match param_names.len() {
$(
$p => self.evaluate_ad_const::<T, $p>(
to_const(param_names),
parameters,
self.inputs().view(),
),
)+
p => unreachable!(
"parameter count {p} is not a member of GRADIENT_SLOTS={:?}",
GRADIENT_SLOTS,
),
}
}
}
};
}
define_dataset_ad!(1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14);