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#[derive(Debug, Clone, Default)]
9#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
10pub enum OrdinalCategories {
11 #[default]
13 Auto,
14 Manual(Vec<Vec<String>>),
17}
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
21#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
22pub enum OrdinalHandleUnknown {
23 #[default]
25 Error,
26 UseNegOne,
28}
29
30#[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 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 pub fn handle_unknown(mut self, h: OrdinalHandleUnknown) -> Self {
62 self.handle_unknown = h;
63 self
64 }
65
66 pub fn categories(&self) -> &[Vec<String>] {
68 &self.category_lists
69 }
70
71 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 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 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 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 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 assert_eq!(enc.categories()[0], &["large", "medium", "small"]);
290 assert_eq!(out.row(0), [2.0]); assert_eq!(out.row(1), [1.0]); assert_eq!(out.row(2), [0.0]); }
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 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 let s = StrMatrix::from_column(["x", "y"]).unwrap();
411 let mut enc = OrdinalEncoder::new(OrdinalCategories::Auto);
412 enc.fit(&s).unwrap();
413 #[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 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 let names = enc.feature_names_out(Some(&["city".into()]));
444 assert_eq!(names, vec!["city", "x1"]);
445 }
446}