use crate::error::{GraphError, GraphResult};
#[derive(Debug, Clone)]
pub struct LouvainConfig {
pub max_passes: usize,
pub min_improvement: f64,
pub max_levels: usize,
}
impl Default for LouvainConfig {
fn default() -> Self {
Self {
max_passes: 10,
min_improvement: 1e-7,
max_levels: 15,
}
}
}
#[derive(Debug)]
pub struct LouvainResult {
pub labels: Vec<usize>,
pub modularity: f64,
pub n_communities: usize,
}
#[derive(Clone, Debug)]
struct Neighbor {
node: usize,
weight: f64,
}
struct WGraph {
n: usize,
adj: Vec<Vec<Neighbor>>,
degree: Vec<f64>,
total_weight: f64,
}
impl WGraph {
fn from_edges(n: usize, edges: &[(usize, usize, f64)]) -> Self {
let mut adj: Vec<Vec<Neighbor>> = vec![Vec::new(); n];
let mut degree = vec![0.0_f64; n];
let mut total_weight = 0.0_f64;
for &(u, v, w) in edges {
adj[u].push(Neighbor { node: v, weight: w });
degree[u] += w;
if u != v {
adj[v].push(Neighbor { node: u, weight: w });
degree[v] += w;
total_weight += w;
} else {
total_weight += w; }
}
Self {
n,
adj,
degree,
total_weight,
}
}
fn modularity(&self, community: &[usize]) -> f64 {
if self.total_weight < 1e-14 {
return 0.0;
}
let two_m = 2.0 * self.total_weight;
let mut q = 0.0_f64;
for u in 0..self.n {
for nb in &self.adj[u] {
if community[u] == community[nb.node] {
q += nb.weight - self.degree[u] * self.degree[nb.node] / two_m;
}
}
}
q / two_m
}
}
fn phase1(graph: &WGraph, community: &mut [usize], config: &LouvainConfig) -> f64 {
let n = graph.n;
let two_m = 2.0 * graph.total_weight;
if two_m < 1e-14 {
return 0.0;
}
let mut sigma_tot = vec![0.0_f64; n];
for i in 0..n {
sigma_tot[community[i]] += graph.degree[i];
}
let mut total_gain = 0.0_f64;
for _pass in 0..config.max_passes {
let mut pass_gain = 0.0_f64;
for i in 0..n {
let current_comm = community[i];
let ki = graph.degree[i];
let mut comm_weights: Vec<(usize, f64)> = Vec::new();
let mut ki_in_current = 0.0_f64;
for nb in &graph.adj[i] {
let c = community[nb.node];
if nb.node != i {
if c == current_comm {
ki_in_current += nb.weight;
}
if let Some(pos) = comm_weights.iter().position(|&(c2, _)| c2 == c) {
comm_weights[pos].1 += nb.weight;
} else {
comm_weights.push((c, nb.weight));
}
}
}
let sigma_d_no_i = sigma_tot[current_comm] - ki;
let base_remove =
-ki_in_current / graph.total_weight + sigma_d_no_i * ki / (two_m * two_m);
let mut best_comm = current_comm;
let mut best_gain = 0.0_f64;
for &(c, ki_c) in &comm_weights {
if c == current_comm {
continue;
}
let delta_add = ki_c / graph.total_weight - sigma_tot[c] * ki / (two_m * two_m);
let gain = delta_add + base_remove;
if gain > best_gain {
best_gain = gain;
best_comm = c;
}
}
if best_comm != current_comm && best_gain > 0.0 {
sigma_tot[current_comm] -= ki;
sigma_tot[best_comm] += ki;
community[i] = best_comm;
pass_gain += best_gain;
}
}
total_gain += pass_gain;
if pass_gain < config.min_improvement {
break;
}
}
total_gain
}
fn phase2(graph: &WGraph, community: &[usize]) -> (WGraph, Vec<usize>) {
let n = graph.n;
let mut comm_map = vec![usize::MAX; n];
let mut k = 0_usize;
let mut old_to_new = vec![usize::MAX; n];
for &c in community.iter() {
if comm_map[c] == usize::MAX {
comm_map[c] = k;
k += 1;
}
}
for i in 0..n {
old_to_new[i] = comm_map[community[i]];
}
let mut edge_map: Vec<(usize, usize, f64)> = Vec::new();
for u in 0..n {
let su = old_to_new[u];
for nb in &graph.adj[u] {
let sv = old_to_new[nb.node];
let mut found = false;
for entry in edge_map.iter_mut() {
if (entry.0 == su && entry.1 == sv) || (entry.0 == sv && entry.1 == su) {
entry.2 += nb.weight * 0.5; found = true;
break;
}
}
if !found {
edge_map.push((su, sv, nb.weight * 0.5));
}
}
}
let super_graph = WGraph::from_edges(k, &edge_map);
(super_graph, old_to_new)
}
pub fn louvain_communities(
n_nodes: usize,
edges: &[(usize, usize, f64)],
config: &LouvainConfig,
) -> GraphResult<LouvainResult> {
if n_nodes == 0 {
return Err(GraphError::EmptyGraph);
}
for &(u, v, _) in edges {
if u >= n_nodes || v >= n_nodes {
return Err(GraphError::InvalidPlan("node_out_of_range".to_owned()));
}
}
let mut node_comm: Vec<usize> = (0..n_nodes).collect();
let graph0 = WGraph::from_edges(n_nodes, edges);
phase1(&graph0, &mut node_comm, config);
let mut current_graph = graph0;
let mut old_to_new: Vec<usize> = node_comm.clone();
for _level in 0..config.max_levels {
let (super_graph, mapping) = phase2(¤t_graph, &old_to_new);
let mut super_comm: Vec<usize> = (0..super_graph.n).collect();
let improvement = phase1(&super_graph, &mut super_comm, config);
for i in 0..n_nodes {
let super_node = mapping[node_comm[i]];
node_comm[i] = super_comm[super_node];
}
old_to_new = super_comm;
current_graph = super_graph;
if improvement < config.min_improvement || current_graph.n <= 1 {
break;
}
}
let mut label_map = vec![usize::MAX; n_nodes];
let mut next_label = 0_usize;
let mut final_labels = vec![0_usize; n_nodes];
for i in 0..n_nodes {
let c = node_comm[i];
if c < n_nodes && label_map[c] == usize::MAX {
label_map[c] = next_label;
next_label += 1;
}
final_labels[i] = if c < n_nodes { label_map[c] } else { 0 };
}
let n_communities = next_label.max(1);
let original_graph = WGraph::from_edges(n_nodes, edges);
let modularity = original_graph.modularity(&final_labels);
Ok(LouvainResult {
labels: final_labels,
modularity,
n_communities,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn clique_edges(nodes: &[usize], weight: f64) -> Vec<(usize, usize, f64)> {
let mut edges = Vec::new();
for i in 0..nodes.len() {
for j in i + 1..nodes.len() {
edges.push((nodes[i], nodes[j], weight));
}
}
edges
}
#[test]
fn empty_graph_error() {
let err = louvain_communities(0, &[], &LouvainConfig::default());
assert!(matches!(err, Err(GraphError::EmptyGraph)), "got: {err:?}");
}
#[test]
fn single_node() {
let r = louvain_communities(1, &[], &LouvainConfig::default())
.expect("value should be present");
assert_eq!(r.n_communities, 1);
assert_eq!(r.labels.len(), 1);
}
#[test]
fn two_isolated_nodes() {
let r = louvain_communities(2, &[], &LouvainConfig::default())
.expect("value should be present");
assert_eq!(r.labels.len(), 2);
assert!(r.n_communities >= 1);
}
#[test]
fn two_cliques_connected_weakly() {
let mut edges = clique_edges(&[0, 1, 2, 3], 10.0);
edges.extend(clique_edges(&[4, 5, 6, 7], 10.0));
edges.push((3, 4, 0.0001)); let r = louvain_communities(8, &edges, &LouvainConfig::default())
.expect("value should be present");
assert_eq!(r.labels.len(), 8);
assert_eq!(r.n_communities, 2);
}
#[test]
fn labels_len() {
let edges = clique_edges(&[0, 1, 2, 3, 4], 1.0);
let r = louvain_communities(5, &edges, &LouvainConfig::default())
.expect("value should be present");
assert_eq!(r.labels.len(), 5);
}
#[test]
fn labels_in_range() {
let edges = clique_edges(&[0, 1, 2], 1.0);
let r = louvain_communities(3, &edges, &LouvainConfig::default())
.expect("value should be present");
for &l in &r.labels {
assert!(
l < r.n_communities,
"label {l} >= n_communities {}",
r.n_communities
);
}
}
#[test]
fn modularity_finite() {
let edges = clique_edges(&[0, 1, 2, 3], 1.0);
let r = louvain_communities(4, &edges, &LouvainConfig::default())
.expect("value should be present");
assert!(r.modularity.is_finite(), "Q = {}", r.modularity);
}
#[test]
fn modularity_nonneg_for_cliques() {
let mut edges = clique_edges(&[0, 1, 2], 1.0);
edges.extend(clique_edges(&[3, 4, 5], 1.0));
edges.push((2, 3, 0.001));
let r = louvain_communities(6, &edges, &LouvainConfig::default())
.expect("value should be present");
assert!(r.modularity >= -1.0, "Q = {} should be >= -1", r.modularity);
}
#[test]
fn all_connected_clique() {
let n = 6;
let nodes: Vec<usize> = (0..n).collect();
let edges = clique_edges(&nodes, 1.0);
let r = louvain_communities(n, &edges, &LouvainConfig::default())
.expect("value should be present");
assert_eq!(r.labels.len(), n);
assert!(r.n_communities >= 1);
}
#[test]
fn n_communities_positive() {
let edges = clique_edges(&[0, 1, 2, 3], 1.0);
let r = louvain_communities(4, &edges, &LouvainConfig::default())
.expect("value should be present");
assert!(r.n_communities >= 1);
}
#[test]
fn node_out_of_range_error() {
let edges = vec![(0, 5, 1.0)]; let err = louvain_communities(4, &edges, &LouvainConfig::default());
assert!(
matches!(err, Err(GraphError::InvalidPlan(_))),
"got: {err:?}"
);
}
#[test]
fn three_cliques_three_communities() {
let mut edges = clique_edges(&[0, 1, 2, 3], 2.0);
edges.extend(clique_edges(&[4, 5, 6, 7], 2.0));
edges.extend(clique_edges(&[8, 9, 10, 11], 2.0));
edges.push((3, 4, 1e-5));
edges.push((7, 8, 1e-5));
let r = louvain_communities(12, &edges, &LouvainConfig::default())
.expect("value should be present");
assert_eq!(r.labels.len(), 12);
assert!(r.n_communities >= 1);
}
}