mod synthetic;
use big_o::Model;
use std::collections::BTreeMap;
const TRIALS: u64 = 25;
const DEEP_TRIALS: u64 = 200;
const NOISE: f64 = 0.05;
const REQUIRED: f64 = 0.95;
fn recovery_rates(trials: u64, sigma: f64) -> BTreeMap<&'static str, (f64, Vec<String>)> {
let mut rates = BTreeMap::new();
for model in MODELS {
let mut recovered = 0u64;
let mut attempted = 0u64;
let mut missed: Vec<String> = Vec::new();
for &share in &synthetic::OFFSET_SHARES {
for trial in 0..trials {
let data = synthetic::with_offset_share(model, trial, sigma, share);
attempted += 1;
match big_o::infer_complexity(&data) {
Ok(inference) if inference.best.model == model => recovered += 1,
Ok(inference) => missed.push(format!(
"offset {share}, trial {trial}: {} ({:?})",
inference.best, inference.best.model
)),
Err(e) => missed.push(format!("offset {share}, trial {trial}: {e}")),
}
}
}
missed.truncate(3);
rates.insert(
model.notation(),
(recovered as f64 / attempted as f64, missed),
);
}
rates
}
const MODELS: [Model; 8] = [
Model::Constant,
Model::Logarithmic,
Model::Linear,
Model::Linearithmic,
Model::Quadratic,
Model::Cubic,
Model::Polynomial,
Model::Exponential,
];
fn assert_recovers(trials: u64, sigma: f64, required: f64) {
let rates = recovery_rates(trials, sigma);
let failures: Vec<String> = rates
.iter()
.filter(|(_, (rate, _))| *rate < required)
.map(|(model, (rate, missed))| {
format!(" {model}: {:.0}% recovered, e.g. {missed:?}", rate * 100.0)
})
.collect();
assert!(
failures.is_empty(),
"at {:.0}% noise, {} of 8 models fell below {:.0}% over {trials} trials per constant-term share:\n{}",
sigma * 100.0,
failures.len(),
required * 100.0,
failures.join("\n")
);
}
#[test]
fn recovers_every_model_from_clean_data() {
assert_recovers(TRIALS, 0.0, 1.0);
}
#[test]
fn recovers_every_model_through_realistic_noise() {
assert_recovers(TRIALS, NOISE, REQUIRED);
}
#[test]
#[ignore = "wider sweep, run on a schedule rather than per commit"]
fn recovers_every_model_over_many_trials() {
assert_recovers(DEEP_TRIALS, 0.0, 1.0);
assert_recovers(DEEP_TRIALS, NOISE, REQUIRED);
}
#[test]
#[ignore = "diagnostic, prints the confusion matrix"]
fn report_confusion_matrix() {
for sigma in [0.0, NOISE] {
println!(
"\nnoise {:.0}%, {DEEP_TRIALS} trials per model per constant-term share {:?}",
sigma * 100.0,
synthetic::OFFSET_SHARES
);
for (model, (rate, missed)) in recovery_rates(DEEP_TRIALS, sigma) {
println!(" {model:<12} {:>5.1}% {missed:?}", rate * 100.0);
}
}
}