Skip to main content

_diffctx/
ppr.rs

1use std::collections::VecDeque;
2
3use rayon;
4use rustc_hash::{FxHashMap, FxHashSet};
5
6use crate::config::limits::PPR;
7use crate::graph::{CsrGraph, Graph};
8use crate::types::FragmentId;
9
10fn init_seed_residuals(
11    csr: &CsrGraph,
12    seeds: &FxHashSet<FragmentId>,
13    seed_weights: Option<&FxHashMap<FragmentId, f64>>,
14) -> Vec<f64> {
15    let n = csr.n;
16    let mut residual = vec![0.0f64; n];
17
18    let valid_seeds: Vec<&FragmentId> = seeds
19        .iter()
20        .filter(|s| csr.node_to_idx.contains_key(*s))
21        .collect();
22
23    if valid_seeds.is_empty() {
24        return residual;
25    }
26
27    if let Some(sw) = seed_weights {
28        let total: f64 = valid_seeds
29            .iter()
30            .map(|s| sw.get(*s).copied().unwrap_or(PPR.default_seed_epsilon))
31            .sum();
32        if total <= 0.0 {
33            return residual;
34        }
35        for s in &valid_seeds {
36            let idx = csr.node_to_idx[*s] as usize;
37            residual[idx] = sw.get(*s).copied().unwrap_or(PPR.default_seed_epsilon) / total;
38        }
39    } else {
40        let weight = 1.0 / valid_seeds.len() as f64;
41        for s in &valid_seeds {
42            let idx = csr.node_to_idx[*s] as usize;
43            residual[idx] = weight;
44        }
45    }
46
47    residual
48}
49
50/// Result of one PPR push pass. `truncated` flags whether the
51/// `max_pushes` budget cut the iteration short — when true, the
52/// returned estimate is biased toward seeds and the absolute scores
53/// are not comparable to a converged run on the same graph. Surfaced
54/// to Python via `LatencyBreakdown.ppr_truncated` so calibration /
55/// final-eval rows can be filtered or flagged for the paper.
56struct PprPushResult {
57    estimate: Vec<f64>,
58    pushes: usize,
59    truncated: bool,
60}
61
62fn ppr_push_csr(
63    csr: &CsrGraph,
64    seeds: &FxHashSet<FragmentId>,
65    alpha: f64,
66    tol: f64,
67    seed_weights: Option<&FxHashMap<FragmentId, f64>>,
68) -> PprPushResult {
69    let n = csr.n;
70    if n == 0 {
71        return PprPushResult {
72            estimate: Vec::new(),
73            pushes: 0,
74            truncated: false,
75        };
76    }
77
78    let restart = 1.0 - alpha;
79    let mut residual = init_seed_residuals(csr, seeds, seed_weights);
80    let mut estimate = vec![0.0f64; n];
81    let mut in_queue = vec![false; n];
82
83    let mut queue: VecDeque<u32> = VecDeque::new();
84    for i in 0..n {
85        if residual[i] >= tol {
86            queue.push_back(i as u32);
87            in_queue[i] = true;
88        }
89    }
90
91    let max_pushes = (n * PPR.push_scale_factor).min(PPR.max_pushes_cap);
92    let mut pushes: usize = 0;
93    let mut truncated = false;
94
95    while let Some(u) = queue.pop_front() {
96        if pushes >= max_pushes {
97            truncated = true;
98            break;
99        }
100        let ui = u as usize;
101        in_queue[ui] = false;
102
103        let r_u = residual[ui];
104        if r_u < tol {
105            continue;
106        }
107
108        estimate[ui] += restart * r_u;
109        residual[ui] = 0.0;
110
111        let total_w = csr.out_weight_sum[ui];
112        if total_w <= 0.0 {
113            pushes += 1;
114            continue;
115        }
116
117        let propagate = alpha * r_u;
118        let start = csr.indptr[ui] as usize;
119        let end = csr.indptr[ui + 1] as usize;
120
121        for k in start..end {
122            let v = csr.indices[k] as usize;
123            let w = csr.weights[k];
124            let delta = propagate * (w / total_w);
125            residual[v] += delta;
126            if !in_queue[v] && residual[v] >= tol {
127                queue.push_back(v as u32);
128                in_queue[v] = true;
129            }
130        }
131
132        pushes += 1;
133    }
134
135    PprPushResult {
136        estimate,
137        pushes,
138        truncated,
139    }
140}
141
142/// Public PPR result. `truncated` is logical-OR of forward + backward
143/// push truncation flags. When true, downstream renormalization
144/// (sum-to-1 across nodes) hides the fact that the iteration was cut
145/// short by `max_pushes_cap`; the absolute relevance scores are
146/// biased toward seeds and not directly comparable across instances.
147/// We surface this to Python so calibration / final-eval rows can be
148/// filtered in post-analysis (PolyBench / Multi-SWE-bench instances
149/// with >20k fragments are the primary suspects).
150pub struct PprResult {
151    pub scores: FxHashMap<FragmentId, f64>,
152    pub truncated: bool,
153    pub forward_pushes: usize,
154    pub backward_pushes: usize,
155}
156
157pub fn personalized_pagerank(
158    graph: &mut Graph,
159    seeds: &FxHashSet<FragmentId>,
160    alpha: f64,
161    tol: f64,
162    forward_blend: f64,
163    seed_weights: Option<&FxHashMap<FragmentId, f64>>,
164) -> PprResult {
165    if graph.node_count() == 0 || seeds.is_empty() {
166        return PprResult {
167            scores: FxHashMap::default(),
168            truncated: false,
169            forward_pushes: 0,
170            backward_pushes: 0,
171        };
172    }
173
174    let (fwd_csr, rev_csr) = graph.to_csr();
175
176    let (forward, backward) = rayon::join(
177        || ppr_push_csr(fwd_csr, seeds, alpha, tol, seed_weights),
178        || ppr_push_csr(rev_csr, seeds, alpha, tol, seed_weights),
179    );
180
181    let n = fwd_csr.n;
182    let mut combined = vec![0.0f64; n];
183    for i in 0..n {
184        combined[i] =
185            forward_blend * forward.estimate[i] + (1.0 - forward_blend) * backward.estimate[i];
186    }
187
188    let total: f64 = combined.iter().sum();
189    if total > 0.0 {
190        for v in &mut combined {
191            *v /= total;
192        }
193    }
194
195    let idx_to_node = &fwd_csr.idx_to_node;
196    let mut scores: FxHashMap<FragmentId, f64> = FxHashMap::default();
197    for i in 0..n {
198        if combined[i] > 0.0 {
199            scores.insert(idx_to_node[i].clone(), combined[i]);
200        }
201    }
202
203    PprResult {
204        scores,
205        truncated: forward.truncated || backward.truncated,
206        forward_pushes: forward.pushes,
207        backward_pushes: backward.pushes,
208    }
209}
210
211#[cfg(test)]
212mod tests {
213    use super::*;
214    use std::sync::Arc;
215
216    fn fid(path: &str, start: u32, end: u32) -> FragmentId {
217        FragmentId::new(Arc::from(path), start, end)
218    }
219
220    #[test]
221    fn ppr_empty_graph() {
222        let mut g = Graph::new();
223        let seeds = FxHashSet::default();
224        let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-4, 0.4, None).scores;
225        assert!(result.is_empty());
226    }
227
228    #[test]
229    fn ppr_single_node() {
230        let mut g = Graph::new();
231        let a = fid("a.rs", 1, 10);
232        g.add_node(a.clone());
233
234        let mut seeds = FxHashSet::default();
235        seeds.insert(a.clone());
236        let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-4, 0.4, None).scores;
237        assert!((result[&a] - 1.0).abs() < 1e-6);
238    }
239
240    #[test]
241    fn ppr_chain_scores_decrease() {
242        let mut g = Graph::new();
243        let a = fid("a.rs", 1, 10);
244        let b = fid("b.rs", 1, 10);
245        let c = fid("c.rs", 1, 10);
246        g.add_node(a.clone());
247        g.add_node(b.clone());
248        g.add_node(c.clone());
249        g.add_edge(a.clone(), b.clone(), 1.0);
250        g.add_edge(b.clone(), c.clone(), 1.0);
251
252        let mut seeds = FxHashSet::default();
253        seeds.insert(a.clone());
254        let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, None).scores;
255
256        assert!(result[&a] > result[&b]);
257        assert!(result[&b] > result[&c]);
258    }
259
260    #[test]
261    fn ppr_normalizes_to_one() {
262        let mut g = Graph::new();
263        let a = fid("a.rs", 1, 10);
264        let b = fid("b.rs", 1, 10);
265        let c = fid("c.rs", 1, 10);
266        g.add_node(a.clone());
267        g.add_node(b.clone());
268        g.add_node(c.clone());
269        g.add_edge(a.clone(), b.clone(), 1.0);
270        g.add_edge(b.clone(), c.clone(), 1.0);
271        g.add_edge(c.clone(), a.clone(), 0.5);
272
273        let mut seeds = FxHashSet::default();
274        seeds.insert(a.clone());
275        let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, None).scores;
276
277        let total: f64 = result.values().sum();
278        assert!((total - 1.0).abs() < 1e-6);
279    }
280
281    #[test]
282    fn ppr_with_seed_weights() {
283        let mut g = Graph::new();
284        let a = fid("a.rs", 1, 10);
285        let b = fid("b.rs", 1, 10);
286        g.add_node(a.clone());
287        g.add_node(b.clone());
288        g.add_edge(a.clone(), b.clone(), 1.0);
289
290        let mut seeds = FxHashSet::default();
291        seeds.insert(a.clone());
292        seeds.insert(b.clone());
293
294        let mut sw = FxHashMap::default();
295        sw.insert(a.clone(), 0.9);
296        sw.insert(b.clone(), 0.1);
297
298        let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, Some(&sw)).scores;
299        assert!(result[&a] > result[&b]);
300    }
301
302    fn build_star_graph() -> (Graph, FragmentId) {
303        let mut g = Graph::new();
304        let center = fid("center.rs", 1, 10);
305        g.add_node(center.clone());
306        for i in 0..5 {
307            let leaf = fid(&format!("leaf_{i}.rs"), 1, 10);
308            g.add_node(leaf.clone());
309            g.add_edge(center.clone(), leaf.clone(), 1.0);
310            g.add_edge(leaf.clone(), center.clone(), 1.0);
311        }
312        (g, center)
313    }
314
315    #[test]
316    fn ppr_is_deterministic_across_calls() {
317        let (mut g1, center) = build_star_graph();
318        let (mut g2, _) = build_star_graph();
319        let seeds: FxHashSet<FragmentId> = std::iter::once(center).collect();
320
321        let r1 = personalized_pagerank(&mut g1, &seeds, 0.6, 1e-6, 0.4, None).scores;
322        let r2 = personalized_pagerank(&mut g2, &seeds, 0.6, 1e-6, 0.4, None).scores;
323
324        assert_eq!(r1.len(), r2.len());
325        for (id, v1) in &r1 {
326            let v2 = r2.get(id).copied().unwrap_or(f64::NAN);
327            assert!((v1 - v2).abs() < 1e-12, "PPR drift at {id}: {v1} vs {v2}");
328        }
329    }
330
331    #[test]
332    fn ppr_converges_under_tighter_tolerance() {
333        let (mut g_loose, center) = build_star_graph();
334        let (mut g_tight, _) = build_star_graph();
335        let seeds: FxHashSet<FragmentId> = std::iter::once(center).collect();
336
337        let loose = personalized_pagerank(&mut g_loose, &seeds, 0.6, 1e-2, 0.4, None).scores;
338        let tight = personalized_pagerank(&mut g_tight, &seeds, 0.6, 1e-6, 0.4, None).scores;
339
340        let max_diff = loose
341            .iter()
342            .map(|(id, v)| (v - tight.get(id).copied().unwrap_or(0.0)).abs())
343            .fold(0.0f64, f64::max);
344        assert!(
345            max_diff < 1e-2,
346            "PPR did not converge: max diff between tol=1e-2 and tol=1e-6 is {max_diff}"
347        );
348    }
349
350    #[test]
351    fn ppr_symmetric_star_assigns_equal_mass_to_leaves() {
352        let (mut g, center) = build_star_graph();
353        let seeds: FxHashSet<FragmentId> = std::iter::once(center.clone()).collect();
354        let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-8, 0.5, None).scores;
355
356        let leaf_scores: Vec<f64> = (0..5)
357            .map(|i| result[&fid(&format!("leaf_{i}.rs"), 1, 10)])
358            .collect();
359        let max_leaf = leaf_scores.iter().cloned().fold(0.0f64, f64::max);
360        let min_leaf = leaf_scores.iter().cloned().fold(f64::INFINITY, f64::min);
361        assert!(
362            (max_leaf - min_leaf) < 1e-6,
363            "Symmetric star should give equal leaf mass; got spread {} (leaves: {leaf_scores:?})",
364            max_leaf - min_leaf
365        );
366        assert!(result[&center] > max_leaf, "Center must dominate leaves");
367    }
368
369    /// Claim 7 (paper §4.4): hub suppression $w'_{uv} = w_{uv} / \ln(1 + \text{in\_deg}(v))$
370    /// reduces PPR mass concentration on hub nodes without removing them.
371    ///
372    /// Two graphs are compared:
373    ///   - Naive: a leaf-to-hub graph built directly via `Graph::add_edge`, no suppression.
374    ///   - Suppressed: same topology built via `build_graph` with non-exempt category,
375    ///     which triggers `apply_hub_suppression` for in-degree above the median.
376    ///
377    /// Expected: hub mass is materially reduced; non-hub mass is largely preserved.
378    #[test]
379    fn claim_7_hub_suppression_reduces_hub_mass_without_removal() {
380        use crate::graph::{EdgeCategory, build_graph};
381        use crate::types::{Fragment, FragmentKind};
382
383        let n_leaves = 20usize;
384        let hub = fid("hub.rs", 1, 10);
385        let leaves: Vec<FragmentId> = (0..n_leaves)
386            .map(|i| fid(&format!("leaf_{i}.rs"), 1, 10))
387            .collect();
388
389        let mut naive = Graph::new();
390        naive.add_node(hub.clone());
391        for leaf in &leaves {
392            naive.add_node(leaf.clone());
393            naive.add_edge(leaf.clone(), hub.clone(), 1.0);
394        }
395        for i in 0..n_leaves - 1 {
396            naive.add_edge(leaves[i].clone(), leaves[i + 1].clone(), 1.0);
397        }
398
399        let fragments: Vec<Fragment> = std::iter::once(hub.clone())
400            .chain(leaves.iter().cloned())
401            .map(|id| Fragment {
402                id,
403                kind: FragmentKind::Function,
404                content: Arc::from(""),
405                identifiers: FxHashSet::default(),
406                token_count: 100,
407                symbol_name: None,
408            })
409            .collect();
410        let mut edges: FxHashMap<(FragmentId, FragmentId), f64> = FxHashMap::default();
411        let mut categories: FxHashMap<(FragmentId, FragmentId), EdgeCategory> =
412            FxHashMap::default();
413        for leaf in &leaves {
414            edges.insert((leaf.clone(), hub.clone()), 1.0);
415            categories.insert((leaf.clone(), hub.clone()), EdgeCategory::Generic);
416        }
417        for i in 0..n_leaves - 1 {
418            edges.insert((leaves[i].clone(), leaves[i + 1].clone()), 1.0);
419            categories.insert(
420                (leaves[i].clone(), leaves[i + 1].clone()),
421                EdgeCategory::Generic,
422            );
423        }
424        let mut suppressed = build_graph(&fragments, edges, categories);
425
426        let seeds: FxHashSet<FragmentId> = leaves.iter().take(3).cloned().collect();
427        let alpha = 0.6;
428        let tol = 1e-8;
429        let blend = 1.0;
430
431        let r_naive = personalized_pagerank(&mut naive, &seeds, alpha, tol, blend, None).scores;
432        let r_suppressed =
433            personalized_pagerank(&mut suppressed, &seeds, alpha, tol, blend, None).scores;
434
435        let hub_naive = r_naive.get(&hub).copied().unwrap_or(0.0);
436        let hub_suppressed = r_suppressed.get(&hub).copied().unwrap_or(0.0);
437
438        assert!(
439            hub_suppressed < hub_naive,
440            "Hub suppression did not reduce hub mass: naive={hub_naive}, suppressed={hub_suppressed}"
441        );
442        let reduction_ratio = hub_naive / hub_suppressed.max(1e-12);
443        assert!(
444            reduction_ratio >= 1.5,
445            "Hub suppression effect too small: only {reduction_ratio:.2}× reduction (want ≥1.5×)"
446        );
447
448        assert!(
449            hub_suppressed > 0.0,
450            "Hub mass should be reduced, not removed; got {hub_suppressed}"
451        );
452
453        let mut leaves_present_after = 0;
454        for leaf in &leaves {
455            if r_suppressed.get(leaf).copied().unwrap_or(0.0) > 0.0 {
456                leaves_present_after += 1;
457            }
458        }
459        assert!(
460            leaves_present_after >= leaves.len() / 2,
461            "Suppression should preserve most leaves in the result, only {leaves_present_after}/{} survived",
462            leaves.len()
463        );
464    }
465
466    /// Claim 10 (paper §4.4 hypothesis): PPR with $\alpha \in [0.5, 0.65]$ ranks nodes
467    /// similarly to ego-graph BFS scoring (hop-decay 1/(1+d)). Spearman rank correlation
468    /// of the two score vectors should be high on a connected graph.
469    #[test]
470    fn claim_10_ppr_and_ego_rankings_correlate_on_synthetic_graph() {
471        let mut g = Graph::new();
472        let nodes: Vec<FragmentId> = (0..30).map(|i| fid(&format!("n_{i}.rs"), 1, 10)).collect();
473        for n in &nodes {
474            g.add_node(n.clone());
475        }
476        let mut rng = 0xC0FFEE_u64;
477        let xorshift = |state: &mut u64| -> u64 {
478            *state ^= *state << 13;
479            *state ^= *state >> 7;
480            *state ^= *state << 17;
481            *state
482        };
483        for i in 0..nodes.len() {
484            for _ in 0..3 {
485                let j = (xorshift(&mut rng) as usize) % nodes.len();
486                if i != j {
487                    g.add_edge(nodes[i].clone(), nodes[j].clone(), 1.0);
488                }
489            }
490        }
491        for i in 0..nodes.len() {
492            let next = (i + 1) % nodes.len();
493            g.add_edge(nodes[i].clone(), nodes[next].clone(), 1.0);
494        }
495
496        let seeds: FxHashSet<FragmentId> = nodes.iter().take(2).cloned().collect();
497        let ppr_scores = personalized_pagerank(&mut g, &seeds, 0.6, 1e-8, 0.5, None).scores;
498        let ego_scores = g.ego_graph(&seeds, 3);
499
500        let common: Vec<FragmentId> = nodes
501            .iter()
502            .filter(|n| ppr_scores.contains_key(n) && ego_scores.contains_key(n))
503            .cloned()
504            .collect();
505        assert!(
506            common.len() >= 10,
507            "Need ≥10 common ranked nodes, got {}",
508            common.len()
509        );
510
511        let ppr_v: Vec<f64> = common.iter().map(|n| ppr_scores[n]).collect();
512        let ego_v: Vec<f64> = common.iter().map(|n| ego_scores[n]).collect();
513
514        let rho = spearman_correlation(&ppr_v, &ego_v);
515        assert!(
516            rho > 0.3,
517            "PPR/EGO Spearman correlation too low: ρ={rho:.3} (paper hypothesizes high correlation, want > 0.3)"
518        );
519    }
520
521    fn rank(values: &[f64]) -> Vec<f64> {
522        let n = values.len();
523        let mut indexed: Vec<(usize, f64)> = values.iter().copied().enumerate().collect();
524        indexed.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
525        let mut ranks = vec![0.0_f64; n];
526        let mut i = 0;
527        while i < n {
528            let mut j = i;
529            while j + 1 < n && indexed[j + 1].1 == indexed[i].1 {
530                j += 1;
531            }
532            let avg_rank = ((i + j) as f64) / 2.0 + 1.0;
533            for k in i..=j {
534                ranks[indexed[k].0] = avg_rank;
535            }
536            i = j + 1;
537        }
538        ranks
539    }
540
541    fn spearman_correlation(x: &[f64], y: &[f64]) -> f64 {
542        assert_eq!(x.len(), y.len());
543        let rx = rank(x);
544        let ry = rank(y);
545        let n = x.len() as f64;
546        let mean_x: f64 = rx.iter().sum::<f64>() / n;
547        let mean_y: f64 = ry.iter().sum::<f64>() / n;
548        let mut cov = 0.0;
549        let mut var_x = 0.0;
550        let mut var_y = 0.0;
551        for i in 0..rx.len() {
552            let dx = rx[i] - mean_x;
553            let dy = ry[i] - mean_y;
554            cov += dx * dy;
555            var_x += dx * dx;
556            var_y += dy * dy;
557        }
558        cov / (var_x.sqrt() * var_y.sqrt()).max(1e-12)
559    }
560
561    /// Claim 6B (paper §4.4): personalized PageRank converges to the closed-form
562    /// stationary distribution π = (1-α)(I - αM)^(-1) p, normalized to sum to 1.
563    ///
564    /// For a symmetric star with center c and N leaves, all bidirectional weight 1:
565    ///     π_center = 1 / (1 + α)
566    ///     π_leaf   = α / (N · (1 + α))
567    /// Derivation:
568    ///     π_c = (1-α) + α · Σ_leaf π_leaf,   π_leaf = α · (1/N) · π_c
569    ///     ⇒ π_c · (1 + α) = 1 after normalization (Σ π = 1).
570    #[test]
571    fn ppr_matches_closed_form_on_symmetric_star() {
572        for &(alpha, n_leaves) in &[(0.6, 5usize), (0.5, 4usize), (0.85, 7usize)] {
573            let mut g = Graph::new();
574            let center = fid("center.rs", 1, 10);
575            g.add_node(center.clone());
576            for i in 0..n_leaves {
577                let leaf = fid(&format!("leaf_{i}.rs"), 1, 10);
578                g.add_node(leaf.clone());
579                g.add_edge(center.clone(), leaf.clone(), 1.0);
580                g.add_edge(leaf.clone(), center.clone(), 1.0);
581            }
582            let seeds: FxHashSet<FragmentId> = std::iter::once(center.clone()).collect();
583            let result = personalized_pagerank(&mut g, &seeds, alpha, 1e-10, 0.5, None).scores;
584
585            let expected_center = 1.0 / (1.0 + alpha);
586            let expected_leaf = alpha / (n_leaves as f64 * (1.0 + alpha));
587
588            let actual_center = result[&center];
589            let center_err = (actual_center - expected_center).abs();
590            assert!(
591                center_err < 1e-3,
592                "α={alpha}, N={n_leaves}: center mass drift |actual - closed-form| = {center_err}; \
593                 expected={expected_center}, got={actual_center}"
594            );
595
596            for i in 0..n_leaves {
597                let leaf_id = fid(&format!("leaf_{i}.rs"), 1, 10);
598                let actual_leaf = result[&leaf_id];
599                let leaf_err = (actual_leaf - expected_leaf).abs();
600                assert!(
601                    leaf_err < 1e-3,
602                    "α={alpha}, N={n_leaves}, leaf_{i}: drift |actual - closed-form| = {leaf_err}; \
603                     expected={expected_leaf}, got={actual_leaf}"
604                );
605            }
606        }
607    }
608}