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    /// [`Self::posterior_sample`] drawn from `seed`: the same seed gives the
21    /// same matrix, whatever the thread count, so a fit trained on it
22    /// replays. Seeded per fixed-width chunk, like
23    /// [`Self::posterior_log_sample`].
24    ///
25    /// The default refuses: an implementation outside this crate that
26    /// predates the method has no seeded draw.
27    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    /// Draw a fresh sample of `log λ` per element. Delta-method
32    /// approximation: `log λ ≈ Normal(posterior_log_mean,
33    /// posterior_log_sd²)`. Caller must have called
34    /// `calibrate_with(CalibrateTarget::All)` first so log_mean / log_sd
35    /// are populated.
36    ///
37    /// `seed` makes the draw reproducible. It used to come from
38    /// `rand::rng()` inside a `map_init`, which is OS-seeded AND partitioned
39    /// by rayon, so a caller could not reproduce its own null distribution
40    /// even with every other seed pinned. Seeding per CHUNK rather than per
41    /// element keeps the draw parallel and independent of how rayon happens
42    /// to split the work.
43    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/// Which posterior quantities to compute during calibration.
50#[derive(Debug, Clone, Copy, PartialEq, Eq)]
51pub enum CalibrateTarget {
52    /// Compute all: mean, sd, log_mean, log_sd
53    All,
54    /// Only posterior mean (a/b)
55    MeanOnly,
56    /// Posterior mean + log mean (digamma(a) - ln(b))
57    MeanAndLogMean,
58}
59
60/// A parameter matrix with two types of statistics
61/// with hyper parameters a0 and b0
62pub 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    /// Calibrate all posterior quantities (mean, sd, log_mean, log_sd)
73    fn calibrate(&mut self) {
74        self.calibrate_with(CalibrateTarget::All);
75    }
76
77    /// Calibrate only the posterior quantities specified by `target`.
78    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
102/// One Gamma draw per `(a, b)` pair, shape `a + ε` and rate `b + ε`, from
103/// `seed`: one generator per fixed-width chunk, so the draw is a function of
104/// the seed and the data shape, not of how rayon splits the work.
105pub(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}