use antecedent_core::ExecutionContext;
use antecedent_kernels::ParCorrWorkspace;
use crate::error::StatsError;
#[derive(Clone, Debug, Default)]
pub struct KnnDependenceWorkspace {
pub index_generation: u64,
pub index_builds: u32,
pub last_dim: usize,
pub last_n: usize,
pub last_fingerprint: u64,
pub features: Vec<f64>,
pub index: Option<crate::matching::MatchingIndex>,
pub perm: Vec<usize>,
pub distances: Vec<f64>,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum SignificanceMethod {
Analytic,
BlockShuffle {
replicates: u32,
block_size: usize,
},
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum ConfidenceMethod {
None,
Analytic {
level: f64,
},
}
impl Default for ConfidenceMethod {
fn default() -> Self {
Self::Analytic { level: 0.95 }
}
}
#[must_use]
pub fn nonparametric_permutation_count(significance: SignificanceMethod) -> usize {
match significance {
SignificanceMethod::Analytic => 49,
SignificanceMethod::BlockShuffle { replicates, .. } => replicates.max(1) as usize,
}
}
#[must_use]
pub fn analytic_confidence_level(confidence: ConfidenceMethod) -> Option<f64> {
match confidence {
ConfidenceMethod::None => None,
ConfidenceMethod::Analytic { level } => Some(level),
}
}
#[derive(Clone, Debug)]
pub struct CiPreparationPlan {
pub significance: SignificanceMethod,
pub confidence: ConfidenceMethod,
}
impl Default for CiPreparationPlan {
fn default() -> Self {
Self { significance: SignificanceMethod::Analytic, confidence: ConfidenceMethod::default() }
}
}
#[derive(Clone, Debug)]
pub struct PreparedCiTest {
pub n: usize,
pub ncols: usize,
pub plan: CiPreparationPlan,
}
impl PreparedCiTest {
pub fn ensure_compatible(&self, request: &CiBatchRequest<'_>) -> Result<(), StatsError> {
let n = request.nrows()?;
if n != self.n {
return Err(StatsError::Shape { message: "CI batch row count differs from prepare()" });
}
if request.columns.len() != self.ncols {
return Err(StatsError::Shape {
message: "CI batch column count differs from prepare()",
});
}
Ok(())
}
#[must_use]
pub fn bind_request<'a>(&self, request: &CiBatchRequest<'a>) -> CiBatchRequest<'a> {
CiBatchRequest {
columns: request.columns,
queries: request.queries,
z_flat: request.z_flat,
significance: self.plan.significance,
confidence: self.plan.confidence,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub struct CiQuery {
pub x: usize,
pub y: usize,
pub z_start: usize,
pub z_len: usize,
}
#[derive(Clone, Debug)]
pub struct CiBatchRequest<'a> {
pub columns: &'a [&'a [f64]],
pub queries: &'a [CiQuery],
pub z_flat: &'a [usize],
pub significance: SignificanceMethod,
pub confidence: ConfidenceMethod,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct CiResult {
pub statistic: f64,
pub p_value: f64,
pub df: f64,
pub ci: Option<(f64, f64)>,
}
#[derive(Clone, Debug, Default)]
pub struct CiBatchResult {
pub results: Vec<CiResult>,
}
pub trait ConditionalIndependenceTest {
fn prepare(
&self,
columns: &[&[f64]],
plan: &CiPreparationPlan,
_ctx: &ExecutionContext,
) -> Result<PreparedCiTest, StatsError> {
if columns.is_empty() {
return Err(StatsError::Shape { message: "no columns" });
}
let n = columns[0].len();
for col in columns {
if col.len() != n {
return Err(StatsError::Shape { message: "column length mismatch" });
}
}
Ok(PreparedCiTest { n, ncols: columns.len(), plan: plan.clone() })
}
fn test(
&self,
prepared: &PreparedCiTest,
columns: &[&[f64]],
query: CiQuery,
z_flat: &[usize],
workspace: &mut CiWorkspace,
ctx: &ExecutionContext,
) -> Result<CiResult, StatsError> {
let req = CiBatchRequest {
columns,
queries: std::slice::from_ref(&query),
z_flat,
significance: prepared.plan.significance,
confidence: prepared.plan.confidence,
};
let out = self.test_batch(prepared, &req, workspace, ctx)?;
out.results
.into_iter()
.next()
.ok_or(StatsError::Shape { message: "CI test returned no results" })
}
fn test_batch(
&self,
prepared: &PreparedCiTest,
request: &CiBatchRequest<'_>,
workspace: &mut CiWorkspace,
ctx: &ExecutionContext,
) -> Result<CiBatchResult, StatsError>;
fn test_batch_adhoc(
&self,
request: &CiBatchRequest<'_>,
workspace: &mut CiWorkspace,
ctx: &ExecutionContext,
) -> Result<CiBatchResult, StatsError> {
let plan = CiPreparationPlan {
significance: request.significance,
confidence: request.confidence,
};
let prepared = self.prepare(request.columns, &plan, ctx)?;
self.test_batch(&prepared, request, workspace, ctx)
}
}
pub use ConditionalIndependenceTest as ConditionalIndependence;
impl CiBatchRequest<'_> {
pub fn nrows(&self) -> Result<usize, StatsError> {
if self.columns.is_empty() {
return Err(StatsError::Shape { message: "no columns" });
}
let n = self.columns[0].len();
for col in self.columns {
if col.len() != n {
return Err(StatsError::Shape { message: "column length mismatch" });
}
}
Ok(n)
}
}
#[derive(Clone, Debug, Default)]
pub struct CiWorkspace {
pub parcorr: ParCorrWorkspace,
pub stats: Vec<Option<f64>>,
pub shuffled: Vec<f64>,
pub block_perm: Vec<usize>,
pub contingency_x_codes: Vec<u32>,
pub contingency_y_codes: Vec<u32>,
pub knn: KnnDependenceWorkspace,
}
impl CiWorkspace {
pub fn prepare_queries(&mut self, n_queries: usize) {
if self.stats.len() < n_queries {
self.stats.resize(n_queries, None);
}
}
}