use std::{
io::{Read, Write},
sync::Arc,
};
use csv::{ReaderBuilder, WriterBuilder};
use ndarray::prelude::*;
use crate::{
datasets::{Dataset, GaussEv, GaussEvT},
io::CsvIO,
models::Labelled,
types::{Error, Labels, Result, Set},
};
pub type GaussType = f64;
pub type GaussSample = Array1<GaussType>;
#[derive(Clone, Debug)]
pub struct GaussTable {
labels: Labels,
values: Array2<GaussType>,
}
pub struct GaussTableEvidenceIter<'a> {
rows: ndarray::iter::LanesIter<'a, GaussType, Ix1>,
labels: &'a Labels,
}
impl<'a> Iterator for GaussTableEvidenceIter<'a> {
type Item = Result<GaussEv>;
fn next(&mut self) -> Option<Self::Item> {
let row = self.rows.next()?;
let evidences = row
.iter()
.enumerate()
.map(|(event, &value)| GaussEvT::CertainPositive { event, value });
Some(GaussEv::new(self.labels.clone(), evidences))
}
}
impl Labelled for GaussTable {
#[inline]
fn labels(&self) -> &Labels {
&self.labels
}
}
impl GaussTable {
pub fn new(mut labels: Labels, mut values: Array2<GaussType>) -> Result<Self> {
if labels.len() != values.ncols() {
return Err(Error::IncompatibleShape(
&labels.len().to_string(),
&values.ncols().to_string(),
));
}
if !labels.is_sorted() {
let mut indices: Vec<usize> = (0..labels.len()).collect();
indices.sort_by_key(|&i| &labels[i]);
labels.sort();
let mut new_values = values.clone();
indices.into_iter().enumerate().for_each(|(i, j)| {
new_values.column_mut(i).assign(&values.column(j));
});
values = new_values;
}
if !values.iter().all(|&x| x.is_finite()) {
return Err(Error::InvalidParameter("values", "must be finite"));
}
Ok(Self { labels, values })
}
}
impl Dataset for GaussTable {
type Values = Array2<GaussType>;
type Evidence = GaussEv;
type EvidenceIter<'a> = GaussTableEvidenceIter<'a>;
#[inline]
fn values(&self) -> &Self::Values {
&self.values
}
fn evidence_iter(&self) -> Self::EvidenceIter<'_> {
GaussTableEvidenceIter {
rows: self.values.rows().into_iter(),
labels: &self.labels,
}
}
#[inline]
fn sample_size(&self) -> f64 {
self.values.nrows() as f64
}
fn select(&self, x: &Set<usize>) -> Result<Self> {
if let Some(&i) = x.iter().find(|&&i| i >= self.values.ncols()) {
return Err(Error::IndexOutOfBounds(i));
}
let labels: Labels = x
.iter()
.map(|&i| {
self.labels
.get_index(i)
.cloned()
.ok_or_else(|| Error::IndexOutOfBounds(i))
})
.collect::<Result<_>>()?;
let mut new_values = Array2::zeros((self.values.nrows(), x.len()));
x.iter().enumerate().for_each(|(j, &i)| {
new_values.column_mut(j).assign(&self.values.column(i));
});
let values = new_values;
Self::new(labels, values)
}
}
impl CsvIO for GaussTable {
fn from_csv_reader<R: Read>(reader: R) -> Result<Self> {
let mut reader = ReaderBuilder::new().has_headers(true).from_reader(reader);
if !reader.has_headers() {
return Err(Error::MissingHeader());
}
let labels: Labels = reader
.headers()?
.into_iter()
.map(|x| x.to_owned())
.collect();
let values: Vec<GaussType> = reader
.into_records()
.enumerate()
.map(|(i, row)| {
let row = row.map_err(|e| Error::Csv(Arc::new(e)))?;
row.into_iter()
.enumerate()
.map(|(j, x)| {
if x.is_empty() {
return Err(Error::MissingValue(i + 1, j + 1));
}
Ok(x.parse::<GaussType>()?)
})
.collect::<Result<Vec<_>>>()
})
.collect::<Result<Vec<_>>>()?
.into_iter()
.flatten()
.collect();
let values = Array1::from_vec(values);
let ncols = labels.len();
let nrows = values.len() / ncols;
let values = values.into_shape_with_order((nrows, ncols))?;
Self::new(labels, values)
}
fn to_csv_writer<W: Write>(&self, writer: W) -> Result<()> {
let mut writer = WriterBuilder::new().has_headers(true).from_writer(writer);
writer.write_record(self.labels.iter())?;
self.values
.rows()
.into_iter()
.try_for_each(|row| -> Result<_> {
let record = row.iter().map(|x| x.to_string());
writer.write_record(record)?;
Ok(())
})?;
Ok(())
}
}