1use super::{DataFormat, DataReaderError};
2use crate::data::{DataPoint, Dataset};
3use std::fs::File;
4use std::io::{BufRead, BufReader};
5use std::path::Path;
6
7pub struct DataReader {
16 format: DataFormat,
17 has_headers: bool,
18 comment_char: Option<char>,
19 label_column: usize,
20}
21
22impl Default for DataReader {
23 fn default() -> Self {
24 Self {
25 format: DataFormat::Space,
26 has_headers: false,
27 comment_char: Some('#'),
28 label_column: 0,
29 }
30 }
31}
32
33impl DataReader {
34 pub fn new() -> Self {
36 Self::default()
37 }
38
39 pub fn with_format(mut self, format: DataFormat) -> Self {
41 self.format = format;
42 self
43 }
44
45 pub fn with_headers(mut self, has_headers: bool) -> Self {
47 self.has_headers = has_headers;
48 self
49 }
50
51 pub fn with_comment_char(mut self, comment_char: Option<char>) -> Self {
53 self.comment_char = comment_char;
54 self
55 }
56
57 pub fn with_label_column(mut self, label_column: usize) -> Self {
59 self.label_column = label_column;
60 self
61 }
62
63 pub fn auto_detect_format(mut self, path: &Path) -> Self {
66 self.format = DataFormat::from_extension(path);
67 self
68 }
69
70 pub fn read_file(&self, path: &Path) -> Result<Dataset, DataReaderError> {
72 let file = File::open(path)?;
73 self.read(BufReader::new(file))
74 }
75
76 pub fn read<R: BufRead>(&self, source: R) -> Result<Dataset, DataReaderError> {
78 let delimiter = self.format.delimiter();
79
80 let mut dataset = Dataset::new();
81 let mut row_idx = 0usize;
82 let mut max_label = 0usize;
83 let mut num_columns: Option<usize> = None;
85 let mut header_pending = self.has_headers;
86
87 for (line_idx, line_result) in source.lines().enumerate() {
88 let line = line_result?;
89 let line = line.trim();
90 let line_no = line_idx + 1;
91
92 if line.is_empty() {
93 continue;
94 }
95 if let Some(comment) = self.comment_char {
96 if line.starts_with(comment) {
97 continue;
98 }
99 }
100 if header_pending {
102 header_pending = false;
103 continue;
104 }
105
106 let tokens: Vec<&str> = if delimiter == ' ' {
107 line.split_whitespace().collect()
108 } else {
109 line.split(delimiter).map(str::trim).collect()
110 };
111
112 match num_columns {
113 None => {
114 if tokens.len() < 2 {
115 return Err(DataReaderError::Format(format!(
116 "line {line_no}: expected a label and at least one feature, found {} \
117 column(s)",
118 tokens.len()
119 )));
120 }
121 if self.label_column >= tokens.len() {
122 return Err(DataReaderError::Format(format!(
123 "line {line_no}: label column {} is out of range, the file has {} \
124 columns",
125 self.label_column,
126 tokens.len()
127 )));
128 }
129 num_columns = Some(tokens.len());
130 }
131 Some(expected) if tokens.len() != expected => {
132 return Err(DataReaderError::Format(format!(
133 "line {line_no}: {} columns, expected {expected}",
134 tokens.len()
135 )));
136 }
137 Some(_) => {}
138 }
139
140 let label_token = tokens[self.label_column];
141 let label: usize = label_token.parse().map_err(|_| {
142 DataReaderError::Parse(format!(
143 "line {line_no}, column {}: `{label_token}` is not a label; labels must be \
144 non-negative integers",
145 self.label_column + 1
146 ))
147 })?;
148 max_label = max_label.max(label);
149
150 for (col_idx, &token) in tokens.iter().enumerate() {
151 if col_idx == self.label_column {
152 continue;
153 }
154
155 let value = token.parse::<f64>().map_err(|_| {
156 DataReaderError::Parse(format!(
157 "line {line_no}, column {}: `{token}` is not a number",
158 col_idx + 1
159 ))
160 })?;
161 if !value.is_finite() {
162 return Err(DataReaderError::Parse(format!(
165 "line {line_no}, column {}: `{token}` is not a finite number",
166 col_idx + 1
167 )));
168 }
169
170 let feature_idx = col_idx - usize::from(col_idx > self.label_column);
171 dataset.insert(DataPoint::new(row_idx, value, label as f64), feature_idx);
172 }
173 row_idx += 1;
174 }
175
176 if row_idx == 0 {
177 return Err(DataReaderError::Format(
178 "no data rows: the file is empty or entirely comments".to_string(),
179 ));
180 }
181
182 if max_label >= row_idx {
186 return Err(DataReaderError::Format(format!(
187 "label {max_label} in a file of {row_idx} rows: labels must be a dense encoding \
188 starting at 0"
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
198#[cfg(test)]
199mod tests {
200 use super::*;
201 use crate::reader::DataFormat;
202 use crate::tests::fixture;
203
204 fn read(text: &str) -> Result<Dataset, DataReaderError> {
205 DataReader::default().read(text.as_bytes())
206 }
207
208 fn read_csv(text: &str) -> Result<Dataset, DataReaderError> {
209 DataReader::default()
210 .with_format(DataFormat::Csv)
211 .read(text.as_bytes())
212 }
213
214 #[test]
215 fn reads_the_label_from_column_zero() {
216 let dataset = read("0 1.0 2.0\n1 3.0 4.0\n").unwrap();
217 assert_eq!(dataset.count(), 2);
218 assert_eq!(dataset.num_features(), 2);
219 assert_eq!(dataset.num_labels(), 2);
220 assert_eq!(dataset[0][0].value(), 1.0);
221 assert_eq!(dataset[1][1].value(), 4.0);
222 assert_eq!(dataset[0][1].label(), 1.0);
223 }
224
225 #[test]
226 fn num_labels_is_max_plus_one_not_the_distinct_count() {
227 let dataset = read("0 1.0\n5 2.0\n0 3.0\n5 4.0\n5 5.0\n0 6.0\n").unwrap();
228 assert_eq!(dataset.num_labels(), 6);
229 }
230
231 #[test]
232 fn a_label_that_cannot_be_a_dense_encoding_is_rejected() {
233 let err = read("0 1.0\n900 2.0\n").unwrap_err();
234 assert!(matches!(err, DataReaderError::Format(_)), "{err}");
235 }
236
237 #[test]
238 fn ragged_rows_are_an_error_not_a_silent_shift() {
239 let wider = read("0 1.0 2.0\n1 3.0 4.0 5.0\n").unwrap_err();
240 assert!(matches!(wider, DataReaderError::Format(_)), "{wider}");
241
242 let narrower = read("0 1.0 2.0\n1 3.0\n").unwrap_err();
243 assert!(matches!(narrower, DataReaderError::Format(_)), "{narrower}");
244 }
245
246 #[test]
247 fn a_missing_csv_field_is_reported_rather_than_deleted() {
248 let err = read_csv("0,1.0,2.0\n1,,3.0\n").unwrap_err();
249 assert!(matches!(err, DataReaderError::Parse(_)), "{err}");
250 }
251
252 #[test]
253 fn non_finite_values_are_rejected() {
254 for text in ["0 1.0\n1 NaN\n", "0 1.0\n1 inf\n"] {
255 let err = read(text).unwrap_err();
256 assert!(matches!(err, DataReaderError::Parse(_)), "{text:?}: {err}");
257 }
258 }
259
260 #[test]
261 fn an_empty_or_all_comment_file_is_an_error() {
262 for text in ["", "\n\n", "# just a comment\n"] {
263 let err = read(text).unwrap_err();
264 assert!(matches!(err, DataReaderError::Format(_)), "{text:?}: {err}");
265 }
266 }
267
268 #[test]
269 fn a_header_after_a_comment_is_still_skipped() {
270 let dataset = DataReader::default()
271 .with_headers(true)
272 .read("# a comment\nlabel f0 f1\n0 1.0 2.0\n1 3.0 4.0\n".as_bytes())
273 .unwrap();
274 assert_eq!(dataset.count(), 2);
275 }
276
277 #[test]
278 fn the_label_column_can_be_moved() {
279 let dataset = DataReader::default()
280 .with_label_column(2)
281 .read("1.0 2.0 0\n3.0 4.0 1\n".as_bytes())
282 .unwrap();
283 assert_eq!(dataset.num_features(), 2);
284 assert_eq!(dataset[0][0].value(), 1.0);
285 assert_eq!(dataset[1][0].value(), 2.0);
286 assert_eq!(dataset[0][0].label(), 0.0);
287 }
288
289 #[test]
290 fn a_label_column_past_the_end_is_an_error() {
291 let err = DataReader::default()
292 .with_label_column(7)
293 .read("0 1.0 2.0\n".as_bytes())
294 .unwrap_err();
295 assert!(matches!(err, DataReaderError::Format(_)), "{err}");
296 }
297
298 #[test]
299 fn the_repository_fixtures_still_load() {
300 for (name, features, labels) in [("iris.txt", 4, 2), ("hepatitis.txt", 68, 2)] {
301 let dataset = DataReader::default().read_file(&fixture(name)).unwrap();
302 assert_eq!(dataset.num_features(), features, "{name}");
303 assert_eq!(dataset.num_labels(), labels, "{name}");
304 assert!(dataset.count() > 0, "{name}");
305 }
306 }
307}