use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use crate::metrics::entropy::{dmi, entropy};
use rand::{RngExt, SeedableRng, rngs::StdRng};
const N_RESAMPLES: usize = 3;
fn data_seed<T: Hash>(x_seq: &[T], y_seq: &[T], seed: u64, frac_idx: u64, resample: u64) -> u64 {
let mut hasher = DefaultHasher::new();
seed.hash(&mut hasher);
frac_idx.hash(&mut hasher);
resample.hash(&mut hasher);
x_seq.hash(&mut hasher); y_seq.hash(&mut hasher);
hasher.finish()
}
fn subsample_indices(n: usize, m: usize, seed: u64) -> Vec<usize> {
let mut rng = StdRng::seed_from_u64(seed);
let mut idx: Vec<usize> = (0..n).collect();
for i in 0..m {
let j = rng.random_range(i..n);
idx.swap(i, j);
}
idx.truncate(m);
idx
}
fn subsampled_mi<T: Eq + Hash + Clone>(x_seq: &[T], y_seq: &[T], m: usize, seed: u64) -> f64 {
let indices = subsample_indices(x_seq.len(), m, seed);
let sub_x: Vec<T> = indices.iter().map(|&i| x_seq[i].clone()).collect();
let sub_y: Vec<T> = indices.iter().map(|&i| y_seq[i].clone()).collect();
dmi(&sub_x, &sub_y)
}
pub fn dmi_qe<T: Eq + Hash + Clone>(x_seq: &[T], y_seq: &[T], seed: u64) -> f64 {
let n = x_seq.len();
if n < 4 || n != y_seq.len() {
return 0.0;
}
let fractions = [1.0, 0.5, 0.25];
let mut estimates: Vec<(f64, f64)> = Vec::new();
for (frac_idx, &frac) in fractions.iter().enumerate() {
let m = (n as f64 * frac) as usize;
if m < 2 {
continue;
}
let estimate = if m == n {
dmi(x_seq, y_seq)
} else {
let mut sum = 0.0;
for r in 0..N_RESAMPLES {
let s = data_seed(x_seq, y_seq, seed, frac_idx as u64, r as u64);
sum += subsampled_mi(x_seq, y_seq, m, s);
}
sum / N_RESAMPLES as f64
};
estimates.push((1.0 / m as f64, estimate));
}
if estimates.len() < 2 {
return dmi(x_seq, y_seq); }
let n_pts = estimates.len() as f64;
let sum_x: f64 = estimates.iter().map(|(x, _)| x).sum();
let sum_y: f64 = estimates.iter().map(|(_, y)| y).sum();
let sum_xy: f64 = estimates.iter().map(|(x, y)| x * y).sum();
let sum_xx: f64 = estimates.iter().map(|(x, _)| x * x).sum();
let denom = n_pts * sum_xx - sum_x * sum_x;
if denom.abs() < 1e-12 {
return dmi(x_seq, y_seq); }
let slope = (n_pts * sum_xy - sum_x * sum_y) / denom;
let intercept = (sum_y - slope * sum_x) / n_pts;
intercept.max(0.0)
}
pub fn nmi_qe<T: Eq + Hash + Clone>(x_seq: &[T], y_seq: &[T], seed: u64) -> f64 {
let mi = dmi_qe(x_seq, y_seq, seed);
if mi == 0.0 {
return 0.0;
}
let h_x = entropy(x_seq);
let h_y = entropy(y_seq);
if h_x == 0.0 || h_y == 0.0 {
return 0.0;
}
(mi / (h_x * h_y).sqrt()).clamp(0.0, 1.0)
}
#[derive(Clone, Copy, Debug)]
pub struct QEShuffleConfig {
pub n_shuffles: usize,
pub seed: u64,
}
impl QEShuffleConfig {
pub fn new(n_shuffles: usize, seed: u64) -> Self {
Self { n_shuffles, seed }
}
}
pub fn shuffle_corrected_qe<T: Eq + Hash + Clone>(
x_seq: &[T],
y_seq: &[T],
config: &QEShuffleConfig,
) -> f64 {
if x_seq.len() < 4 || y_seq.len() < 4 {
return 0.0;
}
let nmi_obs = nmi_qe(x_seq, y_seq, config.seed);
if nmi_obs == 0.0 {
return 0.0;
}
let mut rng = StdRng::seed_from_u64(config.seed);
let mut y_shuffled: Vec<T> = y_seq.to_vec();
let mut nmi_shuffles = Vec::with_capacity(config.n_shuffles);
for _ in 0..config.n_shuffles {
for i in (1..y_shuffled.len()).rev() {
let j = rng.random_range(0..=i);
y_shuffled.swap(i, j);
}
nmi_shuffles.push(nmi_qe(x_seq, &y_shuffled, config.seed));
}
let mean_shuffle: f64 = nmi_shuffles.iter().sum::<f64>() / config.n_shuffles as f64;
(nmi_obs - mean_shuffle).clamp(0.0, 1.0)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_data_seed_deterministic_for_identical_input() {
let x = vec![1u8, 2, 3, 4, 5];
let y = vec![5u8, 4, 3, 2, 1];
let s1 = data_seed(&x, &y, 42, 0, 0);
let s2 = data_seed(&x, &y, 42, 0, 0);
assert_eq!(s1, s2, "identical input must give identical seed");
}
#[test]
fn test_data_seed_differs_for_different_content() {
let x = vec![1u8, 2, 3, 4, 5, 6, 7, 8];
let y1 = vec![1u8, 1, 1, 1, 2, 2, 2, 2];
let y2 = vec![2u8, 2, 2, 2, 1, 1, 1, 1];
let s1 = data_seed(&x, &y1, 42, 0, 0);
let s2 = data_seed(&x, &y2, 42, 0, 0);
assert_ne!(
s1, s2,
"different data content must not collapse to the same subsampling seed"
);
}
#[test]
fn test_data_seed_differs_across_resamples() {
let x = vec![1u8, 2, 3, 4, 5, 6];
let y = vec![6u8, 5, 4, 3, 2, 1];
let s0 = data_seed(&x, &y, 42, 0, 0);
let s1 = data_seed(&x, &y, 42, 0, 1);
assert_ne!(
s0, s1,
"different resample indices should draw different subsamples"
);
}
#[test]
fn test_subsample_indices_size_and_range() {
let idx = subsample_indices(20, 7, 123);
assert_eq!(idx.len(), 7);
assert!(idx.iter().all(|&i| i < 20));
let mut sorted = idx.clone();
sorted.sort_unstable();
sorted.dedup();
assert_eq!(
sorted.len(),
7,
"indices must be distinct (sampling without replacement)"
);
}
#[test]
fn test_dmi_qe_reproducible() {
let x: Vec<u8> = (0..80).map(|i| (i % 5) as u8).collect();
let y: Vec<u8> = (0..80).map(|i| ((i + 2) % 5) as u8).collect();
let a = dmi_qe(&x, &y, 7);
let b = dmi_qe(&x, &y, 7);
assert!(
(a - b).abs() < 1e-12,
"same input and seed must reproduce exactly"
);
}
#[test]
fn test_dmi_qe_independent_across_universes_at_same_delta() {
let shared_seed = 99;
let x: Vec<u8> = (0..60).map(|i| (i % 7) as u8).collect();
let y_universe_a: Vec<u8> = (0..60).map(|i| ((i * 3) % 7) as u8).collect();
let y_universe_b: Vec<u8> = (0..60).map(|i| ((i * 5 + 1) % 7) as u8).collect();
let sa = data_seed(&x, &y_universe_a, shared_seed, 1, 0);
let sb = data_seed(&x, &y_universe_b, shared_seed, 1, 0);
assert_ne!(sa, sb);
let _ = dmi_qe(&x, &y_universe_a, shared_seed);
let _ = dmi_qe(&x, &y_universe_b, shared_seed);
}
#[test]
fn test_qe_reduces_bias() {
let seed: u64 = 42;
let mut rng = StdRng::seed_from_u64(seed);
let x: Vec<u8> = (0..100).map(|_| rng.random_range(0..=7)).collect();
let y: Vec<u8> = (0..100).map(|_| rng.random_range(0..=7)).collect();
let plugin = dmi(&x, &y);
let qe = dmi_qe(&x, &y, seed);
assert!(qe >= 0.0, "QE should be non-negative, got {:.4}", qe);
assert!(
qe <= plugin + 0.1,
"QE ({:.4}) should not greatly exceed plugin ({:.4}) for independent data",
qe,
plugin
);
}
#[test]
fn test_qe_preserves_signal() {
let seed: u64 = 42;
let x = vec![0, 1, 0, 1, 0, 1, 0, 1, 0, 1];
let qe = dmi_qe(&x, &x, seed);
assert!(
qe > 0.3,
"QE should preserve MI for deterministic sequences, got {:.4}",
qe
);
}
#[test]
fn test_nmi_qe_bounded() {
let seed: u64 = 42;
let x = vec![0, 1, 0, 1, 0, 1];
let n = nmi_qe(&x, &x, seed);
assert!(n >= 0.0 && n <= 1.0, "NMI should be in [0,1], got {:.4}", n);
}
}