use super::dist::chi_square_sf_df;
use super::TestResult;
use crate::error::FdarError;
use crate::function_on_scalar::compute_group_means;
use crate::helpers::simpsons_weights;
use crate::matrix::FdMatrix;
pub fn oneway_anova_vstat(
data: &FdMatrix,
groups: &[usize],
argvals: &[f64],
) -> Result<TestResult, FdarError> {
let (n, m) = data.shape();
if m == 0 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "at least 1 column (grid points)".to_string(),
actual: "0 columns".to_string(),
});
}
if groups.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "groups",
expected: format!("{n} elements (matching data rows)"),
actual: format!("{} elements", groups.len()),
});
}
if argvals.len() != m {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: format!("{m} elements (matching data columns)"),
actual: format!("{} elements", argvals.len()),
});
}
if n < 3 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "at least 3 observations".to_string(),
actual: format!("{n} observations"),
});
}
let mut labels: Vec<usize> = groups.to_vec();
labels.sort_unstable();
labels.dedup();
let k = labels.len();
if k < 2 {
return Err(FdarError::InvalidParameter {
parameter: "groups",
message: format!("at least 2 distinct groups required, but only {k} found"),
});
}
let (group_means, overall_mean) = compute_group_means(data, groups, &labels);
let mut counts = vec![0usize; k];
for &g in groups {
let idx = labels.iter().position(|&l| l == g).unwrap_or(0);
counts[idx] += 1;
}
let weights = simpsons_weights(argvals);
let mut v_stat = 0.0;
for t in 0..m {
let mut between_t = 0.0;
for g in 0..k {
let d = group_means[(g, t)] - overall_mean[t];
between_t += counts[g] as f64 * d * d;
}
v_stat += weights[t] * between_t;
}
let df_within = (n as f64 - k as f64).max(1.0);
let mut cov_diag = vec![0.0f64; m];
for t in 0..m {
let mut ss = 0.0;
for i in 0..n {
let g = labels.iter().position(|&l| l == groups[i]).unwrap_or(0);
let d = data[(i, t)] - group_means[(g, t)];
ss += d * d;
}
cov_diag[t] = ss / df_within;
}
let a: f64 = (0..m).map(|t| weights[t] * cov_diag[t]).sum();
let b: f64 = (0..m).map(|t| (weights[t] * cov_diag[t]).powi(2)).sum();
let p_value = if a <= 1e-30 || b <= 1e-30 {
if v_stat > 1e-30 {
0.0
} else {
1.0
}
} else {
let beta = b / a;
let d = (k as f64 - 1.0) * a * a / b;
chi_square_sf_df(v_stat / beta, d)
};
Ok(TestResult {
statistic: v_stat,
p_value,
n_perm: 0,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::function_on_scalar::fanova;
use crate::test_helpers::uniform_grid;
fn noise(seed: &mut u64) -> f64 {
*seed = seed
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let z = (*seed >> 33) as f64 / (1u64 << 31) as f64;
z - 1.0
}
fn make_grouped(
per_group: usize,
argvals: &[f64],
shifts: &[f64],
seed: u64,
) -> (FdMatrix, Vec<usize>) {
let k = shifts.len();
let n = per_group * k;
let m = argvals.len();
let mut mat = FdMatrix::zeros(n, m);
let mut groups = vec![0usize; n];
let mut s = seed;
let mut row = 0;
for (g, &shift) in shifts.iter().enumerate() {
for _ in 0..per_group {
for (j, &t) in argvals.iter().enumerate() {
let base = (2.0 * std::f64::consts::PI * t).sin();
mat[(row, j)] = base + shift + 0.2 * noise(&mut s);
}
groups[row] = g;
row += 1;
}
}
(mat, groups)
}
#[test]
fn vstat_rejects_separated_groups_agrees_with_fanova() {
let argvals = uniform_grid(30);
let (data, groups) = make_grouped(20, &argvals, &[0.0, 1.5, 3.0], 101);
let res = oneway_anova_vstat(&data, &groups, &argvals).unwrap();
assert!(
res.p_value < 0.05,
"separated groups should reject, got p={} (V={})",
res.p_value,
res.statistic
);
let f = fanova(&data, &groups, 499).unwrap();
assert!(
f.p_value < 0.05,
"fanova should also reject separated groups, got p={}",
f.p_value
);
}
#[test]
fn vstat_fails_to_reject_pooled_groups_agrees_with_fanova() {
let argvals = uniform_grid(30);
let (data, groups) = make_grouped(20, &argvals, &[0.0, 0.0, 0.0], 202);
let res = oneway_anova_vstat(&data, &groups, &argvals).unwrap();
assert!(
res.p_value > 0.05,
"pooled groups should not reject, got p={} (V={})",
res.p_value,
res.statistic
);
let f = fanova(&data, &groups, 499).unwrap();
assert!(
f.p_value > 0.05,
"fanova should also fail to reject pooled groups, got p={}",
f.p_value
);
}
#[test]
fn vstat_is_deterministic() {
let argvals = uniform_grid(25);
let (data, groups) = make_grouped(15, &argvals, &[0.0, 1.0], 303);
let a = oneway_anova_vstat(&data, &groups, &argvals).unwrap();
let b = oneway_anova_vstat(&data, &groups, &argvals).unwrap();
assert_eq!(a.statistic, b.statistic);
assert_eq!(a.p_value, b.p_value);
}
#[test]
fn vstat_validates_input() {
let argvals = uniform_grid(20);
let (data, groups) = make_grouped(10, &argvals, &[0.0, 1.0], 404);
let one_group = vec![0usize; groups.len()];
assert!(matches!(
oneway_anova_vstat(&data, &one_group, &argvals),
Err(FdarError::InvalidParameter { .. })
));
let bad_argvals = uniform_grid(15);
assert!(matches!(
oneway_anova_vstat(&data, &groups, &bad_argvals),
Err(FdarError::InvalidDimension { .. })
));
let short_groups = vec![0usize; groups.len() - 1];
assert!(matches!(
oneway_anova_vstat(&data, &short_groups, &argvals),
Err(FdarError::InvalidDimension { .. })
));
}
}