Skip to main content

xlog_induce/
reduce.rs

1//! Deterministic per-topology top-K reduction and tie diagnostics.
2//!
3//! Behaviorally equivalent to the Python reference in
4//! `crates/pyxlog/python/pyxlog/ilp/exact_induce.py`:
5//!
6//! 1. Within each topology, sort scored pairs by lexicographic key
7//!    `(-positives_covered, negatives_covered, left_idx, right_idx)`.
8//! 2. Filter to pairs with `positives_covered > 0`, then keep the first
9//!    `k_per_topology`.
10//! 3. For each kept pair:
11//!    - `local_rank` = 0-indexed position within the positive-filtered list.
12//!    - `next_positives_covered` / `next_negatives_covered` = the next pair
13//!      in the positive-filtered list, or `(0, 0)` if there is none.
14//!    - `tie_class_size` = count of pairs in the FULL sorted list (including
15//!      zero-coverage) sharing the same `(positives_covered, negatives_covered)`.
16//! 4. Output groups candidates by `Topology::ALL` order.
17
18use crate::types::{ScoredCandidate, Topology};
19use xlog_core::RelId;
20
21/// One scored `(topology, left, right)` triple produced by the scoring stage.
22///
23/// Passed into [`reduce_per_topology`] as a flat list; grouping happens inside.
24#[derive(Debug, Clone, Copy, PartialEq, Eq)]
25pub struct ScoredPair {
26    pub topology: Topology,
27    pub left_rel_idx: RelId,
28    pub right_rel_idx: RelId,
29    pub positives_covered: u32,
30    pub negatives_covered: u32,
31}
32
33/// Reduce a flat scored-pair list to the final ordered `ScoredCandidate` list.
34///
35/// Matches the Python reference comparator and diagnostics bit-for-bit.
36pub fn reduce_per_topology(
37    scored_pairs: &[ScoredPair],
38    head_rel_idx: RelId,
39    k_per_topology: u32,
40) -> Vec<ScoredCandidate> {
41    let mut result = Vec::new();
42    let k = k_per_topology as usize;
43
44    for topology in Topology::ALL {
45        // Pull this topology's pairs into a sortable vector.
46        let mut this_topo: Vec<ScoredPair> = scored_pairs
47            .iter()
48            .copied()
49            .filter(|p| p.topology == topology)
50            .collect();
51
52        // Lexicographic sort: max positives, min negatives, min L idx, min R idx.
53        this_topo.sort_by(|a, b| {
54            (
55                std::cmp::Reverse(a.positives_covered),
56                a.negatives_covered,
57                a.left_rel_idx.0,
58                a.right_rel_idx.0,
59            )
60                .cmp(&(
61                    std::cmp::Reverse(b.positives_covered),
62                    b.negatives_covered,
63                    b.left_rel_idx.0,
64                    b.right_rel_idx.0,
65                ))
66        });
67
68        // Positive-coverage filter preserves sorted order.
69        let positives: Vec<ScoredPair> = this_topo
70            .iter()
71            .copied()
72            .filter(|p| p.positives_covered > 0)
73            .collect();
74
75        let kept_n = std::cmp::min(k, positives.len());
76        for (rank, pair) in positives.iter().take(kept_n).enumerate() {
77            // Diagnostics: next candidate in the positive-filtered list.
78            let (next_pos, next_neg) = positives
79                .get(rank + 1)
80                .map(|nxt| (nxt.positives_covered, nxt.negatives_covered))
81                .unwrap_or((0, 0));
82
83            // Tie class counted over the FULL sorted list (including zero-coverage).
84            let tie_count = this_topo
85                .iter()
86                .filter(|s| {
87                    s.positives_covered == pair.positives_covered
88                        && s.negatives_covered == pair.negatives_covered
89                })
90                .count() as u32;
91
92            result.push(ScoredCandidate {
93                topology,
94                head_rel_idx,
95                left_rel_idx: pair.left_rel_idx,
96                right_rel_idx: pair.right_rel_idx,
97                positives_covered: pair.positives_covered,
98                negatives_covered: pair.negatives_covered,
99                local_rank: rank as u32,
100                next_positives_covered: next_pos,
101                next_negatives_covered: next_neg,
102                tie_class_size: tie_count,
103            });
104        }
105    }
106
107    result
108}
109
110/// One kept n-ary pattern after deterministic reduction.
111///
112/// `pattern_idx` points back into the scored pattern batch (canonical
113/// enumeration order), so the caller can recover the full
114/// [`crate::nary::NaryRulePattern`] and its per-atom relation identities.
115#[derive(Debug, Clone, Copy, PartialEq, Eq)]
116pub struct KeptNaryPattern {
117    pub pattern_idx: usize,
118    pub positives_covered: u32,
119    pub negatives_covered: u32,
120    pub local_rank: u32,
121    pub next_positives_covered: u32,
122    pub next_negatives_covered: u32,
123    pub tie_class_size: u32,
124}
125
126/// Reduce per-pattern `(positives_covered, negatives_covered)` counts to
127/// the final ordered top-K.
128///
129/// The binary reduction law with the topology dimension removed and the
130/// canonical enumeration index as the final tie-breaker:
131///
132/// 1. Sort pattern indexes by `(-positives_covered, negatives_covered,
133///    pattern_idx)`. The enumeration is deterministic lexicographic, so
134///    the index tie-break is as stable across calls as the binary
135///    engine's `(left_idx, right_idx)` tie-break.
136/// 2. Filter to `positives_covered > 0`, keep the first `k`.
137/// 3. `local_rank` / `next_*` over the positive-filtered list;
138///    `tie_class_size` over the FULL batch (including zero-coverage).
139pub fn reduce_nary(coverage: &[(u32, u32)], k: u32) -> Vec<KeptNaryPattern> {
140    let mut order: Vec<usize> = (0..coverage.len()).collect();
141    order.sort_by_key(|&i| (std::cmp::Reverse(coverage[i].0), coverage[i].1, i));
142
143    let positives: Vec<usize> = order
144        .iter()
145        .copied()
146        .filter(|&i| coverage[i].0 > 0)
147        .collect();
148
149    // Tie classes are counted ONCE over the whole batch rather than
150    // rescanned per kept pattern: the per-pattern scan is O(k * batch) and
151    // a large n-ary batch makes this diagnostic dominate the reduction it
152    // is only describing.
153    let mut tie_counts: std::collections::HashMap<(u32, u32), u32> =
154        std::collections::HashMap::new();
155    for &pair in coverage {
156        *tie_counts.entry(pair).or_insert(0) += 1;
157    }
158
159    let kept_n = std::cmp::min(k as usize, positives.len());
160    let mut result = Vec::with_capacity(kept_n);
161    for (rank, &idx) in positives.iter().take(kept_n).enumerate() {
162        let (pos, neg) = coverage[idx];
163        let (next_pos, next_neg) = positives
164            .get(rank + 1)
165            .map(|&nxt| coverage[nxt])
166            .unwrap_or((0, 0));
167        let tie_count = tie_counts.get(&(pos, neg)).copied().unwrap_or(0);
168        result.push(KeptNaryPattern {
169            pattern_idx: idx,
170            positives_covered: pos,
171            negatives_covered: neg,
172            local_rank: rank as u32,
173            next_positives_covered: next_pos,
174            next_negatives_covered: next_neg,
175            tie_class_size: tie_count,
176        });
177    }
178    result
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184
185    fn pair(topo: Topology, l: u32, r: u32, pos: u32, neg: u32) -> ScoredPair {
186        ScoredPair {
187            topology: topo,
188            left_rel_idx: RelId(l),
189            right_rel_idx: RelId(r),
190            positives_covered: pos,
191            negatives_covered: neg,
192        }
193    }
194
195    const HEAD: RelId = RelId(100);
196
197    #[test]
198    fn empty_input_yields_empty_output() {
199        let result = reduce_per_topology(&[], HEAD, 2);
200        assert!(result.is_empty());
201    }
202
203    #[test]
204    fn zero_k_yields_empty_output() {
205        let pairs = vec![pair(Topology::Chain, 1, 2, 5, 0)];
206        let result = reduce_per_topology(&pairs, HEAD, 0);
207        assert!(result.is_empty());
208    }
209
210    #[test]
211    fn zero_coverage_pairs_are_excluded_from_kept() {
212        let pairs = vec![
213            pair(Topology::Chain, 1, 2, 0, 0),
214            pair(Topology::Chain, 3, 4, 0, 0),
215        ];
216        let result = reduce_per_topology(&pairs, HEAD, 2);
217        assert!(result.is_empty());
218    }
219
220    #[test]
221    fn single_positive_pair_has_tie_one_and_no_next() {
222        let pairs = vec![pair(Topology::Chain, 1, 2, 5, 0)];
223        let result = reduce_per_topology(&pairs, HEAD, 2);
224        assert_eq!(result.len(), 1);
225        let c = &result[0];
226        assert_eq!(c.topology, Topology::Chain);
227        assert_eq!(c.left_rel_idx, RelId(1));
228        assert_eq!(c.right_rel_idx, RelId(2));
229        assert_eq!(c.positives_covered, 5);
230        assert_eq!(c.negatives_covered, 0);
231        assert_eq!(c.local_rank, 0);
232        assert_eq!(c.next_positives_covered, 0);
233        assert_eq!(c.next_negatives_covered, 0);
234        assert_eq!(c.tie_class_size, 1);
235    }
236
237    #[test]
238    fn max_positives_wins_over_higher_negatives() {
239        // pair A: pos=5, neg=2; pair B: pos=3, neg=0 → A wins despite more negatives.
240        let pairs = vec![
241            pair(Topology::Chain, 1, 2, 3, 0),
242            pair(Topology::Chain, 3, 4, 5, 2),
243        ];
244        let result = reduce_per_topology(&pairs, HEAD, 1);
245        assert_eq!(result.len(), 1);
246        assert_eq!(result[0].positives_covered, 5);
247        assert_eq!(result[0].negatives_covered, 2);
248    }
249
250    #[test]
251    fn min_negatives_breaks_positives_tie() {
252        // Same positives — lower negatives wins.
253        let pairs = vec![
254            pair(Topology::Chain, 1, 2, 5, 3),
255            pair(Topology::Chain, 3, 4, 5, 1),
256        ];
257        let result = reduce_per_topology(&pairs, HEAD, 1);
258        assert_eq!(result.len(), 1);
259        assert_eq!(result[0].left_rel_idx, RelId(3));
260        assert_eq!(result[0].negatives_covered, 1);
261    }
262
263    #[test]
264    fn left_idx_breaks_pos_neg_tie() {
265        // Same positives+negatives — lower left idx wins.
266        let pairs = vec![
267            pair(Topology::Chain, 5, 2, 3, 0),
268            pair(Topology::Chain, 1, 2, 3, 0),
269        ];
270        let result = reduce_per_topology(&pairs, HEAD, 1);
271        assert_eq!(result.len(), 1);
272        assert_eq!(result[0].left_rel_idx, RelId(1));
273    }
274
275    #[test]
276    fn right_idx_breaks_all_other_ties() {
277        // Same pos, neg, left — lower right idx wins.
278        let pairs = vec![
279            pair(Topology::Chain, 1, 5, 3, 0),
280            pair(Topology::Chain, 1, 2, 3, 0),
281        ];
282        let result = reduce_per_topology(&pairs, HEAD, 1);
283        assert_eq!(result.len(), 1);
284        assert_eq!(result[0].right_rel_idx, RelId(2));
285    }
286
287    #[test]
288    fn top_k_truncation_preserves_order() {
289        // Three positive pairs, K=2 → keep the top two.
290        let pairs = vec![
291            pair(Topology::Chain, 1, 2, 3, 0),
292            pair(Topology::Chain, 1, 3, 5, 0),
293            pair(Topology::Chain, 1, 4, 4, 0),
294        ];
295        let result = reduce_per_topology(&pairs, HEAD, 2);
296        assert_eq!(result.len(), 2);
297        assert_eq!(result[0].positives_covered, 5);
298        assert_eq!(result[0].local_rank, 0);
299        assert_eq!(result[1].positives_covered, 4);
300        assert_eq!(result[1].local_rank, 1);
301    }
302
303    #[test]
304    fn next_diagnostics_point_to_rank_plus_one_in_positive_list() {
305        // Three positives; K=2. Top-2 get next_* from positions 1 and 2.
306        let pairs = vec![
307            pair(Topology::Chain, 1, 1, 5, 0), // rank 0
308            pair(Topology::Chain, 2, 2, 4, 0), // rank 1
309            pair(Topology::Chain, 3, 3, 3, 0), // rank 2
310        ];
311        let result = reduce_per_topology(&pairs, HEAD, 2);
312        assert_eq!(result.len(), 2);
313        // rank 0's next_* is rank 1's (pos=4, neg=0)
314        assert_eq!(result[0].next_positives_covered, 4);
315        assert_eq!(result[0].next_negatives_covered, 0);
316        // rank 1's next_* is rank 2's (pos=3, neg=0)
317        assert_eq!(result[1].next_positives_covered, 3);
318        assert_eq!(result[1].next_negatives_covered, 0);
319    }
320
321    #[test]
322    fn next_diagnostics_are_zero_when_no_next() {
323        // Only one positive — next_* should be (0, 0).
324        let pairs = vec![pair(Topology::Chain, 1, 1, 5, 0)];
325        let result = reduce_per_topology(&pairs, HEAD, 2);
326        assert_eq!(result[0].next_positives_covered, 0);
327        assert_eq!(result[0].next_negatives_covered, 0);
328    }
329
330    #[test]
331    fn tie_class_size_counts_same_pos_neg_in_full_sorted_list() {
332        // Three pairs share (pos=5, neg=0); one has different (pos=3, neg=0).
333        let pairs = vec![
334            pair(Topology::Chain, 1, 1, 5, 0),
335            pair(Topology::Chain, 2, 2, 5, 0),
336            pair(Topology::Chain, 3, 3, 5, 0),
337            pair(Topology::Chain, 4, 4, 3, 0),
338        ];
339        let result = reduce_per_topology(&pairs, HEAD, 2);
340        // Top-2 should both have tie_class_size=3 (three pairs share pos=5, neg=0).
341        assert_eq!(result[0].tie_class_size, 3);
342        assert_eq!(result[1].tie_class_size, 3);
343    }
344
345    #[test]
346    fn topologies_output_in_all_order() {
347        // One positive per topology; check output order = chain, star, fanout, fanin.
348        let pairs = vec![
349            pair(Topology::Fanin, 1, 1, 1, 0),
350            pair(Topology::Fanout, 1, 1, 1, 0),
351            pair(Topology::Star, 1, 1, 1, 0),
352            pair(Topology::Chain, 1, 1, 1, 0),
353        ];
354        let result = reduce_per_topology(&pairs, HEAD, 1);
355        assert_eq!(result.len(), 4);
356        assert_eq!(result[0].topology, Topology::Chain);
357        assert_eq!(result[1].topology, Topology::Star);
358        assert_eq!(result[2].topology, Topology::Fanout);
359        assert_eq!(result[3].topology, Topology::Fanin);
360    }
361
362    #[test]
363    fn topology_filtering_is_per_topology() {
364        // Pairs from both chain and star; ranking happens independently.
365        let pairs = vec![
366            pair(Topology::Chain, 1, 1, 3, 0),
367            pair(Topology::Chain, 2, 2, 5, 0),
368            pair(Topology::Star, 3, 3, 2, 0),
369            pair(Topology::Star, 4, 4, 4, 0),
370        ];
371        let result = reduce_per_topology(&pairs, HEAD, 1);
372        assert_eq!(result.len(), 2);
373        // chain winner: pos=5
374        assert_eq!(result[0].topology, Topology::Chain);
375        assert_eq!(result[0].positives_covered, 5);
376        // star winner: pos=4
377        assert_eq!(result[1].topology, Topology::Star);
378        assert_eq!(result[1].positives_covered, 4);
379    }
380
381    #[test]
382    fn k_larger_than_positives_returns_all_positives() {
383        let pairs = vec![
384            pair(Topology::Chain, 1, 1, 5, 0),
385            pair(Topology::Chain, 2, 2, 3, 0),
386        ];
387        let result = reduce_per_topology(&pairs, HEAD, 10);
388        assert_eq!(result.len(), 2);
389    }
390
391    #[test]
392    fn head_rel_idx_propagates_to_every_candidate() {
393        let pairs = vec![
394            pair(Topology::Chain, 1, 1, 5, 0),
395            pair(Topology::Star, 2, 2, 3, 0),
396        ];
397        let head = RelId(42);
398        let result = reduce_per_topology(&pairs, head, 1);
399        for c in &result {
400            assert_eq!(c.head_rel_idx, head);
401        }
402    }
403
404    // ── reduce_nary ─────────────────────────────────────────────────────
405
406    #[test]
407    fn nary_empty_input_yields_empty_output() {
408        assert!(reduce_nary(&[], 2).is_empty());
409    }
410
411    #[test]
412    fn nary_zero_k_yields_empty_output() {
413        assert!(reduce_nary(&[(5, 0)], 0).is_empty());
414    }
415
416    #[test]
417    fn nary_zero_coverage_patterns_are_excluded() {
418        assert!(reduce_nary(&[(0, 0), (0, 3)], 2).is_empty());
419    }
420
421    #[test]
422    fn nary_max_positives_wins_over_higher_negatives() {
423        let kept = reduce_nary(&[(3, 0), (5, 2)], 1);
424        assert_eq!(kept.len(), 1);
425        assert_eq!(kept[0].pattern_idx, 1);
426        assert_eq!(kept[0].positives_covered, 5);
427        assert_eq!(kept[0].negatives_covered, 2);
428    }
429
430    #[test]
431    fn nary_min_negatives_breaks_positives_tie() {
432        let kept = reduce_nary(&[(5, 3), (5, 1)], 1);
433        assert_eq!(kept[0].pattern_idx, 1);
434        assert_eq!(kept[0].negatives_covered, 1);
435    }
436
437    #[test]
438    fn nary_enumeration_index_breaks_full_ties() {
439        let kept = reduce_nary(&[(5, 1), (5, 1)], 1);
440        assert_eq!(kept[0].pattern_idx, 0);
441    }
442
443    #[test]
444    fn nary_top_k_truncation_preserves_order_and_ranks() {
445        let kept = reduce_nary(&[(3, 0), (5, 0), (4, 0)], 2);
446        assert_eq!(kept.len(), 2);
447        assert_eq!(kept[0].pattern_idx, 1);
448        assert_eq!(kept[0].local_rank, 0);
449        assert_eq!(kept[1].pattern_idx, 2);
450        assert_eq!(kept[1].local_rank, 1);
451    }
452
453    #[test]
454    fn nary_next_diagnostics_point_past_kept_prefix() {
455        // Three positives, K=2: rank 0 sees rank 1, rank 1 sees the
456        // UNKEPT rank 2 — next_* reads the positive-filtered list, not
457        // the kept prefix.
458        let kept = reduce_nary(&[(5, 0), (4, 0), (3, 7)], 2);
459        assert_eq!(kept.len(), 2);
460        assert_eq!(
461            (
462                kept[0].next_positives_covered,
463                kept[0].next_negatives_covered
464            ),
465            (4, 0)
466        );
467        assert_eq!(
468            (
469                kept[1].next_positives_covered,
470                kept[1].next_negatives_covered
471            ),
472            (3, 7)
473        );
474    }
475
476    #[test]
477    fn nary_next_diagnostics_are_zero_when_no_next() {
478        let kept = reduce_nary(&[(5, 0)], 2);
479        assert_eq!(kept[0].next_positives_covered, 0);
480        assert_eq!(kept[0].next_negatives_covered, 0);
481    }
482
483    #[test]
484    fn nary_tie_class_counts_full_batch_including_zero_coverage_ties() {
485        // (5,0) appears three times; the fourth pattern differs.
486        let kept = reduce_nary(&[(5, 0), (5, 0), (3, 0), (5, 0)], 2);
487        assert_eq!(kept[0].tie_class_size, 3);
488        assert_eq!(kept[1].tie_class_size, 3);
489    }
490
491    #[test]
492    fn nary_k_larger_than_positives_returns_all_positives() {
493        let kept = reduce_nary(&[(5, 0), (0, 0), (3, 0)], 10);
494        assert_eq!(kept.len(), 2);
495    }
496}