use std::path::Path;
use automl::settings::Algorithm;
use automl::SupervisedModel;
use smartcore::tree::decision_tree_regressor::DecisionTreeRegressorParameters;
use super::metaml::{MetaMLDataset, MetaMLModel};
#[derive(Default)]
pub struct LinearRegressor {
#[allow(dead_code)]
model: Option<SupervisedModel>,
}
impl LinearRegressor {
#[must_use]
pub const fn new() -> Self {
Self { model: None }
}
}
impl MetaMLModel for LinearRegressor {
fn train(&mut self, data: MetaMLDataset) {
let settings = automl::Settings::default_regression().only(Algorithm::Linear);
let mut model = SupervisedModel::new(data, settings);
model.train();
self.model = Some(model);
}
fn predict(&self, features: &[f32; 6]) -> Result<f32, String> {
let Some(model) = self.model.as_ref() else {
return Err("Model must be trained before being saved".to_string());
};
Ok(model.predict(vec![features.to_vec()])[0])
}
fn load(path: &Path) -> Result<Self, String> {
let Some(path_str) = path.to_str() else {
return Err("Failed to convert path to a string".to_string());
};
let model = SupervisedModel::new_from_file(path_str);
Ok(Self { model: Some(model) })
}
fn save(&self, path: &Path) -> Result<(), String> {
let Some(model) = self.model.as_ref() else {
return Err("Model must be trained before being saved".to_string());
};
let Some(path_str) = path.to_str() else {
return Err("Failed to convert path to a string".to_string());
};
model.save(path_str);
Ok(())
}
}
#[derive(Default)]
pub struct DecisionTreeRegressor {
#[allow(dead_code)]
model: Option<SupervisedModel>,
}
impl DecisionTreeRegressor {
#[allow(dead_code)]
const MAX_DEPTH: u16 = 3;
#[must_use]
pub const fn new() -> Self {
Self { model: None }
}
}
impl MetaMLModel for DecisionTreeRegressor {
fn train(&mut self, data: MetaMLDataset) {
let settings = automl::Settings::default_regression()
.only(Algorithm::DecisionTreeRegressor)
.with_decision_tree_regressor_settings(
DecisionTreeRegressorParameters::default().with_max_depth(Self::MAX_DEPTH),
);
let mut model = SupervisedModel::new(data, settings);
model.train();
self.model = Some(model);
}
fn predict(&self, features: &[f32; 6]) -> Result<f32, String> {
let Some(model) = self.model.as_ref() else {
return Err("Model must be trained before being saved".to_string());
};
Ok(model.predict(vec![features.to_vec()])[0])
}
fn load(path: &Path) -> Result<Self, String> {
let Some(path_str) = path.to_str() else {
return Err("Failed to convert path to a string".to_string());
};
let model = SupervisedModel::new_from_file(path_str);
Ok(Self { model: Some(model) })
}
fn save(&self, path: &Path) -> Result<(), String> {
let Some(model) = &self.model.as_ref() else {
return Err("Model must be trained before being saved".to_string());
};
let Some(path_str) = path.to_str() else {
return Err("Failed to convert path to a string".to_string());
};
model.save(path_str);
Ok(())
}
}