diff --git a/.github/workflows/checking.yml b/.github/workflows/checking.yml index 3c3e7dfde..43c142156 100644 --- a/.github/workflows/checking.yml +++ b/.github/workflows/checking.yml @@ -39,7 +39,7 @@ jobs: run: cargo check --workspace --all-targets - name: Run cargo check (with serde) - run: cargo check --workspace --all-targets --features "linfa-clustering/serde linfa-ica/serde linfa-kernel/serde linfa-reduction/serde linfa-svm/serde linfa-elasticnet/serde linfa-pls/serde linfa-trees/serde linfa-nn/serde linfa-linear/serde linfa-preprocessing/serde linfa-bayes/serde linfa-logistic/serde linfa-ftrl/serde" + run: cargo check --workspace --all-targets --features "linfa-clustering/serde linfa-ica/serde linfa-kernel/serde linfa-reduction/serde linfa-svm/serde linfa-elasticnet/serde linfa-pls/serde linfa-trees/serde linfa-nn/serde linfa-linear/serde linfa-preprocessing/serde linfa-bayes/serde linfa-logistic/serde linfa-ftrl/serde linfa-ensemble/serde" check-wasm-browser: name: check-wasm-browser diff --git a/algorithms/linfa-ensemble/Cargo.toml b/algorithms/linfa-ensemble/Cargo.toml index a8e1ffe9c..232393502 100644 --- a/algorithms/linfa-ensemble/Cargo.toml +++ b/algorithms/linfa-ensemble/Cargo.toml @@ -17,7 +17,7 @@ categories = ["algorithms", "mathematics", "science"] [features] default = [] -serde = ["serde_crate", "ndarray/serde"] +serde = ["serde_crate", "ndarray/serde", "linfa-trees/serde"] [dependencies.serde_crate] package = "serde" @@ -38,3 +38,4 @@ linfa-trees = { version = "0.8.1", path = "../linfa-trees" } linfa-datasets = { version = "0.8.1", path = "../../datasets/", features = [ "iris", ] } +serde_json = { version = "1", features = ["float_roundtrip"] } diff --git a/algorithms/linfa-ensemble/src/adaboost.rs b/algorithms/linfa-ensemble/src/adaboost.rs index a8709979f..98d507f5a 100644 --- a/algorithms/linfa-ensemble/src/adaboost.rs +++ b/algorithms/linfa-ensemble/src/adaboost.rs @@ -11,6 +11,9 @@ use ndarray_rand::rand::prelude::*; use ndarray_rand::rand::Rng; use std::{cmp::Eq, collections::HashMap, hash::Hash}; +#[cfg(feature = "serde")] +use serde_crate::{Deserialize, Serialize}; + /// Huge weight used to favor model with perfect predictions const PERFECT_MODEL_WEIGHT: f64 = 1e6; @@ -74,6 +77,11 @@ const PERFECT_MODEL_WEIGHT: f64 = 1e6; /// and an application to boosting. Journal of Computer and System Sciences, 55(1), 119-139. /// * [Scikit-Learn AdaBoost Documentation](https://scikit-learn.org/stable/modules/ensemble.html#adaboost) /// * [An Introduction to Statistical Learning](https://www.statlearning.com/), Chapter 8 +#[cfg_attr( + feature = "serde", + derive(Serialize, Deserialize), + serde(crate = "serde_crate") +)] #[derive(Debug, Clone)] pub struct AdaBoost { /// The fitted base learner models diff --git a/algorithms/linfa-ensemble/src/algorithm.rs b/algorithms/linfa-ensemble/src/algorithm.rs index e2cf531b8..c23a5e3f2 100644 --- a/algorithms/linfa-ensemble/src/algorithm.rs +++ b/algorithms/linfa-ensemble/src/algorithm.rs @@ -10,6 +10,9 @@ use ndarray::{Array2, Axis, Zip}; use rand::Rng; use std::{cmp::Eq, collections::HashMap, hash::Hash}; +#[cfg(feature = "serde")] +use serde_crate::{Deserialize, Serialize}; + /// A fitted ensemble of [Decision Trees](DecisionTree) trained on a random subset of features. /// /// Check out [EnsembleLearner] documentation for more information regarding [RandomForest] interface. @@ -67,6 +70,11 @@ pub type RandomForest = EnsembleLearner>; /// /// * [Scikit-Learn User Guide](https://scikit-learn.org/stable/modules/ensemble.html) /// * [An Introduction to Statistical Learning](https://www.statlearning.com/) +#[cfg_attr( + feature = "serde", + derive(Serialize, Deserialize), + serde(crate = "serde_crate") +)] pub struct EnsembleLearner { pub models: Vec, pub model_features: Vec>, diff --git a/algorithms/linfa-ensemble/src/lib.rs b/algorithms/linfa-ensemble/src/lib.rs index 87f129a78..1350f136c 100644 --- a/algorithms/linfa-ensemble/src/lib.rs +++ b/algorithms/linfa-ensemble/src/lib.rs @@ -337,4 +337,91 @@ mod tests { assert_eq!(classes.len(), 3, "Iris has 3 classes"); assert_eq!(classes, &vec![0, 1, 2], "Classes should be [0, 1, 2]"); } + + #[cfg(feature = "serde")] + #[test] + fn test_ensemble_learner_serde_roundtrip() { + let mut rng = SmallRng::seed_from_u64(42); + let (train, test) = linfa_datasets::iris() + .shuffle(&mut rng) + .split_with_ratio(0.8); + + let model = EnsembleLearnerParams::new_fixed_rng(DecisionTree::params(), rng) + .ensemble_size(20) + .bootstrap_proportion(0.7) + .fit(&train) + .unwrap(); + + let serialized = serde_json::to_string(&model).expect("failed to serialize ensemble"); + let loaded: EnsembleLearner> = + serde_json::from_str(&serialized).expect("failed to deserialize ensemble"); + + assert_eq!(model.predict(&test), loaded.predict(&test)); + assert_eq!(model.models.len(), loaded.models.len()); + assert_eq!(model.model_features, loaded.model_features); + } + + #[cfg(feature = "serde")] + #[test] + fn test_random_forest_serde_roundtrip() { + let mut rng = SmallRng::seed_from_u64(42); + let (train, test) = linfa_datasets::iris() + .shuffle(&mut rng) + .split_with_ratio(0.8); + + let model = RandomForestParams::new_fixed_rng(DecisionTree::params(), rng) + .ensemble_size(20) + .bootstrap_proportion(0.7) + .feature_proportion(0.3) + .fit(&train) + .unwrap(); + + let serialized = serde_json::to_string(&model).expect("failed to serialize random forest"); + let loaded: RandomForest = + serde_json::from_str(&serialized).expect("failed to deserialize random forest"); + + assert_eq!(model.models.len(), loaded.models.len()); + assert_eq!(model.model_features, loaded.model_features); + assert!(model.model_features.iter().any(|feat| feat.len() < 4)); + + // Iterate over the trees by hand, to get reproducible and meaningful comparison + let records = test.records(); + for ((original, reloaded), features) in model + .models + .iter() + .zip(loaded.models.iter()) + .zip(model.model_features.iter()) + { + let subset = records.select(ndarray::Axis(1), features); + assert_eq!( + original.predict(&subset), + reloaded.predict(&subset), + "a tree predicts differently after the round trip" + ); + } + } + + #[cfg(feature = "serde")] + #[test] + fn test_adaboost_serde_roundtrip() { + let mut rng = SmallRng::seed_from_u64(42); + let (train, test) = linfa_datasets::iris() + .shuffle(&mut rng) + .split_with_ratio(0.8); + + let model = AdaBoostParams::new_fixed_rng(DecisionTree::params().max_depth(Some(1)), rng) + .n_estimators(20) + .learning_rate(1.0) + .fit(&train) + .unwrap(); + + let serialized = serde_json::to_string(&model).expect("failed to serialize adaboost"); + let loaded: AdaBoost, usize> = + serde_json::from_str(&serialized).expect("failed to deserialize adaboost"); + + assert_eq!(model.predict(&test), loaded.predict(&test)); + assert_eq!(model.n_estimators(), loaded.n_estimators()); + assert_eq!(model.weights(), loaded.weights()); + assert_eq!(model.classes, loaded.classes); + } }