use crate::causal_discovery::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use crate::causal_discovery::brcd::brcd_gate::{GateConfig, fit_logistic_gate};
use deep_causality_algebra::RealField;
use deep_causality_num::FromPrimitive;
use deep_causality_stats::{
RidgeConfig, fit_ridge as stats_fit_ridge, fit_ridge_streaming as stats_fit_ridge_streaming,
gaussian_log_density,
};
use std::borrow::Cow;
pub const RIDGE_DEFAULT: f64 = 1e-4;
const VARIANCE_FLOOR: f64 = 1e-12;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Transform {
None,
Log,
Log1p,
Yeojohnson,
}
#[derive(Debug, Clone, PartialEq)]
pub struct RidgeFit<T> {
pub beta: Vec<T>,
pub sigma2: T,
}
impl<T: RealField> RidgeFit<T> {
pub fn predict(&self, design_row: &[T]) -> T {
dot(&self.beta, design_row)
}
}
pub fn fit_ridge<T: RealField + FromPrimitive>(
x: &[Vec<T>],
y: &[T],
ridge: T,
) -> Result<RidgeFit<T>, BrcdError> {
let n = x.len();
if n == 0 {
return Err(BrcdError(BrcdErrorEnum::EmptyData));
}
if y.len() != n {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let p = x[0].len();
if p == 0 || x.iter().any(|r| r.len() != p) {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let fit = stats_fit_ridge(x, y, &RidgeConfig::new(ridge)).map_err(ridge_error)?;
Ok(RidgeFit {
beta: fit.beta,
sigma2: floor(fit.sigma2, from_f64::<T>(VARIANCE_FLOOR)),
})
}
fn ridge_error(error: deep_causality_stats::StatsError) -> BrcdError {
match error.kind() {
deep_causality_stats::StatsErrorEnum::EmptyInput(_) => BrcdError(BrcdErrorEnum::EmptyData),
deep_causality_stats::StatsErrorEnum::DimensionMismatch(_) => {
BrcdError(BrcdErrorEnum::DimensionMismatch)
}
_ => BrcdError(BrcdErrorEnum::SingularSystem),
}
}
pub fn effective_transform<T: RealField>(values: &[T], requested: Transform) -> Transform {
let neg_one = -T::one();
let zero = T::zero();
let any_lt_neg1 = values.iter().any(|&v| v < neg_one);
match requested {
Transform::None | Transform::Yeojohnson => requested,
Transform::Log => {
if any_lt_neg1 {
Transform::Yeojohnson
} else if values.iter().any(|&v| v <= zero) {
Transform::Log1p
} else {
Transform::Log
}
}
Transform::Log1p => {
if any_lt_neg1 {
Transform::Yeojohnson
} else {
Transform::Log1p
}
}
}
}
pub fn transform_and_jacobian<T: RealField>(x: T, kind: Transform) -> Result<(T, T), BrcdError> {
match kind {
Transform::None => Ok((x, T::zero())),
Transform::Log => {
if x <= T::zero() {
return Err(BrcdError(BrcdErrorEnum::InvalidTransformDomain));
}
let lx = x.ln();
Ok((lx, -lx))
}
Transform::Log1p => {
if x < -T::one() {
return Err(BrcdError(BrcdErrorEnum::InvalidTransformDomain));
}
let l1p = (T::one() + x).ln();
Ok((l1p, -l1p))
}
Transform::Yeojohnson => Err(BrcdError(BrcdErrorEnum::YeojohnsonUnsupported)),
}
}
pub fn gaussian_single_expert_logdensity<T: RealField + FromPrimitive>(
y: &[T],
parents: &[Vec<T>],
transform: Transform,
ridge: T,
) -> Result<Vec<T>, BrcdError> {
let n = y.len();
if n == 0 {
return Err(BrcdError(BrcdErrorEnum::EmptyData));
}
let p_feat = parents.first().map_or(0, Vec::len);
if !parents.is_empty() && (parents.len() != n || parents.iter().any(|r| r.len() != p_feat)) {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let eff = effective_transform(y, transform);
let mut z = Vec::with_capacity(n);
let mut log_jac = Vec::with_capacity(n);
for &yi in y {
let (zi, ji) = transform_and_jacobian(yi, eff)?;
z.push(zi);
log_jac.push(ji);
}
let (mu, sigma2): (Vec<T>, T) = if p_feat > 0 {
let mut x_fit = Vec::new();
let mut z_fit = Vec::new();
for i in 0..n {
if z[i].is_finite() && parents[i].iter().all(|v| v.is_finite()) {
x_fit.push(design_row(&parents[i]));
z_fit.push(z[i]);
}
}
if x_fit.is_empty() {
(vec![mean(&z); n], variance_ddof1(&z))
} else {
let fit = fit_ridge(&x_fit, &z_fit, ridge)?;
let mus = (0..n)
.map(|i| fit.predict(&design_row(&parents[i])))
.collect();
(mus, fit.sigma2)
}
} else {
(vec![mean(&z); n], variance_ddof1(&z))
};
let logdens = logpdf_rows(&z, &mu, sigma2);
Ok(add_jacobian(logdens, &log_jac))
}
#[derive(Debug, Clone)]
pub struct GaussianFamilyConfig<T> {
pub transform: Transform,
pub transform_parents: bool,
pub ridge: T,
pub gate: GateConfig<T>,
}
impl<T: RealField + FromPrimitive> Default for GaussianFamilyConfig<T> {
fn default() -> Self {
Self {
transform: Transform::None,
transform_parents: false,
ridge: from_f64::<T>(RIDGE_DEFAULT),
gate: GateConfig::default(),
}
}
}
pub fn gaussian_family_logdensity<T: RealField + FromPrimitive>(
node: &[T],
parents: &[Vec<T>],
f: Option<&[bool]>,
f_is_parent: bool,
config: &GaussianFamilyConfig<T>,
) -> Result<Vec<T>, BrcdError> {
let n = node.len();
if n == 0 {
return Err(BrcdError(BrcdErrorEnum::EmptyData));
}
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 let Some(fv) = f
&& fv.len() != n
{
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let eff = effective_transform(node, config.transform);
let mut z = Vec::with_capacity(n);
let mut log_jac = Vec::with_capacity(n);
for &yi in node {
let (zi, ji) = transform_and_jacobian(yi, eff)?;
z.push(zi);
log_jac.push(ji);
}
let parents_t = apply_parent_transform(parents, eff, config.transform_parents)?;
let logdens = match (f, f_is_parent) {
(None, _) => {
let all: Vec<usize> = (0..n).collect();
let (mean_model, var) = fit_expert(&all, &z, &parents_t, config.ridge);
let mu = predict_all(&mean_model, &parents_t, n);
logpdf_rows(&z, &mu, var)
}
(Some(fv), true) => {
let mut out = vec![T::zero(); n];
for regime in [false, true] {
let idxs: Vec<usize> = (0..n).filter(|&i| fv[i] == regime).collect();
if idxs.is_empty() {
continue;
}
let (mean_model, var) = fit_expert_guarded(&idxs, &z, &parents_t, config.ridge);
for &i in &idxs {
let mu = mean_model.predict(parents_t.get(i).map_or(&[][..], Vec::as_slice));
out[i] = single_logpdf(z[i], mu, var);
}
}
out
}
(Some(fv), false) => {
let idx0: Vec<usize> = (0..n).filter(|&i| !fv[i]).collect();
let idx1: Vec<usize> = (0..n).filter(|&i| fv[i]).collect();
let (m0, var0) = fit_expert(&idx0, &z, &parents_t, config.ridge);
let (m1, var1) = fit_expert(&idx1, &z, &parents_t, config.ridge);
let mu0 = predict_all(&m0, &parents_t, n);
let mu1 = predict_all(&m1, &parents_t, n);
let log_n0 = logpdf_rows(&z, &mu0, var0);
let log_n1 = logpdf_rows(&z, &mu1, var1);
let pi1 = gate_probabilities(&parents_t, fv, &config.gate, idx1.len(), n);
let one = T::one();
(0..n)
.map(|i| {
let p1 = clamp_unit(pi1[i]);
logaddexp((one - p1).ln() + log_n0[i], p1.ln() + log_n1[i])
})
.collect()
}
};
Ok(add_jacobian(logdens, &log_jac))
}
enum ExpertMean<T> {
Const(T),
Linear(Vec<T>),
}
impl<T: RealField> ExpertMean<T> {
fn predict(&self, features: &[T]) -> T {
match self {
ExpertMean::Const(m) => *m,
ExpertMean::Linear(beta) => predict_implicit(beta, features),
}
}
}
fn predict_implicit<T: RealField>(beta: &[T], features: &[T]) -> T {
let mut acc = beta[0];
for (b, f) in beta[1..].iter().zip(features.iter()) {
acc += *b * *f;
}
acc
}
fn predict_all<T: RealField>(model: &ExpertMean<T>, parents: &[Vec<T>], n: usize) -> Vec<T> {
(0..n)
.map(|i| model.predict(parents.get(i).map_or(&[][..], Vec::as_slice)))
.collect()
}
fn fit_expert<T: RealField + FromPrimitive>(
idxs: &[usize],
z_all: &[T],
parents_t: &[Vec<T>],
ridge: T,
) -> (ExpertMean<T>, T) {
if idxs.is_empty() {
return (ExpertMean::Const(mean(z_all)), variance_ddof1(z_all));
}
let p = parents_t.first().map_or(0, Vec::len);
if p == 0 {
let ys: Vec<T> = idxs.iter().map(|&i| z_all[i]).collect();
return (ExpertMean::Const(mean(&ys)), variance_ddof1(&ys));
}
match fit_ridge_streaming(idxs, z_all, parents_t, ridge, 0) {
Some(fit) => (ExpertMean::Linear(fit.beta), fit.sigma2),
None => {
let ys: Vec<T> = idxs.iter().map(|&i| z_all[i]).collect();
(ExpertMean::Const(mean(&ys)), variance_ddof1(&ys))
}
}
}
fn fit_expert_guarded<T: RealField + FromPrimitive>(
idxs: &[usize],
z_all: &[T],
parents_t: &[Vec<T>],
ridge: T,
) -> (ExpertMean<T>, T) {
let p = parents_t.first().map_or(0, Vec::len);
if p == 0 {
let ys: Vec<T> = idxs.iter().map(|&i| z_all[i]).collect();
return (ExpertMean::Const(mean(&ys)), variance_ddof1(&ys));
}
match fit_ridge_streaming(idxs, z_all, parents_t, ridge, p + 1) {
Some(fit) => (ExpertMean::Linear(fit.beta), fit.sigma2),
None => {
let ys: Vec<T> = idxs.iter().map(|&i| z_all[i]).collect();
(ExpertMean::Const(mean(&ys)), variance_ddof1(&ys))
}
}
}
fn fit_ridge_streaming<T: RealField + FromPrimitive>(
idxs: &[usize],
z_all: &[T],
parents_t: &[Vec<T>],
ridge: T,
min_finite: usize,
) -> Option<RidgeFit<T>> {
let p = parents_t.first().map_or(0, Vec::len) + 1;
let finite = |i: usize| z_all[i].is_finite() && parents_t[i].iter().all(|v| v.is_finite());
let count = idxs.iter().filter(|&&i| finite(i)).count();
if count <= min_finite {
return None;
}
let rows = idxs.iter().filter(move |&&i| finite(i)).map(move |&i| {
let mut design = Vec::with_capacity(p);
design.push(T::one());
design.extend_from_slice(&parents_t[i]);
(design, z_all[i])
});
let fit = stats_fit_ridge_streaming(rows, &RidgeConfig::new(ridge), p).ok()?;
Some(RidgeFit {
beta: fit.beta,
sigma2: floor(fit.sigma2, from_f64::<T>(VARIANCE_FLOOR)),
})
}
fn gate_probabilities<T: RealField + FromPrimitive>(
parents_t: &[Vec<T>],
f: &[bool],
gate: &GateConfig<T>,
ones: usize,
n: usize,
) -> Vec<T> {
let rows: Vec<Vec<T>> = (0..n)
.map(|i| parents_t.get(i).cloned().unwrap_or_default())
.collect();
match fit_logistic_gate(&rows, f, gate) {
Ok(model) => (0..n).map(|i| model.predict_proba(&rows[i])).collect(),
Err(_) => {
let prior = from_f64::<T>(ones as f64) / from_f64::<T>(n.max(1) as f64);
vec![prior; n]
}
}
}
fn logpdf_rows<T: RealField + FromPrimitive>(z: &[T], mu: &[T], sigma2: T) -> Vec<T> {
let var = density_variance(sigma2);
z.iter()
.zip(mu.iter())
.map(|(&zi, &mi)| logpdf_one(zi, mi, var))
.collect()
}
fn single_logpdf<T: RealField + FromPrimitive>(z: T, mu: T, sigma2: T) -> T {
logpdf_one(z, mu, density_variance(sigma2))
}
fn logpdf_one<T: RealField + FromPrimitive>(z: T, mu: T, var: T) -> T {
match gaussian_log_density(z, mu, var) {
Ok(value) => value,
Err(_) => {
let diff = z - mu;
if diff.is_nan() {
T::nan()
} else {
T::zero().ln()
}
}
}
}
fn density_variance<T: RealField + FromPrimitive>(sigma2: T) -> T {
if sigma2 > T::zero() {
sigma2
} else {
from_f64::<T>(VARIANCE_FLOOR)
}
}
fn add_jacobian<T: RealField>(mut logdens: Vec<T>, log_jac: &[T]) -> Vec<T> {
for (ld, &j) in logdens.iter_mut().zip(log_jac.iter()) {
*ld += j;
}
logdens
}
fn apply_parent_transform<T: RealField + FromPrimitive>(
parents: &[Vec<T>],
eff: Transform,
transform_parents: bool,
) -> Result<Cow<'_, [Vec<T>]>, BrcdError> {
if !transform_parents || eff == Transform::None {
return Ok(Cow::Borrowed(parents));
}
let transformed = parents
.iter()
.map(|row| {
row.iter()
.map(|&v| transform_and_jacobian(v, eff).map(|(z, _)| z))
.collect::<Result<Vec<T>, _>>()
})
.collect::<Result<Vec<Vec<T>>, _>>()?;
Ok(Cow::Owned(transformed))
}
fn logaddexp<T: RealField>(a: T, b: T) -> T {
deep_causality_stats::log_add_exp(a, b)
}
fn clamp_unit<T: RealField + FromPrimitive>(p: T) -> T {
let eps = from_f64::<T>(1e-12);
p.clamp(eps, T::one() - eps)
}
fn design_row<T: RealField>(features: &[T]) -> Vec<T> {
let mut row = Vec::with_capacity(features.len() + 1);
row.push(T::one());
row.extend_from_slice(features);
row
}
fn dot<T: RealField>(a: &[T], b: &[T]) -> T {
let n = a.len().min(b.len());
deep_causality_linear::dot(&a[..n], &b[..n]).unwrap_or_else(|_| T::zero())
}
fn mean<T: RealField + FromPrimitive>(v: &[T]) -> T {
deep_causality_stats::mean(v).unwrap_or_else(|_| T::zero())
}
fn variance_ddof1<T: RealField + FromPrimitive>(v: &[T]) -> T {
deep_causality_stats::variance(v).unwrap_or_else(|_| T::one())
}
fn floor<T: RealField>(x: T, floor: T) -> T {
if x > floor { x } else { floor }
}
fn from_f64<T: FromPrimitive>(x: f64) -> T {
<T as FromPrimitive>::from_f64(x).expect("constant is representable in every RealField")
}