#[must_use]
pub fn conformal_recall_lower_bound(recalls: &[f64], alpha: f64) -> f64 {
if !(0.0..1.0).contains(&alpha) {
return 0.0;
}
let mut sorted: Vec<f64> = recalls.iter().copied().filter(|r| r.is_finite()).collect();
let n = sorted.len();
if n == 0 {
return 0.0;
}
#[allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
let rank = (alpha * (n as f64 + 1.0)).floor() as usize;
if rank == 0 {
return 0.0;
}
sorted.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let idx = (rank - 1).min(n - 1);
sorted[idx].clamp(0.0, 1.0)
}
#[must_use]
pub fn mean_recall_lower_bound(recalls: &[f64], delta: f64) -> f64 {
if !(0.0..1.0).contains(&delta) {
return 0.0;
}
let finite: Vec<f64> = recalls.iter().copied().filter(|r| r.is_finite()).collect();
let n = finite.len();
if n == 0 {
return 0.0;
}
#[allow(clippy::cast_precision_loss)]
let n_f = n as f64;
let mean = finite.iter().sum::<f64>() / n_f;
let radius = ((1.0 / delta).ln() / (2.0 * n_f)).sqrt();
(mean - radius).clamp(0.0, 1.0)
}
#[must_use]
pub fn mean_recall_lower_bound_bernstein(recalls: &[f64], delta: f64) -> f64 {
if !(0.0..1.0).contains(&delta) {
return 0.0;
}
let finite: Vec<f64> = recalls.iter().copied().filter(|r| r.is_finite()).collect();
let n = finite.len();
if n < 2 {
return 0.0;
}
#[allow(clippy::cast_precision_loss)]
let n_f = n as f64;
let mean = finite.iter().sum::<f64>() / n_f;
let var = finite.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / (n_f - 1.0);
let ln = (2.0 / delta).ln();
let bound = mean - (2.0 * var * ln / n_f).sqrt() - 7.0 * ln / (3.0 * (n_f - 1.0));
bound.clamp(0.0, 1.0)
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CertifiedEf {
pub ef_search: usize,
pub certified_recall: f64,
pub meets_target: bool,
}
#[must_use]
pub fn certified_min_ef(
calibration: &[(usize, Vec<f64>)],
target: f64,
alpha: f64,
) -> Option<CertifiedEf> {
let mut best: Option<CertifiedEf> = None;
let mut sorted: Vec<&(usize, Vec<f64>)> = calibration.iter().collect();
sorted.sort_by_key(|(ef, _)| *ef);
for (ef, recalls) in sorted {
let bound = conformal_recall_lower_bound(recalls, alpha);
let candidate = CertifiedEf {
ef_search: *ef,
certified_recall: bound,
meets_target: bound >= target,
};
if candidate.meets_target {
return Some(candidate);
}
let better = best.is_none_or(|b| bound > b.certified_recall);
if better {
best = Some(candidate);
}
}
best
}
#[must_use]
pub fn certified_min_ef_mean(
calibration: &[(usize, Vec<f64>)],
target: f64,
delta: f64,
) -> Option<CertifiedEf> {
let mut best: Option<CertifiedEf> = None;
let mut sorted: Vec<&(usize, Vec<f64>)> = calibration.iter().collect();
sorted.sort_by_key(|(ef, _)| *ef);
for (ef, recalls) in sorted {
let bound = mean_recall_lower_bound_bernstein(recalls, delta);
let candidate = CertifiedEf {
ef_search: *ef,
certified_recall: bound,
meets_target: bound >= target,
};
if candidate.meets_target {
return Some(candidate);
}
let better = best.is_none_or(|b| bound > b.certified_recall);
if better {
best = Some(candidate);
}
}
best
}
#[derive(Debug, Clone, PartialEq)]
pub struct EfCalibration {
pub chosen: CertifiedEf,
pub sweep: Vec<CertifiedEf>,
}
pub fn calibrate_certified_ef(
candidate_efs: &[usize],
mut measure_recall: impl FnMut(usize) -> Vec<f64>,
target: f64,
alpha: f64,
) -> Option<EfCalibration> {
let mut efs: Vec<usize> = candidate_efs.to_vec();
efs.sort_unstable();
efs.dedup();
let mut sweep: Vec<CertifiedEf> = Vec::with_capacity(efs.len());
let mut best: Option<CertifiedEf> = None;
for ef in efs {
let recalls = measure_recall(ef);
let bound = conformal_recall_lower_bound(&recalls, alpha);
let candidate = CertifiedEf {
ef_search: ef,
certified_recall: bound,
meets_target: bound >= target,
};
sweep.push(candidate);
if candidate.meets_target {
return Some(EfCalibration {
chosen: candidate,
sweep,
});
}
let better = best.is_none_or(|b| bound > b.certified_recall);
if better {
best = Some(candidate);
}
}
best.map(|chosen| EfCalibration { chosen, sweep })
}
#[cfg(test)]
mod tests {
use super::*;
struct Lcg(u64);
impl Lcg {
fn next_u64(&mut self) -> u64 {
self.0 = self
.0
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
self.0
}
#[allow(clippy::cast_precision_loss)]
fn unit(&mut self) -> f64 {
((self.next_u64() >> 11) as f64) / ((1u64 << 53) as f64)
}
}
#[test]
#[allow(clippy::float_cmp)] fn conformal_bound_is_trivial_when_sample_too_small() {
let recalls = vec![0.9; 10];
assert_eq!(conformal_recall_lower_bound(&recalls, 0.05), 0.0);
let recalls = vec![0.9; 19];
assert!(conformal_recall_lower_bound(&recalls, 0.05) > 0.0);
assert_eq!(conformal_recall_lower_bound(&[], 0.1), 0.0);
assert_eq!(conformal_recall_lower_bound(&[0.9, 0.8], 0.0), 0.0);
assert_eq!(conformal_recall_lower_bound(&[0.9, 0.8], 1.0), 0.0);
}
#[test]
fn conformal_bound_recovers_the_order_statistic() {
let recalls: Vec<f64> = (1..=99).map(|i| f64::from(i) / 100.0).collect();
let bound = conformal_recall_lower_bound(&recalls, 0.10);
assert!((bound - 0.10).abs() < 1e-9, "got {bound}");
}
#[test]
fn conformal_bound_is_monotone_in_alpha() {
let mut lcg = Lcg(42);
let recalls: Vec<f64> = (0..500).map(|_| lcg.unit()).collect();
let strong = conformal_recall_lower_bound(&recalls, 0.01);
let weak = conformal_recall_lower_bound(&recalls, 0.20);
assert!(strong <= weak, "strong={strong} weak={weak}");
}
#[test]
fn conformal_bound_has_valid_finite_sample_coverage() {
let alpha = 0.10;
let n_cal = 200;
let trials = 4000;
let mut lcg = Lcg(0x5eed);
let mut misses = 0usize;
for _ in 0..trials {
let draw = |lcg: &mut Lcg| -> f64 {
let a = lcg.unit();
let b = lcg.unit();
1.0 - (a * b) * 0.4
};
let cal: Vec<f64> = (0..n_cal).map(|_| draw(&mut lcg)).collect();
let bound = conformal_recall_lower_bound(&cal, alpha);
let fresh = draw(&mut lcg);
if fresh < bound {
misses += 1;
}
}
#[allow(clippy::cast_precision_loss)]
let miss_rate = misses as f64 / f64::from(trials);
assert!(
miss_rate <= alpha + 0.02,
"conformal coverage violated: miss_rate={miss_rate:.4} > alpha={alpha}"
);
}
#[test]
fn mean_bound_lower_bounds_and_tightens_with_n() {
let small = mean_recall_lower_bound(&vec![0.95; 30], 0.05);
let large = mean_recall_lower_bound(&vec![0.95; 3000], 0.05);
assert!(small < large, "small={small} large={large}");
assert!(large <= 0.95 && large > 0.90, "large={large}");
assert!(small >= 0.0);
}
#[test]
fn mean_bound_coverage_holds() {
let delta = 0.05;
let n = 300;
let trials = 3000;
let mut lcg = Lcg(0x00c0_ffee);
let true_mean = 0.9;
let mut misses = 0usize;
for _ in 0..trials {
let cal: Vec<f64> = (0..n)
.map(|_| if lcg.unit() < true_mean { 1.0 } else { 0.0 })
.collect();
if true_mean < mean_recall_lower_bound(&cal, delta) {
misses += 1;
}
}
#[allow(clippy::cast_precision_loss)]
let miss_rate = misses as f64 / f64::from(trials);
assert!(
miss_rate <= delta,
"mean LCB coverage violated: {miss_rate:.4} > {delta}"
);
}
#[test]
fn bernstein_mean_bound_coverage_holds() {
let delta = 0.05;
let n = 300;
let trials = 3000;
let mut lcg = Lcg(0x00b0_bb1e);
let true_mean = 0.97;
let mut misses = 0usize;
for _ in 0..trials {
let cal: Vec<f64> = (0..n)
.map(|_| if lcg.unit() < 0.7 { 1.0 } else { 0.9 })
.collect();
if true_mean < mean_recall_lower_bound_bernstein(&cal, delta) {
misses += 1;
}
}
#[allow(clippy::cast_precision_loss)]
let miss_rate = misses as f64 / f64::from(trials);
assert!(
miss_rate <= delta,
"Bernstein LCB coverage violated: {miss_rate:.4} > {delta}"
);
}
#[test]
fn bernstein_is_tighter_than_hoeffding_on_low_variance_recall() {
let mut recalls = vec![1.0; 950];
recalls.extend((0..50).map(|_| 0.8)); let delta = 0.05;
let hoeffding = mean_recall_lower_bound(&recalls, delta);
let bernstein = mean_recall_lower_bound_bernstein(&recalls, delta);
assert!(
bernstein > hoeffding,
"expected Bernstein tighter: bernstein={bernstein:.4} hoeffding={hoeffding:.4}"
);
}
#[test]
fn mean_mode_certifies_a_cheaper_ef_than_the_per_query_tail_mode() {
let n = 1000;
let ef40: Vec<f64> = (0..n)
.map(|i| if i < 750 { 1.0 } else { 0.95 }) .collect();
let ef100: Vec<f64> = (0..n).map(|i| if i < 940 { 1.0 } else { 0.94 }).collect();
let calibration = vec![(40usize, ef40.clone()), (100usize, ef100)];
let mean_choice = certified_min_ef_mean(&calibration, 0.95, 0.05).unwrap();
assert!(
mean_choice.meets_target,
"mean bound should certify at 0.95"
);
assert_eq!(
mean_choice.ef_search, 40,
"mean mode certifies the cheaper ef=40"
);
assert!(mean_choice.certified_recall >= 0.95);
assert!(mean_recall_lower_bound_bernstein(&ef40, 0.05) >= 0.95);
assert!(mean_recall_lower_bound(&ef40, 0.05) < 0.95);
}
#[test]
fn certified_min_ef_picks_the_cheapest_certified_option() {
let calibration = vec![
(40usize, vec![0.80; 300]),
(100usize, vec![0.99; 300]),
(200usize, vec![0.999; 300]),
];
let choice = certified_min_ef(&calibration, 0.95, 0.05).unwrap();
assert_eq!(choice.ef_search, 100);
assert!(choice.meets_target);
assert!(choice.certified_recall >= 0.95);
}
#[test]
fn certified_min_ef_falls_back_to_best_when_none_meets_target() {
let calibration = vec![(40usize, vec![0.70; 300]), (100usize, vec![0.85; 300])];
let choice = certified_min_ef(&calibration, 0.99, 0.05).unwrap();
assert!(!choice.meets_target);
assert_eq!(choice.ef_search, 100); assert!(certified_min_ef(&[], 0.9, 0.05).is_none());
}
#[test]
fn calibrate_short_circuits_at_the_cheapest_certified_ef() {
let mut measured: Vec<usize> = Vec::new();
let recalls_for = |ef: usize| -> Vec<f64> {
let r = match ef {
20 => 0.80,
40 => 0.99,
_ => 0.999,
};
vec![r; 300]
};
let cal = calibrate_certified_ef(
&[200, 40, 20, 100, 40],
|ef| {
measured.push(ef);
recalls_for(ef)
},
0.95,
0.05,
)
.unwrap();
assert!(cal.chosen.meets_target);
assert_eq!(cal.chosen.ef_search, 40, "cheapest certified ef");
assert_eq!(
measured,
vec![20, 40],
"must stop measuring once ef=40 certifies"
);
assert_eq!(cal.sweep.len(), 2);
}
#[test]
fn calibrate_falls_back_and_measures_all_when_none_certifies() {
let mut count = 0usize;
let cal = calibrate_certified_ef(
&[20, 40, 100],
|ef| {
count += 1;
let r = if ef >= 100 { 0.90 } else { 0.80 };
vec![r; 300]
},
0.99, 0.05,
)
.unwrap();
assert!(!cal.chosen.meets_target);
assert_eq!(cal.chosen.ef_search, 100, "best-certifiable fallback");
assert_eq!(count, 3, "no early stop when nothing certifies");
assert_eq!(cal.sweep.len(), 3);
let mut never = 0usize;
assert!(
calibrate_certified_ef(
&[],
|_| {
never += 1;
vec![1.0]
},
0.9,
0.05
)
.is_none()
);
assert_eq!(never, 0);
}
#[test]
fn certificate_catches_heuristic_overconfidence() {
let heuristic = 0.1_f64
.mul_add((100.0_f64 / 10.0).log2(), 0.9)
.clamp(0.0, 1.0);
assert!(
heuristic >= 0.999,
"heuristic claims ~perfect recall: {heuristic}"
);
let measured = vec![0.85; 500];
let certified = conformal_recall_lower_bound(&measured, 0.05);
assert!(
certified < 0.9,
"certificate should refuse the heuristic's optimism, got {certified}"
);
}
}