Skip to main content

zsfm_tabfm/ensemble/
cat_encoder.rs

1//! `CategoricalOrdinalEncoder` (`cat_encoder_mode="appearance"`, the wrapper's default): assigns
2//! integer codes to a categorical column's values in order of first appearance in the training
3//! data; unknown/missing values (at fit or transform time) map to `-1`.
4
5use std::collections::HashMap;
6
7use serde_json::Value;
8
9pub struct CategoricalOrdinalEncoder {
10    index: HashMap<String, i64>,
11}
12
13impl CategoricalOrdinalEncoder {
14    pub fn fit(column: &[Value]) -> Self {
15        let mut index = HashMap::new();
16        let mut next = 0i64;
17        for v in column {
18            if is_missing(v) {
19                continue;
20            }
21            let key = value_key(v);
22            index.entry(key).or_insert_with(|| {
23                let code = next;
24                next += 1;
25                code
26            });
27        }
28        CategoricalOrdinalEncoder { index }
29    }
30
31    pub fn transform(&self, column: &[Value]) -> Vec<f64> {
32        column
33            .iter()
34            .map(|v| {
35                if is_missing(v) {
36                    return -1.0;
37                }
38                *self.index.get(&value_key(v)).unwrap_or(&-1) as f64
39            })
40            .collect()
41    }
42}
43
44/// Classification label encoder (`CategoricalOrdinalEncoder(mode="alphabetical")` as used on
45/// `y`): unique labels sorted alphabetically (by their string form) get codes `0..n_classes-1`.
46pub struct LabelEncoder {
47    classes: Vec<String>, // sorted; index = class code
48}
49
50impl LabelEncoder {
51    pub fn fit(y: &[Value]) -> Self {
52        let mut uniq: Vec<String> = y.iter().filter(|v| !is_missing(v)).map(value_key).collect();
53        uniq.sort();
54        uniq.dedup();
55        LabelEncoder { classes: uniq }
56    }
57
58    pub fn n_classes(&self) -> usize {
59        self.classes.len()
60    }
61
62    pub fn transform(&self, y: &[Value]) -> Vec<f64> {
63        y.iter()
64            .map(|v| {
65                if is_missing(v) {
66                    return -1.0;
67                }
68                let key = value_key(v);
69                self.classes.iter().position(|c| c == &key).map(|i| i as f64).unwrap_or(-1.0)
70            })
71            .collect()
72    }
73
74    pub fn decode(&self, code: usize) -> &str {
75        &self.classes[code]
76    }
77}
78
79fn is_missing(v: &Value) -> bool {
80    match v {
81        Value::Null => true,
82        Value::String(s) => s.is_empty() || s.eq_ignore_ascii_case("nan"),
83        Value::Number(n) => n.as_f64().map(f64::is_nan).unwrap_or(false),
84        _ => false,
85    }
86}
87
88fn value_key(v: &Value) -> String {
89    match v {
90        Value::String(s) => s.clone(),
91        Value::Number(n) => n.to_string(),
92        Value::Bool(b) => b.to_string(),
93        other => other.to_string(),
94    }
95}
96
97#[cfg(test)]
98mod tests {
99    use super::*;
100
101    #[test]
102    fn test_appearance_order() {
103        let col = vec![
104            Value::String("blue".into()),
105            Value::String("red".into()),
106            Value::String("blue".into()),
107            Value::String("green".into()),
108        ];
109        let enc = CategoricalOrdinalEncoder::fit(&col);
110        assert_eq!(enc.transform(&col), vec![0.0, 1.0, 0.0, 2.0]);
111    }
112
113    #[test]
114    fn test_unknown_at_transform_time() {
115        let train = vec![Value::String("a".into()), Value::String("b".into())];
116        let enc = CategoricalOrdinalEncoder::fit(&train);
117        let test = vec![Value::String("a".into()), Value::String("c".into())];
118        assert_eq!(enc.transform(&test), vec![0.0, -1.0]);
119    }
120
121    #[test]
122    fn test_missing_values() {
123        let col = vec![Value::String("a".into()), Value::Null, Value::String("a".into())];
124        let enc = CategoricalOrdinalEncoder::fit(&col);
125        assert_eq!(enc.transform(&col), vec![0.0, -1.0, 0.0]);
126    }
127}