use crate::causal_discovery::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use crate::causal_discovery::brcd::brcd_linalg::solve_linear;
use deep_causality_num::{FromPrimitive, RealField};
#[derive(Debug, Clone, Copy)]
pub struct GateConfig<T> {
pub ridge: T,
pub max_iter: usize,
pub tol: T,
}
impl<T: RealField + FromPrimitive> Default for GateConfig<T> {
fn default() -> Self {
Self {
ridge: T::one(),
max_iter: 100,
tol: from_f64::<T>(1e-8),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct LogisticGate<T> {
bias: T,
weights: Vec<T>,
}
impl<T: RealField + FromPrimitive> LogisticGate<T> {
pub fn predict_proba(&self, x: &[T]) -> T {
let mut eta = self.bias;
for (w, &xi) in self.weights.iter().zip(x.iter()) {
eta += *w * xi;
}
sigmoid(eta)
}
pub fn bias(&self) -> T {
self.bias
}
pub fn weights(&self) -> &[T] {
&self.weights
}
}
pub fn fit_logistic_gate<T: RealField + FromPrimitive>(
rows: &[Vec<T>],
y: &[bool],
config: &GateConfig<T>,
) -> Result<LogisticGate<T>, BrcdError> {
let n = rows.len();
if n == 0 {
return Err(BrcdError(BrcdErrorEnum::EmptyData));
}
if y.len() != n {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let p = rows[0].len();
if rows.iter().any(|r| r.len() != p) {
return Err(BrcdError(BrcdErrorEnum::DimensionMismatch));
}
let ones = y.iter().filter(|&&v| v).count();
if ones == 0 || ones == n {
let rate = from_f64::<T>(ones as f64) / from_f64::<T>(n as f64);
return Ok(LogisticGate {
bias: logit_clamped(rate),
weights: vec![T::zero(); p],
});
}
let dim = p + 1;
let mut theta = vec![T::zero(); dim];
for _ in 0..config.max_iter {
let mut grad = vec![T::zero(); dim];
let mut hess = vec![T::zero(); dim * dim];
for (row, &label) in rows.iter().zip(y.iter()) {
let eta = theta[0] + dot(&theta[1..], row);
let pi = sigmoid(eta);
let w = pi * (T::one() - pi);
let resid = pi - if label { T::one() } else { T::zero() };
accumulate_row(&mut grad, &mut hess, dim, row, resid, w);
}
for a in 1..dim {
grad[a] += config.ridge * theta[a];
hess[a * dim + a] += config.ridge;
}
solve_linear(&mut hess, &mut grad, dim);
let mut max_step = T::zero();
for a in 0..dim {
theta[a] -= grad[a];
let s = grad[a].abs();
if s > max_step {
max_step = s;
}
}
if theta.iter().any(|t| !t.is_finite()) {
return Err(BrcdError(BrcdErrorEnum::SingularSystem));
}
if max_step < config.tol {
break;
}
}
Ok(LogisticGate {
bias: theta[0],
weights: theta[1..].to_vec(),
})
}
fn accumulate_row<T: RealField>(
grad: &mut [T],
hess: &mut [T],
dim: usize,
row: &[T],
resid: T,
w: T,
) {
let z = |a: usize| if a == 0 { T::one() } else { row[a - 1] };
for a in 0..dim {
let za = z(a);
grad[a] += za * resid;
let zaw = za * w;
for b in 0..dim {
hess[a * dim + b] += zaw * z(b);
}
}
}
fn dot<T: RealField>(a: &[T], b: &[T]) -> T {
a.iter()
.zip(b.iter())
.fold(T::zero(), |acc, (&x, &y)| acc + x * y)
}
fn sigmoid<T: RealField>(x: T) -> T {
let one = T::one();
if x >= T::zero() {
one / (one + (-x).exp())
} else {
let e = x.exp();
e / (one + e)
}
}
fn logit_clamped<T: RealField + FromPrimitive>(p: T) -> T {
let eps = from_f64::<T>(1e-12);
let one = T::one();
let clamped = p.clamp(eps, one - eps);
(clamped / (one - clamped)).ln()
}
fn from_f64<T: FromPrimitive>(x: f64) -> T {
<T as FromPrimitive>::from_f64(x).expect("constant is representable in every RealField")
}