use std::collections::HashMap;
use std::collections::HashSet;
use std::hash::Hash;
use std::error::Error;
use std::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct LengthError(usize, usize);
impl fmt::Display for LengthError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "Dataset lengths must be equal, found {} and {}", self.0, self.1)
}
}
impl Error for LengthError {
fn description(&self) -> &str {
"Dataset lengths do not match"
}
}
pub fn gini<T>(data: &[T]) -> f32
where
T: Eq,
T: Hash,
{
if data.is_empty() {
return 1.0;
}
fn p_squared(count: usize, len: f32) -> f32 {
let p = count as f32 / len;
p * p
}
let len = data.len() as f32;
let mut count = HashMap::new();
for value in data {
*count.entry(value).or_insert(0) += 1;
}
let sum: f32 = count
.into_iter()
.map(|(_, c)| c)
.map(|x| p_squared(x, len))
.sum();
1.0 - sum
}
pub fn categorical_accuracy<T>(pred: &[T], actual: &[T]) -> Result<f32, LengthError>
where
T: Eq,
{
if pred.len() != actual.len(){
return Err(LengthError(pred.len(), actual.len()));
}
let truthy = pred.iter().zip(actual).filter(|(x, y)| x == y).count();
Ok(truthy as f32 / pred.len() as f32)
}
fn class_precision<T>(pred: &[T], actual: &[T], class: &T) -> f32
where
T: Eq,
{
let true_positives = pred
.iter()
.zip(actual)
.filter(|(p, a)| p == a && **p == *class)
.count() as f32;
let all_positives = pred.iter().filter(|p| **p == *class).count() as f32;
if all_positives == 0.0 {
0.0
} else {
true_positives / all_positives
}
}
fn weighted_precision<T>(pred: &[T], actual: &[T]) -> f32
where
T: Eq,
T: Hash,
{
let classes: HashSet<_> = pred.into_iter().collect();
let mut class_weights = HashMap::new();
for value in &classes {
class_weights.insert(
value,
actual.iter().filter(|a| *a == *value).count() as f32 / actual.len() as f32,
);
}
classes
.iter()
.map(|c| class_precision(pred, actual, &c) * class_weights[c])
.sum()
}
fn macro_precision<T>(pred: &[T], actual: &[T]) -> f32
where
T: Eq,
T: Hash,
{
let classes: HashSet<_> = pred.into_iter().collect();
let mut class_weights = HashMap::new();
for value in classes.clone() {
class_weights.insert(value, 1.0 / actual.len() as f32);
}
classes
.iter()
.map(|c| class_precision(pred, actual, c) / classes.len() as f32)
.sum()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum Average {
Macro,
Weighted,
}
impl Default for Average {
fn default() -> Self {
Average::Macro
}
}
pub fn precision<T>(pred: &[T], actual: &[T], average: Average) -> Result<f32, LengthError>
where
T: Eq,
T: Hash,
{
if pred.len() != actual.len(){
return Err(LengthError(pred.len(), actual.len()));
}
match average {
Average::Macro => Ok(macro_precision(pred, actual)),
Average::Weighted => Ok(weighted_precision(pred, actual)),
}
}
fn class_recall<T>(pred: &[T], actual: &[T], class: &T) -> f32
where
T: Eq,
{
let true_positives = pred
.iter()
.zip(actual)
.filter(|(p, a)| p == a && **a == *class)
.count() as f32;
let tp_fn = actual.iter().filter(|a| **a == *class).count() as f32;
if tp_fn == 0.0 {
0.0
} else {
true_positives / tp_fn
}
}
fn weighted_recall<T>(pred: &[T], actual: &[T]) -> f32
where
T: Eq,
T: Hash,
{
let classes: HashSet<_> = pred.into_iter().collect();
let mut class_weights = HashMap::new();
for value in &classes {
class_weights.insert(
value,
actual.iter().filter(|a| **a == **value).count() as f32 / actual.len() as f32,
);
}
classes
.iter()
.map(|c| class_recall(pred, actual, &c) * class_weights[c])
.sum()
}
fn macro_recall<T>(pred: &[T], actual: &[T]) -> f32
where
T: Eq,
T: Hash,
{
let classes: HashSet<_> = pred.into_iter().collect();
classes
.iter()
.map(|c| class_recall(pred, actual, *c) / classes.len() as f32)
.sum()
}
pub fn recall<T>(pred: &[T], actual: &[T], average: Average) -> Result<f32, LengthError>
where
T: Eq,
T: Hash,
{
if pred.len() != actual.len(){
return Err(LengthError(pred.len(), actual.len()));
}
match average {
Average::Macro => Ok(macro_recall(pred, actual)),
Average::Weighted => Ok(weighted_recall(pred, actual)),
}
}
fn macro_f1<T>(pred: &[T], actual: &[T]) -> f32
where
T: Eq,
T: Hash,
{
let recall = macro_recall(pred, actual);
let precision = macro_precision(pred, actual);
2.0 * (recall * precision) / (recall + precision)
}
fn weighted_f1<T>(pred: &[T], actual: &[T]) -> f32
where
T: Eq,
T: Hash,
{
let recall = weighted_recall(pred, actual);
let precision = weighted_precision(pred, actual);
2.0 * (recall * precision) / (recall + precision)
}
pub fn f1_score<T>(pred: &[T], actual: &[T], average: Average) -> Result<f32, LengthError>
where
T: Eq,
T: Hash,
{
if pred.len() != actual.len() {
return Err(LengthError(pred.len(), actual.len()));
}
match average {
Average::Macro => Ok(macro_f1(pred, actual)),
Average::Weighted => Ok(weighted_f1(pred, actual)),
}
}
pub fn hamming_loss<T>(pred: &[T], actual: &[T]) -> Result<f32, LengthError>
where
T: Eq,
{
let cat_acc = categorical_accuracy(pred, actual)?;
Ok(1. - cat_acc)
}
fn macro_fbeta_score<T>(pred: &[T], actual: &[T], beta: f32) -> f32
where
T: Eq,
T: Hash,
{
let precision = macro_precision(pred, actual);
let recall = macro_recall(pred, actual);
let top = (1.0 + beta * beta) * (recall * precision);
let bottom = (beta * beta * precision) + recall;
top / bottom
}
fn weighted_fbeta_score<T>(pred: &[T], actual: &[T], beta: f32) -> f32
where
T: Eq,
T: Hash,
{
let precision = weighted_precision(pred, actual);
let recall = weighted_recall(pred, actual);
let top = (1.0 + beta * beta) * (recall * precision);
let bottom = (beta * beta * precision) + recall;
top / bottom
}
pub fn fbeta_score<T>(pred: &[T], actual: &[T], beta: f32, average: Average) -> Result<f32, LengthError>
where
T: Eq,
T: Hash,
{
if pred.len() != actual.len(){
return Err(LengthError(pred.len(), actual.len()));
}
match average {
Average::Macro => Ok(macro_fbeta_score(pred, actual, beta)),
Average::Weighted => Ok(weighted_fbeta_score(pred, actual, beta)),
}
}
pub fn jaccard_similiarity_score<T>(pred: &[T], actual: &[T]) -> Result<f32, LengthError>
where
T: Eq,
{
categorical_accuracy(pred, actual)
}
#[cfg(test)]
#[macro_use] extern crate approx;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_gini() {
let vec = vec![0, 0, 0, 1];
assert_ulps_eq!(0.375, gini(&vec));
let v2 = vec![0, 0];
assert_ulps_eq!(0.0, gini(&v2));
let mut v3 = vec![0];
v3.pop();
assert_ulps_eq!(1.0, gini(&v3));
}
#[test]
fn test_categorical_accuracy() {
let pred = vec![0, 1, 0, 1, 0, 1];
let real = vec![0, 0, 0, 0, 1, 0];
assert_ulps_eq!(0.33333333, categorical_accuracy(&pred, &real).unwrap());
let pred_short = vec![0];
assert!(categorical_accuracy(&pred_short, &real).is_err());
}
#[test]
fn test_class_precision() {
let actual = vec![0, 1, 2, 0, 1, 2];
let pred = vec![0, 2, 1, 0, 0, 1];
assert_ulps_eq!(0.6666666, class_precision(&pred, &actual, &0));
}
#[test]
fn test_class_recall() {
let actual = vec![0, 1, 2, 0, 0, 0];
let pred = vec![0, 2, 1, 0, 0, 1];
assert_ulps_eq!(0.75, class_recall(&pred, &actual, &0));
}
#[test]
fn test_weighted_precision() {
let actual = vec![0, 1, 2, 0, 1, 2];
let pred = vec![0, 2, 1, 0, 0, 1];
assert_ulps_eq!(0.22222222, weighted_precision(&pred, &actual));
}
#[test]
fn test_macro_precision() {
let actual = vec![0, 1, 2, 0, 1, 2];
let pred = vec![0, 2, 1, 0, 0, 1];
assert_ulps_eq!(0.22222222, macro_precision(&pred, &actual));
}
#[test]
fn test_macro_recall() {
let actual = vec![0, 1, 2, 0, 1, 2];
let pred = vec![0, 2, 1, 0, 0, 1];
assert_ulps_eq!(0.33333333, macro_recall(&pred, &actual));
}
#[test]
fn test_weighted_recall() {
let actual = vec![0, 1, 2, 0, 1, 2];
let pred = vec![0, 2, 1, 0, 0, 1];
assert_ulps_eq!(0.333333333, weighted_recall(&pred, &actual));
}
#[test]
fn test_f1_score() {
let actual = vec![0, 1, 2, 0, 1, 2];
let pred = vec![0, 2, 1, 0, 0, 1];
assert_ulps_eq!(f1_score(&pred, &actual, Average::Macro).unwrap(), 0.26666665);
let pred_short = vec![0];
assert!(f1_score(&pred_short, &actual, Average::Weighted).is_err());
}
}