use crate::causal_discovery::surd::surd_utils;
use crate::causal_discovery::surd::{MaxOrder, SurdResult};
use deep_causality_tensor::CausalTensorError;
use deep_causality_tensor::CausalTensorMathExt;
use deep_causality_tensor::CausalTensorStackExt;
use deep_causality_tensor::{CausalTensor, Tensor};
use std::collections::HashMap;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
pub fn surd_states(
p_raw: &CausalTensor<f64>,
max_order: MaxOrder,
) -> Result<SurdResult<f64>, CausalTensorError> {
if p_raw.is_empty() {
return Err(CausalTensorError::EmptyTensor);
}
let mut p_data = p_raw.as_slice().to_vec();
let total_sum: f64 = p_data.iter().sum();
if total_sum.abs() < 1e-14 {
return Err(CausalTensorError::InvalidOperation);
}
p_data.iter_mut().for_each(|x| *x /= total_sum);
let p = CausalTensor::new(p_data, p_raw.shape().to_vec())?;
let n_total_dims = p.num_dim();
let n_vars = n_total_dims - 1;
let n_target_states = p.shape()[0];
let agent_indices: Vec<usize> = (1..n_total_dims).collect();
let k = max_order.get_k_max_order(n_vars)?;
let h = surd_utils::entropy_nvars(&p, &[0])?;
let hc = surd_utils::cond_entropy(&p, &[0], &agent_indices)?;
let info_leak = if h > 1e-14 {
(hc / h).clamp(0.0, 1.0)
} else {
0.0
};
let mut combs: Vec<Vec<usize>> = Vec::new();
for i in 1..=k.min(n_vars) {
let combinations_for_i = surd_utils::combinations(&agent_indices, i);
combs.extend(combinations_for_i);
}
let p_s = p.sum_axes(&agent_indices)?;
let mut is_map: HashMap<Vec<usize>, CausalTensor<f64>> = HashMap::new();
for j_comb in &combs {
let noj: Vec<usize> = agent_indices
.iter()
.filter(|&ax| !j_comb.contains(ax))
.cloned()
.collect();
let p_as = if noj.is_empty() {
p.clone() } else {
p.sum_axes(&noj)?
};
let p_sources = p.sum_axes(&[0])?;
let noj_mapped: Vec<usize> = noj.iter().map(|ax| ax - 1).collect();
let p_j = if noj_mapped.is_empty() {
p_sources.clone() } else {
p_sources.sum_axes(&noj_mapped)?
};
let p_s_a = p_as.safe_div(&p_j)?;
let p_a_s = p_as.safe_div(&p_s)?;
let mut broadcast_shape = vec![1; p_s_a.num_dim()];
if !broadcast_shape.is_empty() {
broadcast_shape[0] = p_s.shape()[0];
}
let p_s_reshaped = p_s.reshape(&broadcast_shape)?;
let log_diff = p_s_a.surd_log2()? - p_s_reshaped.surd_log2()?;
let specific_info_map = &p_a_s * &log_diff;
let dims_to_sum_for_j_comb: Vec<usize> = (1..=j_comb.len()).collect();
let sum_axes = specific_info_map
.sum_axes(&dims_to_sum_for_j_comb)
.expect("Failed to sum agent axes for specific_info_map");
let ravel = sum_axes.ravel();
is_map.insert(j_comb.clone(), ravel);
}
let mi: HashMap<Vec<usize>, f64> = is_map
.iter()
.map(|(k, v)| Ok((k.clone(), (v * &p_s).as_slice().iter().sum())))
.collect::<Result<_, CausalTensorError>>()?;
let results_per_target: Vec<_> = {
#[cfg(feature = "parallel")]
{
(0..n_target_states)
.into_par_iter()
.map(|t| {
analyze_single_target_state(
t,
n_vars,
&combs,
&is_map,
&p_s,
&p,
&agent_indices,
k,
)
})
.collect::<Result<Vec<_>, _>>()?
}
#[cfg(not(feature = "parallel"))]
{
(0..n_target_states)
.map(|t| {
analyze_single_target_state(
t,
n_vars,
&combs,
&is_map,
&p_s,
&p,
&agent_indices,
k,
)
})
.collect::<Result<Vec<_>, _>>()?
}
};
let mut i_r = HashMap::new();
let mut i_s = HashMap::new();
let mut temp_causal_rd_states: HashMap<Vec<usize>, Vec<CausalTensor<f64>>> = HashMap::new();
let mut temp_causal_un_states: HashMap<Vec<usize>, Vec<CausalTensor<f64>>> = HashMap::new();
let mut temp_causal_sy_states: HashMap<Vec<usize>, Vec<CausalTensor<f64>>> = HashMap::new();
let mut temp_non_causal_rd_states: HashMap<Vec<usize>, Vec<CausalTensor<f64>>> = HashMap::new();
let mut temp_non_causal_un_states: HashMap<Vec<usize>, Vec<CausalTensor<f64>>> = HashMap::new();
let mut temp_non_causal_sy_states: HashMap<Vec<usize>, Vec<CausalTensor<f64>>> = HashMap::new();
for result in results_per_target {
for (k, v) in result.i_r {
*i_r.entry(k).or_insert(0.0) += v;
}
for (k, v) in result.i_s {
*i_s.entry(k).or_insert(0.0) += v;
}
for (k, v) in result.causal_rd_states {
temp_causal_rd_states.entry(k).or_default().push(v);
}
for (k, v) in result.causal_un_states {
temp_causal_un_states.entry(k).or_default().push(v);
}
for (k, v) in result.causal_sy_states {
temp_causal_sy_states.entry(k).or_default().push(v);
}
for (k, v) in result.non_causal_rd_states {
temp_non_causal_rd_states.entry(k).or_default().push(v);
}
for (k, v) in result.non_causal_un_states {
temp_non_causal_un_states.entry(k).or_default().push(v);
}
for (k, v) in result.non_causal_sy_states {
temp_non_causal_sy_states.entry(k).or_default().push(v);
}
}
let causal_redundant_states = temp_causal_rd_states
.into_iter()
.map(|(k, slices)| Ok((k, slices.stack(0)?)))
.collect::<Result<_, _>>()?;
let causal_unique_states = temp_causal_un_states
.into_iter()
.map(|(k, slices)| Ok((k, slices.stack(0)?)))
.collect::<Result<_, _>>()?;
let causal_synergistic_states = temp_causal_sy_states
.into_iter()
.map(|(k, slices)| Ok((k, slices.stack(0)?)))
.collect::<Result<_, _>>()?;
let non_causal_redundant_states = temp_non_causal_rd_states
.into_iter()
.map(|(k, slices)| Ok((k, slices.stack(0)?)))
.collect::<Result<_, _>>()?;
let non_causal_unique_states = temp_non_causal_un_states
.into_iter()
.map(|(k, slices)| Ok((k, slices.stack(0)?)))
.collect::<Result<_, _>>()?;
let non_causal_synergistic_states = temp_non_causal_sy_states
.into_iter()
.map(|(k, slices)| Ok((k, slices.stack(0)?)))
.collect::<Result<_, _>>()?;
Ok(SurdResult::new(
i_r,
i_s,
mi,
info_leak,
causal_redundant_states,
causal_unique_states,
causal_synergistic_states,
non_causal_redundant_states,
non_causal_unique_states,
non_causal_synergistic_states,
))
}
struct PerTargetStateResults {
i_r: HashMap<Vec<usize>, f64>,
i_s: HashMap<Vec<usize>, f64>,
causal_rd_states: HashMap<Vec<usize>, CausalTensor<f64>>,
causal_un_states: HashMap<Vec<usize>, CausalTensor<f64>>,
causal_sy_states: HashMap<Vec<usize>, CausalTensor<f64>>,
non_causal_rd_states: HashMap<Vec<usize>, CausalTensor<f64>>,
non_causal_un_states: HashMap<Vec<usize>, CausalTensor<f64>>,
non_causal_sy_states: HashMap<Vec<usize>, CausalTensor<f64>>,
}
#[allow(clippy::too_many_arguments)]
fn analyze_single_target_state(
t: usize,
n_vars: usize,
combs: &[Vec<usize>],
is_map: &HashMap<Vec<usize>, CausalTensor<f64>>,
p_s: &CausalTensor<f64>,
p: &CausalTensor<f64>,
agent_indices: &[usize],
_k: usize, ) -> Result<PerTargetStateResults, CausalTensorError> {
let mut i_r = HashMap::new();
let mut i_s = HashMap::new();
let mut causal_rd_states = HashMap::new();
let mut causal_un_states = HashMap::new();
let mut causal_sy_states = HashMap::new();
let mut non_causal_rd_states = HashMap::new();
let mut non_causal_un_states = HashMap::new();
let mut non_causal_sy_states = HashMap::new();
let i1_values: Vec<f64> = combs.iter().map(|c| is_map[c].as_slice()[t]).collect();
let i1_sorted_indices = surd_utils::arg_sort(&i1_values);
let lab: Vec<Vec<usize>> = i1_sorted_indices
.iter()
.map(|&i| combs[i].clone())
.collect();
let mut i1_sorted: Vec<f64> = i1_sorted_indices.iter().map(|&i| i1_values[i]).collect();
let lens: Vec<usize> = lab.iter().map(|l| l.len()).collect();
if let Some(&max_len) = lens.iter().max() {
for l in 1..max_len {
let max_prev_level = i1_sorted
.iter()
.zip(&lens)
.filter(|&(_, &len)| len == l)
.map(|(&val, _)| val)
.fold(f64::NEG_INFINITY, f64::max);
if max_prev_level.is_finite() {
i1_sorted
.iter_mut()
.zip(&lens)
.filter(|&(_, &len)| len == l + 1)
.for_each(|(val, _)| {
if *val < max_prev_level {
*val = 0.0;
}
});
}
}
}
let new_sorted_indices = surd_utils::arg_sort(&i1_sorted);
let final_i1: Vec<f64> = new_sorted_indices.iter().map(|&i| i1_sorted[i]).collect();
let final_lab: Vec<Vec<usize>> = new_sorted_indices.iter().map(|&i| lab[i].clone()).collect();
let di_values = surd_utils::diff(&final_i1);
let last_single_var_idx = final_lab
.iter()
.rposition(|lab| lab.len() == 1)
.unwrap_or(usize::MAX);
let mut red_vars: Vec<usize> = (1..=n_vars).collect();
for (i, ll) in final_lab.iter().enumerate() {
let info = di_values[i] * p_s.as_slice()[t];
if info.abs() < 1e-14 {
continue;
}
let prev_ll: &[usize] = if i == 0 { &[] } else { &final_lab[i - 1] };
let (causal_slice, non_causal_slice) =
calculate_state_slice(p, ll, prev_ll, t, agent_indices, n_vars)?;
if ll.len() > 1
{
causal_sy_states.insert(ll.clone(), causal_slice);
non_causal_sy_states.insert(ll.clone(), non_causal_slice);
*i_s.entry(ll.clone()).or_insert(0.0) += info;
} else if ll.len() == 1 {
if i == last_single_var_idx
{
causal_un_states.insert(ll.clone(), causal_slice);
non_causal_un_states.insert(ll.clone(), non_causal_slice);
} else
{
causal_rd_states.insert(red_vars.clone(), causal_slice);
non_causal_rd_states.insert(red_vars.clone(), non_causal_slice);
}
*i_r.entry(red_vars.clone()).or_insert(0.0) += info;
red_vars.retain(|&v| v != ll[0]);
}
}
Ok(PerTargetStateResults {
i_r,
i_s,
causal_rd_states,
causal_un_states,
causal_sy_states,
non_causal_rd_states,
non_causal_un_states,
non_causal_sy_states,
})
}
fn calculate_state_slice(
p: &CausalTensor<f64>,
current_vars: &[usize], prev_vars: &[usize], target_state_index: usize,
agent_indices: &[usize],
n_vars: usize,
) -> Result<(CausalTensor<f64>, CausalTensor<f64>), CausalTensorError> {
let p_slice = p.slice(0, target_state_index)?;
let current_vars_mapped: Vec<usize> = current_vars.iter().map(|&ax| ax - 1).collect();
let prev_vars_mapped: Vec<usize> = prev_vars.iter().map(|&ax| ax - 1).collect();
let all_vars_mapped: Vec<usize> = (0..p_slice.num_dim()).collect();
let p_ti = p_slice.sum_axes(&surd_utils::set_difference(
&all_vars_mapped,
¤t_vars_mapped,
))?;
let source_axes: Vec<usize> = (0..n_vars).collect();
let current_vars_mapped_for_marginal: Vec<usize> =
current_vars.iter().map(|&ax| ax - 1).collect();
let axes_to_sum_for_pi =
surd_utils::set_difference(&source_axes, ¤t_vars_mapped_for_marginal);
let p_i = p.sum_axes(&[0])?.sum_axes(&axes_to_sum_for_pi)?;
let p_target_given_i = p_ti.safe_div(&p_i)?;
let p_target_given_j = if prev_vars.is_empty() {
p.sum_axes(agent_indices)?
} else {
let p_tj = p_slice.sum_axes(&surd_utils::set_difference(
&all_vars_mapped,
&prev_vars_mapped,
))?;
let prev_vars_mapped_for_marginal: Vec<usize> =
prev_vars.iter().map(|&ax| ax - 1).collect();
let axes_to_sum_for_pj =
surd_utils::set_difference(&source_axes, &prev_vars_mapped_for_marginal);
let p_j = p.sum_axes(&[0])?.sum_axes(&axes_to_sum_for_pj)?;
p_tj.safe_div(&p_j)?
};
let log_ratio = (p_target_given_i / p_target_given_j).log2()?;
let mut all_involved_vars = current_vars_mapped.to_vec();
all_involved_vars.extend_from_slice(&prev_vars_mapped);
all_involved_vars.sort();
all_involved_vars.dedup();
let axes_to_sum_out = surd_utils::set_difference(&all_vars_mapped, &all_involved_vars);
let p_tij = p_slice.sum_axes(&axes_to_sum_out)?;
let causal_log_ratio_data: Vec<f64> =
log_ratio.as_slice().iter().map(|&v| v.max(0.0)).collect();
let non_causal_log_ratio_data: Vec<f64> =
log_ratio.as_slice().iter().map(|&v| v.min(0.0)).collect();
let causal_log_ratio = CausalTensor::new(causal_log_ratio_data, log_ratio.shape().to_vec())?;
let non_causal_log_ratio =
CausalTensor::new(non_causal_log_ratio_data, log_ratio.shape().to_vec())?;
let causal_slice = &p_tij * &causal_log_ratio;
let non_causal_slice = &p_tij * &non_causal_log_ratio;
Ok((causal_slice, non_causal_slice))
}