Skip to main content

datarust/encoder/
ordinal.rs

1use std::collections::HashMap;
2
3use crate::error::{DatarustError, Result};
4use crate::matrix::{Matrix, StrMatrix};
5use crate::traits::{default_input_names, CategoricalTransformer, FeatureNames};
6
7/// How to determine categories for ordinal encoding.
8#[derive(Debug, Clone, Default)]
9#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
10pub enum OrdinalCategories {
11    /// Infer categories from the training data (sorted lexicographically per column).
12    #[default]
13    Auto,
14    /// Provide explicit category order per column. If used, every column's categories
15    /// must be specified; each list determines the ordinal mapping.
16    Manual(Vec<Vec<String>>),
17}
18
19/// Strategy for unknown categories during `transform`.
20#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
21#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
22pub enum OrdinalHandleUnknown {
23    /// Raise an error on unknown categories (default).
24    #[default]
25    Error,
26    /// Encode unknown categories as `-1`.
27    UseNegOne,
28}
29
30/// Encode categorical features as ordinal integers (0, 1, 2, …), mirroring
31/// `sklearn.preprocessing.OrdinalEncoder`.
32///
33/// Input is a 2-D [`StrMatrix`]; output is a numeric [`Matrix`] of the same
34/// shape, where each cell is replaced by the ordinal index of its category.
35///
36/// Categories per column are sorted lexicographically by default (sklearn
37/// default), or can be user-specified via [`OrdinalCategories::Manual`].
38#[derive(Debug, Clone)]
39#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
40pub struct OrdinalEncoder {
41    categories: OrdinalCategories,
42    handle_unknown: OrdinalHandleUnknown,
43    category_lists: Vec<Vec<String>>,
44    category_indices: Vec<HashMap<String, usize>>,
45    fitted: bool,
46}
47
48impl OrdinalEncoder {
49    /// Creates a new ordinal encoder with the given category config.
50    pub fn new(categories: OrdinalCategories) -> Self {
51        Self {
52            categories,
53            handle_unknown: OrdinalHandleUnknown::default(),
54            category_lists: vec![],
55            category_indices: vec![],
56            fitted: false,
57        }
58    }
59
60    /// Sets how unknown categories are handled during transform.
61    pub fn handle_unknown(mut self, h: OrdinalHandleUnknown) -> Self {
62        self.handle_unknown = h;
63        self
64    }
65
66    /// Returns the learned categories per column.
67    pub fn categories(&self) -> &[Vec<String>] {
68        &self.category_lists
69    }
70
71    /// Learns the category-to-index mapping per column.
72    pub fn fit(&mut self, x: &StrMatrix) -> Result<()> {
73        let ncols = x.ncols();
74        match &self.categories {
75            OrdinalCategories::Auto => {
76                let mut cat_lists = Vec::with_capacity(ncols);
77                let mut cat_indices = Vec::with_capacity(ncols);
78                for j in 0..ncols {
79                    let col = x.column(j);
80                    let mut set: std::collections::BTreeSet<String> =
81                        std::collections::BTreeSet::new();
82                    for s in &col {
83                        set.insert(s.clone());
84                    }
85                    let list: Vec<String> = set.into_iter().collect();
86                    let idx: HashMap<String, usize> = list
87                        .iter()
88                        .enumerate()
89                        .map(|(i, c)| (c.clone(), i))
90                        .collect();
91                    cat_lists.push(list);
92                    cat_indices.push(idx);
93                }
94                self.category_lists = cat_lists;
95                self.category_indices = cat_indices;
96            }
97            OrdinalCategories::Manual(lists) => {
98                if lists.len() != ncols {
99                    return Err(DatarustError::ShapeMismatch {
100                        expected: format!("{} category lists", ncols),
101                        actual: format!("{} lists", lists.len()),
102                    });
103                }
104                let mut cat_indices = Vec::with_capacity(ncols);
105                for (j, list) in lists.iter().enumerate() {
106                    let idx: HashMap<String, usize> = list
107                        .iter()
108                        .enumerate()
109                        .map(|(i, c)| (c.clone(), i))
110                        .collect();
111                    if idx.len() != list.len() {
112                        return Err(DatarustError::InvalidConfig(format!(
113                            "duplicate category in column {}",
114                            j
115                        )));
116                    }
117                    cat_indices.push(idx);
118                }
119                self.category_lists = lists.clone();
120                self.category_indices = cat_indices;
121            }
122        }
123        self.fitted = true;
124        Ok(())
125    }
126
127    #[allow(clippy::needless_range_loop)]
128    /// Encodes the input as ordinal integer codes.
129    pub fn transform(&self, x: &StrMatrix) -> Result<Matrix> {
130        if !self.fitted {
131            return Err(DatarustError::NotFitted("OrdinalEncoder".into()));
132        }
133        if x.ncols() != self.category_lists.len() {
134            return Err(DatarustError::ShapeMismatch {
135                expected: format!("{} columns", self.category_lists.len()),
136                actual: format!("{} columns", x.ncols()),
137            });
138        }
139        let mut out = vec![vec![0.0; x.ncols()]; x.nrows()];
140
141        #[cfg(feature = "rayon")]
142        {
143            use rayon::prelude::*;
144            let category_indices = &self.category_indices;
145            let handle_unknown = self.handle_unknown;
146            let x_data = &x.data;
147            out.par_iter_mut().enumerate().try_for_each(|(i, row)| {
148                for (j, indices) in category_indices.iter().enumerate() {
149                    row[j] = match indices.get(&x_data[i][j]) {
150                        Some(&idx) => idx as f64,
151                        None => match handle_unknown {
152                            OrdinalHandleUnknown::Error => {
153                                return Err(DatarustError::UnknownCategory(format!(
154                                    "column {} value '{}'",
155                                    j, x_data[i][j]
156                                )))
157                            }
158                            OrdinalHandleUnknown::UseNegOne => -1.0,
159                        },
160                    };
161                }
162                Ok(())
163            })?;
164        }
165
166        #[cfg(not(feature = "rayon"))]
167        {
168            for i in 0..x.nrows() {
169                for j in 0..x.ncols() {
170                    let val = x.get(i, j);
171                    out[i][j] = match self.category_indices[j].get(val) {
172                        Some(&idx) => idx as f64,
173                        None => match self.handle_unknown {
174                            OrdinalHandleUnknown::Error => {
175                                return Err(DatarustError::UnknownCategory(format!(
176                                    "column {} value '{}'",
177                                    j, val
178                                )))
179                            }
180                            OrdinalHandleUnknown::UseNegOne => -1.0,
181                        },
182                    };
183                }
184            }
185        }
186
187        Matrix::new(out)
188    }
189
190    /// Fits the encoder and transforms the input in one step.
191    pub fn fit_transform(&mut self, x: &StrMatrix) -> Result<Matrix> {
192        self.fit(x)?;
193        self.transform(x)
194    }
195
196    #[allow(clippy::needless_range_loop)]
197    /// Decodes ordinal integer codes back to category strings.
198    pub fn inverse_transform(&self, y: &Matrix) -> Result<StrMatrix> {
199        if !self.fitted {
200            return Err(DatarustError::NotFitted("OrdinalEncoder".into()));
201        }
202        if y.ncols() != self.category_lists.len() {
203            return Err(DatarustError::ShapeMismatch {
204                expected: format!("{} columns", self.category_lists.len()),
205                actual: format!("{} columns", y.ncols()),
206            });
207        }
208        let mut out: Vec<Vec<String>> = Vec::with_capacity(y.nrows());
209        for i in 0..y.nrows() {
210            let mut row = Vec::with_capacity(y.ncols());
211            for j in 0..y.ncols() {
212                let v = y.get(i, j);
213                if v.is_nan() {
214                    return Err(DatarustError::InvalidInput(format!(
215                        "NaN value at row {}, column {} in inverse_transform input",
216                        i, j
217                    )));
218                }
219                let idx = v as isize;
220                if idx == -1 {
221                    // Sentinel for unknown categories (UseNegOne)
222                    row.push(String::new());
223                } else if idx < 0 || idx as usize >= self.category_lists[j].len() {
224                    return Err(DatarustError::UnknownLabel(format!(
225                        "index {} out of range for column {}",
226                        idx, j
227                    )));
228                } else {
229                    row.push(self.category_lists[j][idx as usize].clone());
230                }
231            }
232            out.push(row);
233        }
234        StrMatrix::new(out)
235    }
236}
237
238impl CategoricalTransformer for OrdinalEncoder {
239    fn name(&self) -> &'static str {
240        "OrdinalEncoder"
241    }
242
243    fn fit(&mut self, x: &StrMatrix) -> Result<()> {
244        self.fit(x)
245    }
246
247    fn transform(&self, x: &StrMatrix) -> Result<Matrix> {
248        self.transform(x)
249    }
250
251    fn inverse_transform(&self, y: &Matrix) -> Result<StrMatrix> {
252        self.inverse_transform(y)
253    }
254
255    fn is_fitted(&self) -> bool {
256        self.fitted
257    }
258}
259
260impl Default for OrdinalEncoder {
261    fn default() -> Self {
262        Self::new(OrdinalCategories::default())
263    }
264}
265
266impl FeatureNames for OrdinalEncoder {
267    fn feature_names_out(&self, input_features: Option<&[String]>) -> Vec<String> {
268        let n = self.category_lists.len();
269        let names: Vec<String> = match input_features {
270            Some(fs) => (0..n)
271                .map(|i| fs.get(i).cloned().unwrap_or_else(|| format!("x{}", i)))
272                .collect(),
273            None => default_input_names(n),
274        };
275        names
276    }
277}
278
279#[cfg(test)]
280mod tests {
281    use super::*;
282
283    #[test]
284    fn basic_auto_fit() {
285        let s = StrMatrix::from_column(["small", "medium", "large", "small"]).unwrap();
286        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
287        let out = enc.fit_transform(&s).unwrap();
288        // categories sorted: large(0), medium(1), small(2)
289        assert_eq!(enc.categories()[0], &["large", "medium", "small"]);
290        assert_eq!(out.row(0), [2.0]); // small
291        assert_eq!(out.row(1), [1.0]); // medium
292        assert_eq!(out.row(2), [0.0]); // large
293    }
294
295    #[test]
296    fn manual_categories() {
297        let s = StrMatrix::from_column(["small", "medium", "large"]).unwrap();
298        let mut enc = OrdinalEncoder::new(OrdinalCategories::Manual(vec![vec![
299            "small".into(),
300            "medium".into(),
301            "large".into(),
302        ]]));
303        let out = enc.fit_transform(&s).unwrap();
304        assert_eq!(out.row(0), [0.0]);
305        assert_eq!(out.row(1), [1.0]);
306        assert_eq!(out.row(2), [2.0]);
307    }
308
309    #[test]
310    fn inverse_round_trip() {
311        let original = vec!["cat", "dog", "bird", "dog"];
312        let s = StrMatrix::from_column(original.clone()).unwrap();
313        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
314        let encoded = enc.fit_transform(&s).unwrap();
315        let decoded = enc.inverse_transform(&encoded).unwrap();
316        for (i, &orig) in original.iter().enumerate() {
317            assert_eq!(decoded.get(i, 0), orig);
318        }
319    }
320
321    #[test]
322    fn inverse_bad_index_errors() {
323        let s = StrMatrix::from_column(["a", "b"]).unwrap();
324        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
325        enc.fit(&s).unwrap();
326        let bad = Matrix::new(vec![vec![0.0], vec![5.0]]).unwrap();
327        assert!(enc.inverse_transform(&bad).is_err());
328    }
329
330    #[test]
331    fn multi_column() {
332        let s =
333            StrMatrix::from_strings(vec![vec!["a", "x"], vec!["b", "y"], vec!["a", "y"]]).unwrap();
334        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
335        let out = enc.fit_transform(&s).unwrap();
336        assert_eq!(out.ncols(), 2);
337        // col0: a(0), b(1) ; col1: x(0), y(1)
338        assert_eq!(out.row(0), [0.0, 0.0]);
339        assert_eq!(out.row(1), [1.0, 1.0]);
340        assert_eq!(out.row(2), [0.0, 1.0]);
341    }
342
343    #[test]
344    fn handle_unknown_error() {
345        let s = StrMatrix::from_column(["a", "b"]).unwrap();
346        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
347        enc.fit(&s).unwrap();
348        let s2 = StrMatrix::from_column(["a", "z"]).unwrap();
349        assert!(matches!(
350            enc.transform(&s2),
351            Err(DatarustError::UnknownCategory(_))
352        ));
353    }
354
355    #[test]
356    fn handle_unknown_neg_one() {
357        let s = StrMatrix::from_column(["a", "b"]).unwrap();
358        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto)
359            .handle_unknown(OrdinalHandleUnknown::UseNegOne);
360        enc.fit(&s).unwrap();
361        let s2 = StrMatrix::from_column(["a", "z"]).unwrap();
362        let out = enc.transform(&s2).unwrap();
363        assert_eq!(out.row(0), [0.0]);
364        assert_eq!(out.row(1), [-1.0]);
365    }
366
367    #[test]
368    fn manual_column_count_mismatch_errors() {
369        let s = StrMatrix::from_column(["a", "b"]).unwrap();
370        let mut enc = OrdinalEncoder::new(OrdinalCategories::Manual(vec![
371            vec!["a".into()],
372            vec!["b".into()],
373        ]));
374        assert!(enc.fit(&s).is_err());
375    }
376
377    #[test]
378    fn manual_duplicate_category_errors() {
379        let s = StrMatrix::from_column(["a", "b"]).unwrap();
380        let mut enc = OrdinalEncoder::new(OrdinalCategories::Manual(vec![vec![
381            "a".into(),
382            "a".into(),
383        ]]));
384        assert!(enc.fit(&s).is_err());
385    }
386
387    #[test]
388    fn transform_before_fit_errors() {
389        let enc = OrdinalEncoder::new(OrdinalCategories::Auto);
390        let s = StrMatrix::from_column(["a"]).unwrap();
391        assert!(matches!(
392            enc.transform(&s),
393            Err(DatarustError::NotFitted(_))
394        ));
395    }
396
397    #[test]
398    fn inverse_before_fit_errors() {
399        let enc = OrdinalEncoder::new(OrdinalCategories::Auto);
400        let m = Matrix::new(vec![vec![0.0]]).unwrap();
401        assert!(matches!(
402            enc.inverse_transform(&m),
403            Err(DatarustError::NotFitted(_))
404        ));
405    }
406
407    #[test]
408    fn serde_derive() {
409        // compile-test: the struct must derive Serialize/Deserialize when feature is on
410        let s = StrMatrix::from_column(["x", "y"]).unwrap();
411        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
412        enc.fit(&s).unwrap();
413        // no explicit assertion; just ensure the type works under serde feature exists
414        #[cfg(feature = "serde")]
415        {
416            let json = crate::serialize::to_json(&enc).unwrap();
417            let _restored: OrdinalEncoder = crate::serialize::from_json(&json).unwrap();
418        }
419    }
420
421    #[test]
422    fn inverse_transform_sentinel_decodes_to_empty() {
423        let s = StrMatrix::from_column(["cat", "dog"]).unwrap();
424        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto)
425            .handle_unknown(OrdinalHandleUnknown::UseNegOne);
426        enc.fit(&s).unwrap();
427        // unknown value 'fox' encodes to -1.0
428        let x = StrMatrix::from_column(["cat", "fox"]).unwrap();
429        let coded = enc.transform(&x).unwrap();
430        assert_eq!(coded.get(1, 0), -1.0);
431        let decoded = enc.inverse_transform(&coded).unwrap();
432        assert_eq!(decoded.get(0, 0), "cat");
433        assert_eq!(decoded.get(1, 0), "");
434    }
435
436    #[test]
437    fn feature_names_short_input_pads_with_synthetic() {
438        let s =
439            StrMatrix::from_strings(vec![vec!["a", "x"], vec!["b", "y"], vec!["c", "z"]]).unwrap();
440        let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
441        enc.fit(&s).unwrap();
442        // 2 columns but only 1 name provided
443        let names = enc.feature_names_out(Some(&["city".into()]));
444        assert_eq!(names, vec!["city", "x1"]);
445    }
446}