use std::path::Path;
use crate::error::{Error, Result};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Dtype {
Numeric,
Categorical,
}
#[derive(Clone, Debug, PartialEq)]
pub struct Frame {
buf: Vec<f64>,
nrows: usize,
ncols: usize,
columns: Vec<String>,
dtypes: Vec<Dtype>,
}
impl Frame {
pub fn new(buf: Vec<f64>, nrows: usize, ncols: usize, columns: Vec<String>) -> Result<Self> {
if buf.len() != nrows * ncols {
return Err(Error::Shape(format!(
"buffer has {} elements but shape {}x{} needs {}",
buf.len(),
nrows,
ncols,
nrows * ncols
)));
}
if columns.len() != ncols {
return Err(Error::Schema(format!(
"{} column names for {} columns",
columns.len(),
ncols
)));
}
Ok(Frame {
buf,
nrows,
ncols,
columns,
dtypes: vec![Dtype::Numeric; ncols],
})
}
pub fn with_dtypes(mut self, dtypes: Vec<Dtype>) -> Result<Self> {
if dtypes.len() != self.ncols {
return Err(Error::Schema(format!(
"{} dtypes for {} columns",
dtypes.len(),
self.ncols
)));
}
self.dtypes = dtypes;
Ok(self)
}
pub fn from_rows(rows: Vec<Vec<f64>>, columns: Vec<String>) -> Result<Self> {
let nrows = rows.len();
let ncols = columns.len();
let mut buf = Vec::with_capacity(nrows * ncols);
for (r, row) in rows.into_iter().enumerate() {
if row.len() != ncols {
return Err(Error::Shape(format!(
"row {r} has {} values, expected {ncols}",
row.len()
)));
}
buf.extend(row);
}
Frame::new(buf, nrows, ncols, columns)
}
pub fn from_csv(path: impl AsRef<Path>) -> Result<Frame> {
let text = std::fs::read_to_string(path.as_ref())
.map_err(|e| Error::Backend(format!("read csv: {e}")))?;
let mut lines = text.lines().filter(|l| !l.trim().is_empty());
let header = lines
.next()
.ok_or_else(|| Error::Schema("empty CSV".into()))?;
let columns: Vec<String> = header.split(',').map(|s| s.trim().to_string()).collect();
let ncols = columns.len();
let mut rows = Vec::new();
for (i, line) in lines.enumerate() {
let mut row = Vec::with_capacity(ncols);
for cell in line.split(',') {
let t = cell.trim();
row.push(if t.is_empty() {
f64::NAN
} else {
t.parse::<f64>().map_err(|_| {
Error::Schema(format!("row {}: '{t}' is not a number", i + 2))
})?
});
}
rows.push(row);
}
Frame::from_rows(rows, columns)
}
pub fn nrows(&self) -> usize {
self.nrows
}
pub fn ncols(&self) -> usize {
self.ncols
}
pub fn shape(&self) -> (usize, usize) {
(self.nrows, self.ncols)
}
pub fn columns(&self) -> &[String] {
&self.columns
}
pub fn dtypes(&self) -> &[Dtype] {
&self.dtypes
}
pub fn dtype(&self, c: usize) -> Dtype {
self.dtypes[c]
}
pub fn categorical_columns(&self) -> Vec<usize> {
(0..self.ncols)
.filter(|&c| self.dtypes[c] == Dtype::Categorical)
.collect()
}
pub fn buf(&self) -> &[f64] {
&self.buf
}
pub fn column_index(&self, name: &str) -> Option<usize> {
self.columns.iter().position(|c| c == name)
}
pub fn row(&self, r: usize) -> &[f64] {
let start = r * self.ncols;
&self.buf[start..start + self.ncols]
}
pub fn get(&self, r: usize, c: usize) -> f64 {
self.buf[r * self.ncols + c]
}
pub fn column(&self, c: usize) -> Vec<f64> {
(0..self.nrows).map(|r| self.get(r, c)).collect()
}
pub fn as_rows(&self) -> Vec<Vec<f64>> {
(0..self.nrows).map(|r| self.row(r).to_vec()).collect()
}
pub fn select_rows(&self, idx: &[usize]) -> Frame {
let mut buf = Vec::with_capacity(idx.len() * self.ncols);
for &r in idx {
buf.extend_from_slice(self.row(r));
}
Frame {
buf,
nrows: idx.len(),
ncols: self.ncols,
columns: self.columns.clone(),
dtypes: self.dtypes.clone(),
}
}
pub(crate) fn require_columns(&self, expected: &[String]) -> Result<()> {
if self.columns != expected {
return Err(Error::Schema(format!(
"expected columns {expected:?}, got {:?}",
self.columns
)));
}
Ok(())
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct Dataset {
features: Frame,
target: Vec<f64>,
}
impl Dataset {
pub fn new(features: Frame, target: Vec<f64>) -> Result<Self> {
if target.len() != features.nrows() {
return Err(Error::Shape(format!(
"target has {} values but frame has {} rows",
target.len(),
features.nrows()
)));
}
Ok(Dataset { features, target })
}
pub fn features(&self) -> &Frame {
&self.features
}
pub fn target(&self) -> &[f64] {
&self.target
}
pub(crate) fn with_features(&self, features: Frame) -> Dataset {
Dataset {
features,
target: self.target.clone(),
}
}
pub fn select(&self, idx: &[usize]) -> Dataset {
Dataset {
features: self.features.select_rows(idx),
target: idx.iter().map(|&i| self.target[i]).collect(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_mismatched_buffer() {
assert!(Frame::new(vec![1.0, 2.0], 2, 2, vec!["a".into(), "b".into()]).is_err());
}
#[test]
fn round_trips_rows() {
let f = Frame::from_rows(
vec![vec![1.0, 2.0], vec![3.0, 4.0]],
vec!["a".into(), "b".into()],
)
.unwrap();
assert_eq!(f.shape(), (2, 2));
assert_eq!(f.get(1, 0), 3.0);
assert_eq!(f.column(1), vec![2.0, 4.0]);
assert_eq!(f.as_rows(), vec![vec![1.0, 2.0], vec![3.0, 4.0]]);
}
#[test]
fn reads_numeric_csv_with_blanks_as_nan() {
let path = std::env::temp_dir().join("mw_frame_from_csv.csv");
std::fs::write(&path, "a,b\n1,2\n3,\n").unwrap();
let f = Frame::from_csv(&path).unwrap();
assert_eq!(f.shape(), (2, 2));
assert_eq!(f.get(0, 1), 2.0);
assert!(f.get(1, 1).is_nan());
let _ = std::fs::remove_file(&path);
}
#[test]
fn dtypes_default_numeric_and_survive_selection() {
let f = Frame::from_rows(
vec![vec![1.0, 2.0], vec![3.0, 4.0]],
vec!["a".into(), "b".into()],
)
.unwrap();
assert_eq!(f.dtypes(), &[Dtype::Numeric, Dtype::Numeric]);
let typed = f
.with_dtypes(vec![Dtype::Categorical, Dtype::Numeric])
.unwrap();
assert_eq!(typed.categorical_columns(), vec![0]);
assert_eq!(typed.dtype(0), Dtype::Categorical);
assert_eq!(
typed.select_rows(&[1]).dtypes(),
&[Dtype::Categorical, Dtype::Numeric]
);
}
#[test]
fn dataset_checks_lengths() {
let f = Frame::from_rows(vec![vec![1.0], vec![2.0]], vec!["a".into()]).unwrap();
assert!(Dataset::new(f.clone(), vec![0.0]).is_err());
assert!(Dataset::new(f, vec![0.0, 1.0]).is_ok());
}
}