use std::io::{Read, Write};
use csv::{ReaderBuilder, WriterBuilder};
use ndarray::prelude::*;
use crate::{
datasets::Dataset,
io::CsvIO,
models::Labelled,
types::{Labels, Set},
};
pub type GaussType = f64;
pub type GaussSample = Array1<GaussType>;
#[derive(Clone, Debug)]
pub struct GaussTable {
labels: Labels,
values: Array2<GaussType>,
}
impl Labelled for GaussTable {
#[inline]
fn labels(&self) -> &Labels {
&self.labels
}
}
impl GaussTable {
pub fn new(mut labels: Labels, mut values: Array2<GaussType>) -> Self {
assert_eq!(
labels.len(),
values.ncols(),
"Number of labels must match number of columns in values."
);
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;
}
assert!(
values.iter().all(|&x| x.is_finite()),
"Values must have finite values."
);
Self { labels, values }
}
}
impl Dataset for GaussTable {
type Values = Array2<GaussType>;
#[inline]
fn values(&self) -> &Self::Values {
&self.values
}
#[inline]
fn sample_size(&self) -> f64 {
self.values.nrows() as f64
}
fn select(&self, x: &Set<usize>) -> Self {
x.iter().for_each(|&i| {
assert!(
i < self.values.ncols(),
"Index out of bounds in variables selection: \n\
\t expected: index < |columns| , \n\
\t found: index == {} and |columns| == {} .",
i,
self.values.ncols()
);
});
let labels: Labels = x
.iter()
.map(|&i| self.labels.get_index(i).unwrap())
.cloned()
.collect();
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) -> Self {
let mut reader = ReaderBuilder::new().has_headers(true).from_reader(reader);
assert!(reader.has_headers(), "Reader must have headers.");
let labels: Labels = reader
.headers()
.expect("Failed to read the headers.")
.into_iter()
.map(|x| x.to_owned())
.collect();
let values: Array1<_> = reader
.into_records()
.enumerate()
.flat_map(|(i, row)| {
let row = row.unwrap_or_else(|_| panic!("Malformed record on line {}.", i + 1));
let row: Vec<_> = row
.into_iter()
.enumerate()
.map(|(i, x)| {
assert!(!x.is_empty(), "Missing value on line {}.", i + 1);
x.parse::<GaussType>().unwrap()
})
.collect();
row
})
.collect();
let ncols = labels.len();
let nrows = values.len() / ncols;
let values = values
.into_shape_with_order((nrows, ncols))
.expect("Failed to rearrange values to the correct shape.");
Self::new(labels, values)
}
fn to_csv_writer<W: Write>(&self, writer: W) {
let mut writer = WriterBuilder::new().has_headers(true).from_writer(writer);
writer
.write_record(self.labels.iter())
.expect("Failed to write CSV headers.");
self.values.rows().into_iter().for_each(|row| {
let record = row.iter().map(|x| x.to_string());
writer
.write_record(record)
.expect("Failed to write CSV record.");
});
}
}