use csv::{ReaderBuilder, Trim};
use dataset_core::{Dataset, DatasetError, acquire_dataset, download_to};
use ndarray::{Array1, Array2};
use std::fs::File;
type AdultData = (Array2<String>, Array2<f64>, Array1<String>);
const ADULT_DATA_URL: &str =
"https://archive.ics.uci.edu/ml/machine-learning-databases/adult/adult.data";
const ADULT_FILENAME: &str = "adult.csv";
const ADULT_SHA256: &str = "5b00264637dbfec36bdeaab5676b0b309ff9eb788d63554ca0a249491c86603d";
const ADULT_DATASET_NAME: &str = "adult";
const N_SAMPLES: usize = 32_561;
const N_STRING_FEATURES: usize = 8;
const N_NUMERIC_FEATURES: usize = 6;
const N_COLUMNS: usize = 15;
const LABEL_COLUMN: usize = 14;
const STRING_COLUMNS: [(usize, &str); N_STRING_FEATURES] = [
(1, "workclass"),
(3, "education"),
(5, "marital-status"),
(6, "occupation"),
(7, "relationship"),
(8, "race"),
(9, "sex"),
(13, "native-country"),
];
const NUMERIC_COLUMNS: [(usize, &str); N_NUMERIC_FEATURES] = [
(0, "age"),
(2, "fnlwgt"),
(4, "education-num"),
(10, "capital-gain"),
(11, "capital-loss"),
(12, "hours-per-week"),
];
const MISSING_TOKEN: &str = "?";
#[derive(Debug)]
pub struct Adult {
dataset: Dataset<AdultData, DatasetError>,
}
impl Adult {
pub fn new(storage_dir: &str) -> Self {
Adult {
dataset: Dataset::new(storage_dir, Self::load_data),
}
}
fn load_data(dir: &str) -> Result<AdultData, DatasetError> {
let file_path = acquire_dataset(
dir,
ADULT_FILENAME,
ADULT_DATASET_NAME,
Some(ADULT_SHA256),
|temp_path| {
download_to(ADULT_DATA_URL, temp_path, Some(ADULT_FILENAME))?;
Ok(temp_path.join(ADULT_FILENAME))
},
)?;
let file = File::open(&file_path)?;
let mut rdr = ReaderBuilder::new()
.has_headers(false)
.trim(Trim::All)
.from_reader(file);
let mut string_features: Vec<String> = Vec::with_capacity(N_SAMPLES * N_STRING_FEATURES);
let mut numeric_features: Vec<f64> = Vec::with_capacity(N_SAMPLES * N_NUMERIC_FEATURES);
let mut labels: Vec<String> = Vec::with_capacity(N_SAMPLES);
for (idx, result) in rdr.records().enumerate() {
let record = result.map_err(|e| DatasetError::csv_read_error(ADULT_DATASET_NAME, e))?;
let line_num = idx + 1;
if record.iter().all(|f| f.is_empty()) {
continue;
}
if record.len() != N_COLUMNS {
return Err(DatasetError::invalid_column_count(
ADULT_DATASET_NAME,
N_COLUMNS,
record.len(),
line_num,
));
}
for &(col, _name) in STRING_COLUMNS.iter() {
let value = &record[col];
if value == MISSING_TOKEN {
string_features.push(String::new());
} else {
string_features.push(value.to_string());
}
}
for &(col, name) in NUMERIC_COLUMNS.iter() {
let value: f64 = record[col].parse().map_err(|e| {
DatasetError::parse_failed(ADULT_DATASET_NAME, name, line_num, e)
})?;
numeric_features.push(value);
}
let label = &record[LABEL_COLUMN];
if label.is_empty() {
return Err(DatasetError::invalid_value(
ADULT_DATASET_NAME,
"income",
label,
line_num,
));
}
labels.push(label.to_string());
}
let n_samples = labels.len();
if n_samples == 0 {
return Err(DatasetError::empty_dataset(ADULT_DATASET_NAME));
}
let string_array = Array2::from_shape_vec((n_samples, N_STRING_FEATURES), string_features)
.map_err(|e| {
DatasetError::array_shape_error(ADULT_DATASET_NAME, "string_features", e)
})?;
let numeric_array =
Array2::from_shape_vec((n_samples, N_NUMERIC_FEATURES), numeric_features).map_err(
|e| DatasetError::array_shape_error(ADULT_DATASET_NAME, "numeric_features", e),
)?;
let labels_array = Array1::from_vec(labels);
Ok((string_array, numeric_array, labels_array))
}
pub fn features(&self) -> Result<(&Array2<String>, &Array2<f64>), DatasetError> {
let data = self.dataset.load()?;
Ok((&data.0, &data.1))
}
pub fn labels(&self) -> Result<&Array1<String>, DatasetError> {
Ok(&self.dataset.load()?.2)
}
pub fn data(&self) -> Result<&AdultData, DatasetError> {
self.dataset.load()
}
pub fn get_data(&self) -> Option<&AdultData> {
self.dataset.get()
}
pub fn get_data_mut(&mut self) -> Option<&mut AdultData> {
self.dataset.get_mut()
}
pub fn into_data(self) -> Result<AdultData, DatasetError> {
self.dataset.load()?;
Ok(self
.dataset
.into_inner()
.expect("data is present after a successful load"))
}
pub fn take_data(&mut self) -> Result<AdultData, DatasetError> {
self.dataset.load()?;
Ok(self
.dataset
.take()
.expect("data is present after a successful load"))
}
}