use crate::dataframe::DataFrame;
use crate::error::{Error, Result};
use crate::ml::models::{ModelMetrics, SupervisedModel};
use crate::series::Series;
use std::collections::HashMap;
pub trait TunableModel: SupervisedModel + Clone {
fn set_params(&mut self, params: &HashMap<String, String>) -> Result<()>;
}
pub(crate) fn parse_param_bool(key: &str, value: &str) -> Result<bool> {
match value.trim().to_ascii_lowercase().as_str() {
"true" | "1" => Ok(true),
"false" | "0" => Ok(false),
_ => Err(Error::InvalidValue(format!(
"Invalid boolean value for hyperparameter '{}': '{}'",
key, value
))),
}
}
pub(crate) fn parse_param_f64(key: &str, value: &str) -> Result<f64> {
value.trim().parse::<f64>().map_err(|_| {
Error::InvalidValue(format!(
"Invalid floating-point value for hyperparameter '{}': '{}'",
key, value
))
})
}
pub(crate) fn parse_param_usize(key: &str, value: &str) -> Result<usize> {
value.trim().parse::<usize>().map_err(|_| {
Error::InvalidValue(format!(
"Invalid unsigned integer value for hyperparameter '{}': '{}'",
key, value
))
})
}
pub(crate) fn parse_param_u64(key: &str, value: &str) -> Result<u64> {
value.trim().parse::<u64>().map_err(|_| {
Error::InvalidValue(format!(
"Invalid u64 value for hyperparameter '{}': '{}'",
key, value
))
})
}
pub(crate) fn err_on_unknown_params(model_type: &str, unknown: Vec<String>) -> Result<()> {
if unknown.is_empty() {
Ok(())
} else {
let mut unknown = unknown;
unknown.sort();
Err(Error::InvalidValue(format!(
"set_params: unknown hyperparameter(s) {:?} for {}",
unknown, model_type
)))
}
}
#[derive(Debug, Clone)]
pub struct HyperparameterGrid {
pub params: HashMap<String, Vec<String>>,
}
impl HyperparameterGrid {
pub fn new() -> Self {
HyperparameterGrid {
params: HashMap::new(),
}
}
pub fn add_param<T: ToString>(&mut self, name: &str, values: Vec<T>) -> &mut Self {
let string_values = values.into_iter().map(|v| v.to_string()).collect();
self.params.insert(name.to_string(), string_values);
self
}
pub fn parameter_combinations(&self) -> Vec<HashMap<String, String>> {
if self.params.is_empty() {
return vec![HashMap::new()];
}
let keys: Vec<String> = {
let mut k: Vec<String> = self.params.keys().cloned().collect();
k.sort();
k
};
let mut result: Vec<HashMap<String, String>> = vec![HashMap::new()];
for key in &keys {
let values = &self.params[key];
let mut new_result = Vec::with_capacity(result.len() * values.len());
for existing in &result {
for value in values {
let mut combo = existing.clone();
combo.insert(key.clone(), value.clone());
new_result.push(combo);
}
}
result = new_result;
}
result
}
fn sorted_keys(&self) -> Vec<String> {
let mut keys: Vec<String> = self.params.keys().cloned().collect();
keys.sort();
keys
}
}
impl Default for HyperparameterGrid {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
struct CombinationOutcome {
params: HashMap<String, String>,
fold_scores: Vec<f64>,
mean_score: f64,
std_score: f64,
}
fn metric_higher_is_better(metric: &str) -> Result<bool> {
match metric.trim().to_ascii_lowercase().as_str() {
"r2" | "r2_score" | "accuracy" | "precision" | "recall" | "f1" | "auc" | "roc_auc"
| "explained_variance" => Ok(true),
"mse" | "rmse" | "mae" | "mape" | "msle" | "log_loss" | "error" => Ok(false),
other => Err(Error::InvalidInput(format!(
"Unknown scoring metric '{}': cannot determine whether higher or lower is better. \
Supported metrics are r2, accuracy, precision, recall, f1, auc/roc_auc, \
explained_variance (higher is better) and mse, rmse, mae, mape, msle, log_loss \
(lower is better)",
other
))),
}
}
fn sorted_pairs(params: &HashMap<String, String>) -> Vec<String> {
let mut pairs: Vec<String> = params
.iter()
.map(|(key, value)| format!("{}={}", key, value))
.collect();
pairs.sort();
pairs
}
fn fold_score(metrics: &ModelMetrics, scoring: &str, fold_idx: usize) -> Result<f64> {
let value = metrics.get_metric(scoring).copied().ok_or_else(|| {
let mut available: Vec<&String> = metrics.metrics.keys().collect();
available.sort();
Error::InvalidInput(format!(
"Scoring metric '{}' was not reported by the model for CV fold {}; \
metrics reported by this model: {:?}",
scoring, fold_idx, available
))
})?;
if !value.is_finite() {
return Err(Error::Computation(format!(
"Scoring metric '{}' was {} on CV fold {}; a non-finite score cannot be \
compared against other parameter combinations",
scoring, value, fold_idx
)));
}
Ok(value)
}
fn evaluate_combination<T: TunableModel>(
base_model: &T,
params: &HashMap<String, String>,
data: &DataFrame,
target: &str,
cv: usize,
scoring: &str,
) -> Result<CombinationOutcome> {
let mut candidate = base_model.clone();
candidate.set_params(params).map_err(|e| {
Error::InvalidValue(format!(
"parameter combination {:?} could not be applied to the model: {}",
sorted_pairs(params),
e
))
})?;
let fold_metrics = candidate.cross_validate(data, target, cv).map_err(|e| {
Error::Computation(format!(
"parameter combination {:?} failed cross-validation: {}",
sorted_pairs(params),
e
))
})?;
if fold_metrics.len() != cv {
return Err(Error::InvalidOperation(format!(
"parameter combination {:?}: cross_validate returned {} fold result(s) for a \
{}-fold cross-validation; the model's cross_validate implementation must \
produce exactly one ModelMetrics per fold",
sorted_pairs(params),
fold_metrics.len(),
cv
)));
}
let mut fold_scores = Vec::with_capacity(fold_metrics.len());
for (fold_idx, metrics) in fold_metrics.iter().enumerate() {
fold_scores.push(fold_score(metrics, scoring, fold_idx).map_err(|e| {
Error::InvalidInput(format!(
"parameter combination {:?}: {}",
sorted_pairs(params),
e
))
})?);
}
let n = fold_scores.len() as f64;
let mean_score = fold_scores.iter().sum::<f64>() / n;
let variance = fold_scores
.iter()
.map(|&s| (s - mean_score).powi(2))
.sum::<f64>()
/ n;
Ok(CombinationOutcome {
params: params.clone(),
fold_scores,
mean_score,
std_score: variance.sqrt(),
})
}
fn evaluate_all<T: TunableModel>(
base_model: &T,
combinations: &[HashMap<String, String>],
data: &DataFrame,
target: &str,
cv: usize,
scoring: &str,
) -> Result<Vec<CombinationOutcome>> {
metric_higher_is_better(scoring)?;
let mut outcomes = Vec::with_capacity(combinations.len());
for params in combinations {
outcomes.push(evaluate_combination(
base_model, params, data, target, cv, scoring,
)?);
}
Ok(outcomes)
}
fn best_outcome_index(outcomes: &[CombinationOutcome], scoring: &str) -> Result<usize> {
let higher_is_better = metric_higher_is_better(scoring)?;
let mut best_idx = 0usize;
let mut best = outcomes
.first()
.ok_or_else(|| Error::InvalidInput("No parameter combination was evaluated".into()))?
.mean_score;
for (idx, outcome) in outcomes.iter().enumerate().skip(1) {
let better = if higher_is_better {
outcome.mean_score > best
} else {
outcome.mean_score < best
};
if better {
best = outcome.mean_score;
best_idx = idx;
}
}
Ok(best_idx)
}
fn rank_outcomes(outcomes: &[CombinationOutcome], scoring: &str) -> Result<Vec<i64>> {
let higher_is_better = metric_higher_is_better(scoring)?;
let mut ranks = Vec::with_capacity(outcomes.len());
for outcome in outcomes {
let better_count = outcomes
.iter()
.filter(|other| {
if higher_is_better {
other.mean_score > outcome.mean_score
} else {
other.mean_score < outcome.mean_score
}
})
.count();
ranks.push(better_count as i64 + 1);
}
Ok(ranks)
}
fn build_cv_results(
outcomes: &[CombinationOutcome],
param_names: &[String],
scoring: &str,
) -> Result<DataFrame> {
let mut df = DataFrame::new();
for name in param_names {
let column: Vec<String> = outcomes
.iter()
.map(|o| {
o.params.get(name).cloned().ok_or_else(|| {
Error::InvalidOperation(format!(
"cv_results: evaluated combination {:?} is missing searched \
parameter '{}'",
o.params, name
))
})
})
.collect::<Result<Vec<String>>>()?;
let col_name = format!("param_{}", name);
df.add_column(
col_name.clone(),
Series::new(column, Some(col_name.clone()))?,
)?;
}
let n_folds = outcomes.first().map(|o| o.fold_scores.len()).unwrap_or(0);
for fold_idx in 0..n_folds {
let column: Vec<f64> = outcomes
.iter()
.map(|o| o.fold_scores.get(fold_idx).copied().unwrap_or(f64::NAN))
.collect();
let col_name = format!("split{}_test_score", fold_idx);
df.add_column(
col_name.clone(),
Series::new(column, Some(col_name.clone()))?,
)?;
}
let means: Vec<f64> = outcomes.iter().map(|o| o.mean_score).collect();
df.add_column(
"mean_test_score".to_string(),
Series::new(means, Some("mean_test_score".to_string()))?,
)?;
let stds: Vec<f64> = outcomes.iter().map(|o| o.std_score).collect();
df.add_column(
"std_test_score".to_string(),
Series::new(stds, Some("std_test_score".to_string()))?,
)?;
let ranks = rank_outcomes(outcomes, scoring)?;
df.add_column(
"rank_test_score".to_string(),
Series::new(ranks, Some("rank_test_score".to_string()))?,
)?;
Ok(df)
}
pub struct GridSearchCV<T: SupervisedModel> {
pub base_model: T,
pub param_grid: HyperparameterGrid,
pub scoring: String,
pub cv: usize,
pub refit: bool,
pub best_params: Option<HashMap<String, String>>,
pub best_score: Option<f64>,
pub cv_results: Option<DataFrame>,
pub best_estimator_: Option<T>,
}
impl<T: SupervisedModel + Clone> GridSearchCV<T> {
pub fn new(base_model: T, param_grid: HyperparameterGrid, scoring: &str, cv: usize) -> Self {
GridSearchCV {
base_model,
param_grid,
scoring: scoring.to_string(),
cv,
refit: true,
best_params: None,
best_score: None,
cv_results: None,
best_estimator_: None,
}
}
pub fn with_refit(mut self, refit: bool) -> Self {
self.refit = refit;
self
}
pub fn best_estimator(&self) -> Result<T> {
self.best_estimator_
.clone()
.ok_or_else(|| Error::InvalidValue("Grid search not fitted".into()))
}
}
impl<T: TunableModel> GridSearchCV<T> {
pub fn fit(&mut self, data: &DataFrame, target: &str) -> Result<()> {
if !data.has_column(target) {
return Err(Error::InvalidValue(format!(
"Target column '{}' not found",
target
)));
}
if self.cv < 2 {
return Err(Error::InvalidInput(
"Number of CV folds must be at least 2".into(),
));
}
let param_combinations = self.param_grid.parameter_combinations();
if param_combinations.is_empty() {
return Err(Error::InvalidInput(
"No parameter combinations to search".into(),
));
}
let outcomes = evaluate_all(
&self.base_model,
¶m_combinations,
data,
target,
self.cv,
&self.scoring,
)?;
let best_idx = best_outcome_index(&outcomes, &self.scoring)?;
let best = &outcomes[best_idx];
self.best_params = Some(best.params.clone());
self.best_score = Some(best.mean_score);
self.cv_results = Some(build_cv_results(
&outcomes,
&self.param_grid.sorted_keys(),
&self.scoring,
)?);
let mut best_model = self.base_model.clone();
best_model.set_params(&best.params)?;
if self.refit {
best_model.fit(data, target)?;
}
self.best_estimator_ = Some(best_model);
Ok(())
}
}
pub struct RandomizedSearchCV<T: SupervisedModel> {
pub base_model: T,
pub param_grid: HyperparameterGrid,
pub n_iter: usize,
pub scoring: String,
pub cv: usize,
pub random_seed: Option<u64>,
pub refit: bool,
pub best_params: Option<HashMap<String, String>>,
pub best_score: Option<f64>,
pub cv_results: Option<DataFrame>,
pub best_estimator_: Option<T>,
}
impl<T: SupervisedModel + Clone> RandomizedSearchCV<T> {
pub fn new(
base_model: T,
param_grid: HyperparameterGrid,
n_iter: usize,
scoring: &str,
cv: usize,
) -> Self {
RandomizedSearchCV {
base_model,
param_grid,
n_iter,
scoring: scoring.to_string(),
cv,
random_seed: None,
refit: true,
best_params: None,
best_score: None,
cv_results: None,
best_estimator_: None,
}
}
pub fn with_random_seed(mut self, seed: u64) -> Self {
self.random_seed = Some(seed);
self
}
pub fn with_refit(mut self, refit: bool) -> Self {
self.refit = refit;
self
}
pub fn best_estimator(&self) -> Result<T> {
self.best_estimator_
.clone()
.ok_or_else(|| Error::InvalidValue("Randomized search not fitted".into()))
}
fn sample_combinations(&self) -> Vec<HashMap<String, String>> {
let all_combinations = self.param_grid.parameter_combinations();
let n_to_try = self.n_iter.min(all_combinations.len());
if n_to_try >= all_combinations.len() {
return all_combinations;
}
use scirs2_core::random::rngs::StdRng;
use scirs2_core::random::SeedableRng;
use scirs2_core::random::SliceRandom;
let mut rng: StdRng = match self.random_seed {
Some(seed) => StdRng::seed_from_u64(seed),
None => StdRng::seed_from_u64(scirs2_core::random::random::<u64>()),
};
let mut indices: Vec<usize> = (0..all_combinations.len()).collect();
indices.shuffle(&mut rng);
let mut selected: Vec<usize> = indices[..n_to_try].to_vec();
selected.sort_unstable();
selected
.into_iter()
.map(|i| all_combinations[i].clone())
.collect()
}
}
impl<T: TunableModel> RandomizedSearchCV<T> {
pub fn fit(&mut self, data: &DataFrame, target: &str) -> Result<()> {
if !data.has_column(target) {
return Err(Error::InvalidValue(format!(
"Target column '{}' not found",
target
)));
}
if self.cv < 2 {
return Err(Error::InvalidInput(
"Number of CV folds must be at least 2".into(),
));
}
if self.n_iter == 0 {
return Err(Error::InvalidInput(
"n_iter must be at least 1: a randomized search that samples no \
parameter combination cannot report a best configuration"
.into(),
));
}
let selected_combos = self.sample_combinations();
if selected_combos.is_empty() {
return Err(Error::InvalidInput(
"No parameter combinations to search".into(),
));
}
let outcomes = evaluate_all(
&self.base_model,
&selected_combos,
data,
target,
self.cv,
&self.scoring,
)?;
let best_idx = best_outcome_index(&outcomes, &self.scoring)?;
let best = &outcomes[best_idx];
self.best_params = Some(best.params.clone());
self.best_score = Some(best.mean_score);
self.cv_results = Some(build_cv_results(
&outcomes,
&self.param_grid.sorted_keys(),
&self.scoring,
)?);
let mut best_model = self.base_model.clone();
best_model.set_params(&best.params)?;
if self.refit {
best_model.fit(data, target)?;
}
self.best_estimator_ = Some(best_model);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dataframe::DataFrame;
use crate::ml::models::linear::LinearRegression;
use crate::series::Series;
fn make_linear_df(n: usize) -> DataFrame {
let x: Vec<f64> = (0..n).map(|i| i as f64).collect();
let y: Vec<f64> = x.iter().map(|&v| 2.0 * v + 1.0).collect();
let mut df = DataFrame::new();
df.add_column(
"x".to_string(),
Series::new(x, Some("x".to_string())).expect("Series::new"),
)
.expect("add x");
df.add_column(
"y".to_string(),
Series::new(y, Some("y".to_string())).expect("Series::new"),
)
.expect("add y");
df
}
#[test]
fn test_cartesian_product() {
let mut grid = HyperparameterGrid::new();
grid.add_param("a", vec!["1", "2"]);
grid.add_param("b", vec!["x", "y"]);
let combos = grid.parameter_combinations();
assert_eq!(
combos.len(),
4,
"2x2 Cartesian product must yield exactly 4 combinations"
);
for combo in &combos {
assert!(combo.contains_key("a"), "combo missing key 'a'");
assert!(combo.contains_key("b"), "combo missing key 'b'");
}
}
#[test]
fn test_cartesian_empty() {
let grid = HyperparameterGrid::new();
let combos = grid.parameter_combinations();
assert_eq!(
combos.len(),
1,
"empty grid must return exactly one (empty) combination"
);
assert!(combos[0].is_empty(), "the single combination must be empty");
}
#[test]
fn test_gridsearch_cv_real() {
let df = make_linear_df(10);
let model = LinearRegression::new();
let grid = HyperparameterGrid::new(); let mut gs = GridSearchCV::new(model, grid, "r2", 2);
gs.fit(&df, "y").expect("GridSearchCV::fit should succeed");
let best_score = gs.best_score.expect("best_score must be set after fit");
assert!(
best_score > 0.0,
"best_score must be a real positive CV score, got {}",
best_score
);
assert!(gs.cv_results.is_some(), "cv_results must be set after fit");
}
#[test]
fn test_gridsearch_scores_every_combination() {
let df = make_linear_df(12);
let mut grid = HyperparameterGrid::new();
grid.add_param("fit_intercept", vec!["true", "false"]);
let mut gs = GridSearchCV::new(LinearRegression::new(), grid, "r2", 3);
gs.fit(&df, "y").expect("GridSearchCV::fit should succeed");
let results = gs.cv_results.as_ref().expect("cv_results");
assert_eq!(results.nrows(), 2, "one row per parameter combination");
let scores = results
.get_column::<f64>("mean_test_score")
.expect("mean_test_score column")
.values()
.to_vec();
assert!(
(scores[0] - scores[1]).abs() > 1e-9,
"each combination must be scored on its own; got identical scores {:?}",
scores
);
let best = gs.best_params.as_ref().expect("best_params");
assert_eq!(
best.get("fit_intercept").map(String::as_str),
Some("true"),
"the genuinely best combination must be reported as best"
);
}
#[test]
fn test_unknown_scoring_metric_errors() {
let df = make_linear_df(10);
let mut gs = GridSearchCV::new(
LinearRegression::new(),
HyperparameterGrid::new(),
"not_a_metric",
2,
);
assert!(
gs.fit(&df, "y").is_err(),
"an unknown scoring metric must be rejected, not silently ranked"
);
}
#[test]
fn test_unknown_parameter_errors() {
let df = make_linear_df(10);
let mut grid = HyperparameterGrid::new();
grid.add_param("no_such_param", vec!["1", "2"]);
let mut gs = GridSearchCV::new(LinearRegression::new(), grid, "r2", 2);
assert!(
gs.fit(&df, "y").is_err(),
"a parameter the model cannot accept must error, not be skipped"
);
}
#[test]
fn test_randomized_search_samples_and_scores() {
let df = make_linear_df(12);
let mut grid = HyperparameterGrid::new();
grid.add_param("fit_intercept", vec!["true", "false"]);
let mut rs =
RandomizedSearchCV::new(LinearRegression::new(), grid, 2, "r2", 3).with_random_seed(42);
rs.fit(&df, "y")
.expect("RandomizedSearchCV::fit should succeed");
let results = rs.cv_results.as_ref().expect("cv_results");
assert_eq!(results.nrows(), 2, "one row per sampled combination");
assert_eq!(
rs.best_params
.as_ref()
.and_then(|p| p.get("fit_intercept"))
.map(String::as_str),
Some("true")
);
}
}