Skip to main content

proof_engine/graph/
community.rs

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
38/// Compute modularity Q for a given partitioning of the graph.
39/// Q = (1/2m) * sum_ij [ A_ij - k_i*k_j/(2m) ] * delta(c_i, c_j)
40pub fn modularity<N, E>(graph: &Graph<N, E>, communities: &[Community]) -> f32 {
41    // Newman modularity, per community c:
42    //   undirected: Q = sum_c [ L_c / m - (D_c / 2m)^2 ]
43    //   directed:   Q = sum_c [ L_c / m - Dout_c * Din_c / m^2 ]
44    // with L_c the edges inside c and D_c the summed degrees of c. The old
45    // version summed A_ij - k_i k_j / 2m over edges only, leaving out the
46    // expected-edge terms of every non-adjacent pair, so a single community
47    // scored 0.5 instead of 0.
48    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
85/// Louvain method for community detection.
86/// Iteratively moves nodes to communities that maximize modularity gain.
87pub 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    // Each node starts in its own community
105    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    // Precompute degrees and adjacency weights
112    let degrees: HashMap<NodeId, f32> = node_ids.iter()
113        .map(|&nid| (nid, graph.degree(nid) as f32))
114        .collect();
115
116    // Weighted adjacency
117    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    // Sum of weights in each community
127    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            // Compute weights to each neighboring community
145            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            // Remove node from current community
152            *sigma_tot.get_mut(&current_comm).unwrap() -= ki;
153
154            // Find best community
155            let ki_in_current = comm_weights.get(&current_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(&current_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            // Move node to best community
171            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    // Build communities from assignments
188    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
206/// Label propagation community detection.
207/// Each node adopts the label most common among its neighbors.
208pub 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    // Initialize each node with its own label
216    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    // Simple deterministic ordering (could shuffle for randomness)
225    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            // Count label frequencies among neighbors
234            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            // Pick most frequent label (ties broken by smallest label)
241            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    // Build communities from labels
260    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        // Clique 1: 0,1,2
286        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        // Clique 2: 3,4,5
293        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        // Bridge
300        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        // All in one community: modularity should be 0
316        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        // Should find 2 communities
324        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        // No edges => each node is its own community
349        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}