#![allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
clippy::cast_lossless,
clippy::needless_range_loop,
clippy::too_many_arguments,
clippy::similar_names,
clippy::many_single_char_names,
clippy::doc_markdown,
clippy::trivially_copy_pass_by_ref
)]
use std::collections::HashMap;
use antecedent_core::{CausalRng, ExecutionContext, KernelPolicy};
use antecedent_kernels::{shuffle, unbiased_index};
use super::parcorr::PartialCorrelation;
use super::types::{
CiBatchRequest, CiBatchResult, CiResult, CiWorkspace, ConditionalIndependenceTest,
PreparedCiTest, analytic_confidence_level,
};
use crate::special::{gamma_q, normal_ppf};
#[cfg(test)]
use super::types::{CiQuery, ConfidenceMethod, SignificanceMethod};
use crate::error::StatsError;
#[derive(Clone, Debug, Default)]
pub struct GSquared;
impl GSquared {
#[must_use]
pub fn new() -> Self {
Self
}
}
impl ConditionalIndependenceTest for GSquared {
fn test_batch(
&self,
prepared: &PreparedCiTest,
request: &CiBatchRequest<'_>,
workspace: &mut CiWorkspace,
ctx: &ExecutionContext,
) -> Result<CiBatchResult, StatsError> {
prepared.ensure_compatible(request)?;
let request = &prepared.bind_request(request);
let n = request.columns.first().map_or(0, |c| c.len());
if n == 0 {
return Err(StatsError::Shape { message: "no columns" });
}
let level = analytic_confidence_level(request.confidence);
let policy = &ctx.kernel_policy;
let mut results = Vec::with_capacity(request.queries.len());
for (qi, q) in request.queries.iter().enumerate() {
let z = &request.z_flat[q.z_start..q.z_start + q.z_len];
let (g, df) = g_squared_statistic(request.columns, q.x, q.y, z, n, workspace, policy)?;
let (p, ci) = match request.significance {
super::types::SignificanceMethod::Analytic => {
let p = chi2_sf(g, df);
let ci = level.map(|lv| analytic_gsquared_ci(g, df, lv));
(p, ci)
}
super::types::SignificanceMethod::BlockShuffle { replicates, block_size } => {
let n_perm = replicates.max(1) as usize;
let strata = gsq_strata(request.columns, z, n);
let mut y_perm = request.columns[q.y].to_vec();
let mut rng = ctx.rng.stream(0x65C0_u64.wrapping_add(qi as u64));
let mut null_ge = 0u32;
for _ in 0..n_perm {
if block_size > 1 && z.is_empty() {
block_shuffle_y(&mut y_perm, block_size, &mut rng);
} else {
for rows in &strata {
for i in (1..rows.len()).rev() {
let j = unbiased_index(&mut rng, i + 1);
y_perm.swap(rows[i], rows[j]);
}
}
}
let mut cols: Vec<&[f64]> = request.columns.to_vec();
cols[q.y] = &y_perm;
let (g_null, _) =
g_squared_statistic(&cols, q.x, q.y, z, n, workspace, policy)?;
if g_null >= g {
null_ge = null_ge.saturating_add(1);
}
}
let p = (1.0 + f64::from(null_ge)) / (1.0 + n_perm as f64);
(p, None)
}
};
results.push(CiResult { statistic: g, p_value: p, df, ci });
}
Ok(CiBatchResult { results })
}
}
fn analytic_gsquared_ci(g: f64, df: f64, level: f64) -> (f64, f64) {
let z = normal_ppf(0.5 + 0.5 * level.clamp(0.0, 1.0));
let se = (2.0 * df.max(1.0)).sqrt();
((g - z * se).max(0.0), g + z * se)
}
fn gsq_strata(columns: &[&[f64]], z: &[usize], n: usize) -> Vec<Vec<usize>> {
let mut strata: HashMap<u64, Vec<usize>> = HashMap::new();
for r in 0..n {
let key = if z.is_empty() {
0u64
} else {
let mut h = 0xcbf2_9ce4_8422_2325_u64;
for &zc in z {
let v = columns[zc][r].round() as i32;
h ^= u64::from(v as u32);
h = h.wrapping_mul(0x0100_0000_01b3);
}
h
};
strata.entry(key).or_default().push(r);
}
let mut keys: Vec<u64> = strata.keys().copied().collect();
keys.sort_unstable();
keys.into_iter().filter_map(|k| strata.remove(&k)).collect()
}
fn block_shuffle_y(y: &mut [f64], block_size: usize, rng: &mut CausalRng) {
let n = y.len();
let bs = block_size.max(1).min(n);
let n_blocks = n.div_ceil(bs);
let mut order: Vec<usize> = (0..n_blocks).collect();
shuffle(rng, &mut order);
let original = y.to_vec();
let mut dest = 0;
for &bi in &order {
let start = bi * bs;
let end = (start + bs).min(n);
let len = end - start;
y[dest..dest + len].copy_from_slice(&original[start..end]);
dest += len;
}
}
fn g_squared_statistic(
columns: &[&[f64]],
x: usize,
y: usize,
z: &[usize],
n: usize,
workspace: &mut CiWorkspace,
policy: &KernelPolicy,
) -> Result<(f64, f64), StatsError> {
let xi: Vec<i32> = columns[x].iter().map(|v| v.round() as i32).collect();
let yi: Vec<i32> = columns[y].iter().map(|v| v.round() as i32).collect();
let mut strata: HashMap<u64, Vec<usize>> = HashMap::new();
for r in 0..n {
let key = if z.is_empty() {
0u64
} else {
let mut h = 0xcbf2_9ce4_8422_2325_u64;
for &zc in z {
let v = columns[zc][r].round() as i32;
h ^= u64::from(v as u32);
h = h.wrapping_mul(0x0100_0000_01b3);
}
h
};
strata.entry(key).or_default().push(r);
}
let mut g_total = 0.0;
let mut df_total = 0.0;
let mut any = false;
for rows in strata.values() {
if rows.len() < 2 {
continue;
}
let (g, df) = g_squared_on_rows(&xi, &yi, rows, workspace, policy);
g_total += g;
df_total += df;
any = true;
}
if !any {
return Err(StatsError::Shape { message: "empty stratified contingency" });
}
Ok((g_total, df_total.max(1.0)))
}
fn g_squared_on_rows(
xi: &[i32],
yi: &[i32],
rows: &[usize],
workspace: &mut CiWorkspace,
policy: &KernelPolicy,
) -> (f64, f64) {
let mut levels_x: Vec<i32> = rows.iter().map(|&r| xi[r]).collect();
levels_x.sort_unstable();
levels_x.dedup();
let mut levels_y: Vec<i32> = rows.iter().map(|&r| yi[r]).collect();
levels_y.sort_unstable();
levels_y.dedup();
let lx = levels_x.len().max(1);
let ly = levels_y.len().max(1);
let need = lx * ly;
if workspace.shuffled.len() < need {
workspace.shuffled.resize(need, 0.0);
}
for v in &mut workspace.shuffled[..need] {
*v = 0.0;
}
let n_rows = rows.len();
if workspace.contingency_x_codes.len() < n_rows {
workspace.contingency_x_codes.resize(n_rows, 0);
}
if workspace.contingency_y_codes.len() < n_rows {
workspace.contingency_y_codes.resize(n_rows, 0);
}
for (i, &r) in rows.iter().enumerate() {
workspace.contingency_x_codes[i] = levels_x.binary_search(&xi[r]).unwrap_or(0) as u32;
workspace.contingency_y_codes[i] = levels_y.binary_search(&yi[r]).unwrap_or(0) as u32;
}
antecedent_kernels::accumulate_contingency(
policy,
&workspace.contingency_x_codes[..n_rows],
&workspace.contingency_y_codes[..n_rows],
&mut workspace.shuffled[..need],
ly,
);
let mut row_sum = vec![0.0; lx];
let mut col_sum = vec![0.0; ly];
let mut total = 0.0;
for i in 0..lx {
for j in 0..ly {
let o = workspace.shuffled[i * ly + j];
row_sum[i] += o;
col_sum[j] += o;
total += o;
}
}
if total < 1.0 {
return (0.0, 0.0);
}
let mut g = 0.0;
for i in 0..lx {
for j in 0..ly {
let o = workspace.shuffled[i * ly + j];
let e = row_sum[i] * col_sum[j] / total;
if o > 0.0 && e > 0.0 {
g += 2.0 * o * (o / e).ln();
}
}
}
let df = ((lx - 1) * (ly - 1)) as f64;
(g, df)
}
fn chi2_sf(x: f64, df: f64) -> f64 {
if x <= 0.0 {
return 1.0;
}
if df <= 0.0 {
return 0.0;
}
gamma_q(df * 0.5, x * 0.5)
}
#[derive(Clone, Debug, Default)]
pub struct RegressionCi {
inner: PartialCorrelation,
}
impl RegressionCi {
#[must_use]
pub fn new() -> Self {
Self { inner: PartialCorrelation::new() }
}
}
impl ConditionalIndependenceTest for RegressionCi {
fn test_batch(
&self,
prepared: &PreparedCiTest,
request: &CiBatchRequest<'_>,
workspace: &mut CiWorkspace,
ctx: &ExecutionContext,
) -> Result<CiBatchResult, StatsError> {
prepared.ensure_compatible(request)?;
let request = &prepared.bind_request(request);
self.inner.test_batch(prepared, request, workspace, ctx)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gsq_independent_high_p() {
let n = 400usize;
let x: Vec<f64> = (0..n).map(|i| (i % 2) as f64).collect();
let y: Vec<f64> = (0..n).map(|i| ((i / 2) % 2) as f64).collect();
let cols: [&[f64]; 2] = [&x, &y];
let queries = [CiQuery { x: 0, y: 1, z_start: 0, z_len: 0 }];
let req = CiBatchRequest {
columns: &cols,
queries: &queries,
z_flat: &[],
significance: SignificanceMethod::Analytic,
confidence: ConfidenceMethod::default(),
};
let mut ws = CiWorkspace::default();
let ctx = ExecutionContext::for_tests(1);
let out = GSquared::new().test_batch_adhoc(&req, &mut ws, &ctx).unwrap();
assert!(out.results[0].p_value > 0.01, "p={}", out.results[0].p_value);
}
#[test]
fn gsq_dependent_low_p() {
let n = 400usize;
let x: Vec<f64> = (0..n).map(|i| (i % 2) as f64).collect();
let y: Vec<f64> = x.clone();
let cols: [&[f64]; 2] = [&x, &y];
let queries = [CiQuery { x: 0, y: 1, z_start: 0, z_len: 0 }];
let req = CiBatchRequest {
columns: &cols,
queries: &queries,
z_flat: &[],
significance: SignificanceMethod::Analytic,
confidence: ConfidenceMethod::default(),
};
let mut ws = CiWorkspace::default();
let ctx = ExecutionContext::for_tests(2);
let out = GSquared::new().test_batch_adhoc(&req, &mut ws, &ctx).unwrap();
assert!(out.results[0].p_value < 1e-6);
let ci = out.results[0].ci.expect("G² analytic CI");
assert!(ci.0 <= out.results[0].statistic && out.results[0].statistic <= ci.1);
assert!(ci.0 >= 0.0);
}
#[test]
fn chi2_sf_pins_known_values() {
assert!((chi2_sf(3.841_459, 1.0) - 0.05).abs() < 1e-5);
assert!((chi2_sf(5.991_465, 2.0) - 0.05).abs() < 1e-5);
assert!((chi2_sf(10.0, 1.0) - 0.001_565).abs() < 1e-5);
}
#[test]
fn constant_strata_contribute_zero_dof() {
let n = 80usize;
let z: Vec<f64> = (0..n).map(|i| (i % 2) as f64).collect();
let x = z.clone(); let y: Vec<f64> = (0..n).map(|i| ((i / 2) % 2) as f64).collect();
let cols: [&[f64]; 3] = [&x, &y, &z];
let queries = [CiQuery { x: 0, y: 1, z_start: 0, z_len: 1 }];
let z_flat = [2usize];
let req = CiBatchRequest {
columns: &cols,
queries: &queries,
z_flat: &z_flat,
significance: SignificanceMethod::Analytic,
confidence: ConfidenceMethod::default(),
};
let mut ws = CiWorkspace::default();
let ctx = ExecutionContext::for_tests(3);
let out = GSquared::new().test_batch_adhoc(&req, &mut ws, &ctx).unwrap();
assert!((out.results[0].df - 1.0).abs() < 1e-12, "df={}", out.results[0].df);
assert!((out.results[0].p_value - 1.0).abs() < 1e-9);
}
#[test]
fn gsq_reuses_contingency_code_scratch() {
let n = 120usize;
let x: Vec<f64> = (0..n).map(|i| (i % 3) as f64).collect();
let y: Vec<f64> = (0..n).map(|i| (i % 2) as f64).collect();
let cols: [&[f64]; 2] = [&x, &y];
let queries = [CiQuery { x: 0, y: 1, z_start: 0, z_len: 0 }];
let req = CiBatchRequest {
columns: &cols,
queries: &queries,
z_flat: &[],
significance: SignificanceMethod::Analytic,
confidence: ConfidenceMethod::default(),
};
let mut ws = CiWorkspace::default();
let ctx = ExecutionContext::for_tests(4);
let _ = GSquared::new().test_batch_adhoc(&req, &mut ws, &ctx).unwrap();
let x_ptr = ws.contingency_x_codes.as_ptr();
let y_ptr = ws.contingency_y_codes.as_ptr();
let x_cap = ws.contingency_x_codes.capacity();
let y_cap = ws.contingency_y_codes.capacity();
let _ = GSquared::new().test_batch_adhoc(&req, &mut ws, &ctx).unwrap();
assert_eq!(ws.contingency_x_codes.as_ptr(), x_ptr);
assert_eq!(ws.contingency_y_codes.as_ptr(), y_ptr);
assert_eq!(ws.contingency_x_codes.capacity(), x_cap);
assert_eq!(ws.contingency_y_codes.capacity(), y_cap);
}
}