use crate::causal_discovery::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use deep_causality_algebra::RealField;
use deep_causality_num::FromPrimitive;
use std::collections::BTreeMap;
pub const ALPHA_STAR_DEFAULT: f64 = 5.0;
pub fn dirichlet_logdensity<T: RealField + FromPrimitive>(
node: &[usize],
parents: &[Vec<usize>],
cardinality: usize,
alpha_star: T,
) -> Result<Vec<T>, BrcdError> {
let n = node.len();
if n == 0 {
return Err(BrcdError(BrcdErrorEnum::EmptyData));
}
if cardinality == 0 {
return Err(BrcdError(BrcdErrorEnum::ZeroCardinality));
}
if alpha_star <= T::zero() {
return Err(BrcdError(BrcdErrorEnum::NonPositiveConcentration));
}
let p = parents.first().map_or(0, Vec::len);
if !parents.is_empty() && (parents.len() != n || parents.iter().any(|r| r.len() != p)) {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
if node.iter().any(|&x| x >= cardinality) {
return Err(BrcdError(BrcdErrorEnum::StateOutOfRange));
}
let alpha0 = alpha_star / from_usize::<T>(cardinality);
let mut streams: BTreeMap<Vec<usize>, (Vec<usize>, usize)> = BTreeMap::new();
let mut out = Vec::with_capacity(n);
for (i, &x) in node.iter().enumerate() {
let key = parents.get(i).cloned().unwrap_or_default();
let (counts, total) = streams
.entry(key)
.or_insert_with(|| (vec![0usize; cardinality], 0usize));
let num = from_usize::<T>(counts[x]) + alpha0;
let den = from_usize::<T>(*total) + alpha_star;
out.push((num / den).ln());
counts[x] += 1;
*total += 1;
}
Ok(out)
}
fn from_usize<T: FromPrimitive>(n: usize) -> T {
<T as FromPrimitive>::from_usize(n).expect("count is representable in every RealField")
}