pub mod red_wine_quality;
pub mod white_wine_quality;
use crate::table::{Column, ColumnData, Table};
use csv::ReaderBuilder;
use dataset_core::DatasetError;
use ndarray::Array1;
use serde::Deserialize;
const N_FEATURES: usize = 11;
const FEATURE_NAMES: [&str; N_FEATURES] = [
"fixed acidity",
"volatile acidity",
"citric acid",
"residual sugar",
"chlorides",
"free sulfur dioxide",
"total sulfur dioxide",
"density",
"pH",
"sulphates",
"alcohol",
];
const TARGET: &str = "quality";
#[derive(Deserialize)]
struct WineRecord {
fixed_acidity: f64,
volatile_acidity: f64,
citric_acid: f64,
residual_sugar: f64,
chlorides: f64,
free_sulfur_dioxide: f64,
total_sulfur_dioxide: f64,
density: f64,
ph: f64,
sulphates: f64,
alcohol: f64,
quality: f64,
}
fn parse_wine_data_to_table<R: std::io::Read>(
dataset_name: &'static str,
reader: R,
) -> Result<Table, DatasetError> {
let mut rdr = ReaderBuilder::new()
.delimiter(b';')
.has_headers(false)
.from_reader(reader);
let mut fixed_acidity = Vec::new();
let mut volatile_acidity = Vec::new();
let mut citric_acid = Vec::new();
let mut residual_sugar = Vec::new();
let mut chlorides = Vec::new();
let mut free_sulfur_dioxide = Vec::new();
let mut total_sulfur_dioxide = Vec::new();
let mut density = Vec::new();
let mut ph = Vec::new();
let mut sulphates = Vec::new();
let mut alcohol = Vec::new();
let mut quality = Vec::new();
for result in rdr.deserialize::<WineRecord>().skip(1) {
let record = result.map_err(|e| DatasetError::csv_read_error(dataset_name, e))?;
fixed_acidity.push(record.fixed_acidity);
volatile_acidity.push(record.volatile_acidity);
citric_acid.push(record.citric_acid);
residual_sugar.push(record.residual_sugar);
chlorides.push(record.chlorides);
free_sulfur_dioxide.push(record.free_sulfur_dioxide);
total_sulfur_dioxide.push(record.total_sulfur_dioxide);
density.push(record.density);
ph.push(record.ph);
sulphates.push(record.sulphates);
alcohol.push(record.alcohol);
quality.push(record.quality);
}
Table::new(
dataset_name,
vec![
Column::new(
FEATURE_NAMES[0],
ColumnData::Numeric(Array1::from_vec(fixed_acidity)),
),
Column::new(
FEATURE_NAMES[1],
ColumnData::Numeric(Array1::from_vec(volatile_acidity)),
),
Column::new(
FEATURE_NAMES[2],
ColumnData::Numeric(Array1::from_vec(citric_acid)),
),
Column::new(
FEATURE_NAMES[3],
ColumnData::Numeric(Array1::from_vec(residual_sugar)),
),
Column::new(
FEATURE_NAMES[4],
ColumnData::Numeric(Array1::from_vec(chlorides)),
),
Column::new(
FEATURE_NAMES[5],
ColumnData::Numeric(Array1::from_vec(free_sulfur_dioxide)),
),
Column::new(
FEATURE_NAMES[6],
ColumnData::Numeric(Array1::from_vec(total_sulfur_dioxide)),
),
Column::new(
FEATURE_NAMES[7],
ColumnData::Numeric(Array1::from_vec(density)),
),
Column::new(FEATURE_NAMES[8], ColumnData::Numeric(Array1::from_vec(ph))),
Column::new(
FEATURE_NAMES[9],
ColumnData::Numeric(Array1::from_vec(sulphates)),
),
Column::new(
FEATURE_NAMES[10],
ColumnData::Numeric(Array1::from_vec(alcohol)),
),
Column::new(TARGET, ColumnData::Numeric(Array1::from_vec(quality))),
],
)
}