1use crate::matrix::graph::WeightedGraph;
2use crate::matrix::knn::all_pairs::knn_rows_l2;
3use crate::matrix::knn::ivf::{knn_rows_ivf, IvfArgs, DEFAULT_N_PROBE};
4use crate::matrix::knn::{EXACT_THRESHOLD, KNN_SEED};
5use crate::matrix::knn_match::{ColumnDict, SearchScratch};
6
7use indicatif::ParallelProgressIterator;
8use log::info;
9use nalgebra::DMatrix;
10use nalgebra_sparse::CscMatrix;
11use rayon::prelude::*;
12
13const DEFAULT_BLOCK_SIZE: usize = 1000;
14
15pub const ALL_PAIRS_THRESHOLD: usize = 65_536;
20
21pub struct KnnGraph {
22 pub adjacency: CscMatrix<f32>,
24 pub edges: Vec<(usize, usize)>,
26 pub distances: Vec<f32>,
28 pub n_nodes: usize,
30}
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub enum EdgeSource {
35 Primary,
36 Secondary,
37 Both,
40}
41
42#[derive(Debug, Clone, Copy, PartialEq, Eq)]
44pub enum DistanceMerge {
45 Raw,
47 SourceRank,
52}
53
54pub struct KnnGraphArgs {
55 pub knn: usize,
56 pub block_size: usize,
57 pub reciprocal: bool,
60}
61
62impl KnnGraph {
63 pub fn from_columns(points: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
68 Self::from_rows(&points.transpose(), args)
69 }
70
71 pub fn from_rows(data: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
76 let lists = neighbour_lists(data, &args)?;
77 Self::from_neighbours(data.nrows(), &lists, args.reciprocal)
78 }
79
80 pub fn from_rows_fuzzy(
82 data: &DMatrix<f32>,
83 args: KnnGraphArgs,
84 ) -> anyhow::Result<(KnnGraph, Vec<f32>)> {
85 let lists = neighbour_lists(data, &args)?;
86 let graph = Self::from_neighbours(data.nrows(), &lists, args.reciprocal)?;
87 let weights = umap_edge_weights(&graph.edges, &lists);
88 Ok((graph, weights))
89 }
90
91 pub fn from_columns_fuzzy(
105 points: &DMatrix<f32>,
106 args: KnnGraphArgs,
107 ) -> anyhow::Result<(KnnGraph, Vec<f32>)> {
108 Self::from_rows_fuzzy(&points.transpose(), args)
109 }
110
111 fn from_neighbours(
113 nn: usize,
114 lists: &[NeighbourList],
115 reciprocal: bool,
116 ) -> anyhow::Result<KnnGraph> {
117 let n_triplets: usize = lists.iter().map(|(nb, _)| nb.len()).sum();
118 info!("{n_triplets} triplets by kNN matching");
119 if n_triplets == 0 {
120 return Err(anyhow::anyhow!("empty triplets"));
121 }
122
123 let filter_spin =
127 crate::matrix::progress::new_spinner("{spinner} [{elapsed_precise}] {msg}")
128 .with_message("filtering edges");
129 let edges = edges_from_neighbours(lists, reciprocal);
130 filter_spin.finish_and_clear();
131 info!(
132 "{} edges after {} matching",
133 edges.len(),
134 if reciprocal { "reciprocal" } else { "union" }
135 );
136
137 let (edge_pairs, distances): (Vec<_>, Vec<_>) = edges.into_iter().unzip();
138 let adjacency = symmetric_adjacency(nn, &edge_pairs, &distances);
139
140 Ok(KnnGraph {
141 adjacency,
142 edges: edge_pairs,
143 distances,
144 n_nodes: nn,
145 })
146 }
147
148 pub fn union_with(
166 &self,
167 other: &KnnGraph,
168 policy: DistanceMerge,
169 ) -> anyhow::Result<(KnnGraph, Vec<EdgeSource>)> {
170 anyhow::ensure!(
171 self.n_nodes == other.n_nodes,
172 "cannot union graphs over different node counts: {} vs {}",
173 self.n_nodes,
174 other.n_nodes
175 );
176 let n_nodes = self.n_nodes;
177
178 let (a_dist, b_dist) = match policy {
179 DistanceMerge::Raw => (self.distances.clone(), other.distances.clone()),
180 DistanceMerge::SourceRank => (
181 within_source_rank(&self.distances),
182 within_source_rank(&other.distances),
183 ),
184 };
185
186 let canonical = |&(i, j): &(usize, usize)| if i <= j { (i, j) } else { (j, i) };
198 let mut tagged: Vec<TaggedEdge> = Vec::with_capacity(self.edges.len() + other.edges.len());
201 tagged.par_extend(
202 self.edges
203 .par_iter()
204 .zip(a_dist.par_iter())
205 .map(|(e, &d)| (canonical(e), d, 1u8)),
206 );
207 tagged.par_extend(
208 other
209 .edges
210 .par_iter()
211 .zip(b_dist.par_iter())
212 .map(|(e, &d)| (canonical(e), d, 2u8)),
213 );
214 let folded = fold_tagged_edges(tagged);
215 let mut edges = Vec::with_capacity(folded.len());
216 let mut distances = Vec::with_capacity(folded.len());
217 let mut source = Vec::with_capacity(folded.len());
218 for (key, dist, mask) in folded {
219 edges.push(key);
220 distances.push(dist);
221 source.push(match mask {
222 1 => EdgeSource::Primary,
223 2 => EdgeSource::Secondary,
224 _ => EdgeSource::Both,
225 });
226 }
227
228 let adjacency = symmetric_adjacency(n_nodes, &edges, &distances);
230
231 Ok((
232 KnnGraph {
233 adjacency,
234 edges,
235 distances,
236 n_nodes,
237 },
238 source,
239 ))
240 }
241
242 pub fn neighbors(&self, node: usize) -> &[usize] {
244 let offsets = self.adjacency.col_offsets();
245 let start = offsets[node];
246 let end = offsets[node + 1];
247 &self.adjacency.row_indices()[start..end]
248 }
249
250 pub fn num_edges(&self) -> usize {
251 self.edges.len()
252 }
253
254 pub fn num_nodes(&self) -> usize {
255 self.n_nodes
256 }
257
258 pub fn exp_kernel_weights(&self) -> Vec<f32> {
265 if self.distances.is_empty() {
266 return Vec::new();
267 }
268 let sigma = crate::matrix::utils::median(&self.distances);
269 let sigma = if sigma <= 0.0 { 1.0 } else { sigma };
270 info!("exp_kernel_weights: σ (median distance) = {:.4}", sigma);
271 self.distances.iter().map(|&d| (-d / sigma).exp()).collect()
272 }
273
274 pub fn fuzzy_kernel_weights(&self) -> Vec<f32> {
290 if self.distances.is_empty() {
291 return Vec::new();
292 }
293
294 let offsets = self.adjacency.col_offsets();
295 let row_indices = self.adjacency.row_indices();
296 let values = self.adjacency.values();
297
298 let (rho, sigma): (Vec<f32>, Vec<f32>) = (0..self.n_nodes)
300 .into_par_iter()
301 .map(|i| {
302 let start = offsets[i];
303 let end = offsets[i + 1];
304 let dists: Vec<f32> = (start..end).map(|idx| values[idx]).collect();
305 if dists.is_empty() {
306 return (0.0_f32, 1.0_f32);
307 }
308 let rho_i = dists.iter().cloned().fold(f32::INFINITY, f32::min);
309 let target = (dists.len() as f32).log2();
310 let sigma_i = smooth_knn_sigma(&dists, rho_i, target);
311 (rho_i, sigma_i)
312 })
313 .unzip();
314
315 self.edges
318 .par_iter()
319 .map(|&(i, j)| {
320 let d_ij = self.edge_distance_directed(offsets, row_indices, values, i, j);
321 let w_ij = directed_umap_weight(d_ij, rho[i], sigma[i]);
322 let d_ji = self.edge_distance_directed(offsets, row_indices, values, j, i);
323 let w_ji = directed_umap_weight(d_ji, rho[j], sigma[j]);
324 w_ij + w_ji - w_ij * w_ji
326 })
327 .collect()
328 }
329
330 fn edge_distance_directed(
332 &self,
333 offsets: &[usize],
334 row_indices: &[usize],
335 values: &[f32],
336 from: usize,
337 to: usize,
338 ) -> f32 {
339 let start = offsets[from];
340 let end = offsets[from + 1];
341 for idx in start..end {
342 if row_indices[idx] == to {
343 return values[idx];
344 }
345 }
346 f32::INFINITY
347 }
348}
349
350impl WeightedGraph for KnnGraph {
351 fn num_nodes(&self) -> usize {
352 self.n_nodes
353 }
354
355 fn num_edges(&self) -> usize {
356 self.edges.len()
357 }
358
359 fn neighbors_with_weight<'a>(
360 &'a self,
361 node: usize,
362 ) -> Box<dyn Iterator<Item = (usize, f32)> + 'a> {
363 let offsets = self.adjacency.col_offsets();
364 let start = offsets[node];
365 let end = offsets[node + 1];
366 let rows = &self.adjacency.row_indices()[start..end];
367 let vals = &self.adjacency.values()[start..end];
368 Box::new(rows.iter().zip(vals.iter()).map(|(&i, &w)| (i, w)))
369 }
370}
371
372fn smooth_knn_sigma(dists: &[f32], rho: f32, target: f32) -> f32 {
376 const TOLERANCE: f32 = 1e-5;
377 const MAX_ITER: usize = 64;
378
379 let mean_dist: f32 = dists.iter().sum::<f32>() / dists.len().max(1) as f32;
380 let min_sigma = 1e-3 * mean_dist;
381
382 let mut lo = 0.0f32;
383 let mut hi = f32::INFINITY;
384 let mut mid = 1.0f32;
385
386 for _ in 0..MAX_ITER {
387 let mut psum = 0.0f32;
388 for &d in dists {
389 let gap = d - rho;
390 if gap > 0.0 {
391 psum += (-gap / mid).exp();
392 } else {
393 psum += 1.0;
394 }
395 }
396
397 if (psum - target).abs() < TOLERANCE {
398 break;
399 }
400
401 if psum > target {
402 hi = mid;
403 mid = (lo + hi) / 2.0;
404 } else {
405 lo = mid;
406 if hi.is_infinite() {
407 mid *= 2.0;
408 } else {
409 mid = (lo + hi) / 2.0;
410 }
411 }
412 }
413
414 mid.max(min_sigma)
415}
416
417fn directed_umap_weight(d: f32, rho: f32, sigma: f32) -> f32 {
419 if d.is_infinite() || sigma <= 0.0 {
420 return 0.0;
421 }
422 let gap = d - rho;
423 if gap <= 0.0 {
424 1.0
425 } else {
426 (-gap / sigma).exp()
427 }
428}
429
430impl WeightedGraph for crate::leiden::Network {
435 fn num_nodes(&self) -> usize {
436 self.nodes()
437 }
438
439 fn num_edges(&self) -> usize {
440 crate::leiden::Network::edge_count(self)
441 }
442
443 fn neighbors_with_weight<'a>(
444 &'a self,
445 node: usize,
446 ) -> Box<dyn Iterator<Item = (usize, f32)> + 'a> {
447 Box::new(self.neighbors(node).map(|(n, w)| (n, w as f32)))
448 }
449}
450
451#[must_use]
457pub fn modularity_to_cpm_resolution(modularity_gamma: f64, total_edge_weight: f64) -> f64 {
458 modularity_gamma / (2.0 * total_edge_weight).max(1.0)
459}
460
461impl KnnGraph {
462 pub fn to_leiden_network(&self) -> (crate::leiden::Network, f64) {
468 self.to_leiden_network_with(&self.fuzzy_kernel_weights())
469 }
470
471 pub fn to_leiden_network_with(&self, weights: &[f32]) -> (crate::leiden::Network, f64) {
474 let n = self.n_nodes;
475
476 let mut node_degree = vec![0.0f32; n];
477 let mut n_edges = vec![0usize; n];
478 let mut total_edge_weight = 0.0f64;
479 for (&(i, j), &w) in self.edges.iter().zip(weights.iter()) {
480 node_degree[i] += w;
481 node_degree[j] += w;
482 n_edges[i] += 1;
483 n_edges[j] += 1;
484 total_edge_weight += w as f64;
485 }
486
487 let mut network = crate::leiden::Network::with_nodes(&node_degree, &n_edges);
488 for (&(i, j), &w) in self.edges.iter().zip(weights.iter()) {
489 network.add_edge(i, j, w);
490 }
491
492 (network, total_edge_weight)
493 }
494}
495
496pub fn run_leiden(
500 network: &crate::leiden::Network,
501 n: usize,
502 resolution: f64,
503 seed: Option<usize>,
504) -> Vec<usize> {
505 use crate::leiden::clustering::SimpleClustering;
506 use crate::leiden::Clustering;
507 use crate::leiden::Leiden;
508
509 let mut leiden = Leiden::new(resolution, 0.01, seed);
510 let mut clustering = SimpleClustering::init_different_clusters(n);
511
512 for iter in 0..10 {
513 let updated = leiden.iterate(network, &mut clustering);
514 info!(
515 " Leiden iter {}: {} clusters{}",
516 iter + 1,
517 clustering.num_clusters(),
518 if !updated { " (converged)" } else { "" }
519 );
520 if !updated {
521 break;
522 }
523 }
524
525 (0..n).map(|i| clustering.get(i)).collect()
526}
527
528pub fn tune_leiden_resolution(
534 network: &crate::leiden::Network,
535 n: usize,
536 target_k: usize,
537 initial_resolution: f64,
538 seed: Option<usize>,
539) -> Vec<usize> {
540 let mut lo = 1e-6_f64;
541 let mut hi = 10.0_f64;
542 let mut best = run_leiden(network, n, initial_resolution, seed);
543 let best_k = count_distinct(&best);
544
545 info!(
546 " resolution={:.6e} → {} clusters (target {})",
547 initial_resolution, best_k, target_k
548 );
549
550 if best_k == target_k {
551 return best;
552 }
553 if best_k > target_k {
554 hi = initial_resolution;
555 } else {
556 lo = initial_resolution;
557 }
558
559 let mut best_diff = best_k.abs_diff(target_k);
560
561 for _ in 0..20 {
562 let mid = (lo + hi) / 2.0;
563 let result = run_leiden(network, n, mid, seed);
564 let k = count_distinct(&result);
565 info!(" resolution={:.6e} → {} clusters", mid, k);
566
567 if k > target_k {
568 hi = mid;
569 } else {
570 lo = mid;
571 }
572
573 let diff = k.abs_diff(target_k);
574 if diff < best_diff {
575 best = result;
576 best_diff = diff;
577 }
578
579 if k == target_k || (hi - lo) / hi.max(1e-10) < 1e-4 {
580 break;
581 }
582 }
583
584 best
585}
586
587fn count_distinct(labels: &[usize]) -> usize {
589 let max = labels.iter().copied().max().unwrap_or(0);
590 let mut seen = vec![false; max + 1];
591 for &l in labels {
592 seen[l] = true;
593 }
594 seen.iter().filter(|&&s| s).count()
595}
596
597pub fn compact_labels(labels: &mut [usize]) {
599 let max = labels.iter().copied().max().unwrap_or(0);
600 let mut mapping = vec![usize::MAX; max + 1];
601 let mut next = 0usize;
602 for l in labels.iter_mut() {
603 if mapping[*l] == usize::MAX {
604 mapping[*l] = next;
605 next += 1;
606 }
607 *l = mapping[*l];
608 }
609}
610
611fn create_jobs(ntot: usize, block_size: usize) -> Vec<(usize, usize)> {
612 let block_size = if block_size == 0 {
613 DEFAULT_BLOCK_SIZE
614 } else {
615 block_size
616 };
617 let nblock = ntot.div_ceil(block_size);
618 (0..nblock)
619 .map(|block| {
620 let lb = block * block_size;
621 let ub = ((block + 1) * block_size).min(ntot);
622 (lb, ub)
623 })
624 .collect()
625}
626
627#[cfg(test)]
628mod tests;
629
630fn neighbour_lists(data: &DMatrix<f32>, args: &KnnGraphArgs) -> anyhow::Result<Vec<NeighbourList>> {
634 let nn = data.nrows();
635 let n_neighbours = neighbours_per_point(args.knn, nn);
636 Ok(if nn <= EXACT_THRESHOLD {
637 let transposed = data.transpose();
638 let points_vec = transposed.column_iter().collect::<Vec<_>>();
639 let names = (0..nn).collect::<Vec<_>>();
640 let dict = ColumnDict::from_dvector_views(points_vec, names);
641 search_dict(&dict, nn, n_neighbours, args.block_size)?
642 } else {
643 search_rows(data, n_neighbours)
644 })
645}
646
647pub fn umap_memberships(dists: &[f32]) -> Vec<f32> {
652 let Some(rho) = dists.iter().copied().reduce(f32::min) else {
653 return Vec::new();
654 };
655 let target = ((dists.len() + 1) as f32).log2();
656 let sigma = smooth_knn_sigma(dists, rho, target);
657 dists
658 .iter()
659 .map(|&d| directed_umap_weight(d, rho, sigma))
660 .collect()
661}
662
663fn umap_edge_weights(edges: &[(usize, usize)], lists: &[NeighbourList]) -> Vec<f32> {
666 let directed: Vec<Vec<f32>> = lists
668 .par_iter()
669 .map(|(_, dists)| umap_memberships(dists))
670 .collect();
671 let toward = |from: usize, to: usize| -> f32 {
672 let (nb, _) = &lists[from];
673 nb.iter()
674 .position(|&j| j == to)
675 .map_or(0.0, |at| directed[from][at])
676 };
677 edges
678 .par_iter()
679 .map(|&(i, j)| {
680 let (a, b) = (toward(i, j), toward(j, i));
681 a + b - a * b
682 })
683 .collect()
684}
685
686type NeighbourList = (Vec<usize>, Vec<f32>);
687
688type TaggedEdge = ((usize, usize), f32, u8);
691
692fn fold_tagged_edges(mut tagged: Vec<TaggedEdge>) -> Vec<TaggedEdge> {
699 tagged.par_sort_unstable_by_key(|&(key, _, _)| key);
700 tagged.dedup_by(|cur, prev| {
701 if cur.0 == prev.0 {
702 prev.1 = prev.1.min(cur.1);
703 prev.2 |= cur.2;
704 true
705 } else {
706 false
707 }
708 });
709 tagged
710}
711
712fn neighbours_per_point(knn: usize, nn: usize) -> usize {
715 knn.min(nn.saturating_sub(1)).max(1)
716}
717
718fn search_dict(
721 dict: &ColumnDict<usize>,
722 nn: usize,
723 n_neighbours: usize,
724 block_size: usize,
725) -> anyhow::Result<Vec<NeighbourList>> {
726 let jobs = create_jobs(nn, block_size);
727 let search_bar =
731 crate::matrix::progress::new_progress_bar(jobs.len() as u64).with_message("kNN blocks");
732 let result: anyhow::Result<Vec<Vec<NeighbourList>>> = jobs
733 .into_par_iter()
734 .progress_with(search_bar.clone())
735 .map(|(lb, ub)| {
736 let mut scratch = SearchScratch::default();
738 (lb..ub)
739 .map(|i| dict.search_others_reuse(&i, n_neighbours, &mut scratch))
740 .collect()
741 })
742 .collect();
743 search_bar.finish_and_clear();
746 Ok(result?.into_iter().flatten().collect())
747}
748
749fn search_rows(rows: &DMatrix<f32>, n_neighbours: usize) -> Vec<NeighbourList> {
753 let nn = rows.nrows();
754 let (indices, distances) = if nn <= ALL_PAIRS_THRESHOLD {
755 info!("kNN by the exact all-pairs kernel over {nn} points");
756 knn_rows_l2(rows, n_neighbours)
757 } else {
758 info!("kNN by the inverted-file search over {nn} points");
759 knn_rows_ivf(
760 rows,
761 &IvfArgs {
762 k: n_neighbours,
763 n_lists: 0,
764 n_probe: DEFAULT_N_PROBE,
765 seed: KNN_SEED,
766 },
767 )
768 };
769 indices.into_iter().zip(distances).collect()
770}
771
772pub(crate) fn edges_from_neighbours(
777 lists: &[NeighbourList],
778 reciprocal: bool,
779) -> Vec<((usize, usize), f32)> {
780 let tagged: Vec<TaggedEdge> = lists
783 .par_iter()
784 .enumerate()
785 .flat_map_iter(|(i, (nb, ds))| {
786 nb.iter().zip(ds).map(move |(&j, &d)| {
787 if i < j {
788 ((i, j), d, 1u8)
789 } else {
790 ((j, i), d, 2u8)
791 }
792 })
793 })
794 .collect();
795 fold_tagged_edges(tagged)
796 .into_iter()
797 .filter(|&(_, _, mask)| !reciprocal || mask == 3)
798 .map(|(key, dist, _)| (key, dist))
799 .collect()
800}
801
802pub fn symmetric_adjacency(
814 n_nodes: usize,
815 edges: &[(usize, usize)],
816 distances: &[f32],
817) -> CscMatrix<f32> {
818 debug_assert!(
819 edges.iter().all(|&(i, j)| i < j) && edges.windows(2).all(|w| w[0] < w[1]),
820 "symmetric adjacency: edges must be canonical, sorted and unique"
821 );
822 let mut offsets = vec![0usize; n_nodes + 1];
823 for &(i, j) in edges {
824 offsets[i + 1] += 1;
825 offsets[j + 1] += 1;
826 }
827 for c in 0..n_nodes {
828 offsets[c + 1] += offsets[c];
829 }
830 let nnz = offsets[n_nodes];
831 let mut row_indices = vec![0usize; nnz];
832 let mut values = vec![0f32; nnz];
833 let mut cursor = offsets[..n_nodes].to_vec();
834 for (&(i, j), &v) in edges.iter().zip(distances) {
835 row_indices[cursor[i]] = j;
836 values[cursor[i]] = v;
837 cursor[i] += 1;
838 row_indices[cursor[j]] = i;
839 values[cursor[j]] = v;
840 cursor[j] += 1;
841 }
842 CscMatrix::try_from_csc_data(n_nodes, n_nodes, offsets, row_indices, values)
843 .expect("symmetric adjacency: canonical sorted edges fill every column in order")
844}
845
846fn within_source_rank(d: &[f32]) -> Vec<f32> {
852 if d.len() <= 1 {
853 return vec![0.0; d.len()];
854 }
855 let mut order: Vec<usize> = (0..d.len()).collect();
856 order.par_sort_unstable_by(|&a, &b| d[a].total_cmp(&d[b]).then(a.cmp(&b)));
866 let mut out = vec![0.0f32; d.len()];
867 let denom = (d.len() - 1) as f32;
868 for (rank, &idx) in order.iter().enumerate() {
869 out[idx] = rank as f32 / denom;
870 }
871 out
872}