1use 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#[derive(Debug, Clone)]
26pub struct PrincipalGraphArgs {
27 pub n_centroids: usize,
29 pub gamma: f32,
31 pub sigma: f32,
34 pub max_iter: usize,
36 pub tol: f32,
38 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#[derive(Debug, Clone)]
57pub struct PrincipalGraph {
58 pub nodes: DMatrix<f32>,
60 pub edges: Vec<(usize, usize)>,
62 pub edge_weights: Vec<f32>,
64 pub n_iters: usize,
66 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
80pub 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
155pub 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
201fn 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
236pub 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 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#[derive(Debug, Clone, Copy)]
330pub struct CellProjection {
331 pub nearest_edge: usize,
333 pub t: f32,
335 pub sqdist: f32,
337}
338
339pub 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
390fn 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
404pub 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
416pub 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
434pub 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;