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_sample_seeded(&self, seed: u64) -> anyhow::Result<Self::Mat> {
28 let _ = seed;
29 anyhow::bail!("this parameter type has no seeded posterior sample")
30 }
31 fn posterior_log_sample(&self, seed: u64) -> anyhow::Result<Self::Mat>;
44
45 fn nrows(&self) -> usize;
46 fn ncols(&self) -> usize;
47}
48
49#[derive(Debug, Clone, Copy, PartialEq, Eq)]
51pub enum CalibrateTarget {
52 All,
54 MeanOnly,
56 MeanAndLogMean,
58}
59
60pub trait TwoStatParam {
63 type Mat;
64 type Scalar;
65
66 fn new(dims: (usize, usize), a0: Self::Scalar, b0: Self::Scalar) -> Self;
67 fn add_stat(&mut self, add_a: &Self::Mat, add_b: &Self::Mat);
68 fn update_stat(&mut self, update_a: &Self::Mat, update_b: &Self::Mat);
69 fn update_stat_col(&mut self, update_a: &Self::Mat, update_b: &Self::Mat, k: usize);
70 fn reset_stat(&mut self);
71
72 fn calibrate(&mut self) {
74 self.calibrate_with(CalibrateTarget::All);
75 }
76
77 fn calibrate_with(&mut self, target: CalibrateTarget) {
79 match target {
80 CalibrateTarget::All => {
81 self.map_calibrate_mean();
82 self.map_calibrate_log_mean();
83 self.map_calibrate_sd();
84 self.map_calibrate_log_sd();
85 }
86 CalibrateTarget::MeanOnly => {
87 self.map_calibrate_mean();
88 }
89 CalibrateTarget::MeanAndLogMean => {
90 self.map_calibrate_mean();
91 self.map_calibrate_log_mean();
92 }
93 }
94 }
95
96 fn map_calibrate_mean(&mut self);
97 fn map_calibrate_sd(&mut self);
98 fn map_calibrate_log_mean(&mut self);
99 fn map_calibrate_log_sd(&mut self);
100}
101
102pub(crate) fn gamma_sample_seeded(a: &[f32], b: &[f32], seed: u64) -> anyhow::Result<Vec<f32>> {
106 use rand::rngs::SmallRng;
107 use rand::SeedableRng;
108 use rand_distr::{Distribution, Gamma};
109 use rayon::prelude::*;
110 const CHUNK: usize = 1024;
111 let eps = 1e-8;
112 let mut sampled = vec![0.0f32; a.len()];
113 sampled
114 .par_chunks_mut(CHUNK)
115 .enumerate()
116 .try_for_each(|(ci, out)| -> anyhow::Result<()> {
117 let mut rng =
118 SmallRng::seed_from_u64(seed ^ (ci as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
119 let base = ci * CHUNK;
120 for (k, o) in out.iter_mut().enumerate() {
121 let shape = a[base + k] + eps;
122 let scale = (b[base + k] + eps).recip();
123 *o = Gamma::new(shape, scale)?.sample(&mut rng);
124 }
125 Ok(())
126 })?;
127 Ok(sampled)
128}