#![cfg(feature = "dataset")]
use dataset_ml::table::ColumnData;
use dataset_ml::traits::MlDataset;
use dataset_ml::{Digits, Iris, SmsSpam, Titanic};
use std::fs::remove_dir_all;
fn summarize<D: MlDataset>(dataset: &D) -> String {
format!("{} in {}", D::NAME, dataset.storage_dir())
}
#[test]
fn inspection_methods_do_not_load() {
let dataset = Iris::new("./test_traits_no_load");
assert_eq!(dataset.storage_dir(), "./test_traits_no_load");
assert!(!dataset.is_loaded());
assert!(dataset.peek().is_none());
assert_eq!(Iris::NAME, "iris");
assert!(!std::path::Path::new("./test_traits_no_load").exists());
}
#[test]
fn a_generic_function_accepts_every_loader() {
assert_eq!(summarize(&Iris::new("./a")), "iris in ./a");
assert_eq!(summarize(&SmsSpam::new("./b")), "sms_spam in ./b");
assert_eq!(summarize(&Titanic::new("./c")), "titanic in ./c");
assert_eq!(summarize(&Digits::new("./d")), "digits in ./d");
}
#[test]
fn load_peek_and_invalidate_cycle() {
let download_dir = "./test_traits_load_cycle";
let mut dataset = Iris::new(download_dir);
assert!(dataset.peek().is_none());
let table = dataset.load().unwrap();
assert_eq!(table.n_samples(), 150);
assert_eq!(table.n_columns(), 5);
assert!(dataset.is_loaded());
assert!(dataset.peek().is_some());
assert_eq!(dataset.n_samples().unwrap(), 150);
dataset.invalidate();
assert!(!dataset.is_loaded());
assert!(dataset.peek().is_none());
assert_eq!(dataset.n_samples().unwrap(), 150);
assert!(dataset.is_loaded());
remove_dir_all(download_dir).unwrap();
}
#[test]
fn n_samples_agrees_with_every_column() {
let download_dir = "./test_traits_n_samples";
let dataset = Iris::new(download_dir);
let table = dataset.load().unwrap();
assert_eq!(dataset.n_samples().unwrap(), table.n_samples());
for column in table.columns() {
assert_eq!(
column.len(),
table.n_samples(),
"column {} holds a different number of samples",
column.name()
);
}
assert_eq!(Iris::FEATURE_NAMES.len(), 4);
for name in Iris::FEATURE_NAMES {
assert!(
matches!(table.column(name).unwrap().data(), ColumnData::Numeric(_)),
"feature {name} should be numeric"
);
}
assert!(matches!(
table.column(Iris::TARGET).unwrap().data(),
ColumnData::String(_)
));
remove_dir_all(download_dir).unwrap();
}
#[test]
fn load_mut_and_unload_move_data_without_cloning() {
let download_dir = "./test_traits_load_mut_unload";
let mut dataset = Iris::new(download_dir);
let table = dataset.load_mut().unwrap();
if let Some(ColumnData::Numeric(values)) =
table.column_mut("sepal_length").map(|c| c.data_mut())
{
values[0] = 42.0;
}
assert_eq!(
dataset
.load()
.unwrap()
.column("sepal_length")
.unwrap()
.as_numeric()
.unwrap()[0],
42.0
);
let owned = dataset.unload().unwrap();
assert_eq!(
owned.column("sepal_length").unwrap().as_numeric().unwrap()[0],
42.0
);
assert_eq!(owned.n_samples(), 150);
assert!(!dataset.is_loaded());
assert_eq!(
dataset
.load()
.unwrap()
.column("sepal_length")
.unwrap()
.as_numeric()
.unwrap()[0],
5.1
);
dataset.invalidate();
assert!(dataset.unload().is_none());
remove_dir_all(download_dir).unwrap();
}
#[test]
fn into_dataset_yields_the_underlying_container() {
let dataset = Iris::new("./test_traits_into_dataset");
let container = dataset.into_dataset();
assert_eq!(container.storage_dir(), "./test_traits_into_dataset");
assert!(!container.is_loaded());
}