1use std::default::Default;
56use std::fmt::Debug;
57
58#[cfg(feature = "serde")]
59use serde::{Deserialize, Serialize};
60
61use crate::api::{Predictor, SupervisedEstimator};
62use crate::ensemble::base_forest_regressor::{BaseForestRegressor, BaseForestRegressorParameters};
63use crate::error::Failed;
64use crate::linalg::basic::arrays::{Array1, Array2};
65use crate::numbers::basenum::Number;
66use crate::numbers::floatnum::FloatNumber;
67use crate::tree::base_tree_regressor::Splitter;
68
69#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
70#[derive(Debug, Clone)]
71pub struct ExtraTreesRegressorParameters {
74 #[cfg_attr(feature = "serde", serde(default))]
75 pub max_depth: Option<u16>,
77 #[cfg_attr(feature = "serde", serde(default))]
78 pub min_samples_leaf: usize,
80 #[cfg_attr(feature = "serde", serde(default))]
81 pub min_samples_split: usize,
83 #[cfg_attr(feature = "serde", serde(default))]
84 pub n_trees: usize,
86 #[cfg_attr(feature = "serde", serde(default))]
87 pub m: Option<usize>,
89 #[cfg_attr(feature = "serde", serde(default))]
90 pub keep_samples: bool,
92 #[cfg_attr(feature = "serde", serde(default))]
93 pub seed: u64,
95}
96
97#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
99#[derive(Debug)]
100pub struct ExtraTreesRegressor<
101 TX: Number + FloatNumber + PartialOrd,
102 TY: Number,
103 X: Array2<TX>,
104 Y: Array1<TY>,
105> {
106 forest_regressor: Option<BaseForestRegressor<TX, TY, X, Y>>,
107}
108
109impl ExtraTreesRegressorParameters {
110 pub fn with_max_depth(mut self, max_depth: u16) -> Self {
112 self.max_depth = Some(max_depth);
113 self
114 }
115 pub fn with_min_samples_leaf(mut self, min_samples_leaf: usize) -> Self {
117 self.min_samples_leaf = min_samples_leaf;
118 self
119 }
120 pub fn with_min_samples_split(mut self, min_samples_split: usize) -> Self {
122 self.min_samples_split = min_samples_split;
123 self
124 }
125 pub fn with_n_trees(mut self, n_trees: usize) -> Self {
127 self.n_trees = n_trees;
128 self
129 }
130 pub fn with_m(mut self, m: usize) -> Self {
132 self.m = Some(m);
133 self
134 }
135
136 pub fn with_keep_samples(mut self, keep_samples: bool) -> Self {
138 self.keep_samples = keep_samples;
139 self
140 }
141
142 pub fn with_seed(mut self, seed: u64) -> Self {
144 self.seed = seed;
145 self
146 }
147}
148impl Default for ExtraTreesRegressorParameters {
149 fn default() -> Self {
150 ExtraTreesRegressorParameters {
151 max_depth: Option::None,
152 min_samples_leaf: 1,
153 min_samples_split: 2,
154 n_trees: 10,
155 m: Option::None,
156 keep_samples: false,
157 seed: 0,
158 }
159 }
160}
161
162impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>>
163 SupervisedEstimator<X, Y, ExtraTreesRegressorParameters> for ExtraTreesRegressor<TX, TY, X, Y>
164{
165 fn new() -> Self {
166 Self {
167 forest_regressor: Option::None,
168 }
169 }
170
171 fn fit(x: &X, y: &Y, parameters: ExtraTreesRegressorParameters) -> Result<Self, Failed> {
172 ExtraTreesRegressor::fit(x, y, parameters)
173 }
174}
175
176impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>>
177 Predictor<X, Y> for ExtraTreesRegressor<TX, TY, X, Y>
178{
179 fn predict(&self, x: &X) -> Result<Y, Failed> {
180 self.predict(x)
181 }
182}
183
184impl<TX: Number + FloatNumber + PartialOrd, TY: Number, X: Array2<TX>, Y: Array1<TY>>
185 ExtraTreesRegressor<TX, TY, X, Y>
186{
187 pub fn fit(
191 x: &X,
192 y: &Y,
193 parameters: ExtraTreesRegressorParameters,
194 ) -> Result<ExtraTreesRegressor<TX, TY, X, Y>, Failed> {
195 let regressor_params = BaseForestRegressorParameters {
196 max_depth: parameters.max_depth,
197 min_samples_leaf: parameters.min_samples_leaf,
198 min_samples_split: parameters.min_samples_split,
199 n_trees: parameters.n_trees,
200 m: parameters.m,
201 keep_samples: parameters.keep_samples,
202 seed: parameters.seed,
203 bootstrap: false,
204 splitter: Splitter::Random,
205 };
206 let forest_regressor = BaseForestRegressor::fit(x, y, regressor_params)?;
207
208 Ok(ExtraTreesRegressor {
209 forest_regressor: Some(forest_regressor),
210 })
211 }
212
213 pub fn predict(&self, x: &X) -> Result<Y, Failed> {
216 let forest_regressor = self.forest_regressor.as_ref().unwrap();
217 forest_regressor.predict(x)
218 }
219
220 pub fn predict_oob(&self, x: &X) -> Result<Y, Failed> {
222 let forest_regressor = self.forest_regressor.as_ref().unwrap();
223 forest_regressor.predict_oob(x)
224 }
225}
226
227#[cfg(test)]
228mod tests {
229 use super::*;
230 use crate::linalg::basic::matrix::DenseMatrix;
231 use crate::metrics::mean_squared_error;
232
233 #[test]
234 fn test_extra_trees_regressor_fit_predict() {
235 let x = DenseMatrix::from_2d_array(&[
237 &[1., 2.],
238 &[3., 4.],
239 &[5., 6.],
240 &[7., 8.],
241 &[9., 10.],
242 &[11., 12.],
243 &[13., 14.],
244 &[15., 16.],
245 ])
246 .unwrap();
247 let y = vec![1., 2., 3., 4., 5., 6., 7., 8.];
248
249 let parameters = ExtraTreesRegressorParameters::default()
250 .with_n_trees(100)
251 .with_seed(42);
252
253 let regressor = ExtraTreesRegressor::fit(&x, &y, parameters).unwrap();
254 let y_hat = regressor.predict(&x).unwrap();
255
256 assert_eq!(y_hat.len(), y.len());
257 let mse = mean_squared_error(&y, &y_hat);
260 assert!(mse < 1.0);
262 }
263
264 #[test]
265 fn test_fit_predict_higher_dims() {
266 let x = DenseMatrix::from_2d_array(&[
268 &[0., 0., 10., 5., 8., 1., 4., 9., 2., 7.],
270 &[0., 0., 20., 1., 2., 3., 4., 5., 6., 7.],
271 &[0., 0., 30., 7., 6., 5., 4., 3., 2., 1.],
272 &[0., 0., 40., 9., 2., 4., 6., 8., 1., 3.],
273 &[0., 0., 55., 3., 1., 8., 6., 4., 2., 9.],
274 &[0., 0., 65., 2., 4., 7., 5., 3., 1., 8.],
275 ])
276 .unwrap();
277 let y = vec![10., 20., 30., 40., 55., 65.];
278
279 let parameters = ExtraTreesRegressorParameters::default()
280 .with_n_trees(100)
281 .with_seed(42);
282
283 let regressor = ExtraTreesRegressor::fit(&x, &y, parameters).unwrap();
284 let y_hat = regressor.predict(&x).unwrap();
285
286 assert_eq!(y_hat.len(), y.len());
287
288 let mse = mean_squared_error(&y, &y_hat);
289
290 assert!(mse < 1.0);
293 }
294
295 #[test]
296 fn test_reproducibility() {
297 let x = DenseMatrix::from_2d_array(&[
298 &[1., 2.],
299 &[3., 4.],
300 &[5., 6.],
301 &[7., 8.],
302 &[9., 10.],
303 &[11., 12.],
304 ])
305 .unwrap();
306 let y = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
307
308 let params = ExtraTreesRegressorParameters::default().with_seed(42);
309
310 let regressor1 = ExtraTreesRegressor::fit(&x, &y, params.clone()).unwrap();
311 let y_hat1 = regressor1.predict(&x).unwrap();
312
313 let regressor2 = ExtraTreesRegressor::fit(&x, &y, params.clone()).unwrap();
314 let y_hat2 = regressor2.predict(&x).unwrap();
315
316 assert_eq!(y_hat1, y_hat2);
317 }
318}