use scirs2_core::ndarray::{array, Array2};
use sklears::metrics::classification::{accuracy_score, f1_score, precision_score, recall_score};
use sklears::neighbors::KNeighborsClassifier;
use sklears::prelude::*;
#[cfg(feature = "preprocessing")]
use sklears::preprocessing::{MinMaxScaler, StandardScaler};
use sklears::utils::data_generation::make_classification;
#[cfg(feature = "preprocessing")]
#[test]
#[allow(non_snake_case)]
fn test_end_to_end_classification_pipeline() {
let (X, y) = make_classification(100, 4, 3, None, None, 0.0, 1.0, Some(42))
.expect("operation should succeed");
let scaler = StandardScaler::new()
.fit(&X, &())
.expect("StandardScaler fit should succeed");
let X_scaled = scaler
.transform(&X)
.expect("StandardScaler transform should succeed");
let classifier = KNeighborsClassifier::new(3);
let fitted_classifier = classifier
.fit(&X_scaled, &y)
.expect("model fitting should succeed");
let predictions = fitted_classifier
.predict(&X_scaled)
.expect("prediction should succeed");
let accuracy = accuracy_score(&y, &predictions).expect("operation should succeed");
assert!(accuracy > 0.7, "Accuracy should be > 0.7, got {}", accuracy);
assert_eq!(predictions.len(), y.len());
for &pred in predictions.iter() {
assert!(
(0..=2).contains(&pred),
"Predicted class {} should be in [0, 2]",
pred
);
}
}
#[test]
#[allow(non_snake_case)]
fn test_clustering_and_classification() {
let (X, y_true) = make_classification(60, 2, 2, None, None, 0.0, 1.0, Some(42))
.expect("operation should succeed");
let classifier = KNeighborsClassifier::new(5);
let fitted_classifier = classifier
.fit(&X, &y_true)
.expect("model fitting should succeed");
let predictions = fitted_classifier
.predict(&X)
.expect("prediction should succeed");
let accuracy = accuracy_score(&y_true, &predictions).expect("operation should succeed");
assert!(
accuracy > 0.6,
"Accuracy on blob data should be reasonable, got {}",
accuracy
);
}
#[test]
fn test_metrics_consistency() {
let y_true = array![0, 1, 1, 0, 1, 0, 1, 1, 0, 0];
let y_pred = array![0, 1, 0, 0, 1, 1, 1, 1, 0, 1];
let accuracy = accuracy_score(&y_true, &y_pred).expect("operation should succeed");
let precision = precision_score(&y_true, &y_pred, Some(1)).expect("operation should succeed");
let recall = recall_score(&y_true, &y_pred, Some(1)).expect("operation should succeed");
let f1 = f1_score(&y_true, &y_pred, Some(1)).expect("operation should succeed");
assert!((0.0..=1.0).contains(&accuracy));
assert!((0.0..=1.0).contains(&precision));
assert!((0.0..=1.0).contains(&recall));
assert!((0.0..=1.0).contains(&f1));
let expected_f1 = 2.0 * precision * recall / (precision + recall);
assert!((f1 - expected_f1).abs() < 1e-10);
}
#[cfg(feature = "preprocessing")]
#[test]
#[allow(non_snake_case)]
fn test_preprocessing_pipeline() {
let X = Array2::from_shape_vec(
(6, 3),
vec![
1.0, 10.0, -1.0, 2.0, 20.0, -2.0, 3.0, 30.0, -3.0, 4.0, 40.0, -4.0, 5.0, 50.0, -5.0,
6.0, 60.0, -6.0,
],
)
.expect("shape and data length should match");
let std_scaler = StandardScaler::new()
.fit(&X, &())
.expect("StandardScaler fit should succeed");
let X_std = std_scaler
.transform(&X)
.expect("StandardScaler transform should succeed");
assert_eq!(X_std.dim(), X.dim());
for j in 0..X_std.ncols() {
let col = X_std.column(j);
let n = col.len() as f64;
let mean: f64 = col.iter().sum::<f64>() / n;
let variance: f64 = col.iter().map(|&v| (v - mean).powi(2)).sum::<f64>() / n;
assert!(
mean.abs() < 1e-8,
"column {j} mean should be ~0, got {mean}"
);
assert!(
(variance - 1.0).abs() < 1e-8,
"column {j} variance should be ~1, got {variance}"
);
}
let minmax_scaler = MinMaxScaler::new()
.fit(&X, &())
.expect("MinMaxScaler fit should succeed");
let X_minmax = minmax_scaler
.transform(&X)
.expect("MinMaxScaler transform should succeed");
assert_eq!(X_minmax.dim(), X.dim());
let eps = 1e-8;
for &value in X_minmax.iter() {
assert!(
value >= -eps && value <= 1.0 + eps,
"MinMaxScaler output {value} should be within [0, 1]"
);
}
let y = array![0, 0, 0, 1, 1, 1];
let classifier = KNeighborsClassifier::new(3);
let fitted = classifier
.fit(&X_std, &y)
.expect("model fitting should succeed");
let predictions = fitted.predict(&X_std).expect("prediction should succeed");
assert_eq!(predictions.len(), y.len());
}
#[test]
#[allow(non_snake_case)]
fn test_data_generation_consistency() {
let (X1, y1) = make_classification(50, 3, 2, None, None, 0.0, 1.0, Some(42))
.expect("operation should succeed");
let (X2, y2) = make_classification(50, 3, 2, None, None, 0.0, 1.0, Some(42))
.expect("operation should succeed");
assert_eq!(X1, X2);
assert_eq!(y1, y2);
assert_eq!(X1.shape(), &[50, 3]);
assert_eq!(y1.len(), 50);
let mut classes: Vec<i32> = y1.iter().copied().collect();
classes.sort_unstable();
classes.dedup();
assert_eq!(classes.len(), 2);
}
#[test]
#[allow(non_snake_case)]
fn test_cross_crate_type_compatibility() {
let labels = array![0, 1, 2, 1, 0, 2];
let classifier = KNeighborsClassifier::new(3);
let X_test = Array2::from_shape_vec(
(6, 2),
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 fitted = classifier
.fit(&X_test, &labels)
.expect("model fitting should succeed");
let predictions = fitted.predict(&X_test).expect("prediction should succeed");
assert_eq!(predictions.len(), labels.len());
}