From 2b8ee4f6631dfa97aef823430ff3706a665aae92 Mon Sep 17 00:00:00 2001 From: Montana Low Date: Mon, 12 Sep 2022 18:31:34 -0700 Subject: [PATCH 1/4] handle multiclass precision/recall --- src/math/num.rs | 13 +++++++- src/metrics/precision.rs | 66 +++++++++++++++++++++++++------------- src/metrics/recall.rs | 68 ++++++++++++++++++++++++++-------------- 3 files changed, 101 insertions(+), 46 deletions(-) diff --git a/src/math/num.rs b/src/math/num.rs index 71999498..c454b9d0 100644 --- a/src/math/num.rs +++ b/src/math/num.rs @@ -46,8 +46,11 @@ pub trait RealNumber: self * self } - /// Raw transmutation to u64 + /// Raw transmutation to u32 fn to_f32_bits(self) -> u32; + + /// Raw transmutation to u64 + fn to_f64_bits(self) -> u64; } impl RealNumber for f64 { @@ -89,6 +92,10 @@ impl RealNumber for f64 { fn to_f32_bits(self) -> u32 { self.to_bits() as u32 } + + fn to_f64_bits(self) -> u64 { + self.to_bits() + } } impl RealNumber for f32 { @@ -130,6 +137,10 @@ impl RealNumber for f32 { fn to_f32_bits(self) -> u32 { self.to_bits() } + + fn to_f64_bits(self) -> u64 { + self.to_bits() as u64 + } } #[cfg(test)] diff --git a/src/metrics/precision.rs b/src/metrics/precision.rs index a0171aa5..3b1f0b30 100644 --- a/src/metrics/precision.rs +++ b/src/metrics/precision.rs @@ -18,6 +18,8 @@ //! //! //! +use std::collections::HashSet; + #[cfg(feature = "serde")] use serde::{Deserialize, Serialize}; @@ -42,34 +44,35 @@ impl Precision { ); } - let mut tp = 0; - let mut p = 0; - let n = y_true.len(); - for i in 0..n { - if y_true.get(i) != T::zero() && y_true.get(i) != T::one() { - panic!( - "Precision can only be applied to binary classification: {}", - y_true.get(i) - ); - } - - if y_pred.get(i) != T::zero() && y_pred.get(i) != T::one() { - panic!( - "Precision can only be applied to binary classification: {}", - y_pred.get(i) - ); - } - - if y_pred.get(i) == T::one() { - p += 1; + let mut classes = HashSet::new(); + for i in 0..y_true.len() { + classes.insert(y_true.get(i).to_f32_bits()); + } + let classes = classes.len(); - if y_true.get(i) == T::one() { + let mut tp = 0; + let mut fp = 0; + for i in 0..y_true.len() { + if y_pred.get(i) == y_true.get(i) { + if classes == 2 { + if y_true.get(i) == T::one() { + tp += 1; + } + } else { tp += 1; } + } else { + if classes == 2 { + if y_true.get(i) == T::one() { + fp += 1; + } + } else { + fp += 1; + } } } - T::from_i64(tp).unwrap() / T::from_i64(p).unwrap() + T::from_i64(tp).unwrap() / (T::from_i64(tp).unwrap() + T::from_i64(fp).unwrap()) } } @@ -88,5 +91,24 @@ mod tests { assert!((score1 - 0.5).abs() < 1e-8); assert!((score2 - 1.0).abs() < 1e-8); + + let y_pred: Vec = vec![0., 0., 1., 1., 1., 1.]; + let y_true: Vec = vec![0., 1., 1., 0., 1., 0.]; + + let score3: f64 = Precision {}.get_score(&y_pred, &y_true); + assert!((score3 - 0.5).abs() < 1e-8); + } + + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] + #[test] + fn precision_multiclass() { + let y_true: Vec = vec![0., 0., 0., 1., 1., 1., 2., 2., 2.]; + let y_pred: Vec = vec![0., 1., 2., 0., 1., 2., 0., 1., 2.]; + + let score1: f64 = Precision {}.get_score(&y_pred, &y_true); + let score2: f64 = Precision {}.get_score(&y_pred, &y_pred); + + assert!((score1 - 0.333333333).abs() < 1e-8); + assert!((score2 - 1.0).abs() < 1e-8); } } diff --git a/src/metrics/recall.rs b/src/metrics/recall.rs index 18863aee..ed57cecc 100644 --- a/src/metrics/recall.rs +++ b/src/metrics/recall.rs @@ -18,6 +18,9 @@ //! //! //! +use std::collections::HashSet; +use std::convert::TryInto; + #[cfg(feature = "serde")] use serde::{Deserialize, Serialize}; @@ -42,34 +45,34 @@ impl Recall { ); } - let mut tp = 0; - let mut p = 0; - let n = y_true.len(); - for i in 0..n { - if y_true.get(i) != T::zero() && y_true.get(i) != T::one() { - panic!( - "Recall can only be applied to binary classification: {}", - y_true.get(i) - ); - } - - if y_pred.get(i) != T::zero() && y_pred.get(i) != T::one() { - panic!( - "Recall can only be applied to binary classification: {}", - y_pred.get(i) - ); - } - - if y_true.get(i) == T::one() { - p += 1; + let mut classes = HashSet::new(); + for i in 0..y_true.len() { + classes.insert(y_true.get(i).to_f64_bits()); + } + let classes: i64 = classes.len().try_into().unwrap(); - if y_pred.get(i) == T::one() { + let mut tp = 0; + let mut fne = 0; + for i in 0..y_true.len() { + if y_pred.get(i) == y_true.get(i) { + if classes == 2 { + if y_true.get(i) == T::one() { + tp += 1; + } + } else { tp += 1; } + } else { + if classes == 2 { + if y_true.get(i) != T::one() { + fne += 1; + } + } else { + fne += 1; + } } } - - T::from_i64(tp).unwrap() / T::from_i64(p).unwrap() + T::from_i64(tp).unwrap() / (T::from_i64(tp).unwrap() + T::from_i64(fne).unwrap()) } } @@ -88,5 +91,24 @@ mod tests { assert!((score1 - 0.5).abs() < 1e-8); assert!((score2 - 1.0).abs() < 1e-8); + + let y_pred: Vec = vec![0., 0., 1., 1., 1., 1.]; + let y_true: Vec = vec![0., 1., 1., 0., 1., 0.]; + + let score3: f64 = Recall {}.get_score(&y_pred, &y_true); + assert!((score3 - 0.66666666).abs() < 1e-8); + } + + #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)] + #[test] + fn recall_multiclass() { + let y_true: Vec = vec![0., 0., 0., 1., 1., 1., 2., 2., 2.]; + let y_pred: Vec = vec![0., 1., 2., 0., 1., 2., 0., 1., 2.]; + + let score1: f64 = Recall {}.get_score(&y_pred, &y_true); + let score2: f64 = Recall {}.get_score(&y_pred, &y_pred); + + assert!((score1 - 0.333333333).abs() < 1e-8); + assert!((score2 - 1.0).abs() < 1e-8); } } From 8f8a41f98d519ad5eb896d51103b808cfa5762e6 Mon Sep 17 00:00:00 2001 From: Montana Low Date: Mon, 12 Sep 2022 19:12:28 -0700 Subject: [PATCH 2/4] cargo fmt --- src/metrics/precision.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/metrics/precision.rs b/src/metrics/precision.rs index 3b1f0b30..588257cd 100644 --- a/src/metrics/precision.rs +++ b/src/metrics/precision.rs @@ -46,7 +46,7 @@ impl Precision { let mut classes = HashSet::new(); for i in 0..y_true.len() { - classes.insert(y_true.get(i).to_f32_bits()); + classes.insert(y_true.get(i).to_f64_bits()); } let classes = classes.len(); From 128a675d80408241fec853d0d32327604a456e5d Mon Sep 17 00:00:00 2001 From: Montana Low Date: Mon, 12 Sep 2022 19:31:20 -0700 Subject: [PATCH 3/4] cargo fmt --- src/metrics/precision.rs | 4 ++-- src/metrics/recall.rs | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/metrics/precision.rs b/src/metrics/precision.rs index 588257cd..02720e67 100644 --- a/src/metrics/precision.rs +++ b/src/metrics/precision.rs @@ -46,7 +46,7 @@ impl Precision { let mut classes = HashSet::new(); for i in 0..y_true.len() { - classes.insert(y_true.get(i).to_f64_bits()); + classes.insert(y_true.get(i).to_f64_bits()); } let classes = classes.len(); @@ -64,7 +64,7 @@ impl Precision { } else { if classes == 2 { if y_true.get(i) == T::one() { - fp += 1; + fp += 1; } } else { fp += 1; diff --git a/src/metrics/recall.rs b/src/metrics/recall.rs index ed57cecc..a962cb98 100644 --- a/src/metrics/recall.rs +++ b/src/metrics/recall.rs @@ -47,7 +47,7 @@ impl Recall { let mut classes = HashSet::new(); for i in 0..y_true.len() { - classes.insert(y_true.get(i).to_f64_bits()); + classes.insert(y_true.get(i).to_f64_bits()); } let classes: i64 = classes.len().try_into().unwrap(); From 7d43cd1cb3b91be9fa860bb7019831a9bd852d69 Mon Sep 17 00:00:00 2001 From: Montana Low Date: Mon, 12 Sep 2022 19:37:12 -0700 Subject: [PATCH 4/4] clippy --- src/metrics/precision.rs | 10 ++++------ src/metrics/recall.rs | 10 ++++------ 2 files changed, 8 insertions(+), 12 deletions(-) diff --git a/src/metrics/precision.rs b/src/metrics/precision.rs index 02720e67..a2bad30c 100644 --- a/src/metrics/precision.rs +++ b/src/metrics/precision.rs @@ -61,14 +61,12 @@ impl Precision { } else { tp += 1; } - } else { - if classes == 2 { - if y_true.get(i) == T::one() { - fp += 1; - } - } else { + } else if classes == 2 { + if y_true.get(i) == T::one() { fp += 1; } + } else { + fp += 1; } } diff --git a/src/metrics/recall.rs b/src/metrics/recall.rs index a962cb98..48ddeeb2 100644 --- a/src/metrics/recall.rs +++ b/src/metrics/recall.rs @@ -62,14 +62,12 @@ impl Recall { } else { tp += 1; } - } else { - if classes == 2 { - if y_true.get(i) != T::one() { - fne += 1; - } - } else { + } else if classes == 2 { + if y_true.get(i) != T::one() { fne += 1; } + } else { + fne += 1; } } T::from_i64(tp).unwrap() / (T::from_i64(tp).unwrap() + T::from_i64(fne).unwrap())