use crate::matrix::graph::WeightedGraph;
use crate::matrix::knn::all_pairs::knn_rows_l2;
use crate::matrix::knn::ivf::{knn_rows_ivf, IvfArgs, DEFAULT_N_PROBE};
use crate::matrix::knn::{EXACT_THRESHOLD, KNN_SEED};
use crate::matrix::knn_match::{ColumnDict, SearchScratch};
use indicatif::ParallelProgressIterator;
use log::info;
use nalgebra::DMatrix;
use nalgebra_sparse::CscMatrix;
use rayon::prelude::*;
const DEFAULT_BLOCK_SIZE: usize = 1000;
pub const ALL_PAIRS_THRESHOLD: usize = 65_536;
pub struct KnnGraph {
pub adjacency: CscMatrix<f32>,
pub edges: Vec<(usize, usize)>,
pub distances: Vec<f32>,
pub n_nodes: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EdgeSource {
Primary,
Secondary,
Both,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DistanceMerge {
Raw,
SourceRank,
}
pub struct KnnGraphArgs {
pub knn: usize,
pub block_size: usize,
pub reciprocal: bool,
}
impl KnnGraph {
pub fn from_columns(points: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
Self::from_rows(&points.transpose(), args)
}
pub fn from_rows(data: &DMatrix<f32>, args: KnnGraphArgs) -> anyhow::Result<KnnGraph> {
let nn = data.nrows();
let n_neighbours = neighbours_per_point(args.knn, nn);
let lists = if nn <= EXACT_THRESHOLD {
let transposed = data.transpose();
let points_vec = transposed.column_iter().collect::<Vec<_>>();
let names = (0..nn).collect::<Vec<_>>();
let dict = ColumnDict::from_dvector_views(points_vec, names);
search_dict(&dict, nn, n_neighbours, args.block_size)?
} else {
search_rows(data, n_neighbours)
};
Self::from_neighbours(nn, &lists, args.reciprocal)
}
fn from_neighbours(
nn: usize,
lists: &[NeighbourList],
reciprocal: bool,
) -> anyhow::Result<KnnGraph> {
let n_triplets: usize = lists.iter().map(|(nb, _)| nb.len()).sum();
info!("{n_triplets} triplets by kNN matching");
if n_triplets == 0 {
return Err(anyhow::anyhow!("empty triplets"));
}
let filter_spin =
crate::matrix::progress::new_spinner("{spinner} [{elapsed_precise}] {msg}")
.with_message("filtering edges");
let edges = edges_from_neighbours(lists, reciprocal);
filter_spin.finish_and_clear();
info!(
"{} edges after {} matching",
edges.len(),
if reciprocal { "reciprocal" } else { "union" }
);
let (edge_pairs, distances): (Vec<_>, Vec<_>) = edges.into_iter().unzip();
let adjacency = symmetric_adjacency(nn, &edge_pairs, &distances);
Ok(KnnGraph {
adjacency,
edges: edge_pairs,
distances,
n_nodes: nn,
})
}
pub fn union_with(
&self,
other: &KnnGraph,
policy: DistanceMerge,
) -> anyhow::Result<(KnnGraph, Vec<EdgeSource>)> {
anyhow::ensure!(
self.n_nodes == other.n_nodes,
"cannot union graphs over different node counts: {} vs {}",
self.n_nodes,
other.n_nodes
);
let n_nodes = self.n_nodes;
let (a_dist, b_dist) = match policy {
DistanceMerge::Raw => (self.distances.clone(), other.distances.clone()),
DistanceMerge::SourceRank => (
within_source_rank(&self.distances),
within_source_rank(&other.distances),
),
};
let canonical = |&(i, j): &(usize, usize)| if i <= j { (i, j) } else { (j, i) };
let mut tagged: Vec<TaggedEdge> = Vec::with_capacity(self.edges.len() + other.edges.len());
tagged.par_extend(
self.edges
.par_iter()
.zip(a_dist.par_iter())
.map(|(e, &d)| (canonical(e), d, 1u8)),
);
tagged.par_extend(
other
.edges
.par_iter()
.zip(b_dist.par_iter())
.map(|(e, &d)| (canonical(e), d, 2u8)),
);
let folded = fold_tagged_edges(tagged);
let mut edges = Vec::with_capacity(folded.len());
let mut distances = Vec::with_capacity(folded.len());
let mut source = Vec::with_capacity(folded.len());
for (key, dist, mask) in folded {
edges.push(key);
distances.push(dist);
source.push(match mask {
1 => EdgeSource::Primary,
2 => EdgeSource::Secondary,
_ => EdgeSource::Both,
});
}
let adjacency = symmetric_adjacency(n_nodes, &edges, &distances);
Ok((
KnnGraph {
adjacency,
edges,
distances,
n_nodes,
},
source,
))
}
pub fn neighbors(&self, node: usize) -> &[usize] {
let offsets = self.adjacency.col_offsets();
let start = offsets[node];
let end = offsets[node + 1];
&self.adjacency.row_indices()[start..end]
}
pub fn num_edges(&self) -> usize {
self.edges.len()
}
pub fn num_nodes(&self) -> usize {
self.n_nodes
}
pub fn exp_kernel_weights(&self) -> Vec<f32> {
if self.distances.is_empty() {
return Vec::new();
}
let sigma = crate::matrix::utils::median(&self.distances);
let sigma = if sigma <= 0.0 { 1.0 } else { sigma };
info!("exp_kernel_weights: σ (median distance) = {:.4}", sigma);
self.distances.iter().map(|&d| (-d / sigma).exp()).collect()
}
pub fn fuzzy_kernel_weights(&self) -> Vec<f32> {
if self.distances.is_empty() {
return Vec::new();
}
let offsets = self.adjacency.col_offsets();
let row_indices = self.adjacency.row_indices();
let values = self.adjacency.values();
let (rho, sigma): (Vec<f32>, Vec<f32>) = (0..self.n_nodes)
.into_par_iter()
.map(|i| {
let start = offsets[i];
let end = offsets[i + 1];
let dists: Vec<f32> = (start..end).map(|idx| values[idx]).collect();
if dists.is_empty() {
return (0.0_f32, 1.0_f32);
}
let rho_i = dists.iter().cloned().fold(f32::INFINITY, f32::min);
let target = (dists.len() as f32).log2();
let sigma_i = smooth_knn_sigma(&dists, rho_i, target);
(rho_i, sigma_i)
})
.unzip();
self.edges
.par_iter()
.map(|&(i, j)| {
let d_ij = self.edge_distance_directed(offsets, row_indices, values, i, j);
let w_ij = directed_umap_weight(d_ij, rho[i], sigma[i]);
let d_ji = self.edge_distance_directed(offsets, row_indices, values, j, i);
let w_ji = directed_umap_weight(d_ji, rho[j], sigma[j]);
w_ij + w_ji - w_ij * w_ji
})
.collect()
}
fn edge_distance_directed(
&self,
offsets: &[usize],
row_indices: &[usize],
values: &[f32],
from: usize,
to: usize,
) -> f32 {
let start = offsets[from];
let end = offsets[from + 1];
for idx in start..end {
if row_indices[idx] == to {
return values[idx];
}
}
f32::INFINITY
}
}
impl WeightedGraph for KnnGraph {
fn num_nodes(&self) -> usize {
self.n_nodes
}
fn num_edges(&self) -> usize {
self.edges.len()
}
fn neighbors_with_weight<'a>(
&'a self,
node: usize,
) -> Box<dyn Iterator<Item = (usize, f32)> + 'a> {
let offsets = self.adjacency.col_offsets();
let start = offsets[node];
let end = offsets[node + 1];
let rows = &self.adjacency.row_indices()[start..end];
let vals = &self.adjacency.values()[start..end];
Box::new(rows.iter().zip(vals.iter()).map(|(&i, &w)| (i, w)))
}
}
fn smooth_knn_sigma(dists: &[f32], rho: f32, target: f32) -> f32 {
const TOLERANCE: f32 = 1e-5;
const MAX_ITER: usize = 64;
let mean_dist: f32 = dists.iter().sum::<f32>() / dists.len().max(1) as f32;
let min_sigma = 1e-3 * mean_dist;
let mut lo = 0.0f32;
let mut hi = f32::INFINITY;
let mut mid = 1.0f32;
for _ in 0..MAX_ITER {
let mut psum = 0.0f32;
for &d in dists {
let gap = d - rho;
if gap > 0.0 {
psum += (-gap / mid).exp();
} else {
psum += 1.0;
}
}
if (psum - target).abs() < TOLERANCE {
break;
}
if psum > target {
hi = mid;
mid = (lo + hi) / 2.0;
} else {
lo = mid;
if hi.is_infinite() {
mid *= 2.0;
} else {
mid = (lo + hi) / 2.0;
}
}
}
mid.max(min_sigma)
}
fn directed_umap_weight(d: f32, rho: f32, sigma: f32) -> f32 {
if d.is_infinite() || sigma <= 0.0 {
return 0.0;
}
let gap = d - rho;
if gap <= 0.0 {
1.0
} else {
(-gap / sigma).exp()
}
}
impl WeightedGraph for crate::leiden::Network {
fn num_nodes(&self) -> usize {
self.nodes()
}
fn num_edges(&self) -> usize {
crate::leiden::Network::edge_count(self)
}
fn neighbors_with_weight<'a>(
&'a self,
node: usize,
) -> Box<dyn Iterator<Item = (usize, f32)> + 'a> {
Box::new(self.neighbors(node).map(|(n, w)| (n, w as f32)))
}
}
#[must_use]
pub fn modularity_to_cpm_resolution(modularity_gamma: f64, total_edge_weight: f64) -> f64 {
modularity_gamma / (2.0 * total_edge_weight).max(1.0)
}
impl KnnGraph {
pub fn to_leiden_network(&self) -> (crate::leiden::Network, f64) {
let n = self.n_nodes;
let weights = self.fuzzy_kernel_weights();
let mut node_degree = vec![0.0f32; n];
let mut n_edges = vec![0usize; n];
let mut total_edge_weight = 0.0f64;
for (&(i, j), &w) in self.edges.iter().zip(weights.iter()) {
node_degree[i] += w;
node_degree[j] += w;
n_edges[i] += 1;
n_edges[j] += 1;
total_edge_weight += w as f64;
}
let mut network = crate::leiden::Network::with_nodes(&node_degree, &n_edges);
for (&(i, j), &w) in self.edges.iter().zip(weights.iter()) {
network.add_edge(i, j, w);
}
(network, total_edge_weight)
}
}
pub fn run_leiden(
network: &crate::leiden::Network,
n: usize,
resolution: f64,
seed: Option<usize>,
) -> Vec<usize> {
use crate::leiden::clustering::SimpleClustering;
use crate::leiden::Clustering;
use crate::leiden::Leiden;
let mut leiden = Leiden::new(resolution, 0.01, seed);
let mut clustering = SimpleClustering::init_different_clusters(n);
for iter in 0..10 {
let updated = leiden.iterate(network, &mut clustering);
info!(
" Leiden iter {}: {} clusters{}",
iter + 1,
clustering.num_clusters(),
if !updated { " (converged)" } else { "" }
);
if !updated {
break;
}
}
(0..n).map(|i| clustering.get(i)).collect()
}
pub fn tune_leiden_resolution(
network: &crate::leiden::Network,
n: usize,
target_k: usize,
initial_resolution: f64,
seed: Option<usize>,
) -> Vec<usize> {
let mut lo = 1e-6_f64;
let mut hi = 10.0_f64;
let mut best = run_leiden(network, n, initial_resolution, seed);
let best_k = count_distinct(&best);
info!(
" resolution={:.6e} → {} clusters (target {})",
initial_resolution, best_k, target_k
);
if best_k == target_k {
return best;
}
if best_k > target_k {
hi = initial_resolution;
} else {
lo = initial_resolution;
}
let mut best_diff = best_k.abs_diff(target_k);
for _ in 0..20 {
let mid = (lo + hi) / 2.0;
let result = run_leiden(network, n, mid, seed);
let k = count_distinct(&result);
info!(" resolution={:.6e} → {} clusters", mid, k);
if k > target_k {
hi = mid;
} else {
lo = mid;
}
let diff = k.abs_diff(target_k);
if diff < best_diff {
best = result;
best_diff = diff;
}
if k == target_k || (hi - lo) / hi.max(1e-10) < 1e-4 {
break;
}
}
best
}
fn count_distinct(labels: &[usize]) -> usize {
let max = labels.iter().copied().max().unwrap_or(0);
let mut seen = vec![false; max + 1];
for &l in labels {
seen[l] = true;
}
seen.iter().filter(|&&s| s).count()
}
pub fn compact_labels(labels: &mut [usize]) {
let max = labels.iter().copied().max().unwrap_or(0);
let mut mapping = vec![usize::MAX; max + 1];
let mut next = 0usize;
for l in labels.iter_mut() {
if mapping[*l] == usize::MAX {
mapping[*l] = next;
next += 1;
}
*l = mapping[*l];
}
}
fn create_jobs(ntot: usize, block_size: usize) -> Vec<(usize, usize)> {
let block_size = if block_size == 0 {
DEFAULT_BLOCK_SIZE
} else {
block_size
};
let nblock = ntot.div_ceil(block_size);
(0..nblock)
.map(|block| {
let lb = block * block_size;
let ub = ((block + 1) * block_size).min(ntot);
(lb, ub)
})
.collect()
}
#[cfg(test)]
mod tests;
type NeighbourList = (Vec<usize>, Vec<f32>);
type TaggedEdge = ((usize, usize), f32, u8);
fn fold_tagged_edges(mut tagged: Vec<TaggedEdge>) -> Vec<TaggedEdge> {
tagged.par_sort_unstable_by_key(|&(key, _, _)| key);
tagged.dedup_by(|cur, prev| {
if cur.0 == prev.0 {
prev.1 = prev.1.min(cur.1);
prev.2 |= cur.2;
true
} else {
false
}
});
tagged
}
fn neighbours_per_point(knn: usize, nn: usize) -> usize {
knn.min(nn.saturating_sub(1)).max(1)
}
fn search_dict(
dict: &ColumnDict<usize>,
nn: usize,
n_neighbours: usize,
block_size: usize,
) -> anyhow::Result<Vec<NeighbourList>> {
let jobs = create_jobs(nn, block_size);
let search_bar =
crate::matrix::progress::new_progress_bar(jobs.len() as u64).with_message("kNN blocks");
let result: anyhow::Result<Vec<Vec<NeighbourList>>> = jobs
.into_par_iter()
.progress_with(search_bar.clone())
.map(|(lb, ub)| {
let mut scratch = SearchScratch::default();
(lb..ub)
.map(|i| dict.search_others_reuse(&i, n_neighbours, &mut scratch))
.collect()
})
.collect();
search_bar.finish_and_clear();
Ok(result?.into_iter().flatten().collect())
}
fn search_rows(rows: &DMatrix<f32>, n_neighbours: usize) -> Vec<NeighbourList> {
let nn = rows.nrows();
let (indices, distances) = if nn <= ALL_PAIRS_THRESHOLD {
info!("kNN by the exact all-pairs kernel over {nn} points");
knn_rows_l2(rows, n_neighbours)
} else {
info!("kNN by the inverted-file search over {nn} points");
knn_rows_ivf(
rows,
&IvfArgs {
k: n_neighbours,
n_lists: 0,
n_probe: DEFAULT_N_PROBE,
seed: KNN_SEED,
},
)
};
indices.into_iter().zip(distances).collect()
}
pub(crate) fn edges_from_neighbours(
lists: &[NeighbourList],
reciprocal: bool,
) -> Vec<((usize, usize), f32)> {
let tagged: Vec<TaggedEdge> = lists
.par_iter()
.enumerate()
.flat_map_iter(|(i, (nb, ds))| {
nb.iter().zip(ds).map(move |(&j, &d)| {
if i < j {
((i, j), d, 1u8)
} else {
((j, i), d, 2u8)
}
})
})
.collect();
fold_tagged_edges(tagged)
.into_iter()
.filter(|&(_, _, mask)| !reciprocal || mask == 3)
.map(|(key, dist, _)| (key, dist))
.collect()
}
pub fn symmetric_adjacency(
n_nodes: usize,
edges: &[(usize, usize)],
distances: &[f32],
) -> CscMatrix<f32> {
debug_assert!(
edges.iter().all(|&(i, j)| i < j) && edges.windows(2).all(|w| w[0] < w[1]),
"symmetric adjacency: edges must be canonical, sorted and unique"
);
let mut offsets = vec![0usize; n_nodes + 1];
for &(i, j) in edges {
offsets[i + 1] += 1;
offsets[j + 1] += 1;
}
for c in 0..n_nodes {
offsets[c + 1] += offsets[c];
}
let nnz = offsets[n_nodes];
let mut row_indices = vec![0usize; nnz];
let mut values = vec![0f32; nnz];
let mut cursor = offsets[..n_nodes].to_vec();
for (&(i, j), &v) in edges.iter().zip(distances) {
row_indices[cursor[i]] = j;
values[cursor[i]] = v;
cursor[i] += 1;
row_indices[cursor[j]] = i;
values[cursor[j]] = v;
cursor[j] += 1;
}
CscMatrix::try_from_csc_data(n_nodes, n_nodes, offsets, row_indices, values)
.expect("symmetric adjacency: canonical sorted edges fill every column in order")
}
fn within_source_rank(d: &[f32]) -> Vec<f32> {
if d.len() <= 1 {
return vec![0.0; d.len()];
}
let mut order: Vec<usize> = (0..d.len()).collect();
order.par_sort_unstable_by(|&a, &b| d[a].total_cmp(&d[b]).then(a.cmp(&b)));
let mut out = vec![0.0f32; d.len()];
let denom = (d.len() - 1) as f32;
for (rank, &idx) in order.iter().enumerate() {
out[idx] = rank as f32 / denom;
}
out
}