use crate::brcd::brcd_gate::GateConfig;
use crate::brcd::brcd_gaussian::{RIDGE_DEFAULT, Transform};
use deep_causality_algebra::RealField;
use deep_causality_num::FromPrimitive;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FamilyKind {
Continuous,
Discrete,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ConfigStrategy {
#[default]
Full,
MapPrune,
}
#[derive(Debug, Clone)]
pub struct BrcdConfig<T> {
pub seed: u64,
pub family: FamilyKind,
pub node_transform: Transform,
pub transform_parents: bool,
pub num_root_causes: usize,
pub ridge: T,
pub alpha_star: T,
pub gate: GateConfig<T>,
pub config_strategy: ConfigStrategy,
}
impl<T: RealField + FromPrimitive> BrcdConfig<T> {
pub fn continuous(seed: u64) -> Self {
Self {
seed,
family: FamilyKind::Continuous,
node_transform: Transform::None,
transform_parents: false,
num_root_causes: 1,
ridge: from_f64(RIDGE_DEFAULT),
alpha_star: from_f64(5.0),
gate: GateConfig::default(),
config_strategy: ConfigStrategy::Full,
}
}
pub fn discrete(seed: u64) -> Self {
Self {
family: FamilyKind::Discrete,
..Self::continuous(seed)
}
}
}
impl<T: RealField + FromPrimitive> Default for BrcdConfig<T> {
fn default() -> Self {
Self::continuous(0)
}
}
fn from_f64<T: FromPrimitive>(x: f64) -> T {
<T as FromPrimitive>::from_f64(x).expect("constant is representable in every RealField")
}