Skip to main content

legume_numeric/param/
traits.rs

1#[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    // Rows as tensors only when candle is linked (feature `tensor`).
9    #[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    /// Draw a fresh sample of `log λ` per element. Delta-method
21    /// approximation: `log λ ≈ Normal(posterior_log_mean,
22    /// posterior_log_sd²)`. Caller must have called
23    /// `calibrate_with(CalibrateTarget::All)` first so log_mean / log_sd
24    /// are populated.
25    ///
26    /// `seed` makes the draw reproducible. It used to come from
27    /// `rand::rng()` inside a `map_init`, which is OS-seeded AND partitioned
28    /// by rayon, so a caller could not reproduce its own null distribution
29    /// even with every other seed pinned. Seeding per CHUNK rather than per
30    /// element keeps the draw parallel and independent of how rayon happens
31    /// to split the work.
32    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/// Which posterior quantities to compute during calibration.
39#[derive(Debug, Clone, Copy, PartialEq, Eq)]
40pub enum CalibrateTarget {
41    /// Compute all: mean, sd, log_mean, log_sd
42    All,
43    /// Only posterior mean (a/b)
44    MeanOnly,
45    /// Posterior mean + log mean (digamma(a) - ln(b))
46    MeanAndLogMean,
47}
48
49/// A parameter matrix with two types of statistics
50/// with hyper parameters a0 and b0
51pub 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    /// Calibrate all posterior quantities (mean, sd, log_mean, log_sd)
62    fn calibrate(&mut self) {
63        self.calibrate_with(CalibrateTarget::All);
64    }
65
66    /// Calibrate only the posterior quantities specified by `target`.
67    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}