Skip to main content

legume_numeric/matrix/
knn_graph.rs

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
15/// Up to this many points every row's neighbours come from the exact
16/// all-pairs Gram kernel; beyond it from the inverted-file search. Both are
17/// parallel and thread-count independent; the split is where `O(n²)` stops
18/// being affordable.
19pub const ALL_PAIRS_THRESHOLD: usize = 65_536;
20
21pub struct KnnGraph {
22    /// Symmetric CSC adjacency matrix (n_nodes x n_nodes)
23    pub adjacency: CscMatrix<f32>,
24    /// Sorted edge list (i < j), deduplicated
25    pub edges: Vec<(usize, usize)>,
26    /// Edge distances/weights, parallel to `edges`
27    pub distances: Vec<f32>,
28    /// Number of nodes
29    pub n_nodes: usize,
30}
31
32/// Which input graph an edge of a [`KnnGraph::union_with`] came from.
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub enum EdgeSource {
35    Primary,
36    Secondary,
37    /// Present in both inputs. Still a primary-graph edge for any consumer
38    /// filtering on the primary relation, since the primary relation holds.
39    Both,
40}
41
42/// How to reconcile two graphs' `distances` when merging.
43#[derive(Debug, Clone, Copy, PartialEq, Eq)]
44pub enum DistanceMerge {
45    /// Keep raw values. Correct only when both inputs measure the same thing.
46    Raw,
47    /// Replace each side's distances with its own within-source quantile rank
48    /// in `[0, 1]` before merging. Use this whenever the inputs measure
49    /// different things, so the merged column stays comparable across sources
50    /// and monotone within each.
51    SourceRank,
52}
53
54pub struct KnnGraphArgs {
55    pub knn: usize,
56    pub block_size: usize,
57    /// If true, keep only reciprocal edges (i→j AND j→i).
58    /// If false, keep union edges (i→j OR j→i), using min distance.
59    pub reciprocal: bool,
60}
61
62impl KnnGraph {
63    /// Build a KNN graph from column vectors.
64    ///
65    /// * `points` - transposed coordinate matrix (d x n), where each column is a point
66    /// * `args` - KNN graph construction parameters
67    pub fn from_columns(points: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
68        Self::from_rows(&points.transpose(), args)
69    }
70
71    /// Build a KNN graph from row vectors (cells × features).
72    ///
73    /// * `data` - matrix (n x d), where each row is a point
74    /// * `args` - KNN graph construction parameters
75    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    /// [`KnnGraph::from_columns_fuzzy`] over row vectors.
81    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    /// The kNN graph of the columns of `points`, and UMAP's fuzzy membership
92    /// of each edge (parallel to `edges`), as umap-learn and uwot compute it:
93    ///
94    /// 1. each point's weights over its OWN `knn` neighbours only:
95    ///    `exp(-(d - ρ) / σ)`, ρ its nearest distance, σ set so they sum to
96    ///    `log2(knn + 1)` (UMAP counts the point itself among its
97    ///    `n_neighbors`, so `knn` others is `n_neighbors = knn + 1`);
98    /// 2. zero toward a point it did not list;
99    /// 3. the fuzzy union of the two directions, `a + b - a·b`.
100    ///
101    /// [`KnnGraph::fuzzy_kernel_weights`] instead calibrates each point over
102    /// every edge touching it after the union, so a point many others list
103    /// gets a wider kernel and a one-sided edge a weight from both ends.
104    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    /// The graph implied by every point's directed neighbour list.
112    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        // Filtering and the sort ran silent, which on a large pair graph is a
124        // stretch of nothing between the search bar and the next log line,
125        // and reads as a hang.
126        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    /// Merge two graphs over the same nodes, keeping every pair exactly once
149    /// and reporting which input each came from.
150    ///
151    /// Borrows both inputs: a caller that unions a spatial graph with an
152    /// expression one generally still needs the spatial graph afterwards, as
153    /// the topology for anything that reasons about physical adjacency.
154    ///
155    /// Neither input is assumed sorted, nor assumed to store `i < j`. Edge
156    /// order is a constructor invariant here, not a type invariant, and a
157    /// hand-built `KnnGraph` can violate both.
158    ///
159    /// `distances` after a union are NOT a metric. Under
160    /// [`DistanceMerge::SourceRank`] they are within-source quantile ranks,
161    /// which keeps them comparable across sources without pretending the two
162    /// measurements are the same quantity. When an edge is in both inputs the
163    /// smaller value wins, matching the `reciprocal: false` convention in
164    /// `build_from_dict`.
165    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        // Sort then fold, NOT a keyed map. At a few million edges an ordered
187        // map costs a pointer-chasing O(log n) descent and a node allocation
188        // per insert, all of it serial. One flat buffer, one parallel sort and
189        // one linear scan replaces that, and measured an order of magnitude
190        // faster. The allocation count drops too, though only the time was
191        // measured.
192        //
193        // The canonical key is what makes this a set operation: a pair stored
194        // one way round in one input and the other way round in the other must
195        // land on the same key. The dedup has to finish before the COO below,
196        // which SUMS duplicate entries rather than rejecting them.
197        let canonical = |&(i, j): &(usize, usize)| if i <= j { (i, j) } else { (j, i) };
198        // Source as a bitmask, so folding a run is an OR rather than a case
199        // analysis: 1 = primary, 2 = secondary, 3 = both.
200        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        // Derived state, so rebuild rather than merge.
229        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    /// Get neighbors of a node from the CSC adjacency matrix
243    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    /// Convert distances to similarity weights using an exponential kernel:
259    /// `w = exp(-d / σ)` where σ = median distance.
260    ///
261    /// Returns weights parallel to `self.edges`, all in (0, 1].
262    /// Consistent with the softmax(-d) pattern used in counterfactual
263    /// inference (data_beans::alg) but with a global bandwidth.
264    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    /// Adaptive-bandwidth kernel weights with local connectivity.
275    ///
276    /// Per-point sigma calibration (originated in t-SNE, van der Maaten
277    /// & Hinton 2008) ensures every node has the same effective number
278    /// of neighbors, preventing isolated singletons in sparse regions.
279    /// The rho subtraction and fuzzy-union symmetrization follow UMAP
280    /// (McInnes et al. 2018), matching the scanpy default for Leiden.
281    ///
282    /// Algorithm:
283    /// 1. rho_i = distance to nearest neighbor (local connectivity)
284    /// 2. sigma_i via binary search: sum_j exp(-(d_ij - rho_i)/sigma_i) = log2(k)
285    /// 3. Directed weight: w(i→j) = exp(-(d_ij - rho_i) / sigma_i)
286    /// 4. Symmetrize: w_sym = w(i→j) + w(j→i) - w(i→j) * w(j→i)
287    ///
288    /// Returns weights parallel to `self.edges`, all in (0, 1].
289    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        // Step 1-2: compute rho and sigma per node — independent per node.
299        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        // Step 3-4: compute directed weights and symmetrize per edge —
316        // independent per edge, only reads rho/sigma.
317        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                // fuzzy union: P(at least one edge) = P(A) + P(B) - P(A)*P(B)
325                w_ij + w_ji - w_ij * w_ji
326            })
327            .collect()
328    }
329
330    /// Look up the distance from node `from` to node `to` in the CSC adjacency.
331    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
372/// Binary search for per-point sigma (UMAP's smooth_knn_dist).
373///
374/// Finds sigma such that: sum_j exp(-max(0, d_j - rho) / sigma) = target
375fn 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
417/// Compute a single directed UMAP membership weight.
418fn 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
430////////////////////////
431// Leiden integration //
432////////////////////////
433
434impl 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/// Convert a modularity resolution `gamma` to the CPM scale expected by the
452/// Leiden crate, given the total undirected edge weight of the graph.
453///
454/// CPM resolution = `gamma / (2 * total_edge_weight)`. Guards against division
455/// by zero for degenerate graphs by clamping the denominator to at least 1.
456#[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    /// Convert this KNN graph to a Leiden `Network` with modularity objective.
463    ///
464    /// Node weights = weighted degree, edge weights = fuzzy kernel weights.
465    /// Returns `(network, total_edge_weight)`. Pass `total_edge_weight` to
466    /// [`modularity_to_cpm_resolution`] to get a CPM-scale resolution.
467    pub fn to_leiden_network(&self) -> (crate::leiden::Network, f64) {
468        self.to_leiden_network_with(&self.fuzzy_kernel_weights())
469    }
470
471    /// [`KnnGraph::to_leiden_network`] with the edge weights given, parallel
472    /// to `edges` (e.g. from [`KnnGraph::from_rows_fuzzy`]).
473    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
496/// Run Leiden clustering at a fixed (already-scaled) resolution.
497///
498/// Returns cluster labels as `Vec<usize>` (not necessarily contiguous).
499pub 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
528/// Binary search on Leiden resolution to approximate `target_k` clusters.
529///
530/// `initial_resolution` should already be on the CPM scale
531/// (i.e., `modularity_gamma / (2 * total_edge_weight)`).
532/// Returns cluster labels (not necessarily contiguous).
533pub 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
587/// Count distinct values in a label vector.
588fn 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
597/// Remap labels to contiguous 0..k.
598pub 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
630/// One point's neighbours: `(indices, distances)`, nearest first.
631/// Every point's own neighbours (self excluded), exact up to
632/// [`EXACT_THRESHOLD`] points.
633fn 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
647/// UMAP's membership of a point in each of its neighbours, from its
648/// distances to them (itself excluded): `exp(-(d - ρ) / σ)`, ρ the nearest
649/// distance and σ set so they sum to `log2(len + 1)`, as UMAP counts the
650/// point itself among its `n_neighbors`.
651pub 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
663/// UMAP's fuzzy membership of each of `edges` (canonical `i < j`) from the
664/// points' own neighbour lists; see [`KnnGraph::from_columns_fuzzy`].
665fn umap_edge_weights(edges: &[(usize, usize)], lists: &[NeighbourList]) -> Vec<f32> {
666    // Each point's directed weights, parallel to its own list.
667    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
688/// A canonical `(i, j)` key with a distance and a bitmask saying which
689/// inputs listed it.
690type TaggedEdge = ((usize, usize), f32, u8);
691
692/// One flat buffer of canonical keys, one parallel sort and one linear fold:
693/// a run of equal keys keeps the smallest distance and the OR of its masks.
694/// What a keyed map over tens of millions of triplets would do with an
695/// allocation and a contended insert per triplet and a lookup per edge, as a
696/// set operation done in place. The result is sorted by key, and the dedup
697/// has to finish before any CSC build, which SUMS duplicates.
698fn 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
712/// `search_others` returns exactly this many *other* neighbours (self
713/// excluded): the request clamped to the available others, floored at 1.
714fn neighbours_per_point(knn: usize, nn: usize) -> usize {
715    knn.min(nn.saturating_sub(1)).max(1)
716}
717
718/// Every point's neighbours from the exact per-query scan of a
719/// [`ColumnDict`], in parallel blocks — the arm for small point sets.
720fn 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    // Every bar draws through `crate::matrix::progress` so it shares the one
728    // `MultiProgress` the log bridge writes above; a bar built straight
729    // from indicatif registers with neither and corrupts the log.
730    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            // One scratch per block, reused across the block's queries.
737            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    // Clear BEFORE propagating: an error would otherwise leave the bar
744    // ticking over the caller's error output.
745    search_bar.finish_and_clear();
746    Ok(result?.into_iter().flatten().collect())
747}
748
749/// Every row's neighbours among the other rows, exactly by the all-pairs Gram
750/// kernel up to [`ALL_PAIRS_THRESHOLD`] rows and by the inverted-file search
751/// beyond it.
752fn 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
772/// Undirected edges from directed neighbour lists, as `((i, j), distance)`
773/// with `i < j`, sorted. `reciprocal` keeps a pair only when each point
774/// listed the other; otherwise either direction suffices and the smaller of
775/// the two distances is kept.
776pub(crate) fn edges_from_neighbours(
777    lists: &[NeighbourList],
778    reciprocal: bool,
779) -> Vec<((usize, usize), f32)> {
780    // Direction as a bit so a run folds by OR: 1 = listed by the smaller
781    // index, 2 = by the larger.
782    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
802/// The `n x n` symmetric adjacency implied by an undirected edge list, which
803/// must be canonical (`i < j`), sorted and free of duplicates — what every
804/// constructor here produces.
805///
806/// Each edge lands in BOTH endpoints' columns, exactly once per direction.
807/// Built straight from the edge list — degrees, offsets, one fill — rather
808/// than through a `CooMatrix`, whose conversion sorts every entry of the
809/// whole matrix serially. Sorted input is what makes the fill enough: column
810/// `c` receives its partners below `c` in ascending order as the edges
811/// `(i, c)` pass, then its partners above `c` in ascending order from the run
812/// of edges `(c, j)`.
813pub 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
846/// Each value replaced by its rank among the others, scaled to `[0, 1]`.
847///
848/// Ties take distinct adjacent ranks, which is harmless: the point is to put
849/// two incomparable distance scales on one axis, not to be a faithful
850/// empirical CDF.
851fn 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    // Ties break on the index, which is what makes this safe to sort in
857    // parallel: an unstable parallel sort would otherwise put equal distances
858    // in a run-dependent order, and these ranks are written out, so the file
859    // would stop being reproducible. Grid-spaced coordinates tie constantly,
860    // so this is the common case rather than a corner.
861    //
862    // `total_cmp` rather than `partial_cmp().unwrap_or(Equal)`: it is a total
863    // order, so it needs no per-comparison branch and it gives NaN a definite
864    // position instead of making it a wildcard that breaks transitivity.
865    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}