use crate::causal_discovery::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use deep_causality_algebra::RealField;
use deep_causality_num::FromPrimitive;
use deep_causality_stats::{LogisticConfig, Penalisation, fit_logistic, sigmoid};
#[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 design: Vec<Vec<T>> = rows
.iter()
.map(|row| {
let mut z = Vec::with_capacity(p + 1);
z.push(T::one());
z.extend_from_slice(row);
z
})
.collect();
let labels: Vec<T> = y
.iter()
.map(|&v| if v { T::one() } else { T::zero() })
.collect();
let fit = fit_logistic(
&design,
&labels,
&LogisticConfig::new(config.ridge, config.max_iter, config.tol)
.with_penalisation(Penalisation::Excluding(0)),
)
.map_err(|_| BrcdError(BrcdErrorEnum::SingularSystem))?;
Ok(LogisticGate {
bias: fit.beta[0],
weights: fit.beta[1..].to_vec(),
})
}
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")
}