1use std::collections::{HashMap, HashSet};
2use super::graph_core::{Graph, GraphKind, NodeId};
3
4#[derive(Debug, Clone)]
5pub struct Community {
6 pub members: HashSet<NodeId>,
7}
8
9impl Community {
10 pub fn new() -> Self {
11 Self { members: HashSet::new() }
12 }
13
14 pub fn from_members(members: impl IntoIterator<Item = NodeId>) -> Self {
15 Self { members: members.into_iter().collect() }
16 }
17
18 pub fn contains(&self, id: NodeId) -> bool {
19 self.members.contains(&id)
20 }
21
22 pub fn len(&self) -> usize {
23 self.members.len()
24 }
25
26 pub fn is_empty(&self) -> bool {
27 self.members.is_empty()
28 }
29}
30
31#[derive(Debug, Clone)]
32pub struct CommunityResult {
33 pub communities: Vec<Community>,
34 pub modularity: f32,
35 pub iterations: usize,
36}
37
38pub fn modularity<N, E>(graph: &Graph<N, E>, communities: &[Community]) -> f32 {
41 let m = graph.edge_count() as f32;
49 if m == 0.0 { return 0.0; }
50
51 let mut community_of: HashMap<NodeId, usize> = HashMap::new();
52 for (ci, comm) in communities.iter().enumerate() {
53 for &nid in &comm.members {
54 community_of.insert(nid, ci);
55 }
56 }
57
58 let k = communities.len();
59 let mut inside = vec![0.0f32; k];
60 let mut d_out = vec![0.0f32; k];
61 let mut d_in = vec![0.0f32; k];
62 for edge in graph.edges() {
63 let ci = community_of.get(&edge.from).copied();
64 let cj = community_of.get(&edge.to).copied();
65 if let Some(ci) = ci { d_out[ci] += 1.0; }
66 if let Some(cj) = cj { d_in[cj] += 1.0; }
67 if ci.is_some() && ci == cj {
68 inside[ci.unwrap_or(0)] += 1.0;
69 }
70 }
71
72 let mut q = 0.0f32;
73 for c in 0..k {
74 q += inside[c] / m;
75 if graph.kind == GraphKind::Undirected {
76 let d = (d_out[c] + d_in[c]) / (2.0 * m);
77 q -= d * d;
78 } else {
79 q -= d_out[c] * d_in[c] / (m * m);
80 }
81 }
82 q
83}
84
85pub fn louvain<N: Clone, E: Clone>(graph: &Graph<N, E>) -> CommunityResult {
88 let node_ids = graph.node_ids();
89 let n = node_ids.len();
90 if n == 0 {
91 return CommunityResult { communities: Vec::new(), modularity: 0.0, iterations: 0 };
92 }
93
94 let m = graph.edge_count() as f32;
95 if m == 0.0 {
96 let communities: Vec<Community> = node_ids.iter()
97 .map(|&nid| Community::from_members(std::iter::once(nid)))
98 .collect();
99 return CommunityResult { communities, modularity: 0.0, iterations: 0 };
100 }
101
102 let m2 = if graph.kind == GraphKind::Undirected { 2.0 * m } else { m };
103
104 let mut comm_of: HashMap<NodeId, usize> = HashMap::new();
106 for (i, &nid) in node_ids.iter().enumerate() {
107 comm_of.insert(nid, i);
108 }
109 let mut num_communities = n;
110
111 let degrees: HashMap<NodeId, f32> = node_ids.iter()
113 .map(|&nid| (nid, graph.degree(nid) as f32))
114 .collect();
115
116 let mut adj_weights: HashMap<NodeId, Vec<(NodeId, f32)>> = HashMap::new();
118 for &nid in &node_ids {
119 let mut ws = Vec::new();
120 for (nbr, eid) in graph.neighbor_edges(nid) {
121 ws.push((nbr, graph.edge_weight(eid)));
122 }
123 adj_weights.insert(nid, ws);
124 }
125
126 let mut sigma_tot: HashMap<usize, f32> = HashMap::new();
128 for &nid in &node_ids {
129 let c = comm_of[&nid];
130 *sigma_tot.entry(c).or_insert(0.0) += degrees[&nid];
131 }
132
133 let mut iterations = 0;
134 let max_iterations = 100;
135
136 loop {
137 iterations += 1;
138 let mut improved = false;
139
140 for &nid in &node_ids {
141 let current_comm = comm_of[&nid];
142 let ki = degrees[&nid];
143
144 let mut comm_weights: HashMap<usize, f32> = HashMap::new();
146 for &(nbr, w) in adj_weights.get(&nid).unwrap_or(&Vec::new()) {
147 let nc = comm_of[&nbr];
148 *comm_weights.entry(nc).or_insert(0.0) += w;
149 }
150
151 *sigma_tot.get_mut(¤t_comm).unwrap() -= ki;
153
154 let ki_in_current = comm_weights.get(¤t_comm).copied().unwrap_or(0.0);
156 let mut best_comm = current_comm;
157 let mut best_gain = 0.0f32;
158
159 for (&c, &ki_in) in &comm_weights {
160 let st = sigma_tot.get(&c).copied().unwrap_or(0.0);
161 let gain = ki_in / m2 - st * ki / (m2 * m2);
162 let loss = ki_in_current / m2 - sigma_tot.get(¤t_comm).copied().unwrap_or(0.0) * ki / (m2 * m2);
163 let delta_q = gain - loss;
164 if delta_q > best_gain {
165 best_gain = delta_q;
166 best_comm = c;
167 }
168 }
169
170 comm_of.insert(nid, best_comm);
172 *sigma_tot.get_mut(&best_comm).unwrap_or(&mut 0.0) += ki;
173 if !sigma_tot.contains_key(&best_comm) {
174 sigma_tot.insert(best_comm, ki);
175 }
176
177 if best_comm != current_comm {
178 improved = true;
179 }
180 }
181
182 if !improved || iterations >= max_iterations {
183 break;
184 }
185 }
186
187 let mut comm_map: HashMap<usize, Vec<NodeId>> = HashMap::new();
189 for (&nid, &c) in &comm_of {
190 comm_map.entry(c).or_default().push(nid);
191 }
192
193 let communities: Vec<Community> = comm_map.into_values()
194 .map(|members| Community::from_members(members))
195 .collect();
196
197 let mod_val = modularity(graph, &communities);
198
199 CommunityResult {
200 communities,
201 modularity: mod_val,
202 iterations,
203 }
204}
205
206pub fn label_propagation<N, E>(graph: &Graph<N, E>) -> CommunityResult {
209 let node_ids = graph.node_ids();
210 let n = node_ids.len();
211 if n == 0 {
212 return CommunityResult { communities: Vec::new(), modularity: 0.0, iterations: 0 };
213 }
214
215 let mut labels: HashMap<NodeId, u32> = HashMap::new();
217 for (i, &nid) in node_ids.iter().enumerate() {
218 labels.insert(nid, i as u32);
219 }
220
221 let max_iterations = 100;
222 let mut iterations = 0;
223
224 loop {
226 iterations += 1;
227 let mut changed = false;
228
229 for &nid in &node_ids {
230 let neighbors = graph.neighbors(nid);
231 if neighbors.is_empty() { continue; }
232
233 let mut freq: HashMap<u32, usize> = HashMap::new();
235 for nbr in &neighbors {
236 let lbl = labels[nbr];
237 *freq.entry(lbl).or_insert(0) += 1;
238 }
239
240 let max_count = freq.values().copied().max().unwrap_or(0);
242 let best_label = freq.iter()
243 .filter(|(_, &c)| c == max_count)
244 .map(|(&l, _)| l)
245 .min()
246 .unwrap_or(labels[&nid]);
247
248 if labels[&nid] != best_label {
249 labels.insert(nid, best_label);
250 changed = true;
251 }
252 }
253
254 if !changed || iterations >= max_iterations {
255 break;
256 }
257 }
258
259 let mut comm_map: HashMap<u32, Vec<NodeId>> = HashMap::new();
261 for (&nid, &lbl) in &labels {
262 comm_map.entry(lbl).or_default().push(nid);
263 }
264
265 let communities: Vec<Community> = comm_map.into_values()
266 .map(|members| Community::from_members(members))
267 .collect();
268
269 let mod_val = modularity(graph, &communities);
270
271 CommunityResult {
272 communities,
273 modularity: mod_val,
274 iterations,
275 }
276}
277
278#[cfg(test)]
279mod tests {
280 use super::*;
281 use crate::graph::graph_core::GraphKind;
282
283 fn make_two_cliques() -> Graph<(), ()> {
284 let mut g = Graph::new(GraphKind::Undirected);
285 let a = g.add_node(());
287 let b = g.add_node(());
288 let c = g.add_node(());
289 g.add_edge(a, b, ());
290 g.add_edge(b, c, ());
291 g.add_edge(a, c, ());
292 let d = g.add_node(());
294 let e = g.add_node(());
295 let f = g.add_node(());
296 g.add_edge(d, e, ());
297 g.add_edge(e, f, ());
298 g.add_edge(d, f, ());
299 g.add_edge(c, d, ());
301 g
302 }
303
304 #[test]
305 fn test_modularity_single_community() {
306 let mut g = Graph::new(GraphKind::Undirected);
307 let a = g.add_node(());
308 let b = g.add_node(());
309 let c = g.add_node(());
310 g.add_edge(a, b, ());
311 g.add_edge(b, c, ());
312 g.add_edge(a, c, ());
313 let comms = vec![Community::from_members(vec![a, b, c])];
314 let q = modularity(&g, &comms);
315 assert!((q - 0.0).abs() < 0.01);
317 }
318
319 #[test]
320 fn test_louvain_two_cliques() {
321 let g = make_two_cliques();
322 let result = louvain(&g);
323 assert!(result.communities.len() >= 2);
325 assert!(result.modularity >= 0.0);
326 }
327
328 #[test]
329 fn test_label_propagation_two_cliques() {
330 let g = make_two_cliques();
331 let result = label_propagation(&g);
332 assert!(result.communities.len() >= 1);
333 }
334
335 #[test]
336 fn test_louvain_empty() {
337 let g: Graph<(), ()> = Graph::new(GraphKind::Undirected);
338 let result = louvain(&g);
339 assert!(result.communities.is_empty());
340 }
341
342 #[test]
343 fn test_label_propagation_disconnected() {
344 let mut g = Graph::<(), ()>::new(GraphKind::Undirected);
345 let a = g.add_node(());
346 let b = g.add_node(());
347 let c = g.add_node(());
348 let result = label_propagation(&g);
350 assert_eq!(result.communities.len(), 3);
351 }
352
353 #[test]
354 fn test_community_struct() {
355 let c = Community::from_members(vec![NodeId(0), NodeId(1), NodeId(2)]);
356 assert_eq!(c.len(), 3);
357 assert!(c.contains(NodeId(1)));
358 assert!(!c.contains(NodeId(5)));
359 assert!(!c.is_empty());
360 }
361
362 #[test]
363 fn test_louvain_single_node() {
364 let mut g: Graph<(), ()> = Graph::new(GraphKind::Undirected);
365 g.add_node(());
366 let result = louvain(&g);
367 assert_eq!(result.communities.len(), 1);
368 }
369
370 #[test]
371 fn test_modularity_two_perfect_communities() {
372 let mut g = Graph::new(GraphKind::Undirected);
373 let a = g.add_node(());
374 let b = g.add_node(());
375 let c = g.add_node(());
376 let d = g.add_node(());
377 g.add_edge(a, b, ());
378 g.add_edge(c, d, ());
379 let comms = vec![
380 Community::from_members(vec![a, b]),
381 Community::from_members(vec![c, d]),
382 ];
383 let q = modularity(&g, &comms);
384 assert!(q > 0.0, "Modularity should be positive for good partition, got {}", q);
385 }
386}