Skip to main content

model_selection_rs/scoring/
mod.rs

1//! Scoring: a lightweight metric abstraction usable standalone or, optionally,
2//! backed by `smartcore::metrics`.
3//!
4//! The [`Scorer`] trait is deliberately tiny — a single `score` call over
5//! `y_true` / `y_pred`. Built-in scorers (accuracy, MAE, MSE, RMSE, R²) are
6//! implemented directly here so the crate needs no modeling dependency. For
7//! metrics that already exist in mature form elsewhere (F1, ROC-AUC, …), enable
8//! the `smartcore-metrics` feature to get [`smartcore_adapter`] rather than
9//! reimplementing them.
10//!
11//! User-defined metrics are first-class: wrap any closure with
12//! [`make_scorer`].
13
14#[cfg(feature = "smartcore-metrics")]
15pub mod smartcore_adapter;
16
17use ndarray::Array1;
18
19/// A scoring metric over true and predicted target vectors.
20///
21/// Implementors compute a single scalar from aligned `y_true` / `y_pred`.
22/// [`greater_is_better`](Scorer::greater_is_better) tells callers (e.g. model
23/// selection loops) which direction is an improvement.
24pub trait Scorer {
25    /// Compute the score. `y_true` and `y_pred` are assumed equal length.
26    fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64;
27
28    /// A short human-readable name (used to label
29    /// [`CvResults`](crate::evaluate::CvResults) columns).
30    fn name(&self) -> &str;
31
32    /// Whether a larger score is better (e.g. accuracy, R²) or smaller is better
33    /// (e.g. MAE, MSE, RMSE). Defaults to `true`.
34    fn greater_is_better(&self) -> bool {
35        true
36    }
37}
38
39/// Classification accuracy: the fraction of predictions that exactly match.
40///
41/// Comparison is exact-equality on the `f64` values, so encode class labels as
42/// integral `f64`s (`0.0`, `1.0`, …).
43#[derive(Debug, Clone, Copy, Default)]
44pub struct Accuracy;
45
46impl Scorer for Accuracy {
47    fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
48        if y_true.is_empty() {
49            return f64::NAN;
50        }
51        let correct = y_true
52            .iter()
53            .zip(y_pred.iter())
54            .filter(|(t, p)| (**t - **p).abs() < f64::EPSILON)
55            .count();
56        correct as f64 / y_true.len() as f64
57    }
58    fn name(&self) -> &str {
59        "accuracy"
60    }
61}
62
63/// Mean absolute error.
64#[derive(Debug, Clone, Copy, Default)]
65pub struct MeanAbsoluteError;
66
67impl Scorer for MeanAbsoluteError {
68    fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
69        if y_true.is_empty() {
70            return f64::NAN;
71        }
72        let sum: f64 = y_true
73            .iter()
74            .zip(y_pred.iter())
75            .map(|(t, p)| (t - p).abs())
76            .sum();
77        sum / y_true.len() as f64
78    }
79    fn name(&self) -> &str {
80        "mae"
81    }
82    fn greater_is_better(&self) -> bool {
83        false
84    }
85}
86
87/// Mean squared error.
88#[derive(Debug, Clone, Copy, Default)]
89pub struct MeanSquaredError;
90
91impl Scorer for MeanSquaredError {
92    fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
93        if y_true.is_empty() {
94            return f64::NAN;
95        }
96        let sum: f64 = y_true
97            .iter()
98            .zip(y_pred.iter())
99            .map(|(t, p)| (t - p).powi(2))
100            .sum();
101        sum / y_true.len() as f64
102    }
103    fn name(&self) -> &str {
104        "mse"
105    }
106    fn greater_is_better(&self) -> bool {
107        false
108    }
109}
110
111/// Root mean squared error.
112#[derive(Debug, Clone, Copy, Default)]
113pub struct RootMeanSquaredError;
114
115impl Scorer for RootMeanSquaredError {
116    fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
117        MeanSquaredError.score(y_true, y_pred).sqrt()
118    }
119    fn name(&self) -> &str {
120        "rmse"
121    }
122    fn greater_is_better(&self) -> bool {
123        false
124    }
125}
126
127/// Coefficient of determination, R².
128///
129/// Returns `NaN` if `y_true` has zero variance (the score is undefined).
130#[derive(Debug, Clone, Copy, Default)]
131pub struct R2Score;
132
133impl Scorer for R2Score {
134    fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
135        if y_true.is_empty() {
136            return f64::NAN;
137        }
138        let mean = y_true.sum() / y_true.len() as f64;
139        let ss_tot: f64 = y_true.iter().map(|t| (t - mean).powi(2)).sum();
140        if ss_tot == 0.0 {
141            return f64::NAN;
142        }
143        let ss_res: f64 = y_true
144            .iter()
145            .zip(y_pred.iter())
146            .map(|(t, p)| (t - p).powi(2))
147            .sum();
148        1.0 - ss_res / ss_tot
149    }
150    fn name(&self) -> &str {
151        "r2"
152    }
153}
154
155/// A [`Scorer`] backed by an arbitrary closure — the analogue of scikit-learn's
156/// `make_scorer`, so users are never limited to the built-ins.
157pub struct ClosureScorer<F> {
158    name: String,
159    greater_is_better: bool,
160    f: F,
161}
162
163impl<F> Scorer for ClosureScorer<F>
164where
165    F: Fn(&Array1<f64>, &Array1<f64>) -> f64,
166{
167    fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64 {
168        (self.f)(y_true, y_pred)
169    }
170    fn name(&self) -> &str {
171        &self.name
172    }
173    fn greater_is_better(&self) -> bool {
174        self.greater_is_better
175    }
176}
177
178/// Wrap a closure as a [`Scorer`].
179///
180/// ```
181/// use ndarray::array;
182/// use model_selection_rs::scoring::{make_scorer, Scorer};
183///
184/// // Max absolute error, lower is better.
185/// let max_ae = make_scorer("max_ae", false, |t, p| {
186///     t.iter().zip(p.iter()).map(|(a, b)| (a - b).abs()).fold(0.0, f64::max)
187/// });
188/// let s = max_ae.score(&array![1.0, 2.0, 3.0], &array![1.0, 2.0, 5.0]);
189/// assert_eq!(s, 2.0);
190/// assert!(!max_ae.greater_is_better());
191/// ```
192pub fn make_scorer<F>(name: impl Into<String>, greater_is_better: bool, f: F) -> ClosureScorer<F>
193where
194    F: Fn(&Array1<f64>, &Array1<f64>) -> f64,
195{
196    ClosureScorer {
197        name: name.into(),
198        greater_is_better,
199        f,
200    }
201}
202
203#[cfg(test)]
204mod tests {
205    use super::*;
206    use approx::assert_relative_eq;
207    use ndarray::array;
208
209    #[test]
210    fn accuracy_matches_hand_count() {
211        let t = array![0.0, 1.0, 1.0, 0.0];
212        let p = array![0.0, 1.0, 0.0, 0.0];
213        assert_relative_eq!(Accuracy.score(&t, &p), 0.75);
214    }
215
216    #[test]
217    fn regression_metrics_match_hand_values() {
218        let t = array![1.0, 2.0, 3.0];
219        let p = array![1.0, 2.0, 5.0]; // errors: 0, 0, 2
220        assert_relative_eq!(MeanAbsoluteError.score(&t, &p), 2.0 / 3.0);
221        assert_relative_eq!(MeanSquaredError.score(&t, &p), 4.0 / 3.0);
222        assert_relative_eq!(RootMeanSquaredError.score(&t, &p), (4.0f64 / 3.0).sqrt());
223    }
224
225    #[test]
226    fn r2_is_one_for_perfect_fit() {
227        let t = array![1.0, 2.0, 3.0, 4.0];
228        assert_relative_eq!(R2Score.score(&t, &t), 1.0);
229    }
230
231    #[test]
232    fn r2_is_zero_for_mean_predictor() {
233        let t = array![1.0, 2.0, 3.0, 4.0];
234        let mean = array![2.5, 2.5, 2.5, 2.5];
235        assert_relative_eq!(R2Score.score(&t, &mean), 0.0);
236    }
237
238    #[test]
239    fn greater_is_better_flags() {
240        assert!(Accuracy.greater_is_better());
241        assert!(R2Score.greater_is_better());
242        assert!(!MeanAbsoluteError.greater_is_better());
243        assert!(!MeanSquaredError.greater_is_better());
244        assert!(!RootMeanSquaredError.greater_is_better());
245    }
246}