1use crate::data::{DataPoint, Feature};
2use std::ops::{Index, IndexMut};
3
4#[derive(Debug, Default)]
9pub struct Dataset {
10 features: Vec<Feature>,
11 num_labels: usize,
12 sorted: bool,
15 indexed: bool,
16}
17
18impl Dataset {
19 pub fn new() -> Self {
21 Self::default()
22 }
23
24 pub fn is_prepared(&self) -> bool {
28 self.sorted && self.indexed
29 }
30
31 pub fn count(&self) -> usize {
33 self.features.first().map_or(0, Feature::len)
34 }
35
36 pub fn is_empty(&self) -> bool {
38 self.count() == 0
39 }
40
41 pub fn num_features(&self) -> usize {
43 self.features.len()
44 }
45
46 pub fn num_labels(&self) -> usize {
48 self.num_labels
49 }
50
51 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 pub fn set_num_label(&mut self, value: usize) {
62 self.num_labels = value;
63 }
64
65 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 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#[derive(Clone, Debug, PartialEq, Eq)]
103pub enum DatasetError {
104 Empty,
106 ShapeMismatch {
108 values: usize,
109 rows: usize,
110 n_features: usize,
111 },
112 NonFiniteValue { row: usize, feature: usize },
114 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 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}