use crate::matrix::knn_graph::{self, KnnGraph, KnnGraphArgs};
use crate::matrix::traits::MatOps;
use log::info;
use nalgebra::DMatrix;
#[derive(Debug, Clone)]
pub struct KmeansArgs {
pub num_clusters: usize,
pub max_iter: usize,
}
impl Default for KmeansArgs {
fn default() -> Self {
Self {
num_clusters: 1,
max_iter: 100,
}
}
}
impl KmeansArgs {
pub fn with_clusters(num_clusters: usize) -> Self {
Self {
num_clusters,
..Default::default()
}
}
}
pub trait Kmeans {
fn kmeans_columns(&self, args: KmeansArgs) -> Vec<usize>;
fn kmeans_rows(&self, args: KmeansArgs) -> Vec<usize>;
}
impl<T> Kmeans for DMatrix<T>
where
T: Clone + Sync + Send,
Vec<T>: clustering::Elem,
{
fn kmeans_columns(&self, args: KmeansArgs) -> Vec<usize> {
if args.num_clusters <= 1 || self.ncols() == 0 {
return vec![0; self.ncols()];
}
let data: Vec<Vec<T>> = self
.column_iter()
.map(|x| x.iter().cloned().collect())
.collect();
let clust = clustering::kmeans(args.num_clusters, &data, args.max_iter);
clust.membership
}
fn kmeans_rows(&self, args: KmeansArgs) -> Vec<usize> {
if args.num_clusters <= 1 || self.nrows() == 0 {
return vec![0; self.nrows()];
}
let data: Vec<Vec<T>> = self
.row_iter()
.map(|x| x.iter().cloned().collect())
.collect();
let clust = clustering::kmeans(args.num_clusters, &data, args.max_iter);
clust.membership
}
}
pub fn leiden_clustering(
latent: &DMatrix<f32>,
knn: usize,
resolution: f64,
target_clusters: Option<usize>,
seed: Option<u64>,
cosine: bool,
) -> anyhow::Result<Vec<usize>> {
let n = latent.nrows();
let d = latent.ncols();
if n < 2 {
anyhow::bail!("Need at least 2 points for Leiden clustering");
}
info!("Leiden: {n} points x {d} features, knn={knn}, seed={seed:?}, cosine={cosine}");
let mut latent_pre = latent.clone();
if cosine {
for mut row in latent_pre.row_iter_mut() {
let denom = row.norm().max(1e-8);
row /= denom;
}
} else {
latent_pre.scale_columns_inplace();
}
let graph = KnnGraph::from_rows(
&latent_pre,
KnnGraphArgs {
knn,
block_size: 1000,
reciprocal: false,
},
)?;
info!(
"KNN graph: {} nodes, {} edges",
graph.num_nodes(),
graph.num_edges()
);
let (network, total_edge_weight) = graph.to_leiden_network();
let resolution_scaled = knn_graph::modularity_to_cpm_resolution(resolution, total_edge_weight);
let seed_val = seed.map(|s| s as usize);
let mut labels = if let Some(target_k) = target_clusters {
knn_graph::tune_leiden_resolution(&network, n, target_k, resolution_scaled, seed_val)
} else {
knn_graph::run_leiden(&network, n, resolution_scaled, seed_val)
};
knn_graph::compact_labels(&mut labels);
let n_clusters = labels.iter().copied().max().unwrap_or(0) + 1;
info!("Leiden done: {n_clusters} clusters over {n} points");
Ok(labels)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kmeans_columns_single_cluster() {
let mat = DMatrix::from_row_slice(2, 4, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]);
let args = KmeansArgs::with_clusters(1);
let membership = mat.kmeans_columns(args);
assert_eq!(membership.len(), 4);
assert!(membership.iter().all(|&x| x == 0));
}
#[test]
fn test_kmeans_columns_two_clusters() {
let mat = DMatrix::from_row_slice(
2,
6,
&[
0.0, 0.1, 0.2, 10.0, 10.1, 10.2, 0.0, 0.1, 0.0, 10.0, 10.1, 10.2, ],
);
let args = KmeansArgs::with_clusters(2);
let membership = mat.kmeans_columns(args);
assert_eq!(membership.len(), 6);
assert_eq!(membership[0], membership[1]);
assert_eq!(membership[1], membership[2]);
assert_eq!(membership[3], membership[4]);
assert_eq!(membership[4], membership[5]);
assert_ne!(membership[0], membership[3]);
}
#[test]
fn test_kmeans_rows() {
let mat = DMatrix::from_row_slice(
4,
2,
&[
0.0, 0.0, 0.1, 0.1, 10.0, 10.0, 10.1, 10.1, ],
);
let args = KmeansArgs::with_clusters(2);
let membership = mat.kmeans_rows(args);
assert_eq!(membership.len(), 4);
assert_eq!(membership[0], membership[1]);
assert_eq!(membership[2], membership[3]);
assert_ne!(membership[0], membership[2]);
}
#[test]
fn test_kmeans_empty_matrix() {
let mat: DMatrix<f32> = DMatrix::zeros(0, 0);
let col_membership = mat.kmeans_columns(KmeansArgs::with_clusters(2));
let row_membership = mat.kmeans_rows(KmeansArgs::with_clusters(2));
assert!(col_membership.is_empty());
assert!(row_membership.is_empty());
}
}