1use crate::matrix::graph::WeightedGraph;
2use crate::matrix::knn::EXACT_THRESHOLD;
3use crate::matrix::knn_match::{ColumnDict, SearchScratch};
4
5use indicatif::ParallelProgressIterator;
6use log::info;
7use nalgebra::DMatrix;
8use nalgebra_sparse::CscMatrix;
9use rayon::prelude::*;
10
11const DEFAULT_BLOCK_SIZE: usize = 1000;
12
13use crate::matrix::knn::knn_rows;
14pub use crate::matrix::knn::ALL_PAIRS_THRESHOLD;
15
16pub struct KnnGraph {
17 pub adjacency: CscMatrix<f32>,
19 pub edges: Vec<(usize, usize)>,
21 pub distances: Vec<f32>,
23 pub n_nodes: usize,
25}
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum EdgeSource {
30 Primary,
31 Secondary,
32 Both,
35}
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
39pub enum DistanceMerge {
40 Raw,
42 SourceRank,
47}
48
49pub struct KnnGraphArgs {
50 pub knn: usize,
51 pub block_size: usize,
52 pub reciprocal: bool,
55}
56
57impl KnnGraph {
58 pub fn from_columns(points: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
63 Self::from_rows(&points.transpose(), args)
64 }
65
66 pub fn from_rows(data: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
71 let lists = neighbour_lists(data, &args)?;
72 Self::from_neighbours(data.nrows(), &lists, args.reciprocal)
73 }
74
75 pub fn from_rows_fuzzy(
77 data: &DMatrix<f32>,
78 args: KnnGraphArgs,
79 ) -> anyhow::Result<(KnnGraph, Vec<f32>)> {
80 let lists = neighbour_lists(data, &args)?;
81 let graph = Self::from_neighbours(data.nrows(), &lists, args.reciprocal)?;
82 let weights = umap_edge_weights(&graph.edges, &lists);
83 Ok((graph, weights))
84 }
85
86 pub fn from_columns_fuzzy(
100 points: &DMatrix<f32>,
101 args: KnnGraphArgs,
102 ) -> anyhow::Result<(KnnGraph, Vec<f32>)> {
103 Self::from_rows_fuzzy(&points.transpose(), args)
104 }
105
106 fn from_neighbours(
108 nn: usize,
109 lists: &[NeighbourList],
110 reciprocal: bool,
111 ) -> anyhow::Result<KnnGraph> {
112 let n_triplets: usize = lists.iter().map(|(nb, _)| nb.len()).sum();
113 info!("{n_triplets} triplets by kNN matching");
114 if n_triplets == 0 {
115 return Err(anyhow::anyhow!("empty triplets"));
116 }
117
118 let filter_spin =
122 crate::matrix::progress::new_spinner("{spinner} [{elapsed_precise}] {msg}")
123 .with_message("filtering edges");
124 let edges = edges_from_neighbours(lists, reciprocal);
125 filter_spin.finish_and_clear();
126 info!(
127 "{} edges after {} matching",
128 edges.len(),
129 if reciprocal { "reciprocal" } else { "union" }
130 );
131
132 let (edge_pairs, distances): (Vec<_>, Vec<_>) = edges.into_iter().unzip();
133 let adjacency = symmetric_adjacency(nn, &edge_pairs, &distances);
134
135 Ok(KnnGraph {
136 adjacency,
137 edges: edge_pairs,
138 distances,
139 n_nodes: nn,
140 })
141 }
142
143 pub fn union_with(
161 &self,
162 other: &KnnGraph,
163 policy: DistanceMerge,
164 ) -> anyhow::Result<(KnnGraph, Vec<EdgeSource>)> {
165 anyhow::ensure!(
166 self.n_nodes == other.n_nodes,
167 "cannot union graphs over different node counts: {} vs {}",
168 self.n_nodes,
169 other.n_nodes
170 );
171 let n_nodes = self.n_nodes;
172
173 let (a_dist, b_dist) = match policy {
174 DistanceMerge::Raw => (self.distances.clone(), other.distances.clone()),
175 DistanceMerge::SourceRank => (
176 within_source_rank(&self.distances),
177 within_source_rank(&other.distances),
178 ),
179 };
180
181 let canonical = |&(i, j): &(usize, usize)| if i <= j { (i, j) } else { (j, i) };
193 let mut tagged: Vec<TaggedEdge> = Vec::with_capacity(self.edges.len() + other.edges.len());
196 tagged.par_extend(
197 self.edges
198 .par_iter()
199 .zip(a_dist.par_iter())
200 .map(|(e, &d)| (canonical(e), d, 1u8)),
201 );
202 tagged.par_extend(
203 other
204 .edges
205 .par_iter()
206 .zip(b_dist.par_iter())
207 .map(|(e, &d)| (canonical(e), d, 2u8)),
208 );
209 let folded = fold_tagged_edges(tagged);
210 let mut edges = Vec::with_capacity(folded.len());
211 let mut distances = Vec::with_capacity(folded.len());
212 let mut source = Vec::with_capacity(folded.len());
213 for (key, dist, mask) in folded {
214 edges.push(key);
215 distances.push(dist);
216 source.push(match mask {
217 1 => EdgeSource::Primary,
218 2 => EdgeSource::Secondary,
219 _ => EdgeSource::Both,
220 });
221 }
222
223 let adjacency = symmetric_adjacency(n_nodes, &edges, &distances);
225
226 Ok((
227 KnnGraph {
228 adjacency,
229 edges,
230 distances,
231 n_nodes,
232 },
233 source,
234 ))
235 }
236
237 pub fn neighbors(&self, node: usize) -> &[usize] {
239 let offsets = self.adjacency.col_offsets();
240 let start = offsets[node];
241 let end = offsets[node + 1];
242 &self.adjacency.row_indices()[start..end]
243 }
244
245 pub fn num_edges(&self) -> usize {
246 self.edges.len()
247 }
248
249 pub fn num_nodes(&self) -> usize {
250 self.n_nodes
251 }
252
253 pub fn exp_kernel_weights(&self) -> Vec<f32> {
260 if self.distances.is_empty() {
261 return Vec::new();
262 }
263 let sigma = crate::matrix::utils::median(&self.distances);
264 let sigma = if sigma <= 0.0 { 1.0 } else { sigma };
265 info!("exp_kernel_weights: σ (median distance) = {:.4}", sigma);
266 self.distances.iter().map(|&d| (-d / sigma).exp()).collect()
267 }
268
269 pub fn fuzzy_kernel_weights(&self) -> Vec<f32> {
285 if self.distances.is_empty() {
286 return Vec::new();
287 }
288
289 let offsets = self.adjacency.col_offsets();
290 let row_indices = self.adjacency.row_indices();
291 let values = self.adjacency.values();
292
293 let (rho, sigma): (Vec<f32>, Vec<f32>) = (0..self.n_nodes)
295 .into_par_iter()
296 .map(|i| {
297 let start = offsets[i];
298 let end = offsets[i + 1];
299 let dists: Vec<f32> = (start..end).map(|idx| values[idx]).collect();
300 if dists.is_empty() {
301 return (0.0_f32, 1.0_f32);
302 }
303 let rho_i = dists.iter().cloned().fold(f32::INFINITY, f32::min);
304 let target = (dists.len() as f32).log2();
305 let sigma_i = smooth_knn_sigma(&dists, rho_i, target);
306 (rho_i, sigma_i)
307 })
308 .unzip();
309
310 self.edges
313 .par_iter()
314 .map(|&(i, j)| {
315 let d_ij = self.edge_distance_directed(offsets, row_indices, values, i, j);
316 let w_ij = directed_umap_weight(d_ij, rho[i], sigma[i]);
317 let d_ji = self.edge_distance_directed(offsets, row_indices, values, j, i);
318 let w_ji = directed_umap_weight(d_ji, rho[j], sigma[j]);
319 w_ij + w_ji - w_ij * w_ji
321 })
322 .collect()
323 }
324
325 fn edge_distance_directed(
327 &self,
328 offsets: &[usize],
329 row_indices: &[usize],
330 values: &[f32],
331 from: usize,
332 to: usize,
333 ) -> f32 {
334 let start = offsets[from];
335 let end = offsets[from + 1];
336 for idx in start..end {
337 if row_indices[idx] == to {
338 return values[idx];
339 }
340 }
341 f32::INFINITY
342 }
343}
344
345impl WeightedGraph for KnnGraph {
346 fn num_nodes(&self) -> usize {
347 self.n_nodes
348 }
349
350 fn num_edges(&self) -> usize {
351 self.edges.len()
352 }
353
354 fn neighbors_with_weight<'a>(
355 &'a self,
356 node: usize,
357 ) -> Box<dyn Iterator<Item = (usize, f32)> + 'a> {
358 let offsets = self.adjacency.col_offsets();
359 let start = offsets[node];
360 let end = offsets[node + 1];
361 let rows = &self.adjacency.row_indices()[start..end];
362 let vals = &self.adjacency.values()[start..end];
363 Box::new(rows.iter().zip(vals.iter()).map(|(&i, &w)| (i, w)))
364 }
365}
366
367fn smooth_knn_sigma(dists: &[f32], rho: f32, target: f32) -> f32 {
371 const TOLERANCE: f32 = 1e-5;
372 const MAX_ITER: usize = 64;
373
374 let mean_dist: f32 = dists.iter().sum::<f32>() / dists.len().max(1) as f32;
375 let min_sigma = 1e-3 * mean_dist;
376
377 let mut lo = 0.0f32;
378 let mut hi = f32::INFINITY;
379 let mut mid = 1.0f32;
380
381 for _ in 0..MAX_ITER {
382 let mut psum = 0.0f32;
383 for &d in dists {
384 let gap = d - rho;
385 if gap > 0.0 {
386 psum += (-gap / mid).exp();
387 } else {
388 psum += 1.0;
389 }
390 }
391
392 if (psum - target).abs() < TOLERANCE {
393 break;
394 }
395
396 if psum > target {
397 hi = mid;
398 mid = (lo + hi) / 2.0;
399 } else {
400 lo = mid;
401 if hi.is_infinite() {
402 mid *= 2.0;
403 } else {
404 mid = (lo + hi) / 2.0;
405 }
406 }
407 }
408
409 mid.max(min_sigma)
410}
411
412fn directed_umap_weight(d: f32, rho: f32, sigma: f32) -> f32 {
414 if d.is_infinite() || sigma <= 0.0 {
415 return 0.0;
416 }
417 let gap = d - rho;
418 if gap <= 0.0 {
419 1.0
420 } else {
421 (-gap / sigma).exp()
422 }
423}
424
425impl WeightedGraph for crate::leiden::Network {
430 fn num_nodes(&self) -> usize {
431 self.nodes()
432 }
433
434 fn num_edges(&self) -> usize {
435 crate::leiden::Network::edge_count(self)
436 }
437
438 fn neighbors_with_weight<'a>(
439 &'a self,
440 node: usize,
441 ) -> Box<dyn Iterator<Item = (usize, f32)> + 'a> {
442 Box::new(self.neighbors(node).map(|(n, w)| (n, w as f32)))
443 }
444}
445
446#[must_use]
452pub fn modularity_to_cpm_resolution(modularity_gamma: f64, total_edge_weight: f64) -> f64 {
453 modularity_gamma / (2.0 * total_edge_weight).max(1.0)
454}
455
456impl KnnGraph {
457 pub fn to_leiden_network(&self) -> (crate::leiden::Network, f64) {
463 self.to_leiden_network_with(&self.fuzzy_kernel_weights())
464 }
465
466 pub fn to_leiden_network_with(&self, weights: &[f32]) -> (crate::leiden::Network, f64) {
469 let n = self.n_nodes;
470
471 let mut node_degree = vec![0.0f32; n];
472 let mut n_edges = vec![0usize; n];
473 let mut total_edge_weight = 0.0f64;
474 for (&(i, j), &w) in self.edges.iter().zip(weights.iter()) {
475 node_degree[i] += w;
476 node_degree[j] += w;
477 n_edges[i] += 1;
478 n_edges[j] += 1;
479 total_edge_weight += w as f64;
480 }
481
482 let mut network = crate::leiden::Network::with_nodes(&node_degree, &n_edges);
483 for (&(i, j), &w) in self.edges.iter().zip(weights.iter()) {
484 network.add_edge(i, j, w);
485 }
486
487 (network, total_edge_weight)
488 }
489}
490
491pub fn run_leiden(
495 network: &crate::leiden::Network,
496 n: usize,
497 resolution: f64,
498 seed: Option<usize>,
499) -> Vec<usize> {
500 use crate::leiden::clustering::SimpleClustering;
501 use crate::leiden::Clustering;
502 use crate::leiden::Leiden;
503
504 let mut leiden = Leiden::new(resolution, 0.01, seed);
505 let mut clustering = SimpleClustering::init_different_clusters(n);
506
507 for iter in 0..10 {
508 let updated = leiden.iterate(network, &mut clustering);
509 info!(
510 " Leiden iter {}: {} clusters{}",
511 iter + 1,
512 clustering.num_clusters(),
513 if !updated { " (converged)" } else { "" }
514 );
515 if !updated {
516 break;
517 }
518 }
519
520 (0..n).map(|i| clustering.get(i)).collect()
521}
522
523pub fn tune_leiden_resolution(
529 network: &crate::leiden::Network,
530 n: usize,
531 target_k: usize,
532 initial_resolution: f64,
533 seed: Option<usize>,
534) -> Vec<usize> {
535 let mut lo = 1e-6_f64;
536 let mut hi = 10.0_f64;
537 let mut best = run_leiden(network, n, initial_resolution, seed);
538 let best_k = count_distinct(&best);
539
540 info!(
541 " resolution={:.6e} → {} clusters (target {})",
542 initial_resolution, best_k, target_k
543 );
544
545 if best_k == target_k {
546 return best;
547 }
548 if best_k > target_k {
549 hi = initial_resolution;
550 } else {
551 lo = initial_resolution;
552 }
553
554 let mut best_diff = best_k.abs_diff(target_k);
555
556 for _ in 0..20 {
557 let mid = (lo + hi) / 2.0;
558 let result = run_leiden(network, n, mid, seed);
559 let k = count_distinct(&result);
560 info!(" resolution={:.6e} → {} clusters", mid, k);
561
562 if k > target_k {
563 hi = mid;
564 } else {
565 lo = mid;
566 }
567
568 let diff = k.abs_diff(target_k);
569 if diff < best_diff {
570 best = result;
571 best_diff = diff;
572 }
573
574 if k == target_k || (hi - lo) / hi.max(1e-10) < 1e-4 {
575 break;
576 }
577 }
578
579 best
580}
581
582fn count_distinct(labels: &[usize]) -> usize {
584 let max = labels.iter().copied().max().unwrap_or(0);
585 let mut seen = vec![false; max + 1];
586 for &l in labels {
587 seen[l] = true;
588 }
589 seen.iter().filter(|&&s| s).count()
590}
591
592pub fn compact_labels(labels: &mut [usize]) {
594 let max = labels.iter().copied().max().unwrap_or(0);
595 let mut mapping = vec![usize::MAX; max + 1];
596 let mut next = 0usize;
597 for l in labels.iter_mut() {
598 if mapping[*l] == usize::MAX {
599 mapping[*l] = next;
600 next += 1;
601 }
602 *l = mapping[*l];
603 }
604}
605
606fn create_jobs(ntot: usize, block_size: usize) -> Vec<(usize, usize)> {
607 let block_size = if block_size == 0 {
608 DEFAULT_BLOCK_SIZE
609 } else {
610 block_size
611 };
612 let nblock = ntot.div_ceil(block_size);
613 (0..nblock)
614 .map(|block| {
615 let lb = block * block_size;
616 let ub = ((block + 1) * block_size).min(ntot);
617 (lb, ub)
618 })
619 .collect()
620}
621
622#[cfg(test)]
623mod tests;
624
625fn neighbour_lists(data: &DMatrix<f32>, args: &KnnGraphArgs) -> anyhow::Result<Vec<NeighbourList>> {
629 let nn = data.nrows();
630 let n_neighbours = neighbours_per_point(args.knn, nn);
631 Ok(if nn <= EXACT_THRESHOLD {
632 let transposed = data.transpose();
633 let points_vec = transposed.column_iter().collect::<Vec<_>>();
634 let names = (0..nn).collect::<Vec<_>>();
635 let dict = ColumnDict::from_dvector_views(points_vec, names);
636 search_dict(&dict, nn, n_neighbours, args.block_size)?
637 } else {
638 search_rows(data, n_neighbours)
639 })
640}
641
642pub fn umap_memberships(dists: &[f32]) -> Vec<f32> {
647 let Some(rho) = dists.iter().copied().reduce(f32::min) else {
648 return Vec::new();
649 };
650 let target = ((dists.len() + 1) as f32).log2();
651 let sigma = smooth_knn_sigma(dists, rho, target);
652 dists
653 .iter()
654 .map(|&d| directed_umap_weight(d, rho, sigma))
655 .collect()
656}
657
658fn umap_edge_weights(edges: &[(usize, usize)], lists: &[NeighbourList]) -> Vec<f32> {
661 let directed: Vec<Vec<f32>> = lists
663 .par_iter()
664 .map(|(_, dists)| umap_memberships(dists))
665 .collect();
666 let toward = |from: usize, to: usize| -> f32 {
667 let (nb, _) = &lists[from];
668 nb.iter()
669 .position(|&j| j == to)
670 .map_or(0.0, |at| directed[from][at])
671 };
672 edges
673 .par_iter()
674 .map(|&(i, j)| {
675 let (a, b) = (toward(i, j), toward(j, i));
676 a + b - a * b
677 })
678 .collect()
679}
680
681type NeighbourList = (Vec<usize>, Vec<f32>);
682
683type TaggedEdge = ((usize, usize), f32, u8);
686
687fn fold_tagged_edges(mut tagged: Vec<TaggedEdge>) -> Vec<TaggedEdge> {
694 tagged.par_sort_unstable_by_key(|&(key, _, _)| key);
695 tagged.dedup_by(|cur, prev| {
696 if cur.0 == prev.0 {
697 prev.1 = prev.1.min(cur.1);
698 prev.2 |= cur.2;
699 true
700 } else {
701 false
702 }
703 });
704 tagged
705}
706
707fn neighbours_per_point(knn: usize, nn: usize) -> usize {
710 knn.min(nn.saturating_sub(1)).max(1)
711}
712
713fn search_dict(
716 dict: &ColumnDict<usize>,
717 nn: usize,
718 n_neighbours: usize,
719 block_size: usize,
720) -> anyhow::Result<Vec<NeighbourList>> {
721 let jobs = create_jobs(nn, block_size);
722 let search_bar =
726 crate::matrix::progress::new_progress_bar(jobs.len() as u64).with_message("kNN blocks");
727 let result: anyhow::Result<Vec<Vec<NeighbourList>>> = jobs
728 .into_par_iter()
729 .progress_with(search_bar.clone())
730 .map(|(lb, ub)| {
731 let mut scratch = SearchScratch::default();
733 (lb..ub)
734 .map(|i| dict.search_others_reuse(&i, n_neighbours, &mut scratch))
735 .collect()
736 })
737 .collect();
738 search_bar.finish_and_clear();
741 Ok(result?.into_iter().flatten().collect())
742}
743
744fn search_rows(rows: &DMatrix<f32>, n_neighbours: usize) -> Vec<NeighbourList> {
746 let (indices, distances) = knn_rows(rows, n_neighbours);
747 indices.into_iter().zip(distances).collect()
748}
749
750pub(crate) fn edges_from_neighbours(
755 lists: &[NeighbourList],
756 reciprocal: bool,
757) -> Vec<((usize, usize), f32)> {
758 let tagged: Vec<TaggedEdge> = lists
761 .par_iter()
762 .enumerate()
763 .flat_map_iter(|(i, (nb, ds))| {
764 nb.iter().zip(ds).map(move |(&j, &d)| {
765 if i < j {
766 ((i, j), d, 1u8)
767 } else {
768 ((j, i), d, 2u8)
769 }
770 })
771 })
772 .collect();
773 fold_tagged_edges(tagged)
774 .into_iter()
775 .filter(|&(_, _, mask)| !reciprocal || mask == 3)
776 .map(|(key, dist, _)| (key, dist))
777 .collect()
778}
779
780pub fn symmetric_adjacency(
792 n_nodes: usize,
793 edges: &[(usize, usize)],
794 distances: &[f32],
795) -> CscMatrix<f32> {
796 debug_assert!(
797 edges.iter().all(|&(i, j)| i < j) && edges.windows(2).all(|w| w[0] < w[1]),
798 "symmetric adjacency: edges must be canonical, sorted and unique"
799 );
800 let mut offsets = vec![0usize; n_nodes + 1];
801 for &(i, j) in edges {
802 offsets[i + 1] += 1;
803 offsets[j + 1] += 1;
804 }
805 for c in 0..n_nodes {
806 offsets[c + 1] += offsets[c];
807 }
808 let nnz = offsets[n_nodes];
809 let mut row_indices = vec![0usize; nnz];
810 let mut values = vec![0f32; nnz];
811 let mut cursor = offsets[..n_nodes].to_vec();
812 for (&(i, j), &v) in edges.iter().zip(distances) {
813 row_indices[cursor[i]] = j;
814 values[cursor[i]] = v;
815 cursor[i] += 1;
816 row_indices[cursor[j]] = i;
817 values[cursor[j]] = v;
818 cursor[j] += 1;
819 }
820 CscMatrix::try_from_csc_data(n_nodes, n_nodes, offsets, row_indices, values)
821 .expect("symmetric adjacency: canonical sorted edges fill every column in order")
822}
823
824fn within_source_rank(d: &[f32]) -> Vec<f32> {
830 if d.len() <= 1 {
831 return vec![0.0; d.len()];
832 }
833 let mut order: Vec<usize> = (0..d.len()).collect();
834 order.par_sort_unstable_by(|&a, &b| d[a].total_cmp(&d[b]).then(a.cmp(&b)));
844 let mut out = vec![0.0f32; d.len()];
845 let denom = (d.len() - 1) as f32;
846 for (rank, &idx) in order.iter().enumerate() {
847 out[idx] = rank as f32 / denom;
848 }
849 out
850}