legume_numeric/param/
traits.rs1#[cfg(feature = "tensor")]
2use crate::matrix::traits::CandleDataLoaderOps;
3use crate::matrix::traits::{IoOps, MeltOps};
4
5pub trait TwoStatInference: Inference + TwoStatParam {}
6
7pub trait Inference {
8 #[cfg(feature = "tensor")]
10 type Mat: IoOps + MeltOps + CandleDataLoaderOps;
11 #[cfg(not(feature = "tensor"))]
12 type Mat: IoOps + MeltOps;
13 type Scalar: Into<f32>;
14
15 fn posterior_mean(&self) -> &Self::Mat;
16 fn posterior_sd(&self) -> &Self::Mat;
17 fn posterior_log_mean(&self) -> &Self::Mat;
18 fn posterior_log_sd(&self) -> &Self::Mat;
19 fn posterior_sample(&self) -> anyhow::Result<Self::Mat>;
20 fn posterior_log_sample(&self, seed: u64) -> anyhow::Result<Self::Mat>;
33
34 fn nrows(&self) -> usize;
35 fn ncols(&self) -> usize;
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq)]
40pub enum CalibrateTarget {
41 All,
43 MeanOnly,
45 MeanAndLogMean,
47}
48
49pub trait TwoStatParam {
52 type Mat;
53 type Scalar;
54
55 fn new(dims: (usize, usize), a0: Self::Scalar, b0: Self::Scalar) -> Self;
56 fn add_stat(&mut self, add_a: &Self::Mat, add_b: &Self::Mat);
57 fn update_stat(&mut self, update_a: &Self::Mat, update_b: &Self::Mat);
58 fn update_stat_col(&mut self, update_a: &Self::Mat, update_b: &Self::Mat, k: usize);
59 fn reset_stat(&mut self);
60
61 fn calibrate(&mut self) {
63 self.calibrate_with(CalibrateTarget::All);
64 }
65
66 fn calibrate_with(&mut self, target: CalibrateTarget) {
68 match target {
69 CalibrateTarget::All => {
70 self.map_calibrate_mean();
71 self.map_calibrate_log_mean();
72 self.map_calibrate_sd();
73 self.map_calibrate_log_sd();
74 }
75 CalibrateTarget::MeanOnly => {
76 self.map_calibrate_mean();
77 }
78 CalibrateTarget::MeanAndLogMean => {
79 self.map_calibrate_mean();
80 self.map_calibrate_log_mean();
81 }
82 }
83 }
84
85 fn map_calibrate_mean(&mut self);
86 fn map_calibrate_sd(&mut self);
87 fn map_calibrate_log_mean(&mut self);
88 fn map_calibrate_log_sd(&mut self);
89}