causal_hub/estimators/parameters/sufficient_statistics/table/
mixed.rs1use crate::{
2 estimators::{CSSEstimator, ParCSSEstimator, SSE},
3 models::{MixedCPDS, MixedIncTable, MixedTable, MixedWtdTable},
4 types::Set,
5};
6
7macro_rules! impl_css_for_mixed {
8 ($enum:ident) => {
9 impl CSSEstimator<MixedCPDS> for SSE<'_, $enum> {
10 fn fit(&self, x: &Set<usize>, z: &Set<usize>) -> crate::types::Result<MixedCPDS> {
11 match self.dataset {
12 $enum::Categorical(t) => SSE::new(t)
13 .with_missing_method(self.missing_method, self.missing_mechanism.clone())?
14 .fit(x, z)
15 .map(MixedCPDS::Categorical),
16 $enum::Gaussian(t) => SSE::new(t)
17 .with_missing_method(self.missing_method, self.missing_mechanism.clone())?
18 .fit(x, z)
19 .map(|stats| MixedCPDS::Gaussian(Box::new(stats))),
20 }
21 }
22 }
23 impl ParCSSEstimator<MixedCPDS> for SSE<'_, $enum> {
24 fn par_fit(&self, x: &Set<usize>, z: &Set<usize>) -> crate::types::Result<MixedCPDS> {
25 match self.dataset {
26 $enum::Categorical(t) => SSE::new(t)
27 .with_missing_method(self.missing_method, self.missing_mechanism.clone())?
28 .par_fit(x, z)
29 .map(MixedCPDS::Categorical),
30 $enum::Gaussian(t) => SSE::new(t)
31 .with_missing_method(self.missing_method, self.missing_mechanism.clone())?
32 .par_fit(x, z)
33 .map(|stats| MixedCPDS::Gaussian(Box::new(stats))),
34 }
35 }
36 }
37 };
38}
39
40impl_css_for_mixed!(MixedTable);
41impl_css_for_mixed!(MixedIncTable);
42impl_css_for_mixed!(MixedWtdTable);