use super::*;
use scirs2_core::ndarray::{array, Array2};
use sklears_core::traits::{Fit, Predict};
#[test]
fn test_multi_output_classifier() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[0, 1], [1, 0], [1, 1], [0, 0]];
let moc = MultiOutputClassifier::new();
let fitted = moc
.fit(&X.view(), &y)
.expect("model fitting should succeed");
assert_eq!(fitted.n_targets(), 2);
assert_eq!(fitted.classes().len(), 2);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (4, 2));
for target_idx in 0..2 {
let target_classes = &fitted.classes()[target_idx];
for sample_idx in 0..4 {
let pred = predictions[[sample_idx, target_idx]];
assert!(target_classes.contains(&pred));
}
}
}
#[test]
fn test_multi_output_regressor() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[1.5, 2.5], [2.5, 3.5], [2.0, 1.5], [1.0, 1.5]];
let mor = MultiOutputRegressor::new();
let fitted = mor
.fit(&X.view(), &y)
.expect("model fitting should succeed");
assert_eq!(fitted.n_targets(), 2);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (4, 2));
for pred in predictions.iter() {
assert!(pred.is_finite());
}
}
#[test]
fn test_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[0, 1], [1, 0], [0, 1]];
let moc = MultiOutputClassifier::new();
assert!(moc.fit(&X.view(), &y).is_err());
}
#[test]
fn test_empty_targets() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = Array2::<i32>::zeros((2, 0));
let moc = MultiOutputClassifier::new();
assert!(moc.fit(&X.view(), &y).is_err());
}
#[test]
fn test_prediction_shape_mismatch() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[0, 1], [1, 0]];
let moc = MultiOutputClassifier::new();
let fitted = moc
.fit(&X.view(), &y)
.expect("model fitting should succeed");
let X_wrong = array![[1.0, 2.0, 3.0]]; assert!(fitted.predict(&X_wrong.view()).is_err());
}
#[test]
fn test_classifier_chain() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[0, 1], [1, 0], [1, 1], [0, 0]];
let cc = ClassifierChain::new();
let fitted = cc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
assert_eq!(fitted.n_targets(), 2);
assert_eq!(fitted.chain_order(), &[0, 1]);
let predictions = fitted
.predict_simple(&X.view())
.expect("operation should succeed");
assert_eq!(predictions.dim(), (4, 2));
for sample_idx in 0..4 {
for target_idx in 0..2 {
let pred = predictions[[sample_idx, target_idx]];
assert!(pred == 0 || pred == 1);
}
}
}
#[test]
fn test_classifier_chain_custom_order() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[0, 1], [1, 0], [1, 1]];
let cc = ClassifierChain::new().order(vec![1, 0]); let fitted = cc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
assert_eq!(fitted.chain_order(), &[1, 0]);
let predictions = fitted
.predict_simple(&X.view())
.expect("operation should succeed");
assert_eq!(predictions.dim(), (3, 2));
}
#[test]
fn test_classifier_chain_invalid_order() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[0, 1], [1, 0]];
let cc = ClassifierChain::new().order(vec![0, 1, 2]); assert!(cc.fit_simple(&X.view(), &y).is_err());
}
#[test]
fn test_classifier_chain_monte_carlo() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[0, 1], [1, 0], [1, 1], [0, 0]];
let cc = ClassifierChain::new();
let fitted = cc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
let mc_probs = fitted
.predict_monte_carlo(&X.view(), 100, Some(42))
.expect("operation should succeed");
assert_eq!(mc_probs.dim(), (4, 2));
for prob in mc_probs.iter() {
assert!(*prob >= 0.0 && *prob <= 1.0);
}
let mc_labels = fitted
.predict_monte_carlo_labels(&X.view(), 100, Some(42))
.expect("operation should succeed");
assert_eq!(mc_labels.dim(), (4, 2));
for pred in mc_labels.iter() {
assert!(*pred == 0 || *pred == 1);
}
let mc_probs2 = fitted
.predict_monte_carlo(&X.view(), 100, Some(42))
.expect("operation should succeed");
for (i, (&prob1, &prob2)) in mc_probs.iter().zip(mc_probs2.iter()).enumerate() {
assert!(
(prob1 - prob2).abs() < 1e-10,
"Probabilities should be identical with same random state at index {}",
i
);
}
}
#[test]
fn test_classifier_chain_monte_carlo_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[0, 1], [1, 0]];
let cc = ClassifierChain::new();
let fitted = cc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
assert!(fitted.predict_monte_carlo(&X.view(), 0, None).is_err());
let X_wrong = array![[1.0, 2.0, 3.0]]; assert!(fitted
.predict_monte_carlo(&X_wrong.view(), 10, None)
.is_err());
}
#[test]
fn test_regressor_chain() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[1.5, 2.5], [2.5, 3.5], [2.0, 1.5], [1.0, 1.5]];
let rc = RegressorChain::new();
let fitted = rc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
assert_eq!(fitted.n_targets(), 2);
assert_eq!(fitted.chain_order(), &[0, 1]);
let predictions = fitted
.predict_simple(&X.view())
.expect("operation should succeed");
assert_eq!(predictions.dim(), (4, 2));
for pred in predictions.iter() {
assert!(pred.is_finite());
}
}
#[test]
fn test_regressor_chain_custom_order() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[1.5, 2.5], [2.5, 3.5], [2.0, 1.5]];
let rc = RegressorChain::new().order(vec![1, 0]); let fitted = rc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
assert_eq!(fitted.chain_order(), &[1, 0]);
let predictions = fitted
.predict_simple(&X.view())
.expect("operation should succeed");
assert_eq!(predictions.dim(), (3, 2));
for pred in predictions.iter() {
assert!(pred.is_finite());
}
}
#[test]
fn test_regressor_chain_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1.5, 2.5], [2.5, 3.5], [2.0, 1.5]];
let rc = RegressorChain::new();
assert!(rc.fit_simple(&X.view(), &y).is_err());
}
#[test]
fn test_regressor_chain_invalid_order() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1.5, 2.5], [2.5, 3.5]];
let rc = RegressorChain::new().order(vec![0, 1, 2]); assert!(rc.fit_simple(&X.view(), &y).is_err());
}
#[test]
fn test_binary_relevance() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[1, 0], [0, 1], [1, 1], [0, 0]];
let br = BinaryRelevance::new();
let fitted = br.fit(&X.view(), &y).expect("model fitting should succeed");
assert_eq!(fitted.n_labels(), 2);
assert_eq!(fitted.classes().len(), 2);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (4, 2));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
let probabilities = fitted
.predict_proba(&X.view())
.expect("operation should succeed");
assert_eq!(probabilities.dim(), (4, 2));
for prob in probabilities.iter() {
assert!(*prob >= 0.0 && *prob <= 1.0);
}
}
#[test]
fn test_binary_relevance_single_label() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[1], [0], [1]];
let br = BinaryRelevance::new();
let fitted = br.fit(&X.view(), &y).expect("model fitting should succeed");
assert_eq!(fitted.n_labels(), 1);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (3, 1));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
}
#[test]
fn test_binary_relevance_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [0, 1], [1, 1]];
let br = BinaryRelevance::new();
assert!(br.fit(&X.view(), &y).is_err());
}
#[test]
fn test_binary_relevance_non_binary_labels() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[0, 1], [1, 2], [2, 0]];
let br = BinaryRelevance::new();
assert!(br.fit(&X.view(), &y).is_err());
}
#[test]
fn test_binary_relevance_predict_shape_mismatch() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [0, 1]];
let br = BinaryRelevance::new();
let fitted = br.fit(&X.view(), &y).expect("model fitting should succeed");
let X_wrong = array![[1.0, 2.0, 3.0]]; assert!(fitted.predict(&X_wrong.view()).is_err());
assert!(fitted.predict_proba(&X_wrong.view()).is_err());
}
#[test]
fn test_label_powerset() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[1, 0], [0, 1], [1, 1], [0, 0]];
let lp = LabelPowerset::new();
let fitted = lp.fit(&X.view(), &y).expect("model fitting should succeed");
assert_eq!(fitted.n_labels(), 2);
assert_eq!(fitted.n_classes(), 4);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (4, 2));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
let scores = fitted
.decision_function(&X.view())
.expect("operation should succeed");
assert_eq!(scores.dim(), (4, 4));
for score in scores.iter() {
assert!(score.is_finite());
}
}
#[test]
fn test_label_powerset_simple_case() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[1, 0], [0, 1], [1, 0]];
let lp = LabelPowerset::new();
let fitted = lp.fit(&X.view(), &y).expect("model fitting should succeed");
assert_eq!(fitted.n_labels(), 2);
assert_eq!(fitted.n_classes(), 2);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (3, 2));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
}
#[test]
fn test_label_powerset_single_label() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[1], [0], [1]];
let lp = LabelPowerset::new();
let fitted = lp.fit(&X.view(), &y).expect("model fitting should succeed");
assert_eq!(fitted.n_labels(), 1);
assert_eq!(fitted.n_classes(), 2);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (3, 1));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
}
#[test]
fn test_label_powerset_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [0, 1], [1, 1]];
let lp = LabelPowerset::new();
assert!(lp.fit(&X.view(), &y).is_err());
}
#[test]
fn test_label_powerset_non_binary_labels() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[0, 1], [1, 2], [2, 0]];
let lp = LabelPowerset::new();
assert!(lp.fit(&X.view(), &y).is_err());
}
#[test]
fn test_label_powerset_predict_shape_mismatch() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [0, 1]];
let lp = LabelPowerset::new();
let fitted = lp.fit(&X.view(), &y).expect("model fitting should succeed");
let X_wrong = array![[1.0, 2.0, 3.0]]; assert!(fitted.predict(&X_wrong.view()).is_err());
assert!(fitted.decision_function(&X_wrong.view()).is_err());
}
#[test]
fn test_label_powerset_all_same_combination() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0]];
let y = array![[1, 0], [1, 0], [1, 0]];
let lp = LabelPowerset::new();
let fitted = lp.fit(&X.view(), &y).expect("model fitting should succeed");
assert_eq!(fitted.n_labels(), 2);
assert_eq!(fitted.n_classes(), 1);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (3, 2));
for sample_idx in 0..3 {
assert_eq!(predictions[[sample_idx, 0]], 1);
assert_eq!(predictions[[sample_idx, 1]], 0);
}
}
#[test]
fn test_pruned_label_powerset_default_strategy() {
let X = array![
[1.0, 2.0],
[2.0, 3.0],
[3.0, 1.0],
[1.0, 1.0],
[2.0, 2.0],
[3.0, 3.0],
[1.5, 2.5],
[2.5, 1.5]
];
let y = array![
[1, 0],
[0, 1],
[1, 1],
[0, 0], [1, 0],
[0, 1],
[1, 0],
[0, 1], ];
let plp = PrunedLabelPowerset::new()
.min_frequency(2)
.strategy(PruningStrategy::DefaultMapping(vec![0, 0]));
let fitted = plp
.fit(&X.view(), &y)
.expect("model fitting should succeed");
assert!(fitted.n_frequent_classes() <= 4); assert_eq!(fitted.min_frequency(), 2);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (8, 2));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
}
#[test]
fn test_pruned_label_powerset_similarity_strategy() {
let X = array![
[1.0, 2.0],
[2.0, 3.0],
[3.0, 1.0],
[1.0, 1.0],
[2.0, 2.0],
[3.0, 3.0]
];
let y = array![
[1, 0],
[1, 0],
[1, 0], [0, 1],
[0, 1], [1, 1] ];
let plp = PrunedLabelPowerset::new()
.min_frequency(2)
.strategy(PruningStrategy::SimilarityMapping);
let fitted = plp
.fit(&X.view(), &y)
.expect("model fitting should succeed");
assert_eq!(fitted.n_frequent_classes(), 2);
let mapping = fitted.combination_mapping();
let rare_combo = vec![1, 1];
assert!(mapping.contains_key(&rare_combo));
let mapped = mapping.get(&rare_combo).expect("index should be valid");
assert!(mapped == &vec![1, 0] || mapped == &vec![0, 1]);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (6, 2));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
}
#[test]
fn test_pruned_label_powerset_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[0, 1], [1, 0]];
let plp = PrunedLabelPowerset::new().min_frequency(5); assert!(plp.fit(&X.view(), &y).is_err());
let plp = PrunedLabelPowerset::new().strategy(PruningStrategy::DefaultMapping(vec![0, 1, 0])); assert!(plp.fit(&X.view(), &y).is_err());
let y_bad = array![[2, 1], [1, 0]]; let plp = PrunedLabelPowerset::new();
assert!(plp.fit(&X.view(), &y_bad).is_err());
}
#[test]
fn test_pruned_label_powerset_edge_cases() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [1, 0]];
let plp = PrunedLabelPowerset::new().min_frequency(2);
let fitted = plp
.fit(&X.view(), &y)
.expect("model fitting should succeed");
assert!(fitted.n_frequent_classes() >= 1);
assert!(!fitted.frequent_combinations().is_empty());
assert!(fitted.frequent_combinations().contains(&vec![1, 0]));
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (2, 2));
for sample_idx in 0..2 {
assert_eq!(predictions[[sample_idx, 0]], 1);
assert_eq!(predictions[[sample_idx, 1]], 0);
}
}
#[test]
fn test_metrics_hamming_loss() {
let y_true = array![[1, 0, 1], [0, 1, 0], [1, 1, 1]];
let y_pred = array![[1, 0, 0], [0, 1, 1], [1, 0, 1]];
let loss =
metrics::hamming_loss(&y_true.view(), &y_pred.view()).expect("operation should succeed");
assert!((loss - 3.0 / 9.0).abs() < 1e-10);
}
#[test]
fn test_metrics_subset_accuracy() {
let y_true = array![[1, 0, 1], [0, 1, 0], [1, 1, 1]];
let y_pred = array![[1, 0, 1], [0, 1, 1], [1, 0, 1]];
let accuracy =
metrics::subset_accuracy(&y_true.view(), &y_pred.view()).expect("operation should succeed");
assert!((accuracy - 1.0 / 3.0).abs() < 1e-10);
}
#[test]
fn test_metrics_jaccard_score() {
let y_true = array![[1, 0, 1], [0, 1, 0]];
let y_pred = array![[1, 0, 0], [0, 1, 1]];
let score =
metrics::jaccard_score(&y_true.view(), &y_pred.view()).expect("operation should succeed");
assert!((score - 0.5).abs() < 1e-10);
}
#[test]
fn test_metrics_f1_score_micro() {
let y_true = array![[1, 0, 1], [0, 1, 0], [1, 1, 1]];
let y_pred = array![[1, 0, 0], [0, 1, 1], [1, 0, 1]];
let f1 = metrics::f1_score(&y_true.view(), &y_pred.view(), "micro")
.expect("operation should succeed");
assert!((f1 - 0.7272727272727273).abs() < 1e-10);
}
#[test]
fn test_metrics_f1_score_macro() {
let y_true = array![[1, 0], [0, 1], [1, 1]];
let y_pred = array![[1, 0], [0, 1], [1, 0]];
let f1 = metrics::f1_score(&y_true.view(), &y_pred.view(), "macro")
.expect("operation should succeed");
assert!((f1 - 0.8333333333333334).abs() < 1e-10);
}
#[test]
fn test_metrics_f1_score_samples() {
let y_true = array![[1, 0], [0, 1], [1, 1]];
let y_pred = array![[1, 0], [0, 1], [1, 0]];
let f1 = metrics::f1_score(&y_true.view(), &y_pred.view(), "samples")
.expect("sampling should succeed");
assert!((f1 - 0.8888888888888888).abs() < 1e-10);
}
#[test]
fn test_metrics_coverage_error() {
let y_true = array![[1, 0, 1], [0, 1, 0]];
let y_scores = array![[0.9, 0.1, 0.8], [0.2, 0.9, 0.3]];
let coverage = metrics::coverage_error(&y_true.view(), &y_scores.view())
.expect("operation should succeed");
assert!((coverage - 1.5).abs() < 1e-10);
}
#[test]
fn test_metrics_label_ranking_average_precision() {
let y_true = array![[1, 0, 1], [0, 1, 0]];
let y_scores = array![[0.9, 0.1, 0.8], [0.2, 0.9, 0.3]];
let lrap = metrics::label_ranking_average_precision(&y_true.view(), &y_scores.view())
.expect("operation should succeed");
assert!((lrap - 1.0).abs() < 1e-10);
}
#[test]
fn test_metrics_invalid_shapes() {
let y_true = array![[1, 0], [0, 1]];
let y_pred = array![[1, 0, 1]];
assert!(metrics::hamming_loss(&y_true.view(), &y_pred.view()).is_err());
assert!(metrics::subset_accuracy(&y_true.view(), &y_pred.view()).is_err());
assert!(metrics::jaccard_score(&y_true.view(), &y_pred.view()).is_err());
assert!(metrics::f1_score(&y_true.view(), &y_pred.view(), "micro").is_err());
}
#[test]
fn test_metrics_invalid_f1_average() {
let y_true = array![[1, 0], [0, 1]];
let y_pred = array![[1, 0], [0, 1]];
assert!(metrics::f1_score(&y_true.view(), &y_pred.view(), "invalid").is_err());
}
#[test]
fn test_ensemble_of_chains() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[0, 1], [1, 0], [1, 1], [0, 0]];
let eoc = EnsembleOfChains::new().n_chains(3).random_state(42);
let fitted = eoc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
assert_eq!(fitted.n_chains(), 3);
assert_eq!(fitted.n_targets(), 2);
let predictions = fitted
.predict_simple(&X.view())
.expect("operation should succeed");
assert_eq!(predictions.dim(), (4, 2));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
let probabilities = fitted
.predict_proba_simple(&X.view())
.expect("operation should succeed");
assert_eq!(probabilities.dim(), (4, 2));
for prob in probabilities.iter() {
assert!(*prob >= 0.0 && *prob <= 1.0);
}
}
#[test]
fn test_ensemble_of_chains_single_chain() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [0, 1]];
let eoc = EnsembleOfChains::new().n_chains(1);
let fitted = eoc
.fit_simple(&X.view(), &y)
.expect("operation should succeed");
assert_eq!(fitted.n_chains(), 1);
let predictions = fitted
.predict_simple(&X.view())
.expect("operation should succeed");
assert_eq!(predictions.dim(), (2, 2));
}
#[test]
fn test_ensemble_of_chains_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [0, 1], [1, 1]];
let eoc = EnsembleOfChains::new();
assert!(eoc.fit_simple(&X.view(), &y).is_err());
}
#[test]
fn test_one_vs_rest_classifier() {
let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [1.0, 1.0]];
let y = array![[1, 0], [0, 1], [1, 1], [0, 0]];
let ovr = OneVsRestClassifier::new();
let fitted = ovr
.fit(&X.view(), &y)
.expect("model fitting should succeed");
assert_eq!(fitted.n_labels(), 2);
assert_eq!(fitted.classes().len(), 2);
let predictions = fitted
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (4, 2));
for pred in predictions.iter() {
assert!(*pred == 0 || *pred == 1);
}
let probabilities = fitted
.predict_proba(&X.view())
.expect("operation should succeed");
assert_eq!(probabilities.dim(), (4, 2));
for prob in probabilities.iter() {
assert!(*prob >= 0.0 && *prob <= 1.0);
}
let scores = fitted
.decision_function(&X.view())
.expect("operation should succeed");
assert_eq!(scores.dim(), (4, 2));
for score in scores.iter() {
assert!(score.is_finite());
}
}
#[test]
fn test_one_vs_rest_classifier_invalid_input() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[1, 0], [0, 1], [1, 1]];
let ovr = OneVsRestClassifier::new();
assert!(ovr.fit(&X.view(), &y).is_err());
}
#[test]
fn test_metrics_one_error() {
let y_true = array![[1, 0, 0], [0, 1, 0], [0, 0, 1]];
let y_scores = array![[0.9, 0.1, 0.05], [0.1, 0.8, 0.1], [0.05, 0.1, 0.85]];
let one_err =
metrics::one_error(&y_true.view(), &y_scores.view()).expect("operation should succeed");
assert!((one_err - 0.0).abs() < 1e-10);
}
#[test]
fn test_metrics_one_error_with_errors() {
let y_true = array![[1, 0], [0, 1]];
let y_scores = array![[0.3, 0.7], [0.6, 0.4]];
let one_err =
metrics::one_error(&y_true.view(), &y_scores.view()).expect("operation should succeed");
assert!((one_err - 1.0).abs() < 1e-10);
}
#[test]
fn test_metrics_ranking_loss() {
let y_true = array![[1, 0], [0, 1]];
let y_scores = array![[0.8, 0.2], [0.3, 0.7]];
let ranking_loss =
metrics::ranking_loss(&y_true.view(), &y_scores.view()).expect("operation should succeed");
assert!((ranking_loss - 0.0).abs() < 1e-10);
}
#[test]
fn test_metrics_ranking_loss_with_errors() {
let y_true = array![[1, 0], [0, 1]];
let y_scores = array![[0.2, 0.8], [0.7, 0.3]];
let ranking_loss =
metrics::ranking_loss(&y_true.view(), &y_scores.view()).expect("operation should succeed");
assert!((ranking_loss - 1.0).abs() < 1e-10);
}
#[test]
fn test_metrics_average_precision_score() {
let y_true = array![[1, 0, 1], [0, 1, 0]];
let y_scores = array![[0.9, 0.1, 0.8], [0.2, 0.9, 0.3]];
let ap_score = metrics::average_precision_score(&y_true.view(), &y_scores.view())
.expect("operation should succeed");
assert!((ap_score - 1.0).abs() < 1e-10);
}
#[test]
fn test_metrics_precision_recall_micro() {
let y_true = array![[1, 0, 1], [0, 1, 0], [1, 1, 1]];
let y_pred = array![[1, 0, 0], [0, 1, 1], [1, 0, 1]];
let precision = metrics::precision_score_micro(&y_true.view(), &y_pred.view())
.expect("operation should succeed");
let recall = metrics::recall_score_micro(&y_true.view(), &y_pred.view())
.expect("operation should succeed");
assert!((precision - 0.8).abs() < 1e-10);
assert!((recall - 0.6666666666666666).abs() < 1e-10);
}
#[test]
fn test_metrics_invalid_shapes_new_metrics() {
let y_true = array![[1, 0], [0, 1]];
let y_pred = array![[1, 0, 1]]; let y_scores = array![[0.8, 0.2, 0.1]];
assert!(metrics::one_error(&y_true.view(), &y_scores.view()).is_err());
assert!(metrics::ranking_loss(&y_true.view(), &y_scores.view()).is_err());
assert!(metrics::average_precision_score(&y_true.view(), &y_scores.view()).is_err());
assert!(metrics::precision_score_micro(&y_true.view(), &y_pred.view()).is_err());
assert!(metrics::recall_score_micro(&y_true.view(), &y_pred.view()).is_err());
}