From 53049d388357e33e83ebbe451f7a9be82ce31c66 Mon Sep 17 00:00:00 2001 From: dbrundu Date: Mon, 27 Jul 2026 15:58:30 +0200 Subject: [PATCH 1/4] Updated linfa-ensemble Cargo.toml --- algorithms/linfa-ensemble/Cargo.toml | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) 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"] } From d82ad728e33425092086645baae19bb19021dfda Mon Sep 17 00:00:00 2001 From: dbrundu Date: Mon, 27 Jul 2026 15:59:34 +0200 Subject: [PATCH 2/4] Added attributes to derive serde traits on ensemble types --- algorithms/linfa-ensemble/src/adaboost.rs | 8 +++ algorithms/linfa-ensemble/src/algorithm.rs | 8 +++ algorithms/linfa-ensemble/src/lib.rs | 83 ++++++++++++++++++++++ 3 files changed, 99 insertions(+) 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..49993fb2d 100644 --- a/algorithms/linfa-ensemble/src/lib.rs +++ b/algorithms/linfa-ensemble/src/lib.rs @@ -337,4 +337,87 @@ 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, _) = 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.models.len(), loaded.models.len()); + assert_eq!(model.model_features, loaded.model_features); + assert_eq!( + serialized, + serde_json::to_string(&loaded).unwrap(), + "round trip is not lossless" + ); + } + + #[cfg(feature = "serde")] + #[test] + fn test_random_forest_serde_roundtrip() { + let mut rng = SmallRng::seed_from_u64(42); + let (train, _) = 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.model_features, loaded.model_features); + assert!(model.model_features.iter().any(|feat| feat.len() < 4)); + assert_eq!( + serialized, + serde_json::to_string(&loaded).unwrap(), + "round trip is not lossless" + ); + } + + #[cfg(feature = "serde")] + #[test] + fn test_adaboost_serde_roundtrip() { + let mut rng = SmallRng::seed_from_u64(42); + let (train, _) = 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.n_estimators(), loaded.n_estimators()); + assert_eq!(model.weights(), loaded.weights()); + assert_eq!(model.classes, loaded.classes); + assert_eq!( + serialized, + serde_json::to_string(&loaded).unwrap(), + "round trip is not lossless" + ); + } } From fa3df310130ffd2198a8c2e692a4271e54b0f6e8 Mon Sep 17 00:00:00 2001 From: dbrundu Date: Mon, 27 Jul 2026 16:00:37 +0200 Subject: [PATCH 3/4] Added linfa-ensemble/serde to the cargo check github actions --- .github/workflows/checking.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 From 02c433f3a6aa15696ac26de915c214d4c7f3fdb6 Mon Sep 17 00:00:00 2001 From: dbrundu Date: Mon, 27 Jul 2026 17:20:15 +0200 Subject: [PATCH 4/4] Fixed unit tests to compare predictions between serialized and original ensemble models --- algorithms/linfa-ensemble/src/lib.rs | 40 +++++++++++++++------------- 1 file changed, 22 insertions(+), 18 deletions(-) diff --git a/algorithms/linfa-ensemble/src/lib.rs b/algorithms/linfa-ensemble/src/lib.rs index 49993fb2d..1350f136c 100644 --- a/algorithms/linfa-ensemble/src/lib.rs +++ b/algorithms/linfa-ensemble/src/lib.rs @@ -342,7 +342,7 @@ mod tests { #[test] fn test_ensemble_learner_serde_roundtrip() { let mut rng = SmallRng::seed_from_u64(42); - let (train, _) = linfa_datasets::iris() + let (train, test) = linfa_datasets::iris() .shuffle(&mut rng) .split_with_ratio(0.8); @@ -356,20 +356,16 @@ mod tests { 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); - assert_eq!( - serialized, - serde_json::to_string(&loaded).unwrap(), - "round trip is not lossless" - ); } #[cfg(feature = "serde")] #[test] fn test_random_forest_serde_roundtrip() { let mut rng = SmallRng::seed_from_u64(42); - let (train, _) = linfa_datasets::iris() + let (train, test) = linfa_datasets::iris() .shuffle(&mut rng) .split_with_ratio(0.8); @@ -384,20 +380,32 @@ mod tests { 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)); - assert_eq!( - serialized, - serde_json::to_string(&loaded).unwrap(), - "round trip is not lossless" - ); + + // 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, _) = linfa_datasets::iris() + let (train, test) = linfa_datasets::iris() .shuffle(&mut rng) .split_with_ratio(0.8); @@ -411,13 +419,9 @@ mod tests { 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); - assert_eq!( - serialized, - serde_json::to_string(&loaded).unwrap(), - "round trip is not lossless" - ); } }