Skip to main content

legume_numeric/param/
ndarray_gamma.rs

1extern crate special;
2
3use crate::param::io::*;
4use crate::param::traits::*;
5use ndarray::prelude::*;
6use rayon::prelude::*;
7
8#[allow(dead_code)]
9#[derive(Debug, Clone)]
10pub struct GammaMatrix {
11    num_rows: usize,
12    num_columns: usize,
13    //////////////////////
14    // hyper parameters //
15    //////////////////////
16    a0: f32,
17    b0: f32,
18    ///////////////////////////
19    // sufficient statistics //
20    ///////////////////////////
21    a_stat: Array2<f32>,
22    b_stat: Array2<f32>,
23    //////////////////////////
24    // estimated parameters //
25    //////////////////////////
26    estimated_mean: Array2<f32>,
27    estimated_sd: Array2<f32>,
28    estimated_log_mean: Array2<f32>,
29    estimated_log_sd: Array2<f32>,
30}
31
32impl ParamIo for GammaMatrix {
33    type Mat = Array2<f32>;
34}
35
36impl TwoStatParam for GammaMatrix {
37    type Mat = Array2<f32>;
38    type Scalar = f32;
39
40    /// New Poisson-Gamma parameter matrix
41    ///
42    /// ```text
43    /// x[i,j] ~ Poisson(lambda[i,j])
44    /// lambda[i,j] ~ Gamma(a0, b0)
45    /// ```
46    ///
47    /// #Arguments
48    /// * `dims` - dimensions of the matrix (num of rows, num of columns)
49    /// * `a` - hyper parameter a0
50    /// * `b` - hyper parameter b0
51    ///
52    fn new(dims: (usize, usize), a: Self::Scalar, b: Self::Scalar) -> Self {
53        Self {
54            num_rows: dims.0,
55            num_columns: dims.1,
56            a0: a,
57            b0: b,
58            a_stat: Self::Mat::zeros(dims).mapv_into(|x| x + a),
59            b_stat: Self::Mat::zeros(dims).mapv_into(|x| x + b),
60            estimated_mean: Self::Mat::zeros(dims),
61            estimated_sd: Self::Mat::zeros(dims),
62            estimated_log_mean: Self::Mat::zeros(dims),
63            estimated_log_sd: Self::Mat::zeros(dims),
64        }
65    }
66
67    fn add_stat(&mut self, add_a: &Self::Mat, add_b: &Self::Mat) {
68        self.a_stat += add_a;
69        self.b_stat += add_b;
70    }
71
72    fn update_stat(&mut self, add_a: &Self::Mat, add_b: &Self::Mat) {
73        self.reset_stat();
74        self.add_stat(add_a, add_b);
75    }
76
77    fn update_stat_col(&mut self, add_a: &Self::Mat, add_b: &Self::Mat, k: usize) {
78        self.a_stat.column_mut(k).zip_mut_with(add_a, |x, add_x| {
79            *x = self.a0 + add_x;
80        });
81        self.b_stat.column_mut(k).zip_mut_with(add_b, |x, add_x| {
82            *x = self.b0 + add_x;
83        });
84    }
85
86    fn reset_stat(&mut self) {
87        self.a_stat.fill(self.a0);
88        self.b_stat.fill(self.b0);
89    }
90
91    // fn nrows(&self) -> usize {
92    //     self.num_rows
93    // }
94
95    // fn ncols(&self) -> usize {
96    //     self.num_columns
97    // }
98
99    // fn len(&self) -> usize {
100    //     self.num_rows * self.num_columns
101    // }
102    fn map_calibrate_mean(&mut self) {
103        self.estimated_mean = &self.a_stat / &self.b_stat;
104    }
105    fn map_calibrate_sd(&mut self) {
106        self.estimated_sd = &self.a_stat.mapv(|x| x.sqrt()) / &self.b_stat;
107    }
108    fn map_calibrate_log_mean(&mut self) {
109        use special::Gamma;
110        self.estimated_log_mean = &self.a_stat.mapv(Gamma::digamma) - &self.b_stat.mapv(|b| b.ln());
111    }
112    fn map_calibrate_log_sd(&mut self) {
113        // `sd[ln X] = sqrt(trigamma(a))` exactly — see the `dmatrix_gamma`
114        // sibling for why the old `1/sqrt(a - 1)` was wrong wherever counts are
115        // sparse, and why returning 0 below `a = 1` inverted the truth.
116        self.estimated_log_sd = self.a_stat.mapv(|a: f32| {
117            use special::Gamma;
118            a.trigamma().sqrt()
119        });
120    }
121}
122
123impl Inference for GammaMatrix {
124    type Mat = Array2<f32>;
125    type Scalar = f32;
126
127    fn posterior_mean(&self) -> &Self::Mat {
128        &self.estimated_mean
129    }
130
131    fn posterior_sd(&self) -> &Self::Mat {
132        &self.estimated_sd
133    }
134
135    fn posterior_log_mean(&self) -> &Self::Mat {
136        &self.estimated_log_mean
137    }
138
139    fn posterior_log_sd(&self) -> &Self::Mat {
140        &self.estimated_log_sd
141    }
142
143    fn posterior_sample(&self) -> anyhow::Result<Self::Mat> {
144        use rand_distr::{Distribution, Gamma};
145        let eps = 1e-8;
146
147        let a_slice = self
148            .a_stat
149            .as_slice()
150            .ok_or(anyhow::anyhow!("failed to take slice on a_stat"))?;
151        let b_slice = self
152            .b_stat
153            .as_slice()
154            .ok_or(anyhow::anyhow!("failed to take slice on b_stat"))?;
155
156        let sampled = a_slice
157            .par_iter()
158            .zip(b_slice.par_iter())
159            .map_init(rand::rng, |rng, (&a, &b)| -> anyhow::Result<f32> {
160                let shape = a + eps;
161                let scale = (b + eps).recip();
162                let pdf = Gamma::new(shape, scale)?;
163                Ok(pdf.sample(rng))
164            })
165            .collect::<anyhow::Result<Vec<_>>>()?;
166
167        Ok(Self::Mat::from_shape_vec(
168            (self.nrows(), self.ncols()),
169            sampled,
170        )?)
171    }
172
173    fn posterior_log_sample(&self, seed: u64) -> anyhow::Result<Self::Mat> {
174        use rand::rngs::SmallRng;
175        use rand::SeedableRng;
176        use rand_distr::{Distribution, StandardNormal};
177
178        // Fixed chunk width, so which elements share an RNG is a property of
179        // the data shape and not of the thread pool.
180        const CHUNK: usize = 1024;
181
182        let m_slice = self.estimated_log_mean.as_slice().ok_or(anyhow::anyhow!(
183            "failed to take slice on estimated_log_mean"
184        ))?;
185        let s_slice = self
186            .estimated_log_sd
187            .as_slice()
188            .ok_or(anyhow::anyhow!("failed to take slice on estimated_log_sd"))?;
189
190        let mut sampled = vec![0.0f32; m_slice.len()];
191        sampled
192            .par_chunks_mut(CHUNK)
193            .enumerate()
194            .for_each(|(ci, out)| {
195                let mut rng =
196                    SmallRng::seed_from_u64(seed ^ (ci as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
197                let base = ci * CHUNK;
198                for (k, o) in out.iter_mut().enumerate() {
199                    let z: f32 = StandardNormal.sample(&mut rng);
200                    *o = m_slice[base + k] + s_slice[base + k] * z;
201                }
202            });
203
204        Ok(Self::Mat::from_shape_vec(
205            (self.nrows(), self.ncols()),
206            sampled,
207        )?)
208    }
209
210    fn nrows(&self) -> usize {
211        self.num_rows
212    }
213
214    fn ncols(&self) -> usize {
215        self.num_columns
216    }
217}