Skip to main content

legume_numeric/matrix/
knn_graph.rs

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    /// Symmetric CSC adjacency matrix (n_nodes x n_nodes)
18    pub adjacency: CscMatrix<f32>,
19    /// Sorted edge list (i < j), deduplicated
20    pub edges: Vec<(usize, usize)>,
21    /// Edge distances/weights, parallel to `edges`
22    pub distances: Vec<f32>,
23    /// Number of nodes
24    pub n_nodes: usize,
25}
26
27/// Which input graph an edge of a [`KnnGraph::union_with`] came from.
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum EdgeSource {
30    Primary,
31    Secondary,
32    /// Present in both inputs. Still a primary-graph edge for any consumer
33    /// filtering on the primary relation, since the primary relation holds.
34    Both,
35}
36
37/// How to reconcile two graphs' `distances` when merging.
38#[derive(Debug, Clone, Copy, PartialEq, Eq)]
39pub enum DistanceMerge {
40    /// Keep raw values. Correct only when both inputs measure the same thing.
41    Raw,
42    /// Replace each side's distances with its own within-source quantile rank
43    /// in `[0, 1]` before merging. Use this whenever the inputs measure
44    /// different things, so the merged column stays comparable across sources
45    /// and monotone within each.
46    SourceRank,
47}
48
49pub struct KnnGraphArgs {
50    pub knn: usize,
51    pub block_size: usize,
52    /// If true, keep only reciprocal edges (i→j AND j→i).
53    /// If false, keep union edges (i→j OR j→i), using min distance.
54    pub reciprocal: bool,
55}
56
57impl KnnGraph {
58    /// Build a KNN graph from column vectors.
59    ///
60    /// * `points` - transposed coordinate matrix (d x n), where each column is a point
61    /// * `args` - KNN graph construction parameters
62    pub fn from_columns(points: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
63        Self::from_rows(&points.transpose(), args)
64    }
65
66    /// Build a KNN graph from row vectors (cells × features).
67    ///
68    /// * `data` - matrix (n x d), where each row is a point
69    /// * `args` - KNN graph construction parameters
70    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    /// [`KnnGraph::from_columns_fuzzy`] over row vectors.
76    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    /// The kNN graph of the columns of `points`, and UMAP's fuzzy membership
87    /// of each edge (parallel to `edges`), as umap-learn and uwot compute it:
88    ///
89    /// 1. each point's weights over its OWN `knn` neighbours only:
90    ///    `exp(-(d - ρ) / σ)`, ρ its nearest distance, σ set so they sum to
91    ///    `log2(knn + 1)` (UMAP counts the point itself among its
92    ///    `n_neighbors`, so `knn` others is `n_neighbors = knn + 1`);
93    /// 2. zero toward a point it did not list;
94    /// 3. the fuzzy union of the two directions, `a + b - a·b`.
95    ///
96    /// [`KnnGraph::fuzzy_kernel_weights`] instead calibrates each point over
97    /// every edge touching it after the union, so a point many others list
98    /// gets a wider kernel and a one-sided edge a weight from both ends.
99    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    /// The graph implied by every point's directed neighbour list.
107    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        // Filtering and the sort ran silent, which on a large pair graph is a
119        // stretch of nothing between the search bar and the next log line,
120        // and reads as a hang.
121        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    /// Merge two graphs over the same nodes, keeping every pair exactly once
144    /// and reporting which input each came from.
145    ///
146    /// Borrows both inputs: a caller that unions a spatial graph with an
147    /// expression one generally still needs the spatial graph afterwards, as
148    /// the topology for anything that reasons about physical adjacency.
149    ///
150    /// Neither input is assumed sorted, nor assumed to store `i < j`. Edge
151    /// order is a constructor invariant here, not a type invariant, and a
152    /// hand-built `KnnGraph` can violate both.
153    ///
154    /// `distances` after a union are NOT a metric. Under
155    /// [`DistanceMerge::SourceRank`] they are within-source quantile ranks,
156    /// which keeps them comparable across sources without pretending the two
157    /// measurements are the same quantity. When an edge is in both inputs the
158    /// smaller value wins, matching the `reciprocal: false` convention in
159    /// `build_from_dict`.
160    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        // Sort then fold, NOT a keyed map. At a few million edges an ordered
182        // map costs a pointer-chasing O(log n) descent and a node allocation
183        // per insert, all of it serial. One flat buffer, one parallel sort and
184        // one linear scan replaces that, and measured an order of magnitude
185        // faster. The allocation count drops too, though only the time was
186        // measured.
187        //
188        // The canonical key is what makes this a set operation: a pair stored
189        // one way round in one input and the other way round in the other must
190        // land on the same key. The dedup has to finish before the COO below,
191        // which SUMS duplicate entries rather than rejecting them.
192        let canonical = |&(i, j): &(usize, usize)| if i <= j { (i, j) } else { (j, i) };
193        // Source as a bitmask, so folding a run is an OR rather than a case
194        // analysis: 1 = primary, 2 = secondary, 3 = both.
195        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        // Derived state, so rebuild rather than merge.
224        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    /// Get neighbors of a node from the CSC adjacency matrix
238    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    /// Convert distances to similarity weights using an exponential kernel:
254    /// `w = exp(-d / σ)` where σ = median distance.
255    ///
256    /// Returns weights parallel to `self.edges`, all in (0, 1].
257    /// Consistent with the softmax(-d) pattern used in counterfactual
258    /// inference (data_beans::alg) but with a global bandwidth.
259    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    /// Adaptive-bandwidth kernel weights with local connectivity.
270    ///
271    /// Per-point sigma calibration (originated in t-SNE, van der Maaten
272    /// & Hinton 2008) ensures every node has the same effective number
273    /// of neighbors, preventing isolated singletons in sparse regions.
274    /// The rho subtraction and fuzzy-union symmetrization follow UMAP
275    /// (McInnes et al. 2018), matching the scanpy default for Leiden.
276    ///
277    /// Algorithm:
278    /// 1. rho_i = distance to nearest neighbor (local connectivity)
279    /// 2. sigma_i via binary search: sum_j exp(-(d_ij - rho_i)/sigma_i) = log2(k)
280    /// 3. Directed weight: w(i→j) = exp(-(d_ij - rho_i) / sigma_i)
281    /// 4. Symmetrize: w_sym = w(i→j) + w(j→i) - w(i→j) * w(j→i)
282    ///
283    /// Returns weights parallel to `self.edges`, all in (0, 1].
284    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        // Step 1-2: compute rho and sigma per node — independent per node.
294        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        // Step 3-4: compute directed weights and symmetrize per edge —
311        // independent per edge, only reads rho/sigma.
312        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                // fuzzy union: P(at least one edge) = P(A) + P(B) - P(A)*P(B)
320                w_ij + w_ji - w_ij * w_ji
321            })
322            .collect()
323    }
324
325    /// Look up the distance from node `from` to node `to` in the CSC adjacency.
326    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
367/// Binary search for per-point sigma (UMAP's smooth_knn_dist).
368///
369/// Finds sigma such that: sum_j exp(-max(0, d_j - rho) / sigma) = target
370fn 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
412/// Compute a single directed UMAP membership weight.
413fn 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
425////////////////////////
426// Leiden integration //
427////////////////////////
428
429impl 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/// Convert a modularity resolution `gamma` to the CPM scale expected by the
447/// Leiden crate, given the total undirected edge weight of the graph.
448///
449/// CPM resolution = `gamma / (2 * total_edge_weight)`. Guards against division
450/// by zero for degenerate graphs by clamping the denominator to at least 1.
451#[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    /// Convert this KNN graph to a Leiden `Network` with modularity objective.
458    ///
459    /// Node weights = weighted degree, edge weights = fuzzy kernel weights.
460    /// Returns `(network, total_edge_weight)`. Pass `total_edge_weight` to
461    /// [`modularity_to_cpm_resolution`] to get a CPM-scale resolution.
462    pub fn to_leiden_network(&self) -> (crate::leiden::Network, f64) {
463        self.to_leiden_network_with(&self.fuzzy_kernel_weights())
464    }
465
466    /// [`KnnGraph::to_leiden_network`] with the edge weights given, parallel
467    /// to `edges` (e.g. from [`KnnGraph::from_rows_fuzzy`]).
468    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
491/// Run Leiden clustering at a fixed (already-scaled) resolution.
492///
493/// Returns cluster labels as `Vec<usize>` (not necessarily contiguous).
494pub 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
523/// Binary search on Leiden resolution to approximate `target_k` clusters.
524///
525/// `initial_resolution` should already be on the CPM scale
526/// (i.e., `modularity_gamma / (2 * total_edge_weight)`).
527/// Returns cluster labels (not necessarily contiguous).
528pub 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
582/// Count distinct values in a label vector.
583fn 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
592/// Remap labels to contiguous 0..k.
593pub 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
625/// One point's neighbours: `(indices, distances)`, nearest first.
626/// Every point's own neighbours (self excluded), exact up to
627/// [`EXACT_THRESHOLD`] points.
628fn 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
642/// UMAP's membership of a point in each of its neighbours, from its
643/// distances to them (itself excluded): `exp(-(d - ρ) / σ)`, ρ the nearest
644/// distance and σ set so they sum to `log2(len + 1)`, as UMAP counts the
645/// point itself among its `n_neighbors`.
646pub 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
658/// UMAP's fuzzy membership of each of `edges` (canonical `i < j`) from the
659/// points' own neighbour lists; see [`KnnGraph::from_columns_fuzzy`].
660fn umap_edge_weights(edges: &[(usize, usize)], lists: &[NeighbourList]) -> Vec<f32> {
661    // Each point's directed weights, parallel to its own list.
662    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
683/// A canonical `(i, j)` key with a distance and a bitmask saying which
684/// inputs listed it.
685type TaggedEdge = ((usize, usize), f32, u8);
686
687/// One flat buffer of canonical keys, one parallel sort and one linear fold:
688/// a run of equal keys keeps the smallest distance and the OR of its masks.
689/// What a keyed map over tens of millions of triplets would do with an
690/// allocation and a contended insert per triplet and a lookup per edge, as a
691/// set operation done in place. The result is sorted by key, and the dedup
692/// has to finish before any CSC build, which SUMS duplicates.
693fn 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
707/// `search_others` returns exactly this many *other* neighbours (self
708/// excluded): the request clamped to the available others, floored at 1.
709fn neighbours_per_point(knn: usize, nn: usize) -> usize {
710    knn.min(nn.saturating_sub(1)).max(1)
711}
712
713/// Every point's neighbours from the exact per-query scan of a
714/// [`ColumnDict`], in parallel blocks — the arm for small point sets.
715fn 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    // Every bar draws through `crate::matrix::progress` so it shares the one
723    // `MultiProgress` the log bridge writes above; a bar built straight
724    // from indicatif registers with neither and corrupts the log.
725    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            // One scratch per block, reused across the block's queries.
732            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    // Clear BEFORE propagating: an error would otherwise leave the bar
739    // ticking over the caller's error output.
740    search_bar.finish_and_clear();
741    Ok(result?.into_iter().flatten().collect())
742}
743
744/// Every row's neighbour list, by [`knn_rows`].
745fn 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
750/// Undirected edges from directed neighbour lists, as `((i, j), distance)`
751/// with `i < j`, sorted. `reciprocal` keeps a pair only when each point
752/// listed the other; otherwise either direction suffices and the smaller of
753/// the two distances is kept.
754pub(crate) fn edges_from_neighbours(
755    lists: &[NeighbourList],
756    reciprocal: bool,
757) -> Vec<((usize, usize), f32)> {
758    // Direction as a bit so a run folds by OR: 1 = listed by the smaller
759    // index, 2 = by the larger.
760    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
780/// The `n x n` symmetric adjacency implied by an undirected edge list, which
781/// must be canonical (`i < j`), sorted and free of duplicates — what every
782/// constructor here produces.
783///
784/// Each edge lands in BOTH endpoints' columns, exactly once per direction.
785/// Built straight from the edge list — degrees, offsets, one fill — rather
786/// than through a `CooMatrix`, whose conversion sorts every entry of the
787/// whole matrix serially. Sorted input is what makes the fill enough: column
788/// `c` receives its partners below `c` in ascending order as the edges
789/// `(i, c)` pass, then its partners above `c` in ascending order from the run
790/// of edges `(c, j)`.
791pub 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
824/// Each value replaced by its rank among the others, scaled to `[0, 1]`.
825///
826/// Ties take distinct adjacent ranks, which is harmless: the point is to put
827/// two incomparable distance scales on one axis, not to be a faithful
828/// empirical CDF.
829fn 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    // Ties break on the index, which is what makes this safe to sort in
835    // parallel: an unstable parallel sort would otherwise put equal distances
836    // in a run-dependent order, and these ranks are written out, so the file
837    // would stop being reproducible. Grid-spaced coordinates tie constantly,
838    // so this is the common case rather than a corner.
839    //
840    // `total_cmp` rather than `partial_cmp().unwrap_or(Equal)`: it is a total
841    // order, so it needs no per-comparison branch and it gives NaN a definite
842    // position instead of making it a wildcard that breaks transitivity.
843    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}