use crate::brcd::brcd_augment::{augmented_graph, f_node_indicator, get_configurations_multi};
use crate::brcd::brcd_boss_config::BossConfig;
use crate::brcd::brcd_boss_learn::boss_learn;
use crate::brcd::brcd_cache::{FamilyKey, family_key};
use crate::brcd::brcd_config::{BrcdConfig, ConfigStrategy, FamilyKind};
use crate::brcd::brcd_dirichlet::dirichlet_logdensity;
use crate::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use crate::brcd::brcd_gaussian::{GaussianFamilyConfig, gaussian_family_logdensity};
use crate::brcd::brcd_mapconfig::find_map_configs;
use crate::brcd::brcd_result::BrcdResult;
use crate::dag_sampling::{mec_size, representative_dag, sample_dag};
use deep_causality_algebra::RealField;
use deep_causality_num::{FromPrimitive, ToPrimitive};
use deep_causality_par::MaybeParallel;
use deep_causality_rand::Xoshiro256;
use deep_causality_tensor::CausalTensor;
use deep_causality_topology::MixedGraph;
use std::collections::BTreeMap;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
pub fn brcd_run<T, N>(
normal: &CausalTensor<T>,
anomalous: &CausalTensor<T>,
cpdag: Option<&MixedGraph<N>>,
config: &BrcdConfig<T>,
) -> Result<BrcdResult<T>, BrcdError>
where
T: RealField + FromPrimitive + ToPrimitive + MaybeParallel,
N: Clone + MaybeParallel,
{
match cpdag {
Some(graph) => run_with_cpdag(normal, anomalous, graph, config),
None => {
let boss_cfg = BossConfig::<T>::with_seed(config.seed);
let learned = boss_learn(normal, &boss_cfg)?;
run_with_cpdag(normal, anomalous, &learned, config)
}
}
}
fn run_with_cpdag<T, N>(
normal: &CausalTensor<T>,
anomalous: &CausalTensor<T>,
cpdag: &MixedGraph<N>,
config: &BrcdConfig<T>,
) -> Result<BrcdResult<T>, BrcdError>
where
T: RealField + FromPrimitive + ToPrimitive + MaybeParallel,
N: Clone + MaybeParallel,
{
let (n_normal, num_vars) = shape_2d(normal)?;
let (n_anom, num_vars2) = shape_2d(anomalous)?;
if num_vars != num_vars2 || cpdag.num_vertices() != num_vars {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let n_total = n_normal + n_anom;
if n_total == 0 {
return Err(BrcdError(BrcdErrorEnum::EmptyData));
}
let k = config.num_root_causes;
if k == 0 || k > num_vars {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let fnode_idx = num_vars;
let columns = joint_columns(normal, anomalous, num_vars, n_normal, n_anom);
let f_bool = f_node_indicator(n_normal, n_anom);
let (int_columns, cardinalities) = if config.family == FamilyKind::Discrete {
build_discrete(&columns)?
} else {
(Vec::new(), Vec::new())
};
let combos = combinations(num_vars, k);
if combos.is_empty() {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let log_prior = (T::one() / from_usize::<T>(combos.len())).ln();
let ctx = ScoreCtx {
family: config.family,
gaussian_cfg: GaussianFamilyConfig {
transform: config.node_transform,
transform_parents: config.transform_parents,
ridge: config.ridge,
gate: config.gate,
},
alpha_star: config.alpha_star,
columns: &columns,
int_columns: &int_columns,
cardinalities: &cardinalities,
f_bool: &f_bool,
fnode_idx,
n_total,
};
#[cfg(feature = "parallel")]
let plans: Vec<Option<CandidatePlan<T>>> = combos
.par_iter()
.enumerate()
.map(|(i, combo)| {
build_candidate_plan(cpdag, combo, config, &ctx, candidate_seed(config.seed, i))
})
.collect::<Result<Vec<_>, _>>()?;
#[cfg(not(feature = "parallel"))]
let plans: Vec<Option<CandidatePlan<T>>> = combos
.iter()
.enumerate()
.map(|(i, combo)| {
build_candidate_plan(cpdag, combo, config, &ctx, candidate_seed(config.seed, i))
})
.collect::<Result<Vec<_>, _>>()?;
let mut jobs: BTreeMap<FamilyKey, (usize, Vec<usize>)> = BTreeMap::new();
for (dags, _) in plans.iter().flatten() {
for dag in dags {
for node in 0..dag.num_vertices() {
let parents = dag.parents(node);
jobs.entry(family_key(node, &parents))
.or_insert((node, parents));
}
}
}
let scored = score_families(&jobs, &ctx)?;
let mut log_posterior = Vec::with_capacity(combos.len());
for plan in &plans {
let Some((dags, log_p_g)) = plan else {
log_posterior.push(neg_inf::<T>());
continue;
};
if dags.len() == 1 {
let dag = &dags[0];
let mut log_lik = from_usize::<T>(n_total) * log_p_g[0];
for node in 0..dag.num_vertices() {
let key = family_key(node, &dag.parents(node));
log_lik += scored[&key].1;
}
log_posterior.push(log_lik + log_prior);
continue;
}
let mut dag_cols: Vec<Vec<T>> = Vec::with_capacity(dags.len());
for (i, dag) in dags.iter().enumerate() {
let mut log_joint = vec![T::zero(); n_total];
for node in 0..dag.num_vertices() {
let key = family_key(node, &dag.parents(node));
for (acc, &f) in log_joint.iter_mut().zip(scored[&key].0.iter()) {
*acc += f;
}
}
let lg = log_p_g[i];
for acc in log_joint.iter_mut() {
*acc += lg;
}
dag_cols.push(log_joint);
}
let mut log_lik = T::zero();
let mut row_vals = vec![T::zero(); dag_cols.len()];
for r in 0..n_total {
for (slot, col) in row_vals.iter_mut().zip(dag_cols.iter()) {
*slot = col[r];
}
log_lik += logsumexp_slice(&row_vals);
}
log_posterior.push(log_lik + log_prior);
}
Ok(rank(combos, log_posterior))
}
type CandidatePlan<T> = (Vec<MixedGraph<()>>, Vec<T>);
#[inline]
fn candidate_seed(base: u64, index: usize) -> u64 {
let mut z = base ^ (index as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn build_candidate_plan<T, N>(
cpdag: &MixedGraph<N>,
combo: &[usize],
config: &BrcdConfig<T>,
ctx: &ScoreCtx<'_, T>,
seed: u64,
) -> Result<Option<CandidatePlan<T>>, BrcdError>
where
T: RealField + FromPrimitive + ToPrimitive + MaybeParallel,
N: Clone + MaybeParallel,
{
let configs: Vec<MixedGraph<N>> = match config.config_strategy {
ConfigStrategy::Full => get_configurations_multi(cpdag, combo)?,
ConfigStrategy::MapPrune => {
find_map_configs::<T, N, _>(cpdag, combo, |g| config_weight(g, combo, ctx))?.configs
}
};
if configs.is_empty() {
return Ok(None);
}
let mut rng = Xoshiro256::from_seed(seed);
let mut dags: Vec<MixedGraph<()>> = Vec::with_capacity(configs.len());
let mut sizes: Vec<T> = Vec::with_capacity(configs.len());
for cfg in &configs {
let aug = augmented_graph(cfg, combo)?;
sizes.push(mec_size::<T, ()>(&aug));
dags.push(sample_dag::<T, (), _>(&aug, &mut rng)?);
}
let total = sizes.iter().fold(T::zero(), |a, &s| a + s);
let tiny = from_f64::<T>(1e-300);
let log_p_g: Vec<T> = sizes.iter().map(|&s| (s / total + tiny).ln()).collect();
Ok(Some((dags, log_p_g)))
}
fn score_families<T>(
jobs: &BTreeMap<FamilyKey, (usize, Vec<usize>)>,
ctx: &ScoreCtx<'_, T>,
) -> Result<BTreeMap<FamilyKey, (Vec<T>, T)>, BrcdError>
where
T: RealField + FromPrimitive + MaybeParallel,
{
#[cfg(feature = "parallel")]
{
jobs.par_iter()
.map(|(key, (node, parents))| {
let per_row = ctx.score(*node, parents)?;
let total = per_row.iter().fold(T::zero(), |a, &x| a + x);
Ok((key.clone(), (per_row, total)))
})
.collect()
}
#[cfg(not(feature = "parallel"))]
{
jobs.iter()
.map(|(key, (node, parents))| {
let per_row = ctx.score(*node, parents)?;
let total = per_row.iter().fold(T::zero(), |a, &x| a + x);
Ok((key.clone(), (per_row, total)))
})
.collect()
}
}
fn config_weight<T, N>(
config: &MixedGraph<N>,
combo: &[usize],
ctx: &ScoreCtx<'_, T>,
) -> Result<T, BrcdError>
where
T: RealField + FromPrimitive,
N: Clone,
{
let aug = augmented_graph(config, combo)?;
let size = mec_size::<T, ()>(&aug);
let rep = representative_dag::<()>(&aug)?;
let mut log_lik = T::zero();
for node in 0..rep.num_vertices() {
let per_row = ctx.score(node, &rep.parents(node))?;
log_lik += per_row.iter().fold(T::zero(), |a, &x| a + x);
}
let tiny = from_f64::<T>(1e-300);
Ok(log_lik + (size + tiny).ln())
}
struct ScoreCtx<'a, T> {
family: FamilyKind,
gaussian_cfg: GaussianFamilyConfig<T>,
alpha_star: T,
columns: &'a [Vec<T>],
int_columns: &'a [Vec<usize>],
cardinalities: &'a [usize],
f_bool: &'a [bool],
fnode_idx: usize,
n_total: usize,
}
impl<T: RealField + FromPrimitive> ScoreCtx<'_, T> {
fn score(&self, node: usize, parents: &[usize]) -> Result<Vec<T>, BrcdError> {
match self.family {
FamilyKind::Continuous => {
let has_fnode = parents.contains(&self.fnode_idx);
let cont_parents: Vec<usize> = parents
.iter()
.copied()
.filter(|&p| p != self.fnode_idx)
.collect();
let parent_rows = transpose(self.columns, &cont_parents, self.n_total);
let f = if has_fnode { Some(self.f_bool) } else { None };
gaussian_family_logdensity(
&self.columns[node],
&parent_rows,
f,
has_fnode,
&self.gaussian_cfg,
)
}
FamilyKind::Discrete => {
let parent_configs = transpose_int(self.int_columns, parents, self.n_total);
dirichlet_logdensity(
&self.int_columns[node],
&parent_configs,
self.cardinalities[node],
self.alpha_star,
)
}
}
}
}
fn shape_2d<T>(t: &CausalTensor<T>) -> Result<(usize, usize), BrcdError> {
match t.shape() {
[rows, cols] => Ok((*rows, *cols)),
_ => Err(BrcdError(BrcdErrorEnum::DimensionMismatch)),
}
}
fn joint_columns<T: RealField + FromPrimitive>(
normal: &CausalTensor<T>,
anomalous: &CausalTensor<T>,
num_vars: usize,
n_normal: usize,
n_anom: usize,
) -> Vec<Vec<T>> {
let nd = normal.as_slice();
let ad = anomalous.as_slice();
let mut cols = Vec::with_capacity(num_vars + 1);
for j in 0..num_vars {
let mut col = Vec::with_capacity(n_normal + n_anom);
for i in 0..n_normal {
col.push(nd[i * num_vars + j]);
}
for i in 0..n_anom {
col.push(ad[i * num_vars + j]);
}
cols.push(col);
}
let mut f = vec![T::zero(); n_normal];
f.extend(std::iter::repeat_n(T::one(), n_anom));
cols.push(f);
cols
}
fn build_discrete<T: RealField + FromPrimitive + ToPrimitive>(
columns: &[Vec<T>],
) -> Result<(Vec<Vec<usize>>, Vec<usize>), BrcdError> {
let mut ints = Vec::with_capacity(columns.len());
let mut cards = Vec::with_capacity(columns.len());
for col in columns {
let mut ic = Vec::with_capacity(col.len());
let mut max_state = 0usize;
for &v in col {
let rounded = v.round();
if rounded < T::zero() {
return Err(BrcdError(BrcdErrorEnum::StateOutOfRange));
}
let s = rounded
.to_usize()
.ok_or(BrcdError(BrcdErrorEnum::StateOutOfRange))?;
max_state = max_state.max(s);
ic.push(s);
}
ints.push(ic);
cards.push(max_state + 1);
}
Ok((ints, cards))
}
fn transpose<T: RealField>(columns: &[Vec<T>], idxs: &[usize], n: usize) -> Vec<Vec<T>> {
if idxs.is_empty() {
return Vec::new();
}
(0..n)
.map(|i| idxs.iter().map(|&p| columns[p][i]).collect())
.collect()
}
fn transpose_int(columns: &[Vec<usize>], idxs: &[usize], n: usize) -> Vec<Vec<usize>> {
if idxs.is_empty() {
return Vec::new();
}
(0..n)
.map(|i| idxs.iter().map(|&p| columns[p][i]).collect())
.collect()
}
fn combinations(n: usize, k: usize) -> Vec<Vec<usize>> {
if k == 0 || k > n {
return Vec::new();
}
let mut idx: Vec<usize> = (0..k).collect();
let mut out = vec![idx.clone()];
loop {
let mut i = k;
let advanced = loop {
if i == 0 {
break false;
}
i -= 1;
if idx[i] < n - k + i {
break true;
}
};
if !advanced {
break;
}
idx[i] += 1;
for j in (i + 1)..k {
idx[j] = idx[j - 1] + 1;
}
out.push(idx.clone());
}
out
}
fn logsumexp_slice<T: RealField>(vals: &[T]) -> T {
if vals.is_empty() {
return neg_inf::<T>();
}
let max = vals.iter().fold(vals[0], |a, &b| if b > a { b } else { a });
if !max.is_finite() {
return max;
}
let sum = vals.iter().fold(T::zero(), |acc, &v| acc + (v - max).exp());
max + sum.ln()
}
fn rank<T: RealField>(combos: Vec<Vec<usize>>, log_posterior: Vec<T>) -> BrcdResult<T> {
let max = log_posterior
.iter()
.fold(neg_inf::<T>(), |a, &b| if b > a { b } else { a });
let shift = if max.is_finite() { max } else { T::zero() };
let posterior: Vec<T> = log_posterior.iter().map(|&lp| (lp - shift).exp()).collect();
let mut order: Vec<usize> = (0..combos.len()).collect();
order.sort_by(|&a, &b| {
log_posterior[b]
.partial_cmp(&log_posterior[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
BrcdResult::new(
order.iter().map(|&i| combos[i].clone()).collect(),
order.iter().map(|&i| posterior[i]).collect(),
)
}
fn neg_inf<T: RealField>() -> T {
T::zero().ln()
}
fn from_usize<T: FromPrimitive>(n: usize) -> T {
<T as FromPrimitive>::from_usize(n).expect("count is representable in every RealField")
}
fn from_f64<T: FromPrimitive>(x: f64) -> T {
<T as FromPrimitive>::from_f64(x).expect("constant is representable in every RealField")
}