1#![allow(dead_code)]
2
3extern crate special;
4
5use crate::param::io::*;
6use crate::param::traits::*;
7use nalgebra::{DMatrix, DVector};
8use rayon::prelude::*;
9
10#[derive(Debug, Clone)]
11pub struct GammaMatrix {
12 num_rows: usize,
13 num_columns: usize,
14 a0: f32,
18 b0: f32,
19 row_prior: Option<(DVector<f32>, DVector<f32>)>,
22 a_stat: DMatrix<f32>,
26 b_stat: DMatrix<f32>,
27 estimated_mean: DMatrix<f32>,
31 estimated_sd: DMatrix<f32>,
32 estimated_log_mean: DMatrix<f32>,
33 estimated_log_sd: DMatrix<f32>,
34}
35
36impl ParamIo for GammaMatrix {
37 type Mat = DMatrix<f32>;
38}
39
40impl TwoStatParam for GammaMatrix {
41 type Mat = DMatrix<f32>;
42 type Scalar = f32;
43
44 fn new(dims: (usize, usize), a: Self::Scalar, b: Self::Scalar) -> Self {
45 Self {
46 num_rows: dims.0,
47 num_columns: dims.1,
48 a0: a,
49 b0: b,
50 row_prior: None,
51 a_stat: DMatrix::from_element(dims.0, dims.1, a),
52 b_stat: DMatrix::from_element(dims.0, dims.1, b),
53 estimated_mean: DMatrix::zeros(dims.0, dims.1),
57 estimated_sd: DMatrix::zeros(0, 0),
63 estimated_log_mean: DMatrix::zeros(0, 0),
64 estimated_log_sd: DMatrix::zeros(0, 0),
65 }
66 }
67
68 fn add_stat(&mut self, add_a: &Self::Mat, add_b: &Self::Mat) {
69 self.a_stat += add_a;
70 self.b_stat += add_b;
71 }
72 fn update_stat(&mut self, update_a: &Self::Mat, update_b: &Self::Mat) {
73 self.reset_stat();
74 self.add_stat(update_a, update_b);
75 }
76 fn reset_stat(&mut self) {
77 match &self.row_prior {
78 None => {
79 self.a_stat.fill(self.a0);
80 self.b_stat.fill(self.b0);
81 }
82 Some((a0, b0)) => {
83 for mut col in self.a_stat.column_iter_mut() {
86 col.copy_from(a0);
87 }
88 for mut col in self.b_stat.column_iter_mut() {
89 col.copy_from(b0);
90 }
91 }
92 }
93 }
94 fn update_stat_col(&mut self, update_a: &Self::Mat, update_b: &Self::Mat, k: usize) {
95 match &self.row_prior {
96 None => {
97 self.a_stat
98 .column_mut(k)
99 .copy_from(&update_a.map(|x| x + self.a0));
100 self.b_stat
101 .column_mut(k)
102 .copy_from(&update_b.map(|x| x + self.b0));
103 }
104 Some((a0, b0)) => {
105 let mut a = self.a_stat.column_mut(k);
106 a.copy_from(update_a);
107 a += a0;
108 let mut b = self.b_stat.column_mut(k);
109 b.copy_from(update_b);
110 b += b0;
111 }
112 }
113 }
114
115 fn map_calibrate_mean(&mut self) {
127 self.estimated_mean = self.a_stat.zip_map(&self.b_stat, |a, b| a / b);
128 }
129 fn map_calibrate_sd(&mut self) {
130 self.estimated_sd = self.a_stat.zip_map(&self.b_stat, |a, b| a.sqrt() / b);
131 }
132 fn map_calibrate_log_mean(&mut self) {
133 use special::Gamma;
134 self.estimated_log_mean = self
135 .a_stat
136 .zip_map(&self.b_stat, |a, b| a.digamma() - b.ln());
137 }
138 fn map_calibrate_log_sd(&mut self) {
139 use special::Gamma;
152 self.estimated_log_sd = self.a_stat.map(|a| a.trigamma().sqrt());
153 }
154}
155
156impl Inference for GammaMatrix {
157 type Mat = DMatrix<f32>;
158 type Scalar = f32;
159
160 fn posterior_mean(&self) -> &Self::Mat {
161 &self.estimated_mean
162 }
163
164 fn posterior_sd(&self) -> &Self::Mat {
165 &self.estimated_sd
166 }
167
168 fn posterior_log_mean(&self) -> &Self::Mat {
169 &self.estimated_log_mean
170 }
171
172 fn posterior_log_sd(&self) -> &Self::Mat {
173 &self.estimated_log_sd
174 }
175
176 fn posterior_sample(&self) -> anyhow::Result<Self::Mat> {
177 use rand_distr::{Distribution, Gamma};
178 let eps = 1e-8;
179
180 let sampled = self
181 .a_stat
182 .as_slice()
183 .par_iter()
184 .zip(self.b_stat.as_slice().par_iter())
185 .map_init(rand::rng, |rng, (&a, &b)| -> anyhow::Result<f32> {
186 let shape = a + eps;
187 let scale = (b + eps).recip();
188 let pdf = Gamma::new(shape, scale)?;
189 Ok(pdf.sample(rng))
190 })
191 .collect::<anyhow::Result<Vec<_>>>()?;
192
193 Ok(Self::Mat::from_vec(self.nrows(), self.ncols(), sampled))
194 }
195
196 fn posterior_log_sample(&self, seed: u64) -> anyhow::Result<Self::Mat> {
197 use rand::rngs::SmallRng;
198 use rand::SeedableRng;
199 use rand_distr::{Distribution, StandardNormal};
200
201 const CHUNK: usize = 1024;
204 let m_slice = self.estimated_log_mean.as_slice();
205 let s_slice = self.estimated_log_sd.as_slice();
206 let mut sampled = vec![0.0f32; m_slice.len()];
207 sampled
208 .par_chunks_mut(CHUNK)
209 .enumerate()
210 .for_each(|(ci, out)| {
211 let mut rng =
212 SmallRng::seed_from_u64(seed ^ (ci as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
213 let base = ci * CHUNK;
214 for (k, o) in out.iter_mut().enumerate() {
215 let z: f32 = StandardNormal.sample(&mut rng);
216 *o = m_slice[base + k] + s_slice[base + k] * z;
217 }
218 });
219
220 Ok(Self::Mat::from_vec(self.nrows(), self.ncols(), sampled))
221 }
222
223 fn nrows(&self) -> usize {
224 self.num_rows
225 }
226
227 fn ncols(&self) -> usize {
228 self.num_columns
229 }
230}
231
232fn stack_field<F>(
236 blocks: &[GammaMatrix],
237 nrows: usize,
238 ncols: usize,
239 enabled: bool,
240 sel: F,
241) -> DMatrix<f32>
242where
243 F: Fn(&GammaMatrix) -> &DMatrix<f32>,
244{
245 if !enabled || blocks.is_empty() || sel(&blocks[0]).nrows() == 0 {
246 return DMatrix::zeros(0, 0);
247 }
248 let mut out = DMatrix::zeros(nrows, ncols);
249 let mut r0 = 0;
250 for b in blocks {
251 let src = sel(b);
252 out.rows_mut(r0, src.nrows()).copy_from(src);
253 r0 += src.nrows();
254 }
255 out
256}
257
258impl GammaMatrix {
259 pub fn release_stats(&mut self) {
265 self.a_stat = DMatrix::zeros(0, 0);
266 self.b_stat = DMatrix::zeros(0, 0);
267 }
268
269 pub fn sparsify_mean_to_support(&mut self, numerator: &DMatrix<f32>) {
278 debug_assert_eq!(self.estimated_mean.shape(), numerator.shape());
279 self.estimated_mean
280 .iter_mut()
281 .zip(numerator.iter())
282 .for_each(|(m, &n)| {
283 if n == 0.0 {
284 *m = 0.0;
285 }
286 });
287 }
288
289 #[must_use]
298 pub fn has_data_support(&self, row: usize, col: usize) -> bool {
299 self.a_stat[(row, col)] > self.a0_at(row)
300 }
301
302 #[must_use]
314 pub fn evidence_mean(&self, row: usize, col: usize) -> f32 {
315 let a = self.a_stat[(row, col)] - self.a0_at(row);
316 let b = self.b_stat[(row, col)] - self.b0_at(row);
317 if a > 0.0 && b > 0.0 {
318 a / b
319 } else {
320 0.0
321 }
322 }
323
324 #[inline]
326 fn a0_at(&self, row: usize) -> f32 {
327 self.row_prior.as_ref().map_or(self.a0, |(a, _)| a[row])
328 }
329
330 #[inline]
332 fn b0_at(&self, row: usize) -> f32 {
333 self.row_prior.as_ref().map_or(self.b0, |(_, b)| b[row])
334 }
335
336 pub fn with_row_prior(dims: (usize, usize), a0: &DVector<f32>, b0: &DVector<f32>) -> Self {
339 let mut out = Self::new(dims, 0.0, 0.0);
340 out.set_row_prior(a0, b0);
341 out.reset_stat();
342 out
343 }
344
345 pub fn set_row_prior(&mut self, a0: &DVector<f32>, b0: &DVector<f32>) {
348 assert_eq!(a0.len(), self.num_rows, "row prior a0 length != rows");
349 assert_eq!(b0.len(), self.num_rows, "row prior b0 length != rows");
350 self.row_prior = Some((a0.clone(), b0.clone()));
351 }
352
353 #[must_use]
355 pub fn row_prior(&self) -> Option<(&DVector<f32>, &DVector<f32>)> {
356 self.row_prior.as_ref().map(|(a, b)| (a, b))
357 }
358
359 pub fn vconcat(blocks: Vec<GammaMatrix>, stack_stats: bool) -> Self {
369 assert!(!blocks.is_empty(), "vconcat of empty block list");
370 let ncols = blocks[0].num_columns;
371 let a0 = blocks[0].a0;
372 let b0 = blocks[0].b0;
373 let nrows: usize = blocks.iter().map(|b| b.num_rows).sum();
374 let row_prior = if blocks[0].row_prior.is_some() {
375 assert!(
376 blocks.iter().all(|b| b.row_prior.is_some()),
377 "vconcat: blocks mix scalar and row priors"
378 );
379 let a = DVector::from_iterator(
380 nrows,
381 blocks
382 .iter()
383 .flat_map(|b| b.row_prior.as_ref().expect("checked").0.iter().copied()),
384 );
385 let b = DVector::from_iterator(
386 nrows,
387 blocks
388 .iter()
389 .flat_map(|b| b.row_prior.as_ref().expect("checked").1.iter().copied()),
390 );
391 Some((a, b))
392 } else {
393 assert!(
394 blocks
395 .iter()
396 .all(|b| b.row_prior.is_none() && b.a0 == a0 && b.b0 == b0),
397 "vconcat: blocks must share hyper-params"
398 );
399 None
400 };
401 let a_stat = stack_field(&blocks, nrows, ncols, stack_stats, |g| &g.a_stat);
402 let b_stat = stack_field(&blocks, nrows, ncols, stack_stats, |g| &g.b_stat);
403 let estimated_mean = stack_field(&blocks, nrows, ncols, true, |g| &g.estimated_mean);
404 let estimated_sd = stack_field(&blocks, nrows, ncols, true, |g| &g.estimated_sd);
405 let estimated_log_mean =
406 stack_field(&blocks, nrows, ncols, true, |g| &g.estimated_log_mean);
407 let estimated_log_sd = stack_field(&blocks, nrows, ncols, true, |g| &g.estimated_log_sd);
408 Self {
409 num_rows: nrows,
410 num_columns: ncols,
411 a0,
412 b0,
413 row_prior,
414 a_stat,
415 b_stat,
416 estimated_mean,
417 estimated_sd,
418 estimated_log_mean,
419 estimated_log_sd,
420 }
421 }
422}