use nalgebra::{DMatrix, DVector};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum KnockoffS {
Equicorrelated,
Mvr,
Me,
}
pub fn knockoff_s(sigma: &DMatrix<f64>, method: KnockoffS) -> DVector<f64> {
match method {
KnockoffS::Equicorrelated => knockoff_s_equicorrelated(sigma),
KnockoffS::Mvr => knockoff_s_mvr(sigma),
KnockoffS::Me => knockoff_s_me(sigma),
}
}
pub fn knockoff_s_equicorrelated(sigma: &DMatrix<f64>) -> DVector<f64> {
let p = sigma.nrows();
if p == 0 {
return DVector::zeros(0);
}
let lambda_min = min_eig(sigma);
let s = (2.0 * lambda_min).clamp(0.0, 1.0);
DVector::from_element(p, s)
}
pub fn knockoff_s_mvr(sigma: &DMatrix<f64>) -> DVector<f64> {
solve_coordinate(sigma, Objective::Mvr)
}
pub fn knockoff_s_me(sigma: &DMatrix<f64>) -> DVector<f64> {
solve_coordinate(sigma, Objective::Me)
}
#[derive(Clone, Copy)]
enum Objective {
Mvr,
Me,
}
fn solve_coordinate(sigma: &DMatrix<f64>, obj: Objective) -> DVector<f64> {
let p = sigma.nrows();
if p == 0 {
return DVector::zeros(0);
}
let two_sigma = sigma * 2.0;
let lambda_min = min_eig(sigma);
if lambda_min <= 1e-10 {
return knockoff_s_equicorrelated(sigma);
}
let s0 = (2.0 * lambda_min).clamp(1e-6, 1.0) * 0.5;
let mut s = DVector::from_element(p, s0);
const MAX_ITER: usize = 50;
const TOL: f64 = 1e-8;
for _ in 0..MAX_ITER {
let mut m = two_sigma.clone();
for j in 0..p {
m[(j, j)] -= s[j];
}
let mut minv = match m.try_inverse() {
Some(inv) => inv,
None => break, };
let mut max_delta = 0.0f64;
for j in 0..p {
let m_jj = minv[(j, j)];
if m_jj.is_nan() || m_jj <= 1e-12 {
continue;
}
let s_old = s[j];
let s_target = match obj {
Objective::Me => (1.0 + m_jj * s_old) / (2.0 * m_jj),
Objective::Mvr => {
let col = minv.column(j);
let c_j = col.dot(&col);
(1.0 + m_jj * s_old) / (c_j.sqrt() + m_jj)
}
};
let mut delta = s_target - s_old;
let max_step = 0.99 / m_jj;
if delta > max_step {
delta = max_step;
}
if s_old + delta < 1e-8 {
delta = 1e-8 - s_old;
}
if delta.abs() < 1e-15 {
continue;
}
let denom = 1.0 - delta * m_jj;
if denom <= 1e-12 {
continue;
}
let u = minv.column(j).into_owned();
minv.ger(delta / denom, &u, &u, 1.0);
s[j] = s_old + delta;
max_delta = max_delta.max(delta.abs());
}
if max_delta < TOL {
break;
}
}
s
}
fn min_eig(sigma: &DMatrix<f64>) -> f64 {
sigma
.symmetric_eigenvalues()
.iter()
.cloned()
.fold(f64::INFINITY, f64::min)
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::SmallRng;
use rand::SeedableRng;
use rand_distr::{Distribution, StandardNormal};
fn random_corr(p: usize, k: usize, ridge: f64, seed: u64) -> DMatrix<f64> {
let mut rng = SmallRng::seed_from_u64(seed);
let f = DMatrix::<f64>::from_fn(p, k, |_, _| StandardNormal.sample(&mut rng));
let mut cov = &f * f.transpose();
for i in 0..p {
cov[(i, i)] += ridge;
}
let d: Vec<f64> = (0..p).map(|i| cov[(i, i)].sqrt()).collect();
DMatrix::from_fn(p, p, |i, j| cov[(i, j)] / (d[i] * d[j]))
}
fn min_eig_2sigma_minus_d(sigma: &DMatrix<f64>, s: &DVector<f64>) -> f64 {
let p = sigma.nrows();
let mut m = sigma * 2.0;
for j in 0..p {
m[(j, j)] -= s[j];
}
min_eig(&m)
}
#[test]
fn s_vectors_are_feasible() {
for &seed in &[1u64, 2, 3] {
let sigma = random_corr(40, 6, 0.2, seed);
for method in [KnockoffS::Mvr, KnockoffS::Me, KnockoffS::Equicorrelated] {
let s = knockoff_s(&sigma, method);
assert!(s.iter().all(|&v| v > -1e-9), "{method:?}: negative s");
let lam = min_eig_2sigma_minus_d(&sigma, &s);
assert!(
lam > -1e-6,
"{method:?}: 2Σ−D not PSD (λ_min={lam:.3e}) seed={seed}"
);
}
}
}
#[test]
fn mvr_beats_equicorrelated_objective() {
let sigma = random_corr(50, 8, 0.1, 7);
let obj = |s: &DVector<f64>| -> f64 {
let p = sigma.nrows();
let mut m = &sigma * 2.0;
for j in 0..p {
m[(j, j)] -= s[j];
}
let minv = m.try_inverse().unwrap();
minv.diagonal().sum() + s.iter().map(|&v| 1.0 / v).sum::<f64>()
};
let s_mvr = knockoff_s_mvr(&sigma);
let s_equi = knockoff_s_equicorrelated(&sigma);
assert!(
obj(&s_mvr) < obj(&s_equi),
"MVR objective {:.4} not below equicorrelated {:.4}",
obj(&s_mvr),
obj(&s_equi)
);
}
#[test]
fn mvr_outpowers_equicorrelated_with_tight_clusters() {
let p = 20;
let mut sigma = DMatrix::<f64>::identity(p, p);
for &(a, b) in &[(0usize, 1usize), (2, 3)] {
sigma[(a, b)] = 0.985;
sigma[(b, a)] = 0.985;
}
let s_equi = knockoff_s_equicorrelated(&sigma);
let s_mvr = knockoff_s_mvr(&sigma);
assert!(
min_eig_2sigma_minus_d(&sigma, &s_mvr) > -1e-6,
"MVR infeasible"
);
assert!(s_equi[0] < 0.05, "equicorrelated s = {}", s_equi[0]);
let mean_indep_mvr = (4..p).map(|j| s_mvr[j]).sum::<f64>() / (p - 4) as f64;
assert!(
mean_indep_mvr > 0.7,
"MVR independent-feature s too small: {mean_indep_mvr:.3}"
);
assert!(
s_mvr.mean() > 5.0 * s_equi.mean(),
"MVR mean {:.3} not >> equicorrelated {:.3}",
s_mvr.mean(),
s_equi.mean()
);
}
}