Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/checking.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 2 additions & 1 deletion algorithms/linfa-ensemble/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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"] }
8 changes: 8 additions & 0 deletions algorithms/linfa-ensemble/src/adaboost.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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<M, L> {
/// The fitted base learner models
Expand Down
8 changes: 8 additions & 0 deletions algorithms/linfa-ensemble/src/algorithm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -67,6 +70,11 @@ pub type RandomForest<F, L> = EnsembleLearner<DecisionTree<F, L>>;
///
/// * [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<M> {
pub models: Vec<M>,
pub model_features: Vec<Vec<usize>>,
Expand Down
87 changes: 87 additions & 0 deletions algorithms/linfa-ensemble/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<DecisionTree<f64, usize>> =
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<f64, usize> =
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<DecisionTree<f64, usize>, 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);
}
}
Loading