zsfm_tabfm/ensemble/
cat_encoder.rs1use 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
44pub struct LabelEncoder {
47 classes: Vec<String>, }
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}