Skip to main content

contree/data/
dataset.rs

1use crate::data::{DataPoint, Feature};
2use std::ops::{Index, IndexMut};
3
4/// A labelled dataset stored column by column.
5///
6/// Before a search, every column must be sorted and its values indexed; see
7/// [`Dataset::is_prepared`]. [`Dataset::from_rows`] and the file reader do both.
8#[derive(Debug, Default)]
9pub struct Dataset {
10    features: Vec<Feature>,
11    num_labels: usize,
12    // Whether `sort_features` and `compute_unique_feature_values` have run.
13    // `fit` refuses a dataset that is not prepared.
14    sorted: bool,
15    indexed: bool,
16}
17
18impl Dataset {
19    /// An empty dataset.
20    pub fn new() -> Self {
21        Self::default()
22    }
23
24    /// Whether both preparation steps have run. See [`SearchError::UnpreparedDataset`].
25    ///
26    /// [`SearchError::UnpreparedDataset`]: crate::common::SearchError::UnpreparedDataset
27    pub fn is_prepared(&self) -> bool {
28        self.sorted && self.indexed
29    }
30
31    /// Number of instances. Zero for an empty dataset.
32    pub fn count(&self) -> usize {
33        self.features.first().map_or(0, Feature::len)
34    }
35
36    /// Whether the dataset has no instances.
37    pub fn is_empty(&self) -> bool {
38        self.count() == 0
39    }
40
41    /// Number of feature columns.
42    pub fn num_features(&self) -> usize {
43        self.features.len()
44    }
45
46    /// Number of classes; labels are `0..num_labels`.
47    pub fn num_labels(&self) -> usize {
48        self.num_labels
49    }
50
51    /// Appends an observation to column `feature_index`. The observation of
52    /// instance 0 opens a new column.
53    pub fn insert(&mut self, data_point: DataPoint, feature_index: usize) {
54        if data_point.tid() == 0 {
55            self.features.push(Feature::new())
56        }
57        self.features[feature_index].insert(data_point)
58    }
59
60    /// Sets the number of classes.
61    pub fn set_num_label(&mut self, value: usize) {
62        self.num_labels = value;
63    }
64
65    /// Sorts every column by value.
66    pub fn sort_features(&mut self) {
67        for column in self.features.iter_mut() {
68            column.sort();
69        }
70        self.sorted = true;
71    }
72
73    /// Assigns each observation the index of its value among the column's
74    /// distinct values. It may run before or after `sort_features`, and must
75    /// run before any fit.
76    pub fn compute_unique_feature_values(&mut self) {
77        let size = self.count();
78        let mut idx = vec![0; size];
79        for column in &mut self.features {
80            (0..size).for_each(|i| idx[i] = i);
81            idx.sort_unstable_by(|&idx1, &idx2| column[idx1].cmp(&column[idx2]));
82
83            let mut cur_unique = 0;
84            let mut prev: Option<f64> = None;
85            for &index in &idx {
86                let el = &mut column[index];
87                if let Some(prev_val) = prev {
88                    if (el.value - prev_val).abs() >= f64::EPSILON {
89                        cur_unique += 1;
90                    }
91                }
92
93                el.unique_value_idx = cur_unique;
94                prev = Some(el.value);
95            }
96        }
97        self.indexed = true;
98    }
99}
100
101/// Why a dataset could not be built from memory.
102#[derive(Clone, Debug, PartialEq, Eq)]
103pub enum DatasetError {
104    /// No rows or no features.
105    Empty,
106    /// `values.len()` is not `labels.len() * n_features`.
107    ShapeMismatch {
108        values: usize,
109        rows: usize,
110        n_features: usize,
111    },
112    /// A feature value that is NaN or infinite.
113    NonFiniteValue { row: usize, feature: usize },
114    /// A label outside what can be a dense `0..k` encoding.
115    InvalidLabel {
116        row: usize,
117        label: usize,
118        rows: usize,
119    },
120}
121
122impl std::fmt::Display for DatasetError {
123    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
124        match self {
125            DatasetError::Empty => write!(f, "the dataset has no rows"),
126            DatasetError::ShapeMismatch {
127                values,
128                rows,
129                n_features,
130            } => write!(
131                f,
132                "{values} values do not form {rows} rows of {n_features} features"
133            ),
134            DatasetError::NonFiniteValue { row, feature } => {
135                write!(f, "row {row}, feature {feature}: values must be finite")
136            }
137            DatasetError::InvalidLabel { row, label, rows } => write!(
138                f,
139                "row {row}: label {label} in a dataset of {rows} rows; labels must be a dense \
140                 encoding starting at 0"
141            ),
142        }
143    }
144}
145
146impl std::error::Error for DatasetError {}
147
148impl Dataset {
149    /// Builds a dataset from a row-major value buffer and its labels.
150    ///
151    /// This is the in-memory counterpart of
152    /// [`DataReader::read_file`](crate::reader::data_reader::DataReader::read_file),
153    /// and takes numpy's C-order `(n_rows, n_features)` layout directly.
154    ///
155    /// The returned dataset is prepared: its columns are sorted and indexed.
156    pub fn from_rows(
157        values: &[f64],
158        labels: &[usize],
159        n_features: usize,
160    ) -> Result<Self, DatasetError> {
161        let rows = labels.len();
162        if rows == 0 || n_features == 0 {
163            return Err(DatasetError::Empty);
164        }
165        if values.len() != rows * n_features {
166            return Err(DatasetError::ShapeMismatch {
167                values: values.len(),
168                rows,
169                n_features,
170            });
171        }
172
173        let mut max_label = 0;
174        for (row, &label) in labels.iter().enumerate() {
175            if label >= rows {
176                return Err(DatasetError::InvalidLabel { row, label, rows });
177            }
178            max_label = max_label.max(label);
179        }
180
181        let mut dataset = Dataset::new();
182        for (row, (chunk, &label)) in values.chunks_exact(n_features).zip(labels).enumerate() {
183            for (feature, &value) in chunk.iter().enumerate() {
184                if !value.is_finite() {
185                    return Err(DatasetError::NonFiniteValue { row, feature });
186                }
187                dataset.insert(DataPoint::new(row, value, label as f64), feature);
188            }
189        }
190
191        dataset.set_num_label(max_label + 1);
192        dataset.sort_features();
193        dataset.compute_unique_feature_values();
194        Ok(dataset)
195    }
196}
197
198impl Index<usize> for Dataset {
199    type Output = Feature;
200
201    fn index(&self, index: usize) -> &Self::Output {
202        &self.features[index]
203    }
204}
205
206impl IndexMut<usize> for Dataset {
207    fn index_mut(&mut self, index: usize) -> &mut Feature {
208        &mut self.features[index]
209    }
210}
211
212impl<'a> IntoIterator for &'a Dataset {
213    type Item = &'a Feature;
214    type IntoIter = std::slice::Iter<'a, Feature>;
215
216    fn into_iter(self) -> Self::IntoIter {
217        self.features.iter()
218    }
219}