///|
fn class_position(classes : Array[Int], label : Int) -> Int {
for index, value in classes {
if value == label {
return index
}
}
-1
}
///|
fn safe_rate(numerator : Int, denominator : Int) -> Double {
if denominator == 0 {
0.0
} else {
numerator.to_double() / denominator.to_double()
}
}
///|
fn harmonic_f1(precision : Double, recall : Double) -> Double {
if precision + recall == 0.0 {
0.0
} else {
2.0 * precision * recall / (precision + recall)
}
}
///|
/// Computes a confusion matrix and common single-label classification metrics.
pub fn classification_metrics(
actual : Array[Int],
predicted : Array[Int],
) -> Result[ClassificationMetrics, SvmError] {
if actual.is_empty() {
return Err(EmptyMetricInput)
}
if actual.length() != predicted.length() {
return Err(MetricLengthMismatch(actual.length(), predicted.length()))
}
let combined = actual.copy()
for label in predicted {
combined.push(label)
}
let classes = sorted_distinct_labels(combined)
let size = classes.length()
let matrix : Array[Array[Int]] = Array::makei(size, fn(_) {
Array::make(size, 0)
})
for index = 0; index < actual.length(); index = index + 1 {
let row = class_position(classes, actual[index])
let column = class_position(classes, predicted[index])
if row < 0 {
return Err(UnknownClassLabel(actual[index]))
}
if column < 0 {
return Err(UnknownClassLabel(predicted[index]))
}
matrix[row][column] = matrix[row][column] + 1
}
let per_class : Array[ClassMetric] = []
let mut correct = 0
let mut macro_precision = 0.0
let mut macro_recall = 0.0
let mut macro_f1 = 0.0
let mut weighted_precision = 0.0
let mut weighted_recall = 0.0
let mut weighted_f1 = 0.0
for class_index = 0; class_index < size; class_index = class_index + 1 {
let true_positive = matrix[class_index][class_index]
correct = correct + true_positive
let mut support = 0
let mut predicted_count = 0
for other = 0; other < size; other = other + 1 {
support = support + matrix[class_index][other]
predicted_count = predicted_count + matrix[other][class_index]
}
let false_positive = predicted_count - true_positive
let false_negative = support - true_positive
let precision = safe_rate(true_positive, predicted_count)
let recall = safe_rate(true_positive, support)
let f1 = harmonic_f1(precision, recall)
per_class.push({
class_label: classes[class_index],
true_positive_count: true_positive,
false_positive_count: false_positive,
false_negative_count: false_negative,
class_support: support,
precision_value: precision,
recall_value: recall,
f1_value: f1,
})
macro_precision = macro_precision + precision
macro_recall = macro_recall + recall
macro_f1 = macro_f1 + f1
let fraction = support.to_double() / actual.length().to_double()
weighted_precision = weighted_precision + fraction * precision
weighted_recall = weighted_recall + fraction * recall
weighted_f1 = weighted_f1 + fraction * f1
}
macro_precision = macro_precision / size.to_double()
macro_recall = macro_recall / size.to_double()
macro_f1 = macro_f1 / size.to_double()
let accuracy = correct.to_double() / actual.length().to_double()
Ok({
ordered_classes: classes,
confusion_counts: matrix,
class_metrics: per_class,
total_observations: actual.length(),
accuracy_value: accuracy,
error_rate_value: 1.0 - accuracy,
balanced_accuracy_value: macro_recall,
macro_precision_value: macro_precision,
macro_recall_value: macro_recall,
macro_f1_value: macro_f1,
micro_precision_value: accuracy,
micro_recall_value: accuracy,
micro_f1_value: accuracy,
weighted_precision_value: weighted_precision,
weighted_recall_value: weighted_recall,
weighted_f1_value: weighted_f1,
})
}
///|
pub fn ClassMetric::label(self : ClassMetric) -> Int {
self.class_label
}
///|
pub fn ClassMetric::true_positives(self : ClassMetric) -> Int {
self.true_positive_count
}
///|
pub fn ClassMetric::false_positives(self : ClassMetric) -> Int {
self.false_positive_count
}
///|
pub fn ClassMetric::false_negatives(self : ClassMetric) -> Int {
self.false_negative_count
}
///|
pub fn ClassMetric::support(self : ClassMetric) -> Int {
self.class_support
}
///|
pub fn ClassMetric::precision(self : ClassMetric) -> Double {
self.precision_value
}
///|
pub fn ClassMetric::recall(self : ClassMetric) -> Double {
self.recall_value
}
///|
pub fn ClassMetric::f1(self : ClassMetric) -> Double {
self.f1_value
}
///|
pub fn ClassificationMetrics::classes(
self : ClassificationMetrics,
) -> Array[Int] {
self.ordered_classes.copy()
}
///|
pub fn ClassificationMetrics::confusion_matrix(
self : ClassificationMetrics,
) -> Array[Array[Int]] {
let copied : Array[Array[Int]] = []
for row in self.confusion_counts {
copied.push(row.copy())
}
copied
}
///|
pub fn ClassificationMetrics::per_class(
self : ClassificationMetrics,
) -> Array[ClassMetric] {
self.class_metrics.copy()
}
///|
pub fn ClassificationMetrics::observation_count(
self : ClassificationMetrics,
) -> Int {
self.total_observations
}
///|
pub fn ClassificationMetrics::accuracy(self : ClassificationMetrics) -> Double {
self.accuracy_value
}
///|
pub fn ClassificationMetrics::error_rate(
self : ClassificationMetrics,
) -> Double {
self.error_rate_value
}
///|
pub fn ClassificationMetrics::balanced_accuracy(
self : ClassificationMetrics,
) -> Double {
self.balanced_accuracy_value
}
///|
pub fn ClassificationMetrics::macro_precision(
self : ClassificationMetrics,
) -> Double {
self.macro_precision_value
}
///|
pub fn ClassificationMetrics::macro_recall(
self : ClassificationMetrics,
) -> Double {
self.macro_recall_value
}
///|
pub fn ClassificationMetrics::macro_f1(self : ClassificationMetrics) -> Double {
self.macro_f1_value
}
///|
pub fn ClassificationMetrics::micro_precision(
self : ClassificationMetrics,
) -> Double {
self.micro_precision_value
}
///|
pub fn ClassificationMetrics::micro_recall(
self : ClassificationMetrics,
) -> Double {
self.micro_recall_value
}
///|
pub fn ClassificationMetrics::micro_f1(self : ClassificationMetrics) -> Double {
self.micro_f1_value
}
///|
pub fn ClassificationMetrics::weighted_precision(
self : ClassificationMetrics,
) -> Double {
self.weighted_precision_value
}
///|
pub fn ClassificationMetrics::weighted_recall(
self : ClassificationMetrics,
) -> Double {
self.weighted_recall_value
}
///|
pub fn ClassificationMetrics::weighted_f1(
self : ClassificationMetrics,
) -> Double {
self.weighted_f1_value
}