use crate::causal_discovery::surd::surd_utils;
use crate::causal_discovery::surd::surd_utils::surd_utils_cdl;
use crate::causal_discovery::surd::{MaxOrder, SurdResult};
use deep_causality_algebra::RealField;
use deep_causality_num::FromPrimitive;
use deep_causality_tensor::{CausalTensor, CausalTensorError, Tensor};
use std::collections::HashMap;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
pub fn surd_states_cdl<T>(
p_raw: &CausalTensor<Option<T>>,
max_order: MaxOrder,
) -> Result<SurdResult<T>, CausalTensorError>
where
T: RealField + FromPrimitive + Default + Send + Sync,
{
if p_raw.is_empty() {
return Err(CausalTensorError::EmptyTensor);
}
let zero = T::zero();
let one = T::one();
let eps = <T as FromPrimitive>::from_f64(1e-14).expect("1e-14 is representable in RealField");
let total_sum: T = p_raw
.as_slice()
.iter()
.filter_map(|&x| x)
.fold(zero, |acc, v| acc + v);
if total_sum.abs() < eps {
return Err(CausalTensorError::InvalidOperation);
}
let p_data: Vec<Option<T>> = p_raw
.as_slice()
.iter()
.map(|&x| x.map(|val| val / total_sum))
.collect();
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_cdl::entropy_nvars_cdl(&p, &[0])?;
let hc = surd_utils_cdl::cond_entropy_cdl(&p, &[0], &agent_indices)?;
let info_leak = if h > eps {
(hc / h).clamp(zero, one)
} else {
zero
};
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 = surd_utils_cdl::sum_axes_option_f64(&p, &agent_indices)?;
let mut is_map: HashMap<Vec<usize>, CausalTensor<Option<T>>> = 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 {
surd_utils_cdl::sum_axes_option_f64(&p, &noj)?
};
let p_sources = surd_utils_cdl::sum_axes_option_f64(&p, &[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 {
surd_utils_cdl::sum_axes_option_f64(&p_sources, &noj_mapped)?
};
let p_j_broadcasted = surd_utils_cdl::broadcast_to_cdl(&p_j, p_as.shape())?;
let p_s_a = surd_utils_cdl::safe_div_cdl(&p_as, &p_j_broadcasted)?;
let p_s_broadcasted_for_pas = surd_utils_cdl::broadcast_to_cdl(&p_s, p_as.shape())?;
let p_a_s = surd_utils_cdl::safe_div_cdl(&p_as, &p_s_broadcasted_for_pas)?;
let p_s_broadcasted_for_log = surd_utils_cdl::broadcast_to_cdl(&p_s, p_s_a.shape())?;
let log_diff = surd_utils_cdl::sub_cdl(
&surd_utils_cdl::surd_log2_cdl(&p_s_a)?,
&surd_utils_cdl::surd_log2_cdl(&p_s_broadcasted_for_log)?,
)?;
let specific_info_map = surd_utils_cdl::mul_cdl(&p_a_s, &log_diff)?;
let dims_to_sum_for_j_comb: Vec<usize> = (1..=j_comb.len()).collect();
let sum_axes =
surd_utils_cdl::sum_axes_option_f64(&specific_info_map, &dims_to_sum_for_j_comb)?;
let ravel = sum_axes.ravel();
is_map.insert(j_comb.clone(), ravel);
}
let mi: HashMap<Vec<usize>, T> = is_map
.iter()
.map(|(key, v)| {
let multiplied = surd_utils_cdl::mul_cdl(v, &p_s)?;
let sum = multiplied
.as_slice()
.iter()
.filter_map(|&x| x)
.fold(zero, |acc, x| acc + x);
Ok((key.clone(), 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_cdl(
t, n_vars, &combs, &is_map, &p_s, &p, k,
)
})
.collect::<Result<Vec<_>, _>>()?
}
#[cfg(not(feature = "parallel"))]
{
(0..n_target_states)
.map(|t| {
analyze_single_target_state_cdl(
t, n_vars, &combs, &is_map, &p_s, &p, 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<T>>> = HashMap::new();
let mut temp_causal_un_states: HashMap<Vec<usize>, Vec<CausalTensor<T>>> = HashMap::new();
let mut temp_causal_sy_states: HashMap<Vec<usize>, Vec<CausalTensor<T>>> = HashMap::new();
let mut temp_non_causal_rd_states: HashMap<Vec<usize>, Vec<CausalTensor<T>>> = HashMap::new();
let mut temp_non_causal_un_states: HashMap<Vec<usize>, Vec<CausalTensor<T>>> = HashMap::new();
let mut temp_non_causal_sy_states: HashMap<Vec<usize>, Vec<CausalTensor<T>>> = HashMap::new();
for result in results_per_target {
for (key, v) in result.i_r {
*i_r.entry(key).or_insert(zero) += v;
}
for (key, v) in result.i_s {
*i_s.entry(key).or_insert(zero) += v;
}
for (key, v) in result.causal_rd_states {
temp_causal_rd_states.entry(key).or_default().push(v);
}
for (key, v) in result.causal_un_states {
temp_causal_un_states.entry(key).or_default().push(v);
}
for (key, v) in result.causal_sy_states {
temp_causal_sy_states.entry(key).or_default().push(v);
}
for (key, v) in result.non_causal_rd_states {
temp_non_causal_rd_states.entry(key).or_default().push(v);
}
for (key, v) in result.non_causal_un_states {
temp_non_causal_un_states.entry(key).or_default().push(v);
}
for (key, v) in result.non_causal_sy_states {
temp_non_causal_sy_states.entry(key).or_default().push(v);
}
}
let causal_redundant_states = temp_causal_rd_states
.into_iter()
.map(|(key, slices)| Ok((key, CausalTensor::stack(&slices, 0)?)))
.collect::<Result<_, _>>()?;
let causal_unique_states = temp_causal_un_states
.into_iter()
.map(|(key, slices)| Ok((key, CausalTensor::stack(&slices, 0)?)))
.collect::<Result<_, _>>()?;
let causal_synergistic_states = temp_causal_sy_states
.into_iter()
.map(|(key, slices)| Ok((key, CausalTensor::stack(&slices, 0)?)))
.collect::<Result<_, _>>()?;
let non_causal_redundant_states = temp_non_causal_rd_states
.into_iter()
.map(|(key, slices)| Ok((key, CausalTensor::stack(&slices, 0)?)))
.collect::<Result<_, _>>()?;
let non_causal_unique_states = temp_non_causal_un_states
.into_iter()
.map(|(key, slices)| Ok((key, CausalTensor::stack(&slices, 0)?)))
.collect::<Result<_, _>>()?;
let non_causal_synergistic_states = temp_non_causal_sy_states
.into_iter()
.map(|(key, slices)| Ok((key, CausalTensor::stack(&slices, 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<T> {
i_r: HashMap<Vec<usize>, T>,
i_s: HashMap<Vec<usize>, T>,
causal_rd_states: HashMap<Vec<usize>, CausalTensor<T>>,
causal_un_states: HashMap<Vec<usize>, CausalTensor<T>>,
causal_sy_states: HashMap<Vec<usize>, CausalTensor<T>>,
non_causal_rd_states: HashMap<Vec<usize>, CausalTensor<T>>,
non_causal_un_states: HashMap<Vec<usize>, CausalTensor<T>>,
non_causal_sy_states: HashMap<Vec<usize>, CausalTensor<T>>,
}
#[allow(clippy::too_many_arguments)]
fn analyze_single_target_state_cdl<T>(
t: usize,
n_vars: usize,
combs: &[Vec<usize>],
is_map: &HashMap<Vec<usize>, CausalTensor<Option<T>>>,
p_s: &CausalTensor<Option<T>>,
p: &CausalTensor<Option<T>>,
_k: usize,
) -> Result<PerTargetStateResults<T>, CausalTensorError>
where
T: RealField + FromPrimitive + Default + Send + Sync,
{
let zero = T::zero();
#[cfg(not(miri))]
let eps = T::epsilon();
#[cfg(miri)]
let eps = <T as FromPrimitive>::from_f64(1e-9).expect("1e-9 is representable in RealField");
let rank_tol =
<T as FromPrimitive>::from_f64(1e-9).expect("1e-9 is representable in RealField");
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 mut i1_values_with_indices: Vec<(T, usize)> = Vec::new();
for (idx, c) in combs.iter().enumerate() {
if let Some(val) = is_map[c].as_slice()[t] {
i1_values_with_indices.push((val, idx));
}
}
i1_values_with_indices.sort_by(|a, b| {
let a_key = (a.0 / rank_tol).round();
let b_key = (b.0 / rank_tol).round();
a_key
.partial_cmp(&b_key)
.unwrap_or(std::cmp::Ordering::Equal)
});
let i1_sorted_indices: Vec<usize> =
i1_values_with_indices.iter().map(|&(_, idx)| idx).collect();
let i1_sorted_values: Vec<T> = i1_values_with_indices.iter().map(|&(val, _)| val).collect();
let lab: Vec<Vec<usize>> = i1_sorted_indices
.iter()
.map(|&i| combs[i].clone())
.collect();
let mut i1_sorted: Vec<T> = i1_sorted_values.clone();
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(None, |acc: Option<T>, v| match acc {
None => Some(v),
Some(a) => Some(if v > a { v } else { a }),
});
if let Some(max_prev_level) = max_prev_level {
i1_sorted
.iter_mut()
.zip(&lens)
.filter(|&(_, &len)| len == l + 1)
.for_each(|(val, _)| {
if *val < max_prev_level {
*val = max_prev_level;
}
});
}
}
}
let new_sorted_indices = surd_utils::arg_sort_stable(&i1_sorted, rank_tol);
let final_i1: Vec<T> = 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 = if let Some(p_s_val) = p_s.as_slice()[t] {
di_values[i] * p_s_val
} else {
zero };
if info.abs() < eps {
continue;
}
let prev_ll: &[usize] = if i == 0 { &[] } else { &final_lab[i - 1] };
let (causal_slice, non_causal_slice) =
calculate_state_slice_cdl(p, ll, prev_ll, t, 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(zero) += 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(zero) += 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_cdl<T>(
p: &CausalTensor<Option<T>>,
current_vars: &[usize],
prev_vars: &[usize],
target_state_index: usize,
n_vars: usize,
) -> Result<(CausalTensor<T>, CausalTensor<T>), CausalTensorError>
where
T: RealField + FromPrimitive + Default + Send + Sync,
{
let zero = T::zero();
let p_slice_option = 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_option.num_dim()).collect();
let p_ti_option = surd_utils_cdl::sum_axes_option_f64(
&p_slice_option,
&surd_utils_cdl::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_cdl::set_difference(&source_axes, ¤t_vars_mapped_for_marginal);
let p_i_option = surd_utils_cdl::sum_axes_option_f64(
&surd_utils_cdl::sum_axes_option_f64(p, &[0])?,
&axes_to_sum_for_pi,
)?;
if p_ti_option.shape() != p_i_option.shape() {
dbg!("if p_ti_option.shape() != p_i_option.shape() : Tensor ShapeMismatch");
return Err(CausalTensorError::ShapeMismatch);
}
let p_target_given_i_option = surd_utils_cdl::safe_div_cdl(&p_ti_option, &p_i_option)?;
let p_target_given_j_option = if prev_vars.is_empty() {
let p_t_scalar_option =
surd_utils_cdl::sum_axes_option_f64(&p_slice_option, &all_vars_mapped)?;
surd_utils_cdl::broadcast_to_cdl(&p_t_scalar_option, p_target_given_i_option.shape())?
} else {
let p_tj_option = surd_utils_cdl::sum_axes_option_f64(
&p_slice_option,
&surd_utils_cdl::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_cdl::set_difference(&source_axes, &prev_vars_mapped_for_marginal);
let p_j_option = surd_utils_cdl::sum_axes_option_f64(
&surd_utils_cdl::sum_axes_option_f64(p, &[0])?,
&axes_to_sum_for_pj,
)?;
if p_tj_option.shape() != p_j_option.shape() {
dbg!("if p_tj_option.shape() != p_j_option.shape() : Tensor ShapeMismatch");
return Err(CausalTensorError::ShapeMismatch);
}
surd_utils_cdl::safe_div_cdl(&p_tj_option, &p_j_option)?
};
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_cdl::set_difference(&all_vars_mapped, &all_involved_vars);
let p_tij_option = surd_utils_cdl::sum_axes_option_f64(&p_slice_option, &axes_to_sum_out)?;
let p_target_given_i_broadcasted_option =
surd_utils_cdl::broadcast_to_cdl(&p_target_given_i_option, p_tij_option.shape())?;
let p_target_given_j_broadcasted_option =
surd_utils_cdl::broadcast_to_cdl(&p_target_given_j_option, p_tij_option.shape())?;
let log_ratio_option = surd_utils_cdl::sub_cdl(
&surd_utils_cdl::surd_log2_cdl(&p_target_given_i_broadcasted_option)?,
&surd_utils_cdl::surd_log2_cdl(&p_target_given_j_broadcasted_option)?,
)?;
let causal_log_ratio_option_data: Vec<Option<T>> = log_ratio_option
.as_slice()
.iter()
.map(|&v_opt| v_opt.map(|v| if v > zero { v } else { zero }))
.collect();
let non_causal_log_ratio_option_data: Vec<Option<T>> = log_ratio_option
.as_slice()
.iter()
.map(|&v_opt| v_opt.map(|v| if v < zero { v } else { zero }))
.collect();
let causal_log_ratio_tensor = CausalTensor::new(
causal_log_ratio_option_data,
log_ratio_option.shape().to_vec(),
)?;
let non_causal_log_ratio_tensor = CausalTensor::new(
non_causal_log_ratio_option_data,
log_ratio_option.shape().to_vec(),
)?;
let causal_slice_option = surd_utils_cdl::mul_cdl(&p_tij_option, &causal_log_ratio_tensor)?;
let non_causal_slice_option =
surd_utils_cdl::mul_cdl(&p_tij_option, &non_causal_log_ratio_tensor)?;
let causal_slice_data: Vec<T> = causal_slice_option
.as_slice()
.iter()
.map(|&x| x.unwrap_or(zero))
.collect();
let non_causal_slice_data: Vec<T> = non_causal_slice_option
.as_slice()
.iter()
.map(|&x| x.unwrap_or(zero))
.collect();
let causal_slice = CausalTensor::new(causal_slice_data, causal_slice_option.shape().to_vec())?;
let non_causal_slice = CausalTensor::new(
non_causal_slice_data,
non_causal_slice_option.shape().to_vec(),
)?;
Ok((causal_slice, non_causal_slice))
}