#[cfg(feature = "smartcore-metrics")]
pub mod smartcore_adapter;
use ndarray::Array1;
pub trait Scorer {
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64;
fn name(&self) -> &str;
fn greater_is_better(&self) -> bool {
true
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct Accuracy;
impl Scorer for Accuracy {
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
if y_true.is_empty() {
return f64::NAN;
}
let correct = y_true
.iter()
.zip(y_pred.iter())
.filter(|(t, p)| (**t - **p).abs() < f64::EPSILON)
.count();
correct as f64 / y_true.len() as f64
}
fn name(&self) -> &str {
"accuracy"
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct MeanAbsoluteError;
impl Scorer for MeanAbsoluteError {
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
if y_true.is_empty() {
return f64::NAN;
}
let sum: f64 = y_true
.iter()
.zip(y_pred.iter())
.map(|(t, p)| (t - p).abs())
.sum();
sum / y_true.len() as f64
}
fn name(&self) -> &str {
"mae"
}
fn greater_is_better(&self) -> bool {
false
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct MeanSquaredError;
impl Scorer for MeanSquaredError {
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
if y_true.is_empty() {
return f64::NAN;
}
let sum: f64 = y_true
.iter()
.zip(y_pred.iter())
.map(|(t, p)| (t - p).powi(2))
.sum();
sum / y_true.len() as f64
}
fn name(&self) -> &str {
"mse"
}
fn greater_is_better(&self) -> bool {
false
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct RootMeanSquaredError;
impl Scorer for RootMeanSquaredError {
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
MeanSquaredError.score(y_true, y_pred).sqrt()
}
fn name(&self) -> &str {
"rmse"
}
fn greater_is_better(&self) -> bool {
false
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct R2Score;
impl Scorer for R2Score {
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
if y_true.is_empty() {
return f64::NAN;
}
let mean = y_true.sum() / y_true.len() as f64;
let ss_tot: f64 = y_true.iter().map(|t| (t - mean).powi(2)).sum();
if ss_tot == 0.0 {
return f64::NAN;
}
let ss_res: f64 = y_true
.iter()
.zip(y_pred.iter())
.map(|(t, p)| (t - p).powi(2))
.sum();
1.0 - ss_res / ss_tot
}
fn name(&self) -> &str {
"r2"
}
}
pub struct ClosureScorer<F> {
name: String,
greater_is_better: bool,
f: F,
}
impl<F> Scorer for ClosureScorer<F>
where
F: Fn(&Array1<f64>, &Array1<f64>) -> f64,
{
fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
(self.f)(y_true, y_pred)
}
fn name(&self) -> &str {
&self.name
}
fn greater_is_better(&self) -> bool {
self.greater_is_better
}
}
pub fn make_scorer<F>(name: impl Into<String>, greater_is_better: bool, f: F) -> ClosureScorer<F>
where
F: Fn(&Array1<f64>, &Array1<f64>) -> f64,
{
ClosureScorer {
name: name.into(),
greater_is_better,
f,
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
use ndarray::array;
#[test]
fn accuracy_matches_hand_count() {
let t = array![0.0, 1.0, 1.0, 0.0];
let p = array![0.0, 1.0, 0.0, 0.0];
assert_relative_eq!(Accuracy.score(&t, &p), 0.75);
}
#[test]
fn regression_metrics_match_hand_values() {
let t = array![1.0, 2.0, 3.0];
let p = array![1.0, 2.0, 5.0]; assert_relative_eq!(MeanAbsoluteError.score(&t, &p), 2.0 / 3.0);
assert_relative_eq!(MeanSquaredError.score(&t, &p), 4.0 / 3.0);
assert_relative_eq!(RootMeanSquaredError.score(&t, &p), (4.0f64 / 3.0).sqrt());
}
#[test]
fn r2_is_one_for_perfect_fit() {
let t = array![1.0, 2.0, 3.0, 4.0];
assert_relative_eq!(R2Score.score(&t, &t), 1.0);
}
#[test]
fn r2_is_zero_for_mean_predictor() {
let t = array![1.0, 2.0, 3.0, 4.0];
let mean = array![2.5, 2.5, 2.5, 2.5];
assert_relative_eq!(R2Score.score(&t, &mean), 0.0);
}
#[test]
fn greater_is_better_flags() {
assert!(Accuracy.greater_is_better());
assert!(R2Score.greater_is_better());
assert!(!MeanAbsoluteError.greater_is_better());
assert!(!MeanSquaredError.greater_is_better());
assert!(!RootMeanSquaredError.greater_is_better());
}
}