use deep_causality_num::RealField;
use deep_causality_tensor::{CausalTensor, CausalTensorError, Tensor};
use std::cmp::Ordering;
pub(crate) mod surd_utils_cdl;
#[cfg(test)]
mod surd_utils_tests;
pub(crate) fn diff<T: RealField>(slice: &[T]) -> Vec<T> {
let mut result = Vec::new();
if slice.is_empty() {
return result;
}
result.push(slice[0] - T::zero());
for i in 1..slice.len() {
result.push(slice[i] - slice[i - 1]);
}
result
}
pub(crate) fn arg_sort<T: RealField>(slice: &[T]) -> Vec<usize> {
let mut indices: Vec<usize> = (0..slice.len()).collect();
indices.sort_by(|&a_index, &b_index| {
let a_value = slice[a_index];
let b_value = slice[b_index];
a_value.partial_cmp(&b_value).unwrap_or(Ordering::Equal)
});
indices
}
pub(crate) fn set_difference<T: PartialEq + Clone>(a: &[T], b: &[T]) -> Vec<T> {
a.iter()
.filter(|&item| !b.contains(item))
.cloned()
.collect()
}
pub(crate) fn combinations<T: Copy>(pool: &[T], r: usize) -> Vec<Vec<T>> {
if r > pool.len() {
panic!("Cannot choose r elements from a pool smaller than r.");
}
if r == 0 {
return vec![vec![]];
}
let mut result = Vec::new();
let mut indices: Vec<usize> = (0..r).collect();
loop {
result.push(indices.iter().map(|&i| pool[i]).collect());
let mut i = r - 1;
loop {
indices[i] += 1;
if indices[i] < pool.len() - (r - 1 - i) {
for j in (i + 1)..r {
indices[j] = indices[j - 1] + 1;
}
break;
}
if i == 0 {
return result;
}
i -= 1;
}
}
}
pub fn entropy_nvars<T: RealField + Default>(
p: &CausalTensor<T>,
axes: &[usize],
) -> Result<T, CausalTensorError> {
let all_axes: Vec<_> = (0..p.num_dim()).collect();
let axes_to_sum_out: Vec<_> = all_axes
.into_iter()
.filter(|ax| !axes.contains(ax))
.collect();
let zero = T::zero();
if axes_to_sum_out.is_empty() {
let entropy = p.as_slice().iter().fold(zero, |acc, &prob| {
if prob > zero {
acc - prob * prob.log2()
} else {
acc
}
});
Ok(entropy)
} else {
let marginal = p.sum_axes(&axes_to_sum_out)?;
let entropy = marginal.as_slice().iter().fold(zero, |acc, &prob| {
if prob > zero {
acc - prob * prob.log2()
} else {
acc
}
});
Ok(entropy)
}
}
pub fn cond_entropy<T: RealField + Default>(
p: &CausalTensor<T>,
target_axes: &[usize],
cond_axes: &[usize],
) -> Result<T, CausalTensorError> {
let mut joint_axes = target_axes.to_vec();
joint_axes.extend_from_slice(cond_axes);
joint_axes.sort();
joint_axes.dedup();
let h_xy = entropy_nvars(p, &joint_axes)?;
let h_y = entropy_nvars(p, cond_axes)?;
Ok(h_xy - h_y)
}