Skip to main content

legume_numeric/matrix/
embedding_geometry.rs

1//! Read-only geometry of an `[n × h]` embedding table: how many of its `h`
2//! directions are actually in use, and how much its rows share one direction.
3//!
4//! One function serves every table with that shape — a cell embedding, a
5//! per-gene loading, a module dictionary — because the question is the same for
6//! each: "of the `h` dimensions this was given, how many does it use, and is the
7//! apparent answer a genuine collapse or a large mean offset?" The two are
8//! distinguished by centering: a common mode makes the *raw* Gram look
9//! near-rank-1 while the *centered* Gram recovers its rank. Reporting both is
10//! what separates "the fit collapsed" from "the fit has a large mean offset",
11//! which are different problems with different fixes.
12//!
13//! Lives here rather than in any one engine because the number has to mean the
14//! same thing across engines for a comparison to be a comparison — and because
15//! it was first written inside a sampler diagnostic that has since been removed,
16//! where nothing outside that sampler could reach it.
17
18use nalgebra::DMatrix;
19use rayon::prelude::*;
20
21/// Geometry of an embedding table, measured once and handed to the caller to
22/// report. A struct of measured scalars; nothing here decides anything.
23#[derive(Clone, Copy, Debug, Default, PartialEq)]
24pub struct EmbeddingGeometry {
25    /// Rows measured (units: cells, genes, modules).
26    pub n_rows: usize,
27    /// Columns (the embedding dimension `h`).
28    pub h: usize,
29    /// Mean `|cos|` between a row and the mean direction over rows. `→ 1`
30    /// means every row is essentially the same direction plus a small residual;
31    /// `→ 0` means the rows are spread about the origin.
32    pub common_mode_cos: f32,
33    /// SIGNED mean cosine over distinct row pairs, zero-norm rows excluded.
34    /// Computed in closed form from the sum of unit rows, so it is one `O(n·h)`
35    /// pass rather than `n²` pairs. A balanced cloud reads `−1/(n−1)`, not 0.
36    pub mean_pairwise_cos: f32,
37    /// Participation ratio `(Σλ)² / Σλ²` of the **uncentered** Gram `EᵀE/n`,
38    /// in `[1, h]`.
39    ///
40    /// READ THIS AS VARIANCE CONCENTRATION, NOT USEFUL DIMENSIONALITY. A low value
41    /// says the variance is carried by few directions; it does **not** say the
42    /// remaining dims are noise. On real fits this has read under 3 of 16 while
43    /// the dims *beyond* the first few still recovered cell type well above a
44    /// null. Low-variance directions can carry ample signal, and reading this
45    /// field as a capacity estimate is a mistake that has already been made once.
46    pub eff_rank_raw: f32,
47    /// Same, for the **column-centered** Gram. Much larger than
48    /// [`Self::eff_rank_raw`] ⇒ the apparent low rank is a mean offset (a common
49    /// mode), not a genuine collapse.
50    pub eff_rank_centered: f32,
51    /// Same, after DOUBLE centering: each row's mean across its `h` columns
52    /// removed, then each column's mean. A log-simplex dictionary is
53    /// `log β_gk = (per-gene share) + (module/topic content) − (per-topic
54    /// partition)`, and the two offsets dominate its variance: column
55    /// centering leaves the per-gene share, row centering leaves the
56    /// per-topic constant, and either alone reads near 1 with a huge
57    /// [`Self::max_vif`]. Removing both leaves the topic-specific content,
58    /// which is what "how many directions do the topics use" means for such
59    /// a table. For an embedding it removes a per-unit depth-like offset the
60    /// same way. Row centering projects out the all-ones direction, so a
61    /// genuinely full-rank table reads `h − 1` here, not `h`. No VIF is
62    /// reported for it: the row-centered columns sum to zero, so their
63    /// correlation matrix is singular by construction.
64    pub eff_rank_double_centered: f32,
65    /// Largest `|correlation|` between two distinct dims.
66    pub max_abs_corr: f32,
67    /// Largest variance-inflation factor `diag(C⁻¹)` over dims (`1` =
68    /// orthogonal). Above ~5, per-dim readouts stop being trustworthy on their
69    /// own.
70    pub max_vif: f32,
71}
72
73/// Participation ratio `(Σλ)²/Σλ²` of a symmetric PSD matrix's spectrum — a
74/// smooth "how many dims are really in use" that needs no eigenvalue cutoff.
75/// A zero matrix reports `0`, not NaN.
76///
77/// No eigendecomposition: for symmetric `G = QΛQᵀ`, `Σλ = trace(G)` and
78/// `Σλ² = trace(G²) = ‖G‖_F²`, so the ratio is `trace(G)² / ‖G‖_F²` exactly.
79/// That is `O(h²)` off the matrix itself instead of `O(h³)` plus the two `h×h`
80/// copies `symmetric_eigenvalues` makes (it clones internally as well as here).
81fn participation_ratio(gram: &DMatrix<f64>) -> f32 {
82    let (s1, s2) = (gram.trace(), gram.norm_squared());
83    if s2 <= 0.0 {
84        return 0.0;
85    }
86    ((s1 * s1) / s2) as f32
87}
88
89/// `(max |corr|, max VIF)` off a centered Gram.
90///
91/// Correlation from the CENTERED Gram (a correlation is centered by
92/// definition; the raw Gram would report a common mode as collinearity and
93/// conflate the two things this module exists to separate). A constant dim
94/// gets `inv_sd = 0`, so it correlates with nothing.
95///
96/// VIF = diag(C⁻¹), via a Cholesky solve against the identity rather than an
97/// explicit inverse — it matters most exactly here, since near-collinear dims
98/// are the case this measures. A singular `C` means a dim is an exact
99/// combination of the others: infinite inflation, reported as such rather
100/// than as a NaN that would read as "not measured".
101fn corr_and_vif(ctr_gram: &DMatrix<f64>) -> (f32, f32) {
102    let h = ctr_gram.nrows();
103    let inv_sd: Vec<f64> = (0..h)
104        .map(|j| {
105            let sd = ctr_gram[(j, j)].max(0.0).sqrt();
106            if sd > 0.0 {
107                1.0 / sd
108            } else {
109                0.0
110            }
111        })
112        .collect();
113    let corr = DMatrix::<f64>::from_fn(h, h, |i, j| {
114        if i == j {
115            1.0
116        } else {
117            ctr_gram[(i, j)] * inv_sd[i] * inv_sd[j]
118        }
119    });
120    let cr = &corr;
121    let max_abs_corr = (0..h)
122        .flat_map(|i| ((i + 1)..h).map(move |j| cr[(i, j)].abs()))
123        .fold(0.0f64, f64::max);
124    let eye = DMatrix::<f64>::identity(h, h);
125    let max_vif = corr
126        .clone()
127        .cholesky()
128        .map(|c| c.solve(&eye))
129        .or_else(|| corr.lu().solve(&eye))
130        .map_or(f32::INFINITY, |inv| {
131            (0..h).fold(0.0f64, |acc, j| acc.max(inv[(j, j)])) as f32
132        })
133        .max(1.0);
134    (max_abs_corr as f32, max_vif)
135}
136
137/// Per-thread accumulator for the single row pass: the upper triangle of the
138/// raw Gram, the column sums, the sum of unit rows, and how many rows had a
139/// direction at all.
140struct RowPass {
141    gram_upper: Vec<f64>,
142    col_sum: Vec<f64>,
143    unit_sum: Vec<f64>,
144    live: usize,
145}
146
147impl RowPass {
148    fn zero(h: usize) -> Self {
149        Self {
150            gram_upper: vec![0.0; h * h],
151            col_sum: vec![0.0; h],
152            unit_sum: vec![0.0; h],
153            live: 0,
154        }
155    }
156
157    fn add_row(mut self, row: &[f64]) -> Self {
158        let h = row.len();
159        for (i, &ri) in row.iter().enumerate() {
160            self.col_sum[i] += ri;
161            // Row `i` of the upper triangle: entries `(i, i..h)`.
162            let upper = &mut self.gram_upper[i * h + i..(i + 1) * h];
163            for (g, &rj) in upper.iter_mut().zip(&row[i..]) {
164                *g += ri * rj;
165            }
166        }
167        let nrm = row.iter().map(|x| x * x).sum::<f64>().sqrt();
168        if nrm > 0.0 {
169            for (u, &r) in self.unit_sum.iter_mut().zip(row) {
170                *u += r / nrm;
171            }
172            self.live += 1;
173        }
174        self
175    }
176
177    fn merge(mut self, other: Self) -> Self {
178        for (a, b) in self.gram_upper.iter_mut().zip(other.gram_upper) {
179            *a += b;
180        }
181        for (a, b) in self.col_sum.iter_mut().zip(other.col_sum) {
182            *a += b;
183        }
184        for (a, b) in self.unit_sum.iter_mut().zip(other.unit_sum) {
185            *a += b;
186        }
187        self.live += other.live;
188        self
189    }
190}
191
192/// Measure the geometry of `e` (`rows` = units, `cols` = dims). `O(n·h² + h³)`,
193/// with the `O(n·h²)` row pass parallel over rows.
194///
195/// Degenerate inputs return zeros rather than `NaN`: an empty table, or one
196/// whose dims are constant, has no geometry to report.
197#[must_use]
198pub fn embedding_geometry(e: &DMatrix<f32>) -> EmbeddingGeometry {
199    let (n, h) = (e.nrows(), e.ncols());
200    if n == 0 || h == 0 {
201        return EmbeddingGeometry {
202            n_rows: n,
203            h,
204            ..Default::default()
205        };
206    }
207    // Row-major f64 copy: the passes below walk rows, nalgebra stores columns,
208    // and f64 because every quantity here is a sum over all `n` rows. Written
209    // directly rather than via `e.transpose()`, which would materialise a whole
210    // extra `[h, n]` f32 matrix alongside this buffer (~52 MB together on a
211    // 34k × 128 table, for one buffer's worth of data).
212    let mut rm = vec![0.0f64; n * h];
213    rm.par_chunks_mut(h).enumerate().for_each(|(i, dst)| {
214        for (j, d) in dst.iter_mut().enumerate() {
215            *d = f64::from(e[(i, j)]);
216        }
217    });
218    let inv_n = 1.0 / n as f64;
219
220    let pass = rm
221        .par_chunks(h)
222        .fold(|| RowPass::zero(h), RowPass::add_row)
223        .reduce(|| RowPass::zero(h), RowPass::merge);
224
225    let raw_gram = DMatrix::<f64>::from_fn(h, h, |i, j| {
226        let (a, b) = if i <= j { (i, j) } else { (j, i) };
227        pass.gram_upper[a * h + b] * inv_n
228    });
229    let mean: Vec<f64> = pass.col_sum.iter().map(|c| c * inv_n).collect();
230    // Centred Gram in closed form: `Eᶜᵀ Eᶜ / n = EᵀE/n − μμᵀ`. Materialising a
231    // centred copy would cost another `n×h` pass for an `h×h` rank-one update.
232    let ctr_gram = &raw_gram - DMatrix::<f64>::from_fn(h, h, |i, j| mean[i] * mean[j]);
233
234    // Signed mean over distinct pairs of unit rows, from the sum of unit rows:
235    // `Σ_{i≠j} ê_i·ê_j = ‖Σ_i ê_i‖² − m` over the `m` rows that have a direction.
236    let m = pass.live;
237    let mean_pairwise_cos = if m < 2 {
238        0.0
239    } else {
240        let s2: f64 = pass.unit_sum.iter().map(|x| x * x).sum();
241        let m = m as f64;
242        ((s2 - m) / (m * (m - 1.0))) as f32
243    };
244
245    // Mean |cos| to the shared mean direction. A zero-norm row contributes 0 —
246    // it has no direction to agree with, and skipping it would silently change
247    // the denominator.
248    let mean_norm = mean.iter().map(|x| x * x).sum::<f64>().sqrt();
249    let common_mode_cos = if mean_norm <= 0.0 {
250        0.0
251    } else {
252        let acc: f64 = rm
253            .par_chunks(h)
254            .map(|row| {
255                let nrm = row.iter().map(|x| x * x).sum::<f64>().sqrt();
256                if nrm <= 0.0 {
257                    return 0.0;
258                }
259                let dot: f64 = row.iter().zip(&mean).map(|(r, m)| r * m).sum();
260                (dot / (nrm * mean_norm)).abs()
261            })
262            .sum();
263        (acc * inv_n) as f32
264    };
265
266    let (max_abs_corr, max_vif) = corr_and_vif(&ctr_gram);
267
268    // Double centering in closed form. Row centering is `E·P` with the
269    // symmetric idempotent `P = I − 11ᵀ/h`, so the doubly centred Gram is
270    // `P·(EᵀE/n − μμᵀ)·P` = `P·ctr_gram·P` — the column-centred Gram the
271    // block above already holds, projected on both sides. No second pass
272    // over `n`, and no separate row-centred mean to carry.
273    let proj = DMatrix::<f64>::identity(h, h) - DMatrix::<f64>::repeat(h, h, 1.0 / h as f64);
274    let double_gram = &proj * &ctr_gram * &proj;
275
276    EmbeddingGeometry {
277        n_rows: n,
278        h,
279        common_mode_cos,
280        mean_pairwise_cos,
281        eff_rank_raw: participation_ratio(&raw_gram),
282        eff_rank_centered: participation_ratio(&ctr_gram),
283        eff_rank_double_centered: participation_ratio(&double_gram),
284        max_abs_corr,
285        max_vif,
286    }
287}
288
289#[cfg(test)]
290mod tests;