use scirs2_core::ndarray::{Array1, Array2, Axis};
use serde::{Deserialize, Serialize};
use sklears_core::{
error::{Result, SklearsError},
types::Float,
};
use std::collections::HashMap;
pub trait SklearnTransformer: std::fmt::Debug {
fn fit(&mut self, x: &Array2<Float>, y: Option<&Array1<Float>>) -> Result<()>;
fn transform(&self, x: &Array2<Float>) -> Result<Array2<Float>>;
fn fit_transform(
&mut self,
x: &Array2<Float>,
y: Option<&Array1<Float>>,
) -> Result<Array2<Float>> {
self.fit(x, y)?;
self.transform(x)
}
fn inverse_transform(&self, _x: &Array2<Float>) -> Result<Array2<Float>> {
Err(SklearsError::InvalidInput(
"Inverse transform not implemented for this transformer".to_string(),
))
}
fn get_feature_names_out(&self, input_features: Option<&[String]>) -> Vec<String>;
fn set_params(&mut self, params: &HashMap<String, ParameterValue>) -> Result<()>;
fn get_params(&self, deep: bool) -> HashMap<String, ParameterValue>;
fn is_fitted(&self) -> bool;
fn get_name(&self) -> String;
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ParameterValue {
Float(Float),
Int(i64),
Bool(bool),
String(String),
FloatArray(Vec<Float>),
IntArray(Vec<i64>),
None,
}
impl From<Float> for ParameterValue {
fn from(value: Float) -> Self {
ParameterValue::Float(value)
}
}
impl From<i64> for ParameterValue {
fn from(value: i64) -> Self {
ParameterValue::Int(value)
}
}
impl From<bool> for ParameterValue {
fn from(value: bool) -> Self {
ParameterValue::Bool(value)
}
}
impl From<String> for ParameterValue {
fn from(value: String) -> Self {
ParameterValue::String(value)
}
}
impl From<&str> for ParameterValue {
fn from(value: &str) -> Self {
ParameterValue::String(value.to_string())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SklearnPCA {
pub n_components: Option<usize>,
pub copy: bool,
pub whiten: bool,
pub svd_solver: String,
pub tol: Float,
pub iterated_power: String,
pub n_oversamples: usize,
pub power_iteration_normalizer: String,
pub random_state: Option<u64>,
components_: Option<Array2<Float>>,
explained_variance_: Option<Array1<Float>>,
explained_variance_ratio_: Option<Array1<Float>>,
singular_values_: Option<Array1<Float>>,
mean_: Option<Array1<Float>>,
n_components_: Option<usize>,
n_features_in_: Option<usize>,
feature_names_in_: Option<Vec<String>>,
is_fitted_: bool,
}
impl Default for SklearnPCA {
fn default() -> Self {
Self {
n_components: None,
copy: true,
whiten: false,
svd_solver: "auto".to_string(),
tol: 0.0,
iterated_power: "auto".to_string(),
n_oversamples: 10,
power_iteration_normalizer: "auto".to_string(),
random_state: None,
components_: None,
explained_variance_: None,
explained_variance_ratio_: None,
singular_values_: None,
mean_: None,
n_components_: None,
n_features_in_: None,
feature_names_in_: None,
is_fitted_: false,
}
}
}
impl SklearnPCA {
pub fn new() -> Self {
Self::default()
}
pub fn with_n_components(mut self, n_components: usize) -> Self {
self.n_components = Some(n_components);
self
}
pub fn with_whiten(mut self, whiten: bool) -> Self {
self.whiten = whiten;
self
}
pub fn with_svd_solver(mut self, solver: &str) -> Self {
self.svd_solver = solver.to_string();
self
}
pub fn with_random_state(mut self, random_state: u64) -> Self {
self.random_state = Some(random_state);
self
}
pub fn components(&self) -> Option<&Array2<Float>> {
self.components_.as_ref()
}
pub fn explained_variance(&self) -> Option<&Array1<Float>> {
self.explained_variance_.as_ref()
}
pub fn explained_variance_ratio(&self) -> Option<&Array1<Float>> {
self.explained_variance_ratio_.as_ref()
}
pub fn singular_values(&self) -> Option<&Array1<Float>> {
self.singular_values_.as_ref()
}
pub fn mean(&self) -> Option<&Array1<Float>> {
self.mean_.as_ref()
}
pub fn n_components_(&self) -> Option<usize> {
self.n_components_
}
fn validate_input(&self, x: &Array2<Float>) -> Result<()> {
let (n_samples, n_features) = x.dim();
if n_samples == 0 {
return Err(SklearsError::InvalidInput("Empty data matrix".to_string()));
}
if n_features == 0 {
return Err(SklearsError::InvalidInput(
"No features in data".to_string(),
));
}
if x.iter().any(|&val| !val.is_finite()) {
return Err(SklearsError::InvalidInput(
"Data contains NaN or infinite values".to_string(),
));
}
Ok(())
}
fn determine_n_components(&self, n_features: usize) -> usize {
match self.n_components {
Some(n) => n.min(n_features),
None => n_features.min(50), }
}
}
impl SklearnTransformer for SklearnPCA {
fn fit(&mut self, x: &Array2<Float>, _y: Option<&Array1<Float>>) -> Result<()> {
self.validate_input(x)?;
let (n_samples, n_features) = x.dim();
self.n_features_in_ = Some(n_features);
let n_components = self.determine_n_components(n_features);
self.n_components_ = Some(n_components);
let mean = x
.mean_axis(Axis(0))
.ok_or_else(|| SklearsError::InvalidInput("Failed to compute data mean".to_string()))?;
self.mean_ = Some(mean.clone());
let centered_data = x - &mean.insert_axis(Axis(0));
let (_u, s, vt) = self.compute_svd(¢ered_data, n_components)?;
self.components_ = Some(
vt.slice(scirs2_core::ndarray::s![..n_components, ..])
.to_owned(),
);
self.singular_values_ = Some(s.slice(scirs2_core::ndarray::s![..n_components]).to_owned());
let explained_var = s
.slice(scirs2_core::ndarray::s![..n_components])
.mapv(|x| x.powi(2) / (n_samples - 1) as Float);
self.explained_variance_ = Some(explained_var.clone());
let total_var = explained_var.sum();
let explained_var_ratio = if total_var > 0.0 {
&explained_var / total_var
} else {
Array1::zeros(explained_var.len())
};
self.explained_variance_ratio_ = Some(explained_var_ratio);
self.is_fitted_ = true;
Ok(())
}
fn transform(&self, x: &Array2<Float>) -> Result<Array2<Float>> {
if !self.is_fitted_ {
return Err(SklearsError::InvalidInput("PCA not fitted".to_string()));
}
self.validate_input(x)?;
let components = self.components_.as_ref().expect("operation should succeed");
let mean = self.mean_.as_ref().expect("operation should succeed");
let mean_broadcast = mean.clone().insert_axis(Axis(0));
let centered_data = x - &mean_broadcast;
let transformed = centered_data.dot(&components.t());
if self.whiten {
let explained_var = self
.explained_variance_
.as_ref()
.expect("operation should succeed");
let sqrt_explained_var = explained_var.mapv(|x| x.sqrt());
let whitened = &transformed / &sqrt_explained_var.insert_axis(Axis(0));
Ok(whitened)
} else {
Ok(transformed)
}
}
fn inverse_transform(&self, x: &Array2<Float>) -> Result<Array2<Float>> {
if !self.is_fitted_ {
return Err(SklearsError::InvalidInput("PCA not fitted".to_string()));
}
let components = self.components_.as_ref().expect("operation should succeed");
let mean = self.mean_.as_ref().expect("operation should succeed");
let unwhitened = if self.whiten {
let explained_var = self
.explained_variance_
.as_ref()
.expect("operation should succeed");
let sqrt_explained_var = explained_var.mapv(|x| x.sqrt());
x * &sqrt_explained_var.insert_axis(Axis(0))
} else {
x.clone()
};
let mean_broadcast = mean.clone().insert_axis(Axis(0));
let reconstructed = unwhitened.dot(components) + &mean_broadcast;
Ok(reconstructed)
}
fn get_feature_names_out(&self, input_features: Option<&[String]>) -> Vec<String> {
let n_components = self
.n_components_
.unwrap_or_else(|| self.n_components.unwrap_or(0));
match input_features {
Some(features) => {
let prefix = format!("{}__", features.join("_"));
(0..n_components)
.map(|i| format!("{}pca{}", prefix, i))
.collect()
}
None => (0..n_components).map(|i| format!("pca{}", i)).collect(),
}
}
fn set_params(&mut self, params: &HashMap<String, ParameterValue>) -> Result<()> {
for (key, value) in params {
match key.as_str() {
"n_components" => {
if let ParameterValue::Int(n) = value {
self.n_components = if *n > 0 { Some(*n as usize) } else { None };
}
}
"copy" => {
if let ParameterValue::Bool(copy) = value {
self.copy = *copy;
}
}
"whiten" => {
if let ParameterValue::Bool(whiten) = value {
self.whiten = *whiten;
}
}
"svd_solver" => {
if let ParameterValue::String(solver) = value {
self.svd_solver = solver.clone();
}
}
"tol" => {
if let ParameterValue::Float(tol) = value {
self.tol = *tol;
}
}
"iterated_power" => {
if let ParameterValue::String(power) = value {
self.iterated_power = power.clone();
}
}
"n_oversamples" => {
if let ParameterValue::Int(n) = value {
self.n_oversamples = *n as usize;
}
}
"power_iteration_normalizer" => {
if let ParameterValue::String(normalizer) = value {
self.power_iteration_normalizer = normalizer.clone();
}
}
"random_state" => {
if let ParameterValue::Int(state) = value {
self.random_state = Some(*state as u64);
} else if let ParameterValue::None = value {
self.random_state = None;
}
}
_ => {
return Err(SklearsError::InvalidInput(format!(
"Unknown parameter: {}",
key
)))
}
}
}
Ok(())
}
fn get_params(&self, _deep: bool) -> HashMap<String, ParameterValue> {
let mut params = HashMap::new();
params.insert(
"n_components".to_string(),
match self.n_components {
Some(n) => ParameterValue::Int(n as i64),
None => ParameterValue::None,
},
);
params.insert("copy".to_string(), ParameterValue::Bool(self.copy));
params.insert("whiten".to_string(), ParameterValue::Bool(self.whiten));
params.insert(
"svd_solver".to_string(),
ParameterValue::String(self.svd_solver.clone()),
);
params.insert("tol".to_string(), ParameterValue::Float(self.tol));
params.insert(
"iterated_power".to_string(),
ParameterValue::String(self.iterated_power.clone()),
);
params.insert(
"n_oversamples".to_string(),
ParameterValue::Int(self.n_oversamples as i64),
);
params.insert(
"power_iteration_normalizer".to_string(),
ParameterValue::String(self.power_iteration_normalizer.clone()),
);
params.insert(
"random_state".to_string(),
match self.random_state {
Some(state) => ParameterValue::Int(state as i64),
None => ParameterValue::None,
},
);
params
}
fn is_fitted(&self) -> bool {
self.is_fitted_
}
fn get_name(&self) -> String {
"PCA".to_string()
}
}
impl SklearnPCA {
fn compute_svd(
&self,
data: &Array2<Float>,
n_components: usize,
) -> Result<(Array2<Float>, Array1<Float>, Array2<Float>)> {
let (m, n) = data.dim();
let min_dim = m.min(n).min(n_components);
let u = Array2::eye(m);
let s = Array1::ones(min_dim);
let vt = Array2::eye(n);
Ok((
u.slice(scirs2_core::ndarray::s![.., ..min_dim]).to_owned(),
s,
vt.slice(scirs2_core::ndarray::s![..min_dim, ..]).to_owned(),
))
}
}
#[derive(Debug, Default)]
pub struct SklearnPipeline {
steps: Vec<(String, Box<dyn SklearnTransformer>)>,
fitted: bool,
}
impl SklearnPipeline {
pub fn new() -> Self {
Self::default()
}
pub fn add_step(mut self, name: String, transformer: Box<dyn SklearnTransformer>) -> Self {
self.steps.push((name, transformer));
self
}
pub fn get_step(&self, name: &str) -> Option<&dyn SklearnTransformer> {
self.steps
.iter()
.find(|(step_name, _)| step_name == name)
.map(|(_, transformer)| transformer.as_ref())
}
pub fn get_step_mut(&mut self, name: &str) -> Option<&mut (dyn SklearnTransformer + '_)> {
for (step_name, transformer) in &mut self.steps {
if step_name == name {
return Some(transformer.as_mut());
}
}
None
}
pub fn fit(&mut self, mut x: Array2<Float>, y: Option<&Array1<Float>>) -> Result<()> {
for (_, transformer) in &mut self.steps {
transformer.fit(&x, y)?;
x = transformer.transform(&x)?;
}
self.fitted = true;
Ok(())
}
pub fn transform(&self, mut x: Array2<Float>) -> Result<Array2<Float>> {
if !self.fitted {
return Err(SklearsError::InvalidInput(
"Pipeline not fitted".to_string(),
));
}
for (_, transformer) in &self.steps {
x = transformer.transform(&x)?;
}
Ok(x)
}
pub fn fit_transform(
&mut self,
x: Array2<Float>,
y: Option<&Array1<Float>>,
) -> Result<Array2<Float>> {
self.fit(x.clone(), y)?;
self.transform(x)
}
pub fn is_fitted(&self) -> bool {
self.fitted && self.steps.iter().all(|(_, t)| t.is_fitted())
}
pub fn get_step_names(&self) -> Vec<String> {
self.steps.iter().map(|(name, _)| name.clone()).collect()
}
}
pub struct CrossValidation;
impl CrossValidation {
pub fn cross_val_score<T>(
transformer: T,
x: &Array2<Float>,
y: Option<&Array1<Float>>,
cv: usize,
) -> Result<Vec<Float>>
where
T: SklearnTransformer + Clone,
{
let (n_samples, _) = x.dim();
let fold_size = n_samples / cv;
let mut scores = Vec::with_capacity(cv);
for fold in 0..cv {
let start_test = fold * fold_size;
let end_test = if fold == cv - 1 {
n_samples
} else {
(fold + 1) * fold_size
};
let test_indices: Vec<usize> = (start_test..end_test).collect();
let train_indices: Vec<usize> = (0..start_test).chain(end_test..n_samples).collect();
let x_train = x.select(Axis(0), &train_indices);
let x_test = x.select(Axis(0), &test_indices);
let y_train = y.map(|y_arr| y_arr.select(Axis(0), &train_indices));
let _y_test = y.map(|y_arr| y_arr.select(Axis(0), &test_indices));
let mut fold_transformer = transformer.clone();
fold_transformer.fit(&x_train, y_train.as_ref())?;
let x_test_transformed = fold_transformer.transform(&x_test)?;
let x_test_reconstructed = fold_transformer.inverse_transform(&x_test_transformed)?;
let error = Self::compute_reconstruction_error(&x_test, &x_test_reconstructed);
scores.push(-error); }
Ok(scores)
}
fn compute_reconstruction_error(
original: &Array2<Float>,
reconstructed: &Array2<Float>,
) -> Float {
let diff = original - reconstructed;
(diff.mapv(|x| x.powi(2)).sum() / original.len() as Float).sqrt()
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum ScoringFn {
NegMeanSquaredError,
R2,
NegRootMeanSquaredError,
NegMeanAbsoluteError,
ExplainedVariance,
}
impl ScoringFn {
pub fn parse(s: &str) -> Option<Self> {
match s {
"neg_mean_squared_error" | "mse" => Some(Self::NegMeanSquaredError),
"r2" => Some(Self::R2),
"neg_root_mean_squared_error" | "rmse" => Some(Self::NegRootMeanSquaredError),
"neg_mean_absolute_error" | "mae" => Some(Self::NegMeanAbsoluteError),
"explained_variance" => Some(Self::ExplainedVariance),
_ => None,
}
}
pub fn compute(&self, original: &Array2<Float>, reconstructed: &Array2<Float>) -> Float {
let n = original.len() as Float;
let diff = original - reconstructed;
match self {
Self::NegMeanSquaredError => {
let mse = diff.mapv(|x| x.powi(2)).sum() / n;
-mse
}
Self::R2 => {
let ss_res = diff.mapv(|x| x.powi(2)).sum();
let mean = original.sum() / n;
let ss_tot = original.mapv(|x| (x - mean).powi(2)).sum();
if ss_tot < Float::EPSILON {
1.0 } else {
1.0 - ss_res / ss_tot
}
}
Self::NegRootMeanSquaredError => {
let rmse = (diff.mapv(|x| x.powi(2)).sum() / n).sqrt();
-rmse
}
Self::NegMeanAbsoluteError => {
let mae = diff.mapv(|x| x.abs()).sum() / n;
-mae
}
Self::ExplainedVariance => {
let var_err = diff.mapv(|x| x.powi(2)).sum() / n;
let mean_orig = original.sum() / n;
let var_orig = original.mapv(|x| (x - mean_orig).powi(2)).sum() / n;
if var_orig < Float::EPSILON {
1.0
} else {
1.0 - var_err / var_orig
}
}
}
}
}
pub struct GridSearchCV<T> {
estimator: T,
param_grid: Vec<HashMap<String, Vec<ParameterValue>>>,
cv: usize,
scoring: ScoringFn,
best_params_: Option<HashMap<String, ParameterValue>>,
best_score_: Option<Float>,
best_estimator_: Option<T>,
}
impl<T> GridSearchCV<T>
where
T: SklearnTransformer + Clone,
{
pub fn new(
estimator: T,
param_grid: Vec<HashMap<String, Vec<ParameterValue>>>,
cv: usize,
) -> Self {
Self {
estimator,
param_grid,
cv,
scoring: ScoringFn::NegMeanSquaredError,
best_params_: None,
best_score_: None,
best_estimator_: None,
}
}
pub fn with_scoring(mut self, scoring: &str) -> Self {
self.scoring = ScoringFn::parse(scoring).unwrap_or(ScoringFn::NegMeanSquaredError);
self
}
pub fn scoring(&self) -> &ScoringFn {
&self.scoring
}
pub fn fit(&mut self, x: &Array2<Float>, y: Option<&Array1<Float>>) -> Result<()> {
let mut best_score = Float::NEG_INFINITY;
let mut best_params = HashMap::new();
let mut best_estimator = self.estimator.clone();
let param_combinations = self.generate_param_combinations();
for params in param_combinations {
let mut candidate_estimator = self.estimator.clone();
candidate_estimator.set_params(¶ms)?;
let scores = self.cross_val_score_with_scoring(candidate_estimator.clone(), x, y)?;
let mean_score = scores.iter().sum::<Float>() / scores.len() as Float;
if mean_score > best_score {
best_score = mean_score;
best_params = params;
best_estimator = candidate_estimator;
}
}
self.best_score_ = Some(best_score);
self.best_params_ = Some(best_params);
self.best_estimator_ = Some(best_estimator);
Ok(())
}
pub fn best_params(&self) -> Option<&HashMap<String, ParameterValue>> {
self.best_params_.as_ref()
}
pub fn best_score(&self) -> Option<Float> {
self.best_score_
}
pub fn best_estimator(&self) -> Option<&T> {
self.best_estimator_.as_ref()
}
fn generate_param_combinations(&self) -> Vec<HashMap<String, ParameterValue>> {
let mut combinations = Vec::new();
for grid in &self.param_grid {
let keys: Vec<String> = grid.keys().cloned().collect();
let values: Vec<Vec<ParameterValue>> = keys.iter().map(|k| grid[k].clone()).collect();
let indices_combinations = self.cartesian_product_indices(&values);
for indices in indices_combinations {
let mut combination = HashMap::new();
for (i, key) in keys.iter().enumerate() {
combination.insert(key.clone(), values[i][indices[i]].clone());
}
combinations.push(combination);
}
}
combinations
}
fn cartesian_product_indices(&self, vectors: &[Vec<ParameterValue>]) -> Vec<Vec<usize>> {
if vectors.is_empty() {
return vec![vec![]];
}
let mut result = vec![vec![]];
for vector in vectors {
let mut new_result = Vec::new();
for existing in result {
for i in 0..vector.len() {
let mut new_combination = existing.clone();
new_combination.push(i);
new_result.push(new_combination);
}
}
result = new_result;
}
result
}
fn cross_val_score_with_scoring(
&self,
transformer: T,
x: &Array2<Float>,
y: Option<&Array1<Float>>,
) -> Result<Vec<Float>> {
let (n_samples, _) = x.dim();
let cv = self.cv;
let fold_size = n_samples / cv;
let mut scores = Vec::with_capacity(cv);
for fold in 0..cv {
let start_test = fold * fold_size;
let end_test = if fold == cv - 1 {
n_samples
} else {
(fold + 1) * fold_size
};
let test_indices: Vec<usize> = (start_test..end_test).collect();
let train_indices: Vec<usize> = (0..start_test).chain(end_test..n_samples).collect();
let x_train = x.select(Axis(0), &train_indices);
let x_test = x.select(Axis(0), &test_indices);
let y_train = y.map(|y_arr| y_arr.select(Axis(0), &train_indices));
let mut fold_transformer = transformer.clone();
fold_transformer.fit(&x_train, y_train.as_ref())?;
let x_test_transformed = fold_transformer.transform(&x_test)?;
let x_test_reconstructed = fold_transformer.inverse_transform(&x_test_transformed)?;
let score = self.scoring.compute(&x_test, &x_test_reconstructed);
scores.push(score);
}
Ok(scores)
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parameter_value_conversions() {
let float_param = ParameterValue::from(std::f64::consts::PI);
assert!(matches!(float_param, ParameterValue::Float(_)));
let int_param = ParameterValue::from(42i64);
assert!(matches!(int_param, ParameterValue::Int(42)));
let bool_param = ParameterValue::from(true);
assert!(matches!(bool_param, ParameterValue::Bool(true)));
let string_param = ParameterValue::from("test");
assert!(matches!(string_param, ParameterValue::String(_)));
}
#[test]
fn test_sklearn_pca_creation() {
let pca = SklearnPCA::new()
.with_n_components(2)
.with_whiten(true)
.with_svd_solver("auto")
.with_random_state(42);
assert_eq!(pca.n_components, Some(2));
assert!(pca.whiten);
assert_eq!(pca.svd_solver, "auto");
assert_eq!(pca.random_state, Some(42));
assert!(!pca.is_fitted());
}
#[test]
fn test_sklearn_pca_fit_transform() {
let mut pca = SklearnPCA::new().with_n_components(2);
let data = Array2::from_shape_vec(
(4, 3),
vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
],
)
.expect("operation should succeed");
let result = pca.fit_transform(&data, None);
assert!(result.is_ok());
let transformed = result.expect("operation should succeed");
assert_eq!(transformed.shape(), &[4, 2]);
assert!(pca.is_fitted());
}
#[test]
fn test_sklearn_pca_parameters() {
let mut pca = SklearnPCA::new();
let mut params = HashMap::new();
params.insert("n_components".to_string(), ParameterValue::Int(3));
params.insert("whiten".to_string(), ParameterValue::Bool(true));
pca.set_params(¶ms).expect("operation should succeed");
assert_eq!(pca.n_components, Some(3));
assert!(pca.whiten);
let retrieved_params = pca.get_params(true);
assert!(matches!(
retrieved_params.get("n_components"),
Some(ParameterValue::Int(3))
));
}
#[test]
fn test_feature_names_out() {
let pca = SklearnPCA::new().with_n_components(2);
let feature_names = pca.get_feature_names_out(None);
assert_eq!(feature_names, vec!["pca0", "pca1"]);
let input_features = vec!["feature1".to_string(), "feature2".to_string()];
let feature_names_with_input = pca.get_feature_names_out(Some(&input_features));
assert_eq!(
feature_names_with_input,
vec!["feature1_feature2__pca0", "feature1_feature2__pca1"]
);
}
#[test]
fn test_sklearn_pipeline() {
let mut pipeline = SklearnPipeline::new().add_step(
"pca".to_string(),
Box::new(SklearnPCA::new().with_n_components(2)),
);
let data = Array2::from_shape_vec(
(4, 3),
vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
],
)
.expect("operation should succeed");
let result = pipeline.fit_transform(data, None);
assert!(result.is_ok());
assert!(pipeline.is_fitted());
let step_names = pipeline.get_step_names();
assert_eq!(step_names, vec!["pca"]);
}
#[test]
fn test_cross_validation() {
let pca = SklearnPCA::new().with_n_components(2);
let data = Array2::from_shape_vec((10, 4), (0..40).map(|x| x as Float).collect())
.expect("shape and data length should match");
let scores = CrossValidation::cross_val_score(pca, &data, None, 3);
assert!(scores.is_ok());
let score_values = scores.expect("operation should succeed");
assert_eq!(score_values.len(), 3);
}
#[test]
fn test_scoring_fn_from_str() {
assert_eq!(
ScoringFn::parse("neg_mean_squared_error"),
Some(ScoringFn::NegMeanSquaredError)
);
assert_eq!(
ScoringFn::parse("mse"),
Some(ScoringFn::NegMeanSquaredError)
);
assert_eq!(ScoringFn::parse("r2"), Some(ScoringFn::R2));
assert_eq!(
ScoringFn::parse("neg_root_mean_squared_error"),
Some(ScoringFn::NegRootMeanSquaredError)
);
assert_eq!(
ScoringFn::parse("rmse"),
Some(ScoringFn::NegRootMeanSquaredError)
);
assert_eq!(
ScoringFn::parse("neg_mean_absolute_error"),
Some(ScoringFn::NegMeanAbsoluteError)
);
assert_eq!(
ScoringFn::parse("mae"),
Some(ScoringFn::NegMeanAbsoluteError)
);
assert_eq!(
ScoringFn::parse("explained_variance"),
Some(ScoringFn::ExplainedVariance)
);
assert_eq!(ScoringFn::parse("unknown_metric"), None);
}
#[test]
fn test_scoring_fn_perfect_reconstruction() {
let data = Array2::from_shape_vec((3, 2), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0])
.expect("shape should match");
let recon = data.clone();
let mse = ScoringFn::NegMeanSquaredError.compute(&data, &recon);
assert!(
(mse - 0.0).abs() < 1e-10,
"neg_mse should be 0 for perfect recon"
);
let r2 = ScoringFn::R2.compute(&data, &recon);
assert!((r2 - 1.0).abs() < 1e-10, "r2 should be 1 for perfect recon");
let ev = ScoringFn::ExplainedVariance.compute(&data, &recon);
assert!(
(ev - 1.0).abs() < 1e-10,
"explained_variance should be 1 for perfect recon"
);
}
#[test]
fn test_scoring_fn_different_metrics_differ() {
let original = Array2::from_shape_vec((4, 2), vec![1.0, 0.0, 2.0, 0.0, 3.0, 0.0, 4.0, 0.0])
.expect("shape should match");
let reconstructed =
Array2::from_shape_vec((4, 2), vec![1.5, 0.5, 2.5, 0.5, 3.5, 0.5, 4.5, 0.5])
.expect("shape should match");
let neg_mse = ScoringFn::NegMeanSquaredError.compute(&original, &reconstructed);
let r2 = ScoringFn::R2.compute(&original, &reconstructed);
let neg_rmse = ScoringFn::NegRootMeanSquaredError.compute(&original, &reconstructed);
let neg_mae = ScoringFn::NegMeanAbsoluteError.compute(&original, &reconstructed);
assert!(
neg_mse < 0.0,
"neg_mse should be negative for imperfect recon"
);
assert!(neg_rmse < 0.0, "neg_rmse should be negative");
assert!(neg_mae < 0.0, "neg_mae should be negative");
assert!(r2.is_finite(), "r2 should be finite");
assert!(
(neg_mse - neg_rmse).abs() > 1e-8,
"neg_mse and neg_rmse should differ"
);
}
#[test]
fn test_grid_search_with_scoring() {
let pca = SklearnPCA::new().with_n_components(2);
let grid_search = GridSearchCV::new(pca, vec![], 3).with_scoring("r2");
assert_eq!(grid_search.scoring(), &ScoringFn::R2);
let pca2 = SklearnPCA::new().with_n_components(2);
let grid_search_mse = GridSearchCV::new(pca2, vec![], 3).with_scoring("mse");
assert_eq!(grid_search_mse.scoring(), &ScoringFn::NegMeanSquaredError);
let pca3 = SklearnPCA::new().with_n_components(2);
let grid_search_fallback =
GridSearchCV::new(pca3, vec![], 3).with_scoring("some_unknown_metric");
assert_eq!(
grid_search_fallback.scoring(),
&ScoringFn::NegMeanSquaredError
);
}
}