use crate::causal_discovery::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use crate::causal_discovery::brcd::brcd_gate::{GateConfig, fit_logistic_gate};
use crate::causal_discovery::brcd::brcd_linalg::solve_linear;
use deep_causality_algebra::RealField;
use deep_causality_num::FromPrimitive;
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 mut xtx = vec![T::zero(); p * p];
let mut xty = vec![T::zero(); p];
for (row, &yi) in x.iter().zip(y.iter()) {
for a in 0..p {
xty[a] += row[a] * yi;
let ra = row[a];
for b in 0..p {
xtx[a * p + b] += ra * row[b];
}
}
}
for a in 0..p {
xtx[a * p + a] += ridge;
}
solve_linear(&mut xtx, &mut xty, p);
let beta = xty;
let mut rss = T::zero();
for (row, &yi) in x.iter().zip(y.iter()) {
let r = yi - dot(&beta, row);
rss += r * r;
}
let dof = t_usize::<T>(n.saturating_sub(p).max(1));
let sigma2 = floor(rss / dof, from_f64::<T>(VARIANCE_FLOOR));
Ok(RidgeFit { beta, sigma2 })
}
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 mut xtx = vec![T::zero(); p * p];
let mut xty = vec![T::zero(); p];
let mut design = vec![T::zero(); p];
design[0] = T::one();
let mut count = 0usize;
for &i in idxs {
if z_all[i].is_finite() && parents_t[i].iter().all(|v| v.is_finite()) {
design[1..].copy_from_slice(&parents_t[i]);
let yi = z_all[i];
for a in 0..p {
xty[a] += design[a] * yi;
let ra = design[a];
for b in 0..p {
xtx[a * p + b] += ra * design[b];
}
}
count += 1;
}
}
if count <= min_finite {
return None;
}
for a in 0..p {
xtx[a * p + a] += ridge;
}
solve_linear(&mut xtx, &mut xty, p);
let beta = xty;
let mut rss = T::zero();
for &i in idxs {
if z_all[i].is_finite() && parents_t[i].iter().all(|v| v.is_finite()) {
design[1..].copy_from_slice(&parents_t[i]);
let r = z_all[i] - dot(&beta, &design);
rss += r * r;
}
}
let dof = t_usize::<T>(count.saturating_sub(p).max(1));
let sigma2 = floor(rss / dof, from_f64::<T>(VARIANCE_FLOOR));
Some(RidgeFit { beta, sigma2 })
}
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);
let half = from_f64::<T>(0.5);
let two = from_f64::<T>(2.0);
let log_two_pi_var = (two * T::pi() * var).ln();
z.iter()
.zip(mu.iter())
.map(|(&zi, &mi)| {
let diff = zi - mi;
-half * (log_two_pi_var + (diff * diff) / var)
})
.collect()
}
fn single_logpdf<T: RealField + FromPrimitive>(z: T, mu: T, sigma2: T) -> T {
let var = density_variance(sigma2);
let half = from_f64::<T>(0.5);
let two = from_f64::<T>(2.0);
let log_two_pi_var = (two * T::pi() * var).ln();
let diff = z - mu;
-half * (log_two_pi_var + (diff * diff) / var)
}
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 {
let m = if a >= b { a } else { b };
if !m.is_finite() {
return m;
}
m + ((a - m).exp() + (b - m).exp()).ln()
}
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 {
a.iter()
.zip(b.iter())
.fold(T::zero(), |acc, (&x, &y)| acc + x * y)
}
fn mean<T: RealField + FromPrimitive>(v: &[T]) -> T {
if v.is_empty() {
return T::zero();
}
v.iter().fold(T::zero(), |acc, &x| acc + x) / t_usize::<T>(v.len())
}
fn variance_ddof1<T: RealField + FromPrimitive>(v: &[T]) -> T {
if v.len() < 2 {
return T::one();
}
let m = mean(v);
let ss = v.iter().fold(T::zero(), |acc, &x| {
let d = x - m;
acc + d * d
});
ss / t_usize::<T>(v.len() - 1)
}
fn floor<T: RealField>(x: T, floor: T) -> T {
if x > floor { x } else { floor }
}
fn t_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")
}