use crate::error::RillError;
use crate::sparse::SparseFeatures;
pub trait OnlineRegressor {
fn feature_count(&self) -> usize;
fn samples_seen(&self) -> u64;
fn predict(&self, features: &[f64]) -> Result<f64, RillError>;
fn learn(&mut self, features: &[f64], target: f64) -> Result<(), RillError>;
fn reset(&mut self);
}
pub trait OnlineBinaryClassifier {
fn feature_count(&self) -> usize;
fn samples_seen(&self) -> u64;
fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError>;
fn predict(&self, features: &[f64]) -> Result<bool, RillError> {
Ok(self.predict_proba(features)? >= 0.5)
}
fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError>;
fn reset(&mut self);
}
pub trait Transformer {
fn input_dim(&self) -> usize;
fn output_dim(&self) -> usize;
fn transform(&self, features: &[f64]) -> Result<Vec<f64>, RillError>;
fn update(&mut self, features: &[f64]) -> Result<(), RillError>;
fn samples_seen(&self) -> u64;
fn reset(&mut self);
}
pub trait Metric {
type Truth;
type Prediction;
fn update(&mut self, truth: Self::Truth, prediction: Self::Prediction)
-> Result<(), RillError>;
fn value(&self) -> Option<f64>;
fn samples_seen(&self) -> u64;
fn reset(&mut self);
}
pub trait OnlineStatistic {
fn update(&mut self, value: f64) -> Result<(), RillError>;
fn samples_seen(&self) -> u64;
fn reset(&mut self);
}
pub trait SparseRegressor {
fn samples_seen(&self) -> u64;
fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError>;
fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError>;
fn reset(&mut self);
}
pub trait SparseClassifier {
fn samples_seen(&self) -> u64;
fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError>;
fn predict(&self, features: &SparseFeatures) -> Result<bool, RillError> {
Ok(self.predict_proba(features)? >= 0.5)
}
fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError>;
fn reset(&mut self);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::metrics::{Accuracy, F1Score, LogLoss, Mae, Mse, Precision, R2, Recall, Rmse};
fn assert_samples_seen_contract<T: Metric>(
metric: &mut T,
n: u64,
truth: impl Fn(u64) -> T::Truth,
prediction: impl Fn(u64) -> T::Prediction,
) {
assert_eq!(metric.samples_seen(), 0, "fresh metric must start at 0");
for i in 0..n {
metric.update(truth(i), prediction(i)).unwrap_or_else(|e| {
panic!("update {i} must succeed in trait contract test: {e:?}")
});
assert_eq!(
metric.samples_seen(),
i + 1,
"samples_seen must equal successful update count after {} updates",
i + 1
);
}
assert_eq!(metric.samples_seen(), n);
metric.reset();
assert_eq!(
metric.samples_seen(),
0,
"reset must clear samples_seen for {}",
std::any::type_name::<T>()
);
}
#[test]
fn mae_samples_seen_equals_successful_updates() {
let mut m = Mae::new();
assert_samples_seen_contract(&mut m, 5, |i| i as f64, |i| (i as f64) + 0.5);
}
#[test]
fn mse_samples_seen_equals_successful_updates() {
let mut m = Mse::new();
assert_samples_seen_contract(&mut m, 5, |i| i as f64, |i| (i as f64) + 0.5);
}
#[test]
fn rmse_samples_seen_equals_successful_updates() {
let mut m = Rmse::new();
assert_samples_seen_contract(&mut m, 5, |i| i as f64, |i| (i as f64) + 0.5);
}
#[test]
fn r2_samples_seen_equals_successful_updates() {
let mut m = R2::new();
assert_samples_seen_contract(&mut m, 5, |i| (i as f64) + 1.0, |i| (i as f64) + 1.1);
}
#[test]
fn accuracy_samples_seen_equals_successful_updates() {
let mut m = Accuracy::default();
assert_samples_seen_contract(&mut m, 5, |i| i % 2 == 0, |i| i % 3 == 0);
}
#[test]
fn precision_samples_seen_equals_successful_updates() {
let mut m = Precision::default();
assert_samples_seen_contract(&mut m, 4, |i| i % 2 == 0, |i| i % 3 == 0);
}
#[test]
fn recall_samples_seen_equals_successful_updates() {
let mut m = Recall::default();
assert_samples_seen_contract(&mut m, 4, |i| i % 2 == 0, |i| i % 3 == 0);
}
#[test]
fn f1_samples_seen_equals_successful_updates() {
let mut m = F1Score::default();
assert_samples_seen_contract(&mut m, 4, |i| i % 2 == 0, |i| i % 3 == 0);
}
#[test]
fn log_loss_samples_seen_equals_successful_updates() {
let mut m = LogLoss::default();
assert_samples_seen_contract(&mut m, 5, |i| i % 2 == 0, |i| 0.3 + 0.1 * (i as f64));
}
#[test]
fn failed_update_does_not_advance_samples_seen() {
let mut m = Mae::new();
m.update(1.0, 2.0).unwrap();
assert_eq!(m.samples_seen(), 1);
assert!(m.update(f64::NAN, 1.0).is_err());
assert_eq!(m.samples_seen(), 1, "failed update must not advance count");
}
}