Skip to main content

legume_numeric/param/
dmatrix_gamma.rs

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    //////////////////////
15    // hyper parameters //
16    //////////////////////
17    a0: f32,
18    b0: f32,
19    /// Per-row prior `(a0, b0)`, overriding the scalar pair when set.
20    /// Lengths equal `num_rows`. See [`Self::with_row_prior`].
21    row_prior: Option<(DVector<f32>, DVector<f32>)>,
22    ///////////////////////////
23    // sufficient statistics //
24    ///////////////////////////
25    a_stat: DMatrix<f32>,
26    b_stat: DMatrix<f32>,
27    //////////////////////////
28    // estimated parameters //
29    //////////////////////////
30    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` is eager: the coordinate descent reads
54            // `posterior_mean()` before the first calibration (relying on a
55            // zero start), and it gets allocated immediately anyway.
56            estimated_mean: DMatrix::zeros(dims.0, dims.1),
57            // The sd / log_mean / log_sd planes are lazily allocated by
58            // `map_calibrate_*` (via `calibrate_with`). An iterative fit that
59            // only reads `posterior_mean()` (calibrating `MeanOnly`) never
60            // pays for them — they're materialized only when output needs
61            // them (a calibrate with `All` / `MeanAndLogMean`).
62            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                // Column-wise copies are contiguous in column-major storage and
84                // are a no-op on a released (0 x 0) plane.
85                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 nrows(&self) -> usize {
116    //     self.num_rows
117    // }
118
119    // fn ncols(&self) -> usize {
120    //     self.num_columns
121    // }
122
123    // fn len(&self) -> usize {
124    //     self.num_rows * self.num_columns
125    // }
126    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        // `sd[ln X] = sqrt(trigamma(a))` exactly, for `X ~ Gamma(a, b)` — note it
140        // does not depend on the rate, which is why `b_stat` plays no part here.
141        //
142        // This replaced `1/sqrt(a - 1)`, the large-`a` asymptote, which was
143        // wrong in the regime that dominates sparse count data: 46% high at
144        // `a = 1.5`, 21% high at `a = 2.2` (a typical detected feature), and
145        // agreeing only past `a ~ 100`. Below `a = 1` it has no real value at
146        // all, and the old code returned 0 there — i.e. it reported PERFECT
147        // certainty for a feature with no counts, whose posterior is the prior
148        // and whose true `sd` is the largest in the matrix (1.283 at `a = 1`).
149        // Anything reading `log_sd` as a precision was being handed the
150        // inversion of the truth.
151        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        // Fixed chunk width, so which elements share an RNG is a property of
202        // the data shape and not of how rayon happened to split the work.
203        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
232/// Row-stack one plane across blocks. Returns an empty matrix (and skips
233/// the work) when `enabled` is false or the plane is lazily-unallocated in
234/// the first block, so empty/dropped planes never get materialized.
235fn 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    /// Drop the sufficient-stat planes (`a_stat` / `b_stat`) after
260    /// calibration, keeping only the posterior estimates. Use when the
261    /// consumer reads posterior means / log-means but never
262    /// `posterior_sample` (which is the only reader of `a_stat`/`b_stat`).
263    /// Halves the resident footprint of a calibrated parameter.
264    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    /// Zero every `estimated_mean` entry whose corresponding `numerator` is
270    /// zero, collapsing the per-column Gamma prior baseline (`a0/denom`,
271    /// present at *every* unobserved cell) to exact zero. This lets a
272    /// downstream triplet-ization of the mean be **sparse** — only the
273    /// observed support survives. It's the lossy-but-correct choice for
274    /// count-based consumers (the baseline is a regularization floor, not
275    /// signal). `numerator` must match the mean's shape; only meaningful
276    /// after a mean calibration.
277    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    /// Whether the posterior at `(row, col)` carries any data beyond the
290    /// prior: `a_stat > a0`. The read-only counterpart of
291    /// [`Self::sparsify_mean_to_support`], for consumers that serialize the
292    /// mean without owning the numerator — an unsupported entry's mean is the
293    /// prior floor `a0 / (b0 + denom)`, which is regularization, not signal.
294    /// Writing those floors out turns a sparse posterior dense: a carried
295    /// pseudobulk reference measured 100.0% dense (34M of 34M entries) before
296    /// its writer checked this.
297    #[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    /// The unregularized rate at `(row, col)`: `(a_stat − a0) / (b_stat − b0)`
303    /// — data sum over data denominator, no prior in either. Zero when the
304    /// entry has no data support.
305    ///
306    /// This is what a *serialized* posterior should usually store: paired with
307    /// its denominator, it is a bijection of the sufficient statistics, so a
308    /// consumer reconstructs `a_stat`/`b_stat` exactly. The posterior mean
309    /// `(a0 + sum)/(b0 + n)` is the right *estimate* but the wrong *carrier* —
310    /// its prior shrinkage (1.85× at `sum = 1, n = 12`) gets re-ingested as if
311    /// it were data, and a second posterior forms around an already-shrunk
312    /// value.
313    #[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    /// Prior shape for `row`: its row prior when set, else the scalar `a0`.
325    #[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    /// Prior rate for `row`: its row prior when set, else the scalar `b0`.
331    #[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    /// A matrix whose row `d` has prior `Gamma(a0[d], b0[d])`. The statistics
337    /// start at the prior, as with [`TwoStatParam::new`].
338    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    /// Replace the per-row prior. Accumulated statistics are left as they
346    /// are; the new prior applies at the next `reset_stat` / `update_stat`.
347    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    /// The per-row prior, if one is set.
354    #[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    /// Row-stack per-feature-block parameters (from a gene-blocked fit)
360    /// into one `[Σrowsᵢ × K]` parameter. All blocks must share the column
361    /// count and either share the scalar hyper-params or all carry a row
362    /// prior, in which case the row priors are concatenated. Calibrated
363    /// planes present in the first block
364    /// are stacked; lazily-empty planes stay empty. `stack_stats` controls
365    /// whether `a_stat`/`b_stat` are carried through — pass `false` when the
366    /// output only needs posterior estimates, so the heavy sufficient-stat
367    /// planes are never assembled at full width.
368    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}