Skip to main content

legume_numeric/matrix/
principal_graph.rs

1//! Principal graph fitting (`SimplePPT`) on a low-dimensional embedding.
2//!
3//! Faithful port of Mao et al. 2015, "`SimplePPT`: A Simple Principal Tree
4//! Algorithm" — the same fitter Monocle 3 uses inside `learn_graph()` once
5//! cells have been embedded in a low-dim space (UMAP for Monocle 3,
6//! topic θ / SVD components for senna).
7//!
8//! The objective is
9//!
10//!   L(Y, R) = `Σ_n` `Σ_k` `r_nk` ‖`z_n` − `y_k‖²`
11//!           + σ `Σ_n` `Σ_k` `r_nk` log `r_nk`
12//!           + (γ/2) Σ_(j,k)∈E(Y) ‖`y_j` − `y_k‖²`
13//!
14//! and is solved by alternating soft-assignment, MST recomputation, and
15//! a Laplacian-regularized linear solve `(diag(R^T 1) + γL) Y = R^T Z`.
16
17use nalgebra::DMatrix;
18use petgraph::algo::{dijkstra, min_spanning_tree};
19use petgraph::data::FromElements;
20use petgraph::graph::{NodeIndex, UnGraph};
21use petgraph::visit::EdgeRef;
22use rayon::prelude::*;
23
24/// Configuration for [`fit_principal_graph`].
25#[derive(Debug, Clone)]
26pub struct PrincipalGraphArgs {
27    /// Number of centroids K (graph nodes). Monocle 3 default ≈ 200.
28    pub n_centroids: usize,
29    /// Tree-smoothing strength γ. Higher = stiffer / fewer wiggles.
30    pub gamma: f32,
31    /// Soft-assignment bandwidth σ (in the same units as ‖z‖²).
32    /// Set ≤ 0 to use an adaptive σ = mean of per-cell nearest-centroid dist².
33    pub sigma: f32,
34    /// Maximum outer `SimplePPT` iterations.
35    pub max_iter: usize,
36    /// Relative objective change for early stop.
37    pub tol: f32,
38    /// k-means iterations for centroid initialization.
39    pub kmeans_max_iter: usize,
40}
41
42impl Default for PrincipalGraphArgs {
43    fn default() -> Self {
44        Self {
45            n_centroids: 200,
46            gamma: 10.0,
47            sigma: -1.0,
48            max_iter: 25,
49            tol: 1e-4,
50            kmeans_max_iter: 100,
51        }
52    }
53}
54
55/// A fitted principal graph (tree) over the latent space.
56#[derive(Debug, Clone)]
57pub struct PrincipalGraph {
58    /// K × D centroid coordinates in the latent space.
59    pub nodes: DMatrix<f32>,
60    /// MST edges as `(j, k)` with `j < k`. Length = K − 1 for a connected K-tree.
61    pub edges: Vec<(usize, usize)>,
62    /// Euclidean edge weights, parallel to `edges`.
63    pub edge_weights: Vec<f32>,
64    /// Outer iterations consumed.
65    pub n_iters: usize,
66    /// Final objective value.
67    pub final_objective: f32,
68}
69
70impl PrincipalGraph {
71    pub fn n_nodes(&self) -> usize {
72        self.nodes.nrows()
73    }
74
75    pub fn n_edges(&self) -> usize {
76        self.edges.len()
77    }
78}
79
80/// Fit a principal tree to the rows of `z` (cells × D).
81pub fn fit_principal_graph(
82    z: &DMatrix<f32>,
83    args: &PrincipalGraphArgs,
84) -> anyhow::Result<PrincipalGraph> {
85    anyhow::ensure!(args.n_centroids >= 2, "need at least 2 centroids");
86    anyhow::ensure!(
87        z.nrows() >= args.n_centroids,
88        "fewer cells ({}) than requested centroids ({})",
89        z.nrows(),
90        args.n_centroids
91    );
92    anyhow::ensure!(args.gamma >= 0.0, "gamma must be ≥ 0");
93
94    let k = args.n_centroids;
95    let d = z.ncols();
96
97    let mut y = kmeans_centroids(z, k, args.kmeans_max_iter).0;
98
99    let mut prev_obj = f32::INFINITY;
100    let mut edges: Vec<(usize, usize)> = Vec::new();
101    let mut edge_weights: Vec<f32> = Vec::new();
102    let mut n_iters = 0usize;
103
104    for iter in 0..args.max_iter {
105        let dist_nk = pairwise_sqdist_rows_to_rows(z, &y);
106
107        let sigma = if args.sigma > 0.0 {
108            args.sigma
109        } else {
110            adaptive_sigma(&dist_nk)
111        };
112
113        let r_nk = softmin_rows(&dist_nk, sigma);
114
115        let dist_kk = pairwise_sqdist_rows_to_rows(&y, &y);
116        let (mst_edges, mst_weights) = mst_from_sqdist(&dist_kk);
117
118        let s = r_nk.row_sum();
119        let mut sys = laplacian(k, &mst_edges) * args.gamma;
120        for i in 0..k {
121            sys[(i, i)] += s[i] + 1e-6;
122        }
123        let rhs = r_nk.transpose() * z;
124        y = solve_spd(&sys, &rhs)?;
125
126        let obj = objective(&dist_nk, &r_nk, &y, &mst_edges, sigma, args.gamma);
127        n_iters = iter + 1;
128        edges = mst_edges;
129        edge_weights = mst_weights;
130
131        let denom = prev_obj.abs().max(1.0);
132        let rel = (prev_obj - obj).abs() / denom;
133        log::debug!("SimplePPT iter {iter}: obj={obj:.4} (Δrel={rel:.2e}, σ={sigma:.4})");
134        if rel < args.tol {
135            prev_obj = obj;
136            break;
137        }
138        prev_obj = obj;
139    }
140
141    debug_assert_eq!(y.nrows(), k);
142    debug_assert_eq!(y.ncols(), d);
143
144    Ok(PrincipalGraph {
145        nodes: y,
146        edges,
147        edge_weights,
148        n_iters,
149        final_objective: prev_obj,
150    })
151}
152
153pub use crate::matrix::kmeans::{kmeans_centroids, kmeans_centroids_seeded};
154
155/// `(N×D, K×D) → N×K` matrix of squared Euclidean distances. Fills a
156/// row-major flat buffer in parallel via `par_chunks_exact_mut` (one
157/// alloc) before handing it to `DMatrix::from_row_slice`.
158pub fn pairwise_sqdist_rows_to_rows(a: &DMatrix<f32>, b: &DMatrix<f32>) -> DMatrix<f32> {
159    let n = a.nrows();
160    let k = b.nrows();
161    let d = a.ncols();
162    debug_assert_eq!(b.ncols(), d);
163
164    let mut buf = vec![0f32; n * k];
165    buf.par_chunks_exact_mut(k)
166        .enumerate()
167        .for_each(|(i, row)| {
168            for kk in 0..k {
169                let mut s = 0f32;
170                for j in 0..d {
171                    let v = a[(i, j)] - b[(kk, j)];
172                    s += v * v;
173                }
174                row[kk] = s;
175            }
176        });
177    DMatrix::from_row_slice(n, k, &buf)
178}
179
180fn adaptive_sigma(dist_nk: &DMatrix<f32>) -> f32 {
181    let n = dist_nk.nrows();
182    if n == 0 {
183        return 1.0;
184    }
185    let sum: f32 = (0..n)
186        .into_par_iter()
187        .map(|i| {
188            let mut m = f32::INFINITY;
189            for k in 0..dist_nk.ncols() {
190                if dist_nk[(i, k)] < m {
191                    m = dist_nk[(i, k)];
192                }
193            }
194            m
195        })
196        .sum();
197    let mean_min = sum / n as f32;
198    mean_min.max(1e-8)
199}
200
201/// Row-wise softmax of `−d/σ` with subtraction of per-row min for
202/// numerical stability.
203fn softmin_rows(d: &DMatrix<f32>, sigma: f32) -> DMatrix<f32> {
204    let n = d.nrows();
205    let k = d.ncols();
206    let mut buf = vec![0f32; n * k];
207    buf.par_chunks_exact_mut(k)
208        .enumerate()
209        .for_each(|(i, row)| {
210            let mut min_d = f32::INFINITY;
211            for kk in 0..k {
212                if d[(i, kk)] < min_d {
213                    min_d = d[(i, kk)];
214                }
215            }
216            let mut zsum = 0f32;
217            for kk in 0..k {
218                let v = (-(d[(i, kk)] - min_d) / sigma).exp();
219                row[kk] = v;
220                zsum += v;
221            }
222            if zsum > 0.0 {
223                for v in row.iter_mut() {
224                    *v /= zsum;
225                }
226            } else {
227                let u = 1.0 / k as f32;
228                for v in row.iter_mut() {
229                    *v = u;
230                }
231            }
232        });
233    DMatrix::from_row_slice(n, k, &buf)
234}
235
236/// MST over the K centroids using petgraph's `min_spanning_tree`. Edge
237/// weights for ranking are the squared distances in `dist_kk`; the
238/// returned weights are the Euclidean (sqrt) distances so downstream
239/// geodesic distances are in latent-space units.
240pub fn mst_from_sqdist(dist_kk: &DMatrix<f32>) -> (Vec<(usize, usize)>, Vec<f32>) {
241    let k = dist_kk.nrows();
242    if k <= 1 {
243        return (Vec::new(), Vec::new());
244    }
245    // Build a complete undirected graph on K nodes; petgraph runs
246    // Kruskal's MST internally on PartialOrd edge weights.
247    let mut g: UnGraph<(), f32> = UnGraph::with_capacity(k, k * (k - 1) / 2);
248    let nodes: Vec<NodeIndex> = (0..k).map(|_| g.add_node(())).collect();
249    for a in 0..k {
250        for b in (a + 1)..k {
251            g.add_edge(nodes[a], nodes[b], dist_kk[(a, b)].max(0.0));
252        }
253    }
254    let mst: UnGraph<(), f32> = UnGraph::from_elements(min_spanning_tree(&g));
255
256    let mut edges = Vec::with_capacity(k - 1);
257    let mut weights = Vec::with_capacity(k - 1);
258    for e in mst.edge_references() {
259        let a = e.source().index();
260        let b = e.target().index();
261        let (lo, hi) = if a < b { (a, b) } else { (b, a) };
262        edges.push((lo, hi));
263        weights.push(e.weight().max(0.0).sqrt());
264    }
265    (edges, weights)
266}
267
268fn laplacian(k: usize, edges: &[(usize, usize)]) -> DMatrix<f32> {
269    let mut l = DMatrix::<f32>::zeros(k, k);
270    for &(a, b) in edges {
271        l[(a, a)] += 1.0;
272        l[(b, b)] += 1.0;
273        l[(a, b)] -= 1.0;
274        l[(b, a)] -= 1.0;
275    }
276    l
277}
278
279fn solve_spd(a: &DMatrix<f32>, b: &DMatrix<f32>) -> anyhow::Result<DMatrix<f32>> {
280    let chol = a
281        .clone()
282        .cholesky()
283        .ok_or_else(|| anyhow::anyhow!("Cholesky failed on principal-graph M-step system"))?;
284    Ok(chol.solve(b))
285}
286
287fn objective(
288    dist_nk: &DMatrix<f32>,
289    r_nk: &DMatrix<f32>,
290    y: &DMatrix<f32>,
291    edges: &[(usize, usize)],
292    sigma: f32,
293    gamma: f32,
294) -> f32 {
295    let n = dist_nk.nrows();
296    let k = dist_nk.ncols();
297    let (data_term, entropy) = (0..n)
298        .into_par_iter()
299        .map(|i| {
300            let mut d_i = 0f32;
301            let mut e_i = 0f32;
302            for kk in 0..k {
303                let r = r_nk[(i, kk)];
304                d_i += r * dist_nk[(i, kk)];
305                if r > 1e-12 {
306                    e_i += r * r.ln();
307                }
308            }
309            (d_i, e_i)
310        })
311        .reduce(|| (0f32, 0f32), |a, b| (a.0 + b.0, a.1 + b.1));
312    let mut tree_term = 0f32;
313    for &(a, b) in edges {
314        let mut s = 0f32;
315        for j in 0..y.ncols() {
316            let v = y[(a, j)] - y[(b, j)];
317            s += v * v;
318        }
319        tree_term += s;
320    }
321    data_term + sigma * entropy + 0.5 * gamma * tree_term
322}
323
324///////////////////////////////////////////
325// Cell projection + geodesic pseudotime //
326///////////////////////////////////////////
327
328/// Per-cell projection onto the principal graph.
329#[derive(Debug, Clone, Copy)]
330pub struct CellProjection {
331    /// Index into `graph.edges` of the closest edge.
332    pub nearest_edge: usize,
333    /// Parameter t ∈ [0, 1] along that edge from `(j → k)`.
334    pub t: f32,
335    /// Squared distance from cell to its projected point.
336    pub sqdist: f32,
337}
338
339/// Project each row of `z` to its nearest point on the principal graph.
340pub fn project_cells_to_graph(z: &DMatrix<f32>, graph: &PrincipalGraph) -> Vec<CellProjection> {
341    let n = z.nrows();
342    let d = z.ncols();
343    let nodes = &graph.nodes;
344    let edges = &graph.edges;
345
346    (0..n)
347        .into_par_iter()
348        .map(|i| {
349            let mut best = CellProjection {
350                nearest_edge: 0,
351                t: 0.0,
352                sqdist: f32::INFINITY,
353            };
354            for (eidx, &(j, k)) in edges.iter().enumerate() {
355                let mut dot = 0f32;
356                let mut len2 = 0f32;
357                for dd in 0..d {
358                    let yj = nodes[(j, dd)];
359                    let yk = nodes[(k, dd)];
360                    let zi = z[(i, dd)];
361                    dot += (zi - yj) * (yk - yj);
362                    len2 += (yk - yj) * (yk - yj);
363                }
364                let t = if len2 > 1e-12 {
365                    (dot / len2).clamp(0.0, 1.0)
366                } else {
367                    0.0
368                };
369                let mut sd = 0f32;
370                for dd in 0..d {
371                    let yj = nodes[(j, dd)];
372                    let yk = nodes[(k, dd)];
373                    let proj = yj + t * (yk - yj);
374                    let v = z[(i, dd)] - proj;
375                    sd += v * v;
376                }
377                if sd < best.sqdist {
378                    best = CellProjection {
379                        nearest_edge: eidx,
380                        t,
381                        sqdist: sd,
382                    };
383                }
384            }
385            best
386        })
387        .collect()
388}
389
390/// Build a petgraph view of the principal tree (nodes = centroid index,
391/// edges = MST with Euclidean weights). The graph is reconstructed lazily
392/// because it's tiny (K ≈ 200) and avoids forcing `PrincipalGraph` itself
393/// to carry a non-Clone graph type.
394fn build_petgraph(graph: &PrincipalGraph) -> (UnGraph<(), f32>, Vec<NodeIndex>) {
395    let k = graph.n_nodes();
396    let mut g: UnGraph<(), f32> = UnGraph::with_capacity(k, graph.n_edges());
397    let nodes: Vec<NodeIndex> = (0..k).map(|_| g.add_node(())).collect();
398    for (&(a, b), &w) in graph.edges.iter().zip(&graph.edge_weights) {
399        g.add_edge(nodes[a], nodes[b], w);
400    }
401    (g, nodes)
402}
403
404/// Geodesic distances from `root` to every node of the principal graph,
405/// computed via `petgraph::algo::dijkstra`.
406pub fn node_geodesic_from(graph: &PrincipalGraph, root: usize) -> Vec<f32> {
407    let (g, nodes) = build_petgraph(graph);
408    let dist_map = dijkstra(&g, nodes[root], None, |e| *e.weight());
409    let mut out = vec![f32::INFINITY; graph.n_nodes()];
410    for (nid, d) in dist_map {
411        out[nid.index()] = d;
412    }
413    out
414}
415
416/// Per-cell pseudotime: geodesic distance from `root_node` to each cell's
417/// projection on the principal graph.
418pub fn pseudotime_from_root(
419    graph: &PrincipalGraph,
420    projections: &[CellProjection],
421    root_node: usize,
422) -> Vec<f32> {
423    let node_dist = node_geodesic_from(graph, root_node);
424    projections
425        .iter()
426        .map(|p| {
427            let (j, k) = graph.edges[p.nearest_edge];
428            let w = graph.edge_weights[p.nearest_edge];
429            (node_dist[j] + p.t * w).min(node_dist[k] + (1.0 - p.t) * w)
430        })
431        .collect()
432}
433
434/// Find the centroid index closest to row `row` of `z`.
435pub fn closest_node_to_row(z: &DMatrix<f32>, row: usize, graph: &PrincipalGraph) -> usize {
436    let d = z.ncols();
437    let mut best = 0usize;
438    let mut best_sd = f32::INFINITY;
439    for k in 0..graph.n_nodes() {
440        let mut s = 0f32;
441        for dd in 0..d {
442            let v = z[(row, dd)] - graph.nodes[(k, dd)];
443            s += v * v;
444        }
445        if s < best_sd {
446            best_sd = s;
447            best = k;
448        }
449    }
450    best
451}
452
453#[cfg(test)]
454mod tests;