///|
/// Group-level diagnostics for online classification services.
pub struct GroupMetric {
group : String
mut count : Int
mut positives : Int
mut true_positives : Int
mut false_positives : Int
mut false_negatives : Int
}
///|
pub fn GroupMetric::new(group : String) -> GroupMetric {
{
group,
count: 0,
positives: 0,
true_positives: 0,
false_positives: 0,
false_negatives: 0,
}
}
///|
pub fn GroupMetric::observe(
self : GroupMetric,
prediction : Double,
label : Double,
threshold? : Double = 0.5,
) -> Unit {
let predicted_positive = prediction >= threshold
let actual_positive = label >= 0.5
self.count += 1
if actual_positive {
self.positives += 1
}
if predicted_positive && actual_positive {
self.true_positives += 1
}
if predicted_positive && !actual_positive {
self.false_positives += 1
}
if !predicted_positive && actual_positive {
self.false_negatives += 1
}
}
///|
pub fn GroupMetric::group(self : GroupMetric) -> String {
self.group
}
///|
pub fn GroupMetric::count(self : GroupMetric) -> Int {
self.count
}
///|
pub fn GroupMetric::positive_rate(self : GroupMetric) -> Double {
if self.count == 0 {
0.0
} else {
(self.true_positives + self.false_positives).to_double() /
self.count.to_double()
}
}
///|
pub fn GroupMetric::recall(self : GroupMetric) -> Double {
if self.positives == 0 {
0.0
} else {
self.true_positives.to_double() / self.positives.to_double()
}
}
///|
pub fn GroupMetric::precision(self : GroupMetric) -> Double {
let predicted = self.true_positives + self.false_positives
if predicted == 0 {
0.0
} else {
self.true_positives.to_double() / predicted.to_double()
}
}
///|
pub fn GroupMetric::false_positive_rate(self : GroupMetric) -> Double {
let negatives = self.count - self.positives
if negatives <= 0 {
0.0
} else {
self.false_positives.to_double() / negatives.to_double()
}
}
///|
pub fn GroupMetric::reset(self : GroupMetric) -> Unit {
self.count = 0
self.positives = 0
self.true_positives = 0
self.false_positives = 0
self.false_negatives = 0
}
///|
pub struct FairnessMonitor {
groups : Map[String, GroupMetric]
mut observations : Int
}
///|
pub fn FairnessMonitor::new() -> FairnessMonitor {
{ groups: {}, observations: 0 }
}
///|
pub fn FairnessMonitor::observe(
self : FairnessMonitor,
group : String,
prediction : Double,
label : Double,
threshold? : Double = 0.5,
) -> Unit {
let metric = self.groups.get(group).unwrap_or(GroupMetric::new(group))
metric.observe(prediction, label, threshold~)
self.groups[group] = metric
self.observations += 1
}
///|
pub fn FairnessMonitor::group(
self : FairnessMonitor,
group : String,
) -> GroupMetric? {
self.groups.get(group)
}
///|
pub fn FairnessMonitor::groups(self : FairnessMonitor) -> Array[String] {
self.groups.keys().to_array()
}
///|
pub fn FairnessMonitor::observations(self : FairnessMonitor) -> Int {
self.observations
}
///|
pub fn FairnessMonitor::max_positive_rate_gap(self : FairnessMonitor) -> Double {
let names = self.groups.keys().to_array()
if names.length() < 2 {
0.0
} else {
let mut minimum = 1.0
let mut maximum = 0.0
for name in names {
let rate = self.groups
.get(name)
.map(metric => metric.positive_rate())
.unwrap_or(0.0)
if rate < minimum {
minimum = rate
}
if rate > maximum {
maximum = rate
}
}
maximum - minimum
}
}
///|
pub fn FairnessMonitor::max_recall_gap(self : FairnessMonitor) -> Double {
let names = self.groups.keys().to_array()
if names.length() < 2 {
0.0
} else {
let mut minimum = 1.0
let mut maximum = 0.0
for name in names {
let recall = self.groups
.get(name)
.map(metric => metric.recall())
.unwrap_or(0.0)
if recall < minimum {
minimum = recall
}
if recall > maximum {
maximum = recall
}
}
maximum - minimum
}
}
///|
pub fn FairnessMonitor::reset(self : FairnessMonitor) -> Unit {
self.groups.clear()
self.observations = 0
}