legume_numeric/param/
ndarray_gamma.rs1extern 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 a0: f32,
17 b0: f32,
18 a_stat: Array2<f32>,
22 b_stat: Array2<f32>,
23 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 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 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 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 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}