dtrees_rs/cover/
similarities.rs1use crate::cover::reversible_cover::{Difference, ShallowBitset, SparseBitset};
4
5#[derive(Debug)]
9pub struct SimilarityCover {
10 covers: [Option<ShallowBitset>; 2],
11 errors: [f64; 2],
12}
13
14impl Default for SimilarityCover {
15 fn default() -> Self {
16 Self::new()
17 }
18}
19
20impl SimilarityCover {
21 pub fn new() -> Self {
23 Self {
24 covers: [None, None],
25 errors: [f64::INFINITY; 2],
26 }
27 }
28
29 pub fn update(&mut self, cover: &SparseBitset, error: f64) {
32 let shallow_cover: ShallowBitset = cover.into();
33
34 match (self.covers[0].as_ref(), self.covers[1].as_ref()) {
35 (None, _) => {
36 self.covers[0] = Some(shallow_cover);
37 self.errors[0] = error;
38 return;
39 }
40 (_, None) => {
41 self.covers[1] = Some(shallow_cover);
42 self.errors[1] = error;
43 return;
44 }
45 _ => {}
46 }
47
48 let differences: Vec<Difference> = self
49 .covers
50 .iter()
51 .map(|sim_cover| sim_cover.as_ref().map(|c| cover - c).unwrap_or_default())
52 .collect();
53
54 let min_idx = differences
55 .iter()
56 .enumerate()
57 .min_by_key(|&(_, diff)| diff)
58 .map(|(idx, _)| idx);
59 if let Some(idx) = min_idx {
60 self.covers[idx] = Some(shallow_cover);
61 self.errors[idx] = error;
62 }
63 }
64
65 pub fn compute_similarity(&self, cover: &SparseBitset) -> f64 {
67 self.covers
68 .iter()
69 .enumerate()
70 .filter_map(|(i, cover_opt)| {
71 cover_opt.as_ref().map(|sim_cover| {
72 let diff = cover - sim_cover;
73 self.errors[i] - diff.out_count as f64
74 })
75 })
76 .fold(0.0, f64::max)
77 }
78}
79
80#[cfg(test)]
81mod tests {
82 use super::SimilarityCover;
83 use crate::bitsets::{BitCollection, Bitset, BitsetInit};
84 use crate::cover::reversible_cover::SparseBitset;
85
86 fn cover(rows: &[usize]) -> SparseBitset {
88 let mut feature = Bitset::new(BitsetInit::Empty(8));
89 for &row in rows {
90 feature.set(row);
91 }
92 let mut cover = SparseBitset::new(8);
93 cover.intersect_with(&feature, false);
94 cover
95 }
96
97 #[test]
98 fn a_replaced_reference_set_takes_its_new_error() {
99 let mut similarity = SimilarityCover::new();
100 similarity.update(&cover(&[0, 1, 2, 3]), 3.0);
101 similarity.update(&cover(&[4, 5, 6, 7]), 3.0);
102 similarity.update(&cover(&[0, 1, 2]), 1.0);
104
105 assert_eq!(similarity.compute_similarity(&cover(&[0, 1, 2])), 1.0);
108 }
109}