use approx::assert_abs_diff_eq;
use ndarray::{Array1, Array2, array};
use rustyml::error::Error;
use rustyml::machine_learning::IsolationForest;
use rustyml::machine_learning::KMeans;
use rustyml::machine_learning::LinearRegression;
use rustyml::machine_learning::traits::{Fit, Predict};
use rustyml::types::{DistanceCalculationMetric, KernelType};
use crate::common::assert_allclose;
use rustyml::machine_learning::DBSCAN;
use rustyml::machine_learning::DistanceCalculationMetric as Metric;
use rustyml::machine_learning::LDA;
use rustyml::machine_learning::MeanShift;
use rustyml::machine_learning::SVC;
use rustyml::machine_learning::{Algorithm, DecisionTree};
use rustyml::machine_learning::{KNN, WeightingStrategy};
use rustyml::machine_learning::{LinearSVC, RegularizationType};
#[test]
fn error_dimension_mismatch_constructor_fields() {
let err = Error::dimension_mismatch(3, 2);
match err {
Error::DimensionMismatch { expected, found } => {
assert_eq!(expected, 3, "expected field should be 3");
assert_eq!(found, 2, "found field should be 2");
}
other => panic!("expected DimensionMismatch, got {other:?}"),
}
}
#[test]
fn error_dimension_mismatch_zero_expected() {
let err = Error::dimension_mismatch(0, 5);
assert!(
matches!(
err,
Error::DimensionMismatch {
expected: 0,
found: 5
}
),
"expected DimensionMismatch{{0,5}}, got {err:?}"
);
}
#[test]
fn error_invalid_parameter_constructor_fields() {
let err = Error::invalid_parameter("lr", "must be positive");
match err {
Error::InvalidParameter {
ref name,
ref reason,
} => {
assert_eq!(name, "lr");
assert!(
reason.contains("must be positive"),
"reason should contain 'must be positive', got: {reason}"
);
}
other => panic!("expected InvalidParameter, got {other:?}"),
}
}
#[test]
fn error_non_finite_constructor_carries_context() {
let err = Error::non_finite("weights");
match err {
Error::NonFinite(ref ctx) => {
assert!(
ctx.contains("weights"),
"context should mention 'weights', got: {ctx}"
);
}
other => panic!("expected NonFinite, got {other:?}"),
}
}
#[test]
fn error_empty_input_constructor_carries_what() {
let err = Error::empty_input("dataset");
match err {
Error::EmptyInput(ref what) => {
assert!(
what.contains("dataset"),
"payload should mention 'dataset', got: {what}"
);
}
other => panic!("expected EmptyInput, got {other:?}"),
}
}
#[test]
fn error_not_fitted_constructor_carries_model_name() {
let err = Error::not_fitted("KMeans");
match err {
Error::NotFitted(name) => {
assert_eq!(name, "KMeans", "model name should be 'KMeans'");
}
other => panic!("expected NotFitted, got {other:?}"),
}
}
#[test]
fn rustyml_result_alias_is_result_t_error() {
use rustyml::error::RustymlResult;
let r: RustymlResult<i32> = Ok(42);
assert!(matches!(r, Ok(42)));
let e: RustymlResult<i32> = Err(Error::empty_input("test"));
assert!(matches!(e, Err(Error::EmptyInput(_))));
}
#[test]
fn error_computation_constructor_no_source() {
let err = Error::computation("overflow");
match err {
Error::Computation {
ref context,
ref source,
} => {
assert!(
context.contains("overflow"),
"context should contain 'overflow'"
);
assert!(
source.is_none(),
"source should be None for Error::computation"
);
}
other => panic!("expected Computation, got {other:?}"),
}
}
#[test]
fn error_shape_mismatch_constructor_fields() {
let err = Error::shape_mismatch(vec![2usize, 3], vec![2usize, 4]);
match err {
Error::ShapeMismatch {
ref expected,
ref found,
} => {
assert_eq!(expected, &vec![2usize, 3]);
assert_eq!(found, &vec![2usize, 4]);
}
other => panic!("expected ShapeMismatch, got {other:?}"),
}
}
#[test]
fn error_invalid_input_constructor() {
let err = Error::invalid_input("unexpected rank");
match err {
Error::InvalidInput(ref msg) => {
assert!(
msg.contains("unexpected rank"),
"message should mention the supplied text"
);
}
other => panic!("expected InvalidInput, got {other:?}"),
}
}
#[test]
fn kernel_linear_orthogonal_is_zero() {
let k = KernelType::Linear;
let x1 = array![1.0_f64, 0.0];
let x2 = array![0.0_f64, 1.0];
assert_abs_diff_eq!(k.compute(x1.view(), x2.view()), 0.0, epsilon = 1e-12);
}
#[test]
fn kernel_linear_general_is_eleven() {
let k = KernelType::Linear;
let x1 = array![1.0_f64, 2.0];
let x2 = array![3.0_f64, 4.0];
assert_abs_diff_eq!(k.compute(x1.view(), x2.view()), 11.0, epsilon = 1e-12);
}
#[test]
fn kernel_rbf_same_vector_is_one() {
let k = KernelType::RBF { gamma: 1.0 };
let x = array![1.0_f64, 0.0];
assert_abs_diff_eq!(k.compute(x.view(), x.view()), 1.0, epsilon = 1e-12);
}
#[test]
fn kernel_rbf_orthonormal_pair_is_exp_minus_two() {
let k = KernelType::RBF { gamma: 1.0 };
let x1 = array![1.0_f64, 0.0];
let x2 = array![0.0_f64, 1.0];
let expected = (-2.0_f64).exp();
assert_abs_diff_eq!(k.compute(x1.view(), x2.view()), expected, epsilon = 1e-12);
}
#[test]
fn kernel_poly_degree2_orthogonal_is_zero() {
let k = KernelType::Poly {
degree: 2,
gamma: 1.0,
coef0: 0.0,
};
let x1 = array![1.0_f64, 0.0];
let x2 = array![0.0_f64, 1.0];
assert_abs_diff_eq!(k.compute(x1.view(), x2.view()), 0.0, epsilon = 1e-12);
}
#[test]
fn kernel_poly_degree2_with_coef0_one() {
let k = KernelType::Poly {
degree: 2,
gamma: 1.0,
coef0: 1.0,
};
let x = array![1.0_f64, 1.0];
assert_abs_diff_eq!(k.compute(x.view(), x.view()), 9.0, epsilon = 1e-12);
}
#[test]
fn kernel_sigmoid_unit_vector_is_tanh_one() {
let k = KernelType::Sigmoid {
gamma: 1.0,
coef0: 0.0,
};
let x = array![1.0_f64, 0.0];
let expected = 1.0_f64.tanh();
assert_abs_diff_eq!(k.compute(x.view(), x.view()), expected, epsilon = 1e-12);
}
#[test]
fn kernel_cosine_zero_vector_is_zero() {
let k = KernelType::Cosine;
let zero = array![0.0_f64, 0.0];
let other = array![1.0_f64, 2.0];
assert_abs_diff_eq!(k.compute(zero.view(), other.view()), 0.0, epsilon = 1e-12);
}
#[test]
fn kernel_cosine_identical_vector_is_one() {
let k = KernelType::Cosine;
let x = array![3.0_f64, 4.0]; assert_abs_diff_eq!(k.compute(x.view(), x.view()), 1.0, epsilon = 1e-12);
}
#[test]
fn distance_euclidean_3_4_triangle_is_5() {
let metric = DistanceCalculationMetric::Euclidean;
let a = array![0.0_f64, 0.0];
let b = array![3.0_f64, 4.0];
assert_abs_diff_eq!(metric.distance(a.view(), b.view()), 5.0, epsilon = 1e-12);
}
#[test]
fn distance_euclidean_is_symmetric() {
let metric = DistanceCalculationMetric::Euclidean;
let a = array![0.0_f64, 0.0];
let b = array![3.0_f64, 4.0];
assert_abs_diff_eq!(
metric.distance(a.view(), b.view()),
metric.distance(b.view(), a.view()),
epsilon = 1e-12
);
}
#[test]
fn distance_manhattan_3_4_is_7() {
let metric = DistanceCalculationMetric::Manhattan;
let a = array![0.0_f64, 0.0];
let b = array![3.0_f64, 4.0];
assert_abs_diff_eq!(metric.distance(a.view(), b.view()), 7.0, epsilon = 1e-12);
}
#[test]
fn distance_minkowski_p3_is_91_cbrt() {
let metric = DistanceCalculationMetric::Minkowski(3.0);
let a = array![0.0_f64, 0.0];
let b = array![3.0_f64, 4.0];
let expected = 91.0_f64.powf(1.0 / 3.0);
assert_abs_diff_eq!(
metric.distance(a.view(), b.view()),
expected,
epsilon = 1e-9
);
}
#[test]
fn distance_euclidean_self_is_zero() {
let metric = DistanceCalculationMetric::Euclidean;
let a = array![3.0_f64, 4.0];
assert_abs_diff_eq!(metric.distance(a.view(), a.view()), 0.0, epsilon = 1e-12);
}
#[test]
fn distance_manhattan_self_is_zero() {
let metric = DistanceCalculationMetric::Manhattan;
let a = array![3.0_f64, 4.0];
assert_abs_diff_eq!(metric.distance(a.view(), a.view()), 0.0, epsilon = 1e-12);
}
#[test]
fn distance_minkowski_p2_equals_euclidean() {
let euclidean = DistanceCalculationMetric::Euclidean;
let mink2 = DistanceCalculationMetric::Minkowski(2.0);
let a = array![0.0_f64, 0.0];
let b = array![3.0_f64, 4.0];
assert_abs_diff_eq!(mink2.distance(a.view(), b.view()), 5.0, epsilon = 1e-9);
assert_abs_diff_eq!(
mink2.distance(a.view(), b.view()),
euclidean.distance(a.view(), b.view()),
epsilon = 1e-9
);
}
fn train_and_predict_count<M>(
model: &mut M,
x_train: &Array2<f64>,
y_train: &Array1<f64>,
x_test: &Array2<f64>,
) -> usize
where
M: for<'a> Fit<(&'a Array2<f64>, &'a Array1<f64>)>
+ for<'a> Predict<&'a Array2<f64>, Output = Array1<f64>>,
{
Fit::fit(model, (x_train, y_train)).expect("fit through trait should succeed");
let preds = Predict::predict(model, x_test).expect("predict through trait should succeed");
preds.len()
}
#[test]
fn generic_fit_predict_with_linear_regression() {
let x_train = Array2::from_shape_vec((5, 1), vec![1.0, 2.0, 3.0, 4.0, 5.0]).unwrap();
let y_train = Array1::from_vec(vec![3.0, 5.0, 7.0, 9.0, 11.0]);
let x_test = Array2::from_shape_vec((3, 1), vec![6.0, 7.0, 8.0]).unwrap();
let mut model = LinearRegression::new(true, 0.01, 10_000, 1e-10).unwrap();
let n = train_and_predict_count(&mut model, &x_train, &y_train, &x_test);
assert_eq!(n, 3, "expected 3 predictions for 3 test points");
}
#[test]
fn generic_fit_trait_with_kmeans_unsupervised() {
let data = Array2::from_shape_vec(
(6, 2),
vec![
0.0, 0.0, 0.1, 0.0, 0.0, 0.1, 10.0, 0.0, 10.1, 0.0, 10.0, 0.1, ],
)
.unwrap();
let mut km = KMeans::new(2, 200, 1e-4).unwrap().with_random_state(42);
Fit::fit(&mut km, &data).expect("fit via Fit trait should succeed");
let labels = Predict::predict(&km, &data).expect("predict via Predict trait should succeed");
assert_eq!(labels.len(), 6, "should produce one label per sample");
}
#[test]
fn trait_predictions_match_inherent_method_predictions() {
let x_train = Array2::from_shape_vec((5, 1), vec![1.0, 2.0, 3.0, 4.0, 5.0]).unwrap();
let y_train = Array1::from_vec(vec![3.0, 5.0, 7.0, 9.0, 11.0]);
let x_test = Array2::from_shape_vec((1, 1), vec![6.0]).unwrap();
let mut model_trait = LinearRegression::new(true, 0.01, 10_000, 1e-10).unwrap();
Fit::fit(&mut model_trait, (&x_train, &y_train)).unwrap();
let preds_trait = Predict::predict(&model_trait, &x_test).unwrap();
let mut model_direct = LinearRegression::new(true, 0.01, 10_000, 1e-10).unwrap();
model_direct.fit(&x_train, &y_train).unwrap();
let preds_direct = model_direct.predict(&x_test).unwrap();
assert_abs_diff_eq!(preds_trait[0], preds_direct[0], epsilon = 0.0);
assert_abs_diff_eq!(preds_trait[0], 13.0, epsilon = 5e-3);
}
#[test]
fn linear_regression_save_load_round_trip_predictions_identical() {
let x_train = Array2::from_shape_vec((5, 1), vec![1.0, 2.0, 3.0, 4.0, 5.0]).unwrap();
let y_train = Array1::from_vec(vec![3.0, 5.0, 7.0, 9.0, 11.0]);
let x_test = Array2::from_shape_vec((2, 1), vec![6.0, 0.0]).unwrap();
let mut model = LinearRegression::new(true, 0.01, 10_000, 1e-10).unwrap();
model.fit(&x_train, &y_train).unwrap();
let preds_before = model.predict(&x_test).unwrap();
assert_abs_diff_eq!(preds_before[0], 13.0, epsilon = 5e-3);
assert_abs_diff_eq!(preds_before[1], 1.0, epsilon = 5e-3);
let path = "/tmp/rustyml_ml_infra_linreg_round_trip.json";
model
.save_to_path(path)
.expect("save_to_path should succeed");
let loaded = LinearRegression::load_from_path(path).expect("load_from_path should succeed");
let preds_after = loaded.predict(&x_test).unwrap();
assert_allclose(&preds_before, &preds_after, 0.0);
let _ = std::fs::remove_file(path);
}
#[test]
fn linear_regression_save_load_preserves_hyperparameters() {
let x_train = Array2::from_shape_vec((5, 1), vec![1.0, 2.0, 3.0, 4.0, 5.0]).unwrap();
let y_train = Array1::from_vec(vec![3.0, 5.0, 7.0, 9.0, 11.0]);
let mut model = LinearRegression::new(false, 0.005, 8_000, 1e-9).unwrap();
model.fit(&x_train, &y_train).unwrap();
let path = "/tmp/rustyml_ml_infra_linreg_hyperparams.json";
model.save_to_path(path).unwrap();
let loaded = LinearRegression::load_from_path(path).unwrap();
assert_eq!(loaded.get_fit_intercept(), model.get_fit_intercept());
assert_abs_diff_eq!(
loaded.get_learning_rate(),
model.get_learning_rate(),
epsilon = 1e-15
);
assert_eq!(loaded.get_max_iterations(), model.get_max_iterations());
let orig = model.get_coefficients().unwrap();
let load = loaded.get_coefficients().unwrap();
assert_allclose(orig, load, 0.0);
let _ = std::fs::remove_file(path);
}
fn three_blob_data_for_round_trip() -> Array2<f64> {
Array2::from_shape_vec(
(15, 2),
vec![
-0.05, 0.03, 0.04, -0.02, 0.01, 0.05, -0.03, -0.04, 0.02, 0.01,
99.95, 0.03, 100.04, -0.02, 100.01, 0.05, 99.97, -0.04, 100.02, 0.01,
49.95, 100.03, 50.04, 99.98, 50.01, 100.05, 49.97, 99.96, 50.02, 100.01,
],
)
.unwrap()
}
#[test]
fn kmeans_save_load_round_trip_predictions_identical() {
let data = three_blob_data_for_round_trip();
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km.fit(&data).unwrap();
let preds_before = km.predict(&data).unwrap();
let path = "/tmp/rustyml_ml_infra_kmeans_round_trip.json";
km.save_to_path(path).expect("save_to_path should succeed");
let km_loaded = KMeans::load_from_path(path).expect("load_from_path should succeed");
let preds_after = km_loaded.predict(&data).unwrap();
assert_eq!(preds_before.len(), preds_after.len());
for (o, l) in preds_before.iter().zip(preds_after.iter()) {
assert_eq!(o, l, "prediction mismatch after save/load round-trip");
}
let c_orig = km.get_centroids().unwrap();
let c_load = km_loaded.get_centroids().unwrap();
assert_allclose(c_orig, c_load, 0.0_f64);
let _ = std::fs::remove_file(path);
}
#[test]
fn kmeans_save_load_preserves_hyperparameters() {
let data = three_blob_data_for_round_trip();
let mut km = KMeans::new(3, 200, 1e-5).unwrap().with_random_state(7);
km.fit(&data).unwrap();
let path = "/tmp/rustyml_ml_infra_kmeans_hyperparams.json";
km.save_to_path(path).unwrap();
let loaded = KMeans::load_from_path(path).unwrap();
assert_eq!(loaded.get_n_clusters(), km.get_n_clusters());
assert_eq!(loaded.get_max_iterations(), km.get_max_iterations());
assert_abs_diff_eq!(loaded.get_tolerance(), km.get_tolerance(), epsilon = 1e-15);
assert_eq!(loaded.get_random_state(), km.get_random_state());
let _ = std::fs::remove_file(path);
}
#[test]
fn linear_regression_predict_before_fit_is_not_fitted() {
let model = LinearRegression::new(true, 0.01, 100, 1e-6).unwrap();
let x = array![[1.0, 2.0]];
let result = model.predict(&x);
assert!(
matches!(result, Err(Error::NotFitted(_))),
"expected NotFitted, got {result:?}"
);
}
#[test]
fn kmeans_predict_before_fit_is_not_fitted() {
let km = KMeans::new(3, 100, 1e-4).unwrap().with_random_state(42);
let x = array![[1.0, 2.0]];
let result = km.predict(&x);
assert!(
matches!(result, Err(Error::NotFitted(_))),
"expected NotFitted, got {result:?}"
);
}
#[test]
fn isolation_forest_predict_before_fit_is_not_fitted() {
let forest = IsolationForest::new(10, 32).unwrap().with_random_state(42);
let x = array![[1.0, 2.0]];
let result = forest.predict(&x);
assert!(
matches!(result, Err(Error::NotFitted(_))),
"expected NotFitted, got {result:?}"
);
}
#[test]
fn linear_regression_default_predict_before_fit_is_not_fitted() {
let model = LinearRegression::default();
let x = array![[1.0]];
let result = model.predict(&x);
assert!(
matches!(result, Err(Error::NotFitted(_))),
"expected NotFitted, got {result:?}"
);
}
#[test]
fn kmeans_default_predict_before_fit_is_not_fitted() {
let km = KMeans::default();
let x = array![[1.0, 2.0]];
let result = km.predict(&x);
assert!(
matches!(result, Err(Error::NotFitted(_))),
"expected NotFitted, got {result:?}"
);
}
#[test]
fn isolation_forest_default_predict_before_fit_is_not_fitted() {
let forest = IsolationForest::default();
let x = array![[1.0, 2.0]];
let result = forest.predict(&x);
assert!(
matches!(result, Err(Error::NotFitted(_))),
"expected NotFitted, got {result:?}"
);
}
#[test]
fn generic_fit_predict_isolation_forest_outputs_f64_scores() {
let data = Array2::from_shape_vec(
(5, 2),
vec![0.0, 0.0, 0.1, 0.0, 0.0, 0.1, 0.1, 0.1, 50.0, 50.0],
)
.unwrap();
let mut forest = IsolationForest::new(20, 32).unwrap().with_random_state(42);
Fit::fit(&mut forest, &data).expect("fit via Fit trait should succeed");
let scores: Array1<f64> =
Predict::predict(&forest, &data).expect("predict via Predict trait should succeed");
assert_eq!(scores.len(), 5, "one anomaly score per sample");
for (i, &s) in scores.iter().enumerate() {
assert!((0.0..=1.0).contains(&s), "score[{i}] = {s} not in [0,1]");
}
}
#[test]
fn generic_fit_predict_dbscan_outputs_isize_labels() {
let data = Array2::from_shape_vec(
(6, 2),
vec![
0.0, 0.0, 0.1, 0.0, 0.0, 0.1, 10.0, 10.0, 10.1, 10.0, 10.0, 10.1, ],
)
.unwrap();
let mut db = DBSCAN::new(0.5, 2).unwrap();
Fit::fit(&mut db, &data).expect("fit via Fit trait should succeed");
let labels: Array1<isize> =
Predict::predict(&db, &data).expect("predict via Predict trait should succeed");
assert_eq!(labels.len(), 6, "one label per sample");
for (i, &l) in labels.iter().enumerate() {
assert!(
l >= 0,
"label[{i}] = {l} should be a non-negative cluster id"
);
}
assert_ne!(labels[0], labels[3], "the two separated blobs must differ");
}
#[test]
fn generic_fit_predict_knn_outputs_generic_labels() {
let x_train = Array2::from_shape_vec((2, 2), vec![0.0, 0.0, 10.0, 0.0]).unwrap();
let y_train = Array1::from_vec(vec![0_i32, 1]);
let mut knn = KNN::<i32>::new(1)
.unwrap()
.with_weighting_strategy(WeightingStrategy::Uniform)
.with_metric(Metric::Euclidean)
.unwrap();
Fit::fit(&mut knn, (&x_train, &y_train)).expect("fit via Fit trait should succeed");
let x_test = Array2::from_shape_vec((2, 2), vec![0.5, 0.0, 9.5, 0.0]).unwrap();
let preds: Array1<i32> =
Predict::predict(&knn, &x_test).expect("predict via Predict trait should succeed");
assert_eq!(preds.len(), 2, "one label per test point");
assert_eq!(preds[0], 0, "point near anchor 0 must get label 0");
assert_eq!(preds[1], 1, "point near anchor 1 must get label 1");
}
#[test]
fn generic_fit_predict_lda_outputs_i32_labels() {
let x_train = Array2::from_shape_vec((6, 1), vec![1.0, 2.0, 3.0, 7.0, 8.0, 9.0]).unwrap();
let y_train = Array1::from_vec(vec![0_i32, 0, 0, 1, 1, 1]);
let mut lda = LDA::new(1).unwrap();
Fit::fit(&mut lda, (&x_train, &y_train)).expect("fit via Fit trait should succeed");
let preds: Array1<i32> =
Predict::predict(&lda, &x_train).expect("predict via Predict trait should succeed");
assert_eq!(preds.len(), 6, "one label per sample");
for (i, (&p, &t)) in preds.iter().zip(y_train.iter()).enumerate() {
assert_eq!(
p, t,
"sample {i}: LDA via trait predicted {p}, expected {t}"
);
}
}
#[test]
fn generic_fit_predict_decision_tree_outputs_f64_labels() {
let x_train = Array2::from_shape_vec((6, 1), vec![0.0, 0.1, 0.2, 1.0, 1.1, 1.2]).unwrap();
let y_train = Array1::from_vec(vec![0.0, 0.0, 0.0, 1.0, 1.0, 1.0]);
let mut tree = DecisionTree::new(Algorithm::CART, true).unwrap();
Fit::fit(&mut tree, (&x_train, &y_train)).expect("fit via Fit trait should succeed");
let preds: Array1<f64> =
Predict::predict(&tree, &x_train).expect("predict via Predict trait should succeed");
assert_eq!(preds.len(), 6, "one label per sample");
for (i, (&p, &t)) in preds.iter().zip(y_train.iter()).enumerate() {
assert_abs_diff_eq!(p, t, epsilon = 1e-9);
let _ = i;
}
}
#[test]
fn generic_fit_predict_linear_svc_outputs_f64_labels() {
let x_train = Array2::from_shape_vec(
(8, 2),
vec![
-5.0, 0.0, -6.0, 0.0, -7.0, 0.0, -4.0, 0.0, 5.0, 0.0, 6.0, 0.0, 7.0, 0.0, 4.0, 0.0, ],
)
.unwrap();
let y_train = Array1::from_vec(vec![0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0]);
let mut svc = LinearSVC::new(10_000, 0.01, RegularizationType::L2(0.01), true, 1e-6).unwrap();
Fit::fit(&mut svc, (&x_train, &y_train)).expect("fit via Fit trait should succeed");
let preds: Array1<f64> =
Predict::predict(&svc, &x_train).expect("predict via Predict trait should succeed");
assert_eq!(preds.len(), 8, "one label per sample");
let expected = [0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0];
for (i, (&p, &t)) in preds.iter().zip(expected.iter()).enumerate() {
assert_eq!(
p, t,
"sample {i}: LinearSVC via trait predicted {p}, expected {t}"
);
}
}
#[test]
fn generic_fit_predict_mean_shift_outputs_usize_labels() {
let data = Array2::from_shape_vec(
(6, 2),
vec![
-0.1, 0.0, 0.1, 0.0, 0.0, 0.0, 19.9, 20.0, 20.1, 20.0, 20.0, 20.0, ],
)
.unwrap();
let mut ms = MeanShift::new(2.0)
.unwrap()
.with_max_iter(300)
.unwrap()
.with_tolerance(1e-5)
.unwrap()
.with_bin_seeding(true)
.with_cluster_all(true);
Fit::fit(&mut ms, &data).expect("fit via Fit trait should succeed");
let labels: Array1<usize> =
Predict::predict(&ms, &data).expect("predict via Predict trait should succeed");
assert_eq!(labels.len(), 6, "one label per sample");
assert_eq!(labels[0], labels[1], "blob A samples must share a cluster");
assert_eq!(labels[3], labels[4], "blob B samples must share a cluster");
assert_ne!(
labels[0], labels[3],
"the two far-apart blobs must be different clusters"
);
}
#[test]
fn generic_fit_predict_svc_outputs_pm1_labels() {
let x_train = Array2::from_shape_vec(
(8, 2),
vec![
2.0, 2.0, 3.0, 2.0, 2.0, 3.0, 3.0, 3.0, -2.0, -2.0, -3.0, -2.0, -2.0, -3.0, -3.0, -3.0, ],
)
.unwrap();
let y_train = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, -1.0, -1.0, -1.0, -1.0]);
let mut svc = SVC::new(KernelType::Linear, 10.0, 1e-3, 1000)
.unwrap()
.with_random_state(42);
Fit::fit(&mut svc, (&x_train, &y_train)).expect("fit via Fit trait should succeed");
let preds: Array1<f64> =
Predict::predict(&svc, &x_train).expect("predict via Predict trait should succeed");
assert_eq!(preds.len(), 8, "one label per sample");
for (i, (&p, &t)) in preds.iter().zip(y_train.iter()).enumerate() {
assert_eq!(
p, t,
"sample {i}: SVC via trait predicted {p}, expected {t}"
);
}
}