use dataset_core::DatasetError;
use ndarray::{Array1, Array2};
#[derive(Debug, Clone, PartialEq)]
pub enum ColumnData {
Numeric(Array1<f64>),
Integer(Array1<i64>),
String(Array1<String>),
Bytes(Array2<u8>),
}
impl ColumnData {
pub fn len(&self) -> usize {
match self {
ColumnData::Numeric(values) => values.len(),
ColumnData::Integer(values) => values.len(),
ColumnData::String(values) => values.len(),
ColumnData::Bytes(values) => values.nrows(),
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn kind(&self) -> &'static str {
match self {
ColumnData::Numeric(_) => "numeric",
ColumnData::Integer(_) => "integer",
ColumnData::String(_) => "string",
ColumnData::Bytes(_) => "bytes",
}
}
pub fn width(&self) -> usize {
match self {
ColumnData::Bytes(values) => values.ncols(),
_ => 1,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Column {
name: &'static str,
data: ColumnData,
}
impl Column {
pub fn new(name: &'static str, data: ColumnData) -> Self {
Column { name, data }
}
pub fn name(&self) -> &'static str {
self.name
}
pub fn data(&self) -> &ColumnData {
&self.data
}
pub fn data_mut(&mut self) -> &mut ColumnData {
&mut self.data
}
pub fn len(&self) -> usize {
self.data.len()
}
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
pub fn as_numeric(&self) -> Option<&Array1<f64>> {
match &self.data {
ColumnData::Numeric(values) => Some(values),
_ => None,
}
}
pub fn as_integer(&self) -> Option<&Array1<i64>> {
match &self.data {
ColumnData::Integer(values) => Some(values),
_ => None,
}
}
pub fn as_string(&self) -> Option<&Array1<String>> {
match &self.data {
ColumnData::String(values) => Some(values),
_ => None,
}
}
pub fn as_bytes(&self) -> Option<&Array2<u8>> {
match &self.data {
ColumnData::Bytes(values) => Some(values),
_ => None,
}
}
pub fn to_numeric(&self) -> Option<Array1<f64>> {
match &self.data {
ColumnData::Numeric(values) => Some(values.clone()),
ColumnData::Integer(values) => Some(values.mapv(|value| value as f64)),
_ => None,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Table {
name: &'static str,
columns: Vec<Column>,
n_samples: usize,
}
impl Table {
pub fn new(name: &'static str, columns: Vec<Column>) -> Result<Self, DatasetError> {
let Some(first) = columns.first() else {
return Err(DatasetError::empty_dataset(name));
};
let n_samples = first.len();
if n_samples == 0 {
return Err(DatasetError::empty_dataset(name));
}
for column in &columns {
if column.len() != n_samples {
return Err(DatasetError::length_mismatch(
name,
column.name(),
n_samples,
column.len(),
));
}
}
for (index, column) in columns.iter().enumerate() {
if columns[..index]
.iter()
.any(|other| other.name() == column.name())
{
return Err(DatasetError::invalid_value(
name,
"column name",
column.name(),
index + 1,
));
}
}
Ok(Table {
name,
columns,
n_samples,
})
}
pub fn name(&self) -> &'static str {
self.name
}
pub fn n_samples(&self) -> usize {
self.n_samples
}
pub fn n_columns(&self) -> usize {
self.columns.len()
}
pub fn columns(&self) -> &[Column] {
&self.columns
}
pub fn columns_mut(&mut self) -> &mut [Column] {
&mut self.columns
}
pub fn names(&self) -> impl Iterator<Item = &'static str> + '_ {
self.columns.iter().map(Column::name)
}
pub fn column(&self, name: &str) -> Option<&Column> {
self.columns.iter().find(|column| column.name() == name)
}
pub fn column_mut(&mut self, name: &str) -> Option<&mut Column> {
self.columns.iter_mut().find(|column| column.name() == name)
}
pub fn numeric_matrix(&self, names: &[&str]) -> Result<Array2<f64>, DatasetError> {
if names.is_empty() {
return Err(DatasetError::length_mismatch(
self.name,
"requested columns",
1,
0,
));
}
let mut selected = Vec::with_capacity(names.len());
for name in names {
let Some(column) = self.column(name) else {
return Err(DatasetError::unknown_column(self.name, name));
};
selected.push(column);
}
let width: usize = selected.iter().map(|column| column.data().width()).sum();
let mut values: Vec<f64> = Vec::with_capacity(self.n_samples * width);
for row in 0..self.n_samples {
for column in &selected {
match column.data() {
ColumnData::Numeric(source) => values.push(source[row]),
ColumnData::Integer(source) => values.push(source[row] as f64),
ColumnData::Bytes(source) => {
for col in 0..source.ncols() {
values.push(f64::from(source[[row, col]]));
}
}
other => {
return Err(DatasetError::column_type_mismatch(
self.name,
column.name(),
"numeric",
other.kind(),
));
}
}
}
}
Array2::from_shape_vec((self.n_samples, width), values)
.map_err(|e| DatasetError::array_shape_error(self.name, "numeric matrix", e))
}
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
fn numeric(name: &'static str, values: [f64; 3]) -> Column {
Column::new(name, ColumnData::Numeric(Array1::from_vec(values.to_vec())))
}
fn sample_table() -> Table {
Table::new(
"sample",
vec![
numeric("a", [1.0, 2.0, 3.0]),
numeric("b", [4.0, 5.0, 6.0]),
Column::new(
"label",
ColumnData::String(array!["x".into(), "y".into(), "x".into()]),
),
],
)
.unwrap()
}
#[test]
fn new_rejects_an_empty_column_list() {
assert!(Table::new("t", vec![]).is_err());
}
#[test]
fn new_rejects_zero_samples() {
let column = Column::new("a", ColumnData::Numeric(Array1::zeros(0)));
assert!(Table::new("t", vec![column]).is_err());
}
#[test]
fn new_rejects_columns_of_different_lengths() {
let short = Column::new("a", ColumnData::Numeric(array![1.0, 2.0]));
let long = Column::new("b", ColumnData::Numeric(array![1.0, 2.0, 3.0]));
let error = Table::new("t", vec![short, long]).unwrap_err().to_string();
assert!(error.contains("expected 2"), "{error}");
}
#[test]
fn new_rejects_a_repeated_name() {
let one = numeric("a", [1.0, 2.0, 3.0]);
let two = numeric("a", [4.0, 5.0, 6.0]);
assert!(Table::new("t", vec![one, two]).is_err());
}
#[test]
fn a_table_reports_its_name_shape_and_column_names() {
let table = sample_table();
assert_eq!(table.name(), "sample");
assert_eq!(table.n_samples(), 3);
assert_eq!(table.n_columns(), 3);
assert_eq!(table.names().collect::<Vec<_>>(), vec!["a", "b", "label"]);
}
#[test]
fn column_lookup_is_by_name_not_position() {
let table = sample_table();
assert_eq!(table.column("b").unwrap().as_numeric().unwrap()[0], 4.0);
assert!(table.column("missing").is_none());
}
#[test]
fn numeric_matrix_keeps_the_requested_order() {
let table = sample_table();
let matrix = table.numeric_matrix(&["b", "a"]).unwrap();
assert_eq!(matrix.shape(), &[3, 2]);
assert_eq!(matrix.row(0).to_vec(), vec![4.0, 1.0]);
assert_eq!(matrix.row(2).to_vec(), vec![6.0, 3.0]);
}
#[test]
fn numeric_matrix_takes_a_subset() {
let table = sample_table();
let matrix = table.numeric_matrix(&["a"]).unwrap();
assert_eq!(matrix.shape(), &[3, 1]);
}
#[test]
fn numeric_matrix_repeats_a_repeated_name() {
let table = sample_table();
let matrix = table.numeric_matrix(&["a", "a"]).unwrap();
assert_eq!(matrix.shape(), &[3, 2]);
assert_eq!(matrix.row(1).to_vec(), vec![2.0, 2.0]);
}
#[test]
fn numeric_matrix_converts_integers() {
let table = Table::new(
"t",
vec![
Column::new("count", ColumnData::Integer(array![1, 2])),
Column::new("when", ColumnData::Integer(array![10, 20])),
],
)
.unwrap();
let matrix = table.numeric_matrix(&["count", "when"]).unwrap();
assert_eq!(matrix.row(1).to_vec(), vec![2.0, 20.0]);
}
#[test]
fn numeric_matrix_expands_a_bytes_column_to_its_width() {
let pixels = Array2::from_shape_vec((2, 3), vec![1u8, 2, 3, 4, 5, 6]).unwrap();
let table =
Table::new("t", vec![Column::new("pixels", ColumnData::Bytes(pixels))]).unwrap();
let matrix = table.numeric_matrix(&["pixels"]).unwrap();
assert_eq!(matrix.shape(), &[2, 3]);
assert_eq!(matrix.row(1).to_vec(), vec![4.0, 5.0, 6.0]);
}
#[test]
fn numeric_matrix_rejects_an_empty_request() {
let table = sample_table();
let error = table.numeric_matrix(&[]).unwrap_err().to_string();
assert!(error.contains("requested columns"), "{error}");
}
#[test]
fn numeric_matrix_rejects_an_unknown_name() {
let table = sample_table();
let error = table.numeric_matrix(&["a", "missing"]).unwrap_err();
let message = error.to_string();
assert!(message.contains("no column named `missing`"), "{message}");
}
#[test]
fn numeric_matrix_rejects_a_string_column() {
let table = sample_table();
let message = table.numeric_matrix(&["label"]).unwrap_err().to_string();
assert!(message.contains("`label`"), "{message}");
assert!(message.contains("expected `numeric`"), "{message}");
}
#[test]
fn numeric_matrix_names_the_string_column_wherever_it_sits() {
let table = sample_table();
let message = table
.numeric_matrix(&["a", "label", "b"])
.unwrap_err()
.to_string();
assert!(message.contains("`label`"), "{message}");
assert!(message.contains("`string`"), "{message}");
}
#[test]
fn to_numeric_reads_numbers_and_refuses_the_rest() {
let pixels = Array2::<u8>::zeros((3, 4));
let table = Table::new(
"t",
vec![
numeric("a", [1.0, 2.0, 3.0]),
Column::new("count", ColumnData::Integer(array![1, 2, 3])),
Column::new(
"label",
ColumnData::String(Array1::from_vec(vec!["x".to_string(); 3])),
),
Column::new("pixels", ColumnData::Bytes(pixels)),
],
)
.unwrap();
assert_eq!(table.column("a").unwrap().to_numeric().unwrap()[0], 1.0);
assert_eq!(table.column("count").unwrap().to_numeric().unwrap()[2], 3.0);
assert!(table.column("label").unwrap().to_numeric().is_none());
assert!(table.column("pixels").unwrap().to_numeric().is_none());
}
#[test]
fn column_mut_edits_in_place() {
let mut table = sample_table();
if let Some(ColumnData::Numeric(values)) = table.column_mut("a").map(Column::data_mut) {
values[0] = 99.0;
}
assert_eq!(table.column("a").unwrap().as_numeric().unwrap()[0], 99.0);
}
#[test]
fn width_is_one_except_for_bytes() {
assert_eq!(ColumnData::Numeric(array![1.0]).width(), 1);
assert_eq!(ColumnData::Integer(array![1]).width(), 1);
assert_eq!(ColumnData::String(array!["x".to_string()]).width(), 1);
let pixels = Array2::<u8>::zeros((1, 5));
assert_eq!(ColumnData::Bytes(pixels).width(), 5);
}
}