model_selection_rs/scoring/
mod.rs1#[cfg(feature = "smartcore-metrics")]
15pub mod smartcore_adapter;
16
17use ndarray::Array1;
18
19pub trait Scorer {
25 fn score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> f64;
27
28 fn name(&self) -> &str;
31
32 fn greater_is_better(&self) -> bool {
35 true
36 }
37}
38
39#[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#[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#[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#[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#[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
155pub 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
178pub 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]; 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}