Skip to main content

contree/reader/
data_reader.rs

1use super::{DataFormat, DataReaderError};
2use crate::data::{DataPoint, Dataset};
3use std::fs::File;
4use std::io::{BufRead, BufReader};
5use std::path::Path;
6
7/// Reads a delimited text file into a [`Dataset`].
8///
9/// The default format is one instance per line, whitespace separated, the
10/// label in column 0, `#` starting a comment, and no header row.
11///
12/// Labels must be non-negative integers. They are treated as a dense encoding,
13/// so `num_labels` is `max_label + 1` and a file whose labels are `{0, 2}`
14/// simply has an empty class 1.
15pub 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    /// A reader with the default format.
35    pub fn new() -> Self {
36        Self::default()
37    }
38
39    /// Sets the column delimiter.
40    pub fn with_format(mut self, format: DataFormat) -> Self {
41        self.format = format;
42        self
43    }
44
45    /// Whether the first data line is a header to skip.
46    pub fn with_headers(mut self, has_headers: bool) -> Self {
47        self.has_headers = has_headers;
48        self
49    }
50
51    /// Lines starting with this character are ignored.
52    pub fn with_comment_char(mut self, comment_char: Option<char>) -> Self {
53        self.comment_char = comment_char;
54        self
55    }
56
57    /// Which column holds the label. Defaults to 0.
58    pub fn with_label_column(mut self, label_column: usize) -> Self {
59        self.label_column = label_column;
60        self
61    }
62
63    /// Picks the delimiter from the file extension (`.csv`, `.tsv`, otherwise
64    /// whitespace).
65    pub fn auto_detect_format(mut self, path: &Path) -> Self {
66        self.format = DataFormat::from_extension(path);
67        self
68    }
69
70    /// Reads a file. The returned dataset is sorted and indexed, ready to fit.
71    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    /// Reads from any line source. `read_file` is this over a file.
77    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        // Set by the first data row; every later row must have as many.
84        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            // The header is the first non-comment, non-blank line.
101            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                    // NaN and infinities parse as `f64` but break the ordering
163                    // of the column.
164                    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        // Class histograms are indexed by label, so `num_labels` is `max + 1`
183        // rather than the number of distinct labels. A label at or above the
184        // row count cannot come from a dense encoding.
185        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}