pub fn loss_recovered(l_clean: f64, l_recon: f64, l_ablate: f64) -> f64 {
const EPS: f64 = 1e-12;
let denom = l_ablate - l_clean;
if denom.abs() < EPS {
return f64::NAN;
}
(l_ablate - l_recon) / denom
}
pub fn r2_score(clean: &[f64], approx: &[f64], n_rows: usize, n_cols: usize) -> f64 {
assert_eq!(clean.len(), n_rows * n_cols, "clean shape mismatch");
assert_eq!(approx.len(), n_rows * n_cols, "approx shape mismatch");
if n_rows == 0 || n_cols == 0 {
return 0.0;
}
let mut col_mean = vec![0.0_f64; n_cols];
for row in 0..n_rows {
let base = row * n_cols;
for col in 0..n_cols {
col_mean[col] += clean[base + col];
}
}
let inv_rows = 1.0 / n_rows as f64;
for value in col_mean.iter_mut() {
*value *= inv_rows;
}
let mut rss = 0.0_f64;
let mut tss = 0.0_f64;
for row in 0..n_rows {
let base = row * n_cols;
for col in 0..n_cols {
let c = clean[base + col];
let residual = c - approx[base + col];
rss += residual * residual;
let centered = c - col_mean[col];
tss += centered * centered;
}
}
if tss > 0.0 {
1.0 - rss / tss
} else {
0.0
}
}
pub fn kl_categorical_rows(
clean_logprobs: &[f64],
other_logprobs: &[f64],
n_rows: usize,
n_cols: usize,
) -> f64 {
assert_eq!(clean_logprobs.len(), n_rows * n_cols, "clean logprobs shape mismatch");
assert_eq!(other_logprobs.len(), n_rows * n_cols, "other logprobs shape mismatch");
if n_rows == 0 {
return 0.0;
}
let mut total = 0.0_f64;
for row in 0..n_rows {
let base = row * n_cols;
let mut row_kl = 0.0_f64;
for col in 0..n_cols {
let logp_clean = clean_logprobs[base + col];
let p = logp_clean.exp();
row_kl += p * (logp_clean - other_logprobs[base + col]);
}
total += row_kl;
}
total / n_rows as f64
}
pub fn distortion_floor_r2(
r2s: &[f64],
loss_recovered: &[f64],
tol_frac: f64,
) -> Option<f64> {
assert_eq!(r2s.len(), loss_recovered.len(), "parallel arrays length mismatch");
if r2s.is_empty() {
return None;
}
let mut order: Vec<usize> = (0..r2s.len()).collect();
order.sort_by(|&a, &b| {
r2s[b]
.partial_cmp(&r2s[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
let plateau = loss_recovered[order[0]];
let threshold = plateau - tol_frac * plateau.abs();
for &idx in &order {
if loss_recovered[idx] < threshold {
return Some(r2s[idx]);
}
}
Some(r2s[order[order.len() - 1]])
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn loss_recovered_endpoints_and_degenerate() {
assert!((loss_recovered(1.0, 1.0, 3.0) - 1.0).abs() < 1e-15);
assert!((loss_recovered(1.0, 3.0, 3.0) - 0.0).abs() < 1e-15);
assert!((loss_recovered(1.0, 2.0, 3.0) - 0.5).abs() < 1e-15);
assert!(loss_recovered(2.0, 2.0, 2.0).is_nan());
}
#[test]
fn r2_perfect_and_mean_baselines() {
let clean = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
assert!((r2_score(&clean, &clean, 3, 2) - 1.0).abs() < 1e-15);
let mean_pred = [3.0, 4.0, 3.0, 4.0, 3.0, 4.0];
assert!(r2_score(&clean, &mean_pred, 3, 2).abs() < 1e-15);
}
#[test]
fn r2_constant_clean_is_zero() {
let clean = [7.0, 7.0, 7.0, 7.0];
let approx = [1.0, 2.0, 3.0, 4.0];
assert_eq!(r2_score(&clean, &approx, 2, 2), 0.0);
}
#[test]
fn kl_identical_is_zero_and_positive_otherwise() {
let p = [0.5_f64, 0.25, 0.25, 0.1, 0.6, 0.3];
let logp: Vec<f64> = p.iter().map(|v| v.ln()).collect();
assert!(kl_categorical_rows(&logp, &logp, 2, 3).abs() < 1e-15);
let uniform = vec![(1.0_f64 / 3.0).ln(); 6];
let kl = kl_categorical_rows(&logp, &uniform, 2, 3);
let mut expect = 0.0;
for row in 0..2 {
let mut h = 0.0;
for col in 0..3 {
let pv = p[row * 3 + col];
h += pv * pv.ln();
}
expect += h + (3.0_f64).ln();
}
expect /= 2.0;
assert!((kl - expect).abs() < 1e-15);
}
#[test]
fn floor_flat_plateau_returns_coarsest() {
let r2s = [0.99, 0.90, 0.50];
let lr = [1.0, 1.0, 1.0];
assert_eq!(distortion_floor_r2(&r2s, &lr, 0.05), Some(0.50));
}
#[test]
fn floor_detects_drop_point() {
let r2s = [0.99, 0.95, 0.60, 0.30];
let lr = [1.0, 0.99, 0.80, 0.40];
assert_eq!(distortion_floor_r2(&r2s, &lr, 0.05), Some(0.60));
}
#[test]
fn floor_empty_is_none() {
assert_eq!(distortion_floor_r2(&[], &[], 0.05), None);
}
}