Skip to main content

dtrees_rs/cover/
reversible_cover.rs

1//! A reversible sparse bitset of instances.
2
3use crate::bitsets::Bitset;
4use search_trail::{
5    ReversibleU64, ReversibleUsize, SaveAndRestore, StateManager, U64Manager, UsizeManager,
6};
7use std::cmp::Ordering;
8use std::ops::Sub;
9
10/// A set of instances that supports intersection and constant-time undo.
11///
12/// Words that become zero are moved out of the active part of
13/// `non_zero_words`, so later operations only visit non-empty words.
14pub struct SparseBitset {
15    words: Vec<ReversibleU64>,
16    non_zero_words: Vec<usize>,
17    nb_non_zero: ReversibleUsize,
18
19    state_manager: StateManager,
20}
21
22/// A non-reversible snapshot of a [`SparseBitset`]: every word, with the
23/// inactive ones as 0.
24#[derive(Debug)]
25pub struct ShallowBitset {
26    words: Vec<u64>,
27}
28
29/// Sizes of the two set differences between a [`SparseBitset`] and a
30/// [`ShallowBitset`].
31#[derive(Default)]
32pub struct Difference {
33    pub(crate) in_count: usize,
34    pub(crate) out_count: usize,
35}
36
37impl SparseBitset {
38    /// The full set `0..n`.
39    pub fn new(n: usize) -> Self {
40        let mut state_manager = StateManager::default();
41
42        let nb_words = n.div_ceil(64);
43        let mut words = Vec::with_capacity(nb_words);
44        for _ in 0..nb_words {
45            words.push(state_manager.manage_u64(u64::MAX));
46        }
47
48        let mask = if n % 64 == 0 {
49            u64::MAX
50        } else {
51            (1u64 << (n % 64)) - 1
52        };
53        if let Some(last) = words.last_mut() {
54            state_manager.set_u64(*last, mask);
55        }
56        let non_zero_words = (0..nb_words).collect();
57        let nb_non_zero = state_manager.manage_usize(nb_words);
58
59        state_manager.save_state();
60
61        Self {
62            words,
63            non_zero_words,
64            nb_non_zero,
65            state_manager,
66        }
67    }
68
69    /// Number of elements.
70    pub fn count(&self) -> usize {
71        let mut count = 0;
72        let nb_non_zero = self.state_manager.get_usize(self.nb_non_zero);
73        for i in (0..nb_non_zero).rev() {
74            count += self
75                .state_manager
76                .get_u64(self.words[self.non_zero_words[i]])
77                .count_ones();
78        }
79        count as usize
80    }
81
82    /// Whether the set is empty.
83    pub fn is_empty(&self) -> bool {
84        self.state_manager.get_usize(self.nb_non_zero) == 0
85    }
86
87    /// Intersects with `other` (with its complement when `invert`), saving the
88    /// current state for [`Self::restore`]. Returns the new size.
89    pub fn intersect_with(&mut self, other: &Bitset, invert: bool) -> usize {
90        self.state_manager.save_state();
91
92        let mut size = self.state_manager.get_usize(self.nb_non_zero);
93        let mut count = 0;
94        for i in (0..size).rev() {
95            let idx = self.non_zero_words[i];
96            let intersect = self.state_manager.get_u64(self.words[idx])
97                & if invert { !other[idx] } else { other[idx] };
98            if intersect == 0 {
99                size -= 1;
100                self.non_zero_words[i] = self.non_zero_words[size];
101                self.non_zero_words[size] = idx;
102            } else {
103                self.state_manager.set_u64(self.words[idx], intersect);
104                count += intersect.count_ones();
105            }
106        }
107        self.state_manager.set_usize(self.nb_non_zero, size);
108
109        count as usize
110    }
111
112    /// Size of the intersection with `other` (or its complement), without
113    /// changing the set.
114    pub fn count_intersect_with(&self, other: &Bitset, invert: bool) -> usize {
115        let size = self.state_manager.get_usize(self.nb_non_zero);
116        let mut count = 0;
117        for i in (0..size).rev() {
118            let idx = self.non_zero_words[i];
119            let intersect = self.state_manager.get_u64(self.words[idx])
120                & if invert { !other[idx] } else { other[idx] };
121            if intersect != 0 {
122                count += intersect.count_ones();
123            }
124        }
125        count as usize
126    }
127
128    /// Size of the intersection with each of `others`.
129    pub fn count_intersect_with_many(&self, others: &[Bitset]) -> Vec<usize> {
130        let mut counts = vec![0; others.len()];
131        let size = self.state_manager.get_usize(self.nb_non_zero);
132        for i in (0..size).rev() {
133            let idx = self.non_zero_words[i];
134            let word = self.state_manager.get_u64(self.words[idx]);
135            for (bid, other) in others.iter().enumerate() {
136                counts[bid] += (word & other[idx]).count_ones() as usize
137            }
138        }
139        counts
140    }
141
142    /// The elements of the set.
143    pub fn to_vec(&self) -> Vec<usize> {
144        let mut result = Vec::new();
145        let nb_non_zero = self.state_manager.get_usize(self.nb_non_zero);
146
147        for i in 0..nb_non_zero {
148            let word_idx = self.non_zero_words[i];
149            let mut word = self.state_manager.get_u64(self.words[word_idx]);
150            let base_idx = word_idx * 64;
151
152            while word != 0 {
153                let bit_pos = word.trailing_zeros() as usize;
154                result.push(base_idx + bit_pos);
155                word &= word - 1;
156            }
157        }
158        result
159    }
160
161    /// Undoes the last intersection.
162    #[inline]
163    pub fn restore(&mut self) {
164        self.state_manager.restore_state();
165    }
166
167    /// Every word of the set, with the inactive ones as 0.
168    ///
169    /// A word that becomes empty is moved out of the active part of
170    /// `non_zero_words` without being overwritten, so its stored value is
171    /// stale and must not be read.
172    fn active_words(&self) -> Vec<u64> {
173        let mut words = vec![0; self.words.len()];
174        let nb_non_zero = self.state_manager.get_usize(self.nb_non_zero);
175        for &idx in &self.non_zero_words[..nb_non_zero] {
176            words[idx] = self.state_manager.get_u64(self.words[idx]);
177        }
178        words
179    }
180}
181
182impl From<&SparseBitset> for ShallowBitset {
183    fn from(val: &SparseBitset) -> Self {
184        ShallowBitset {
185            words: val.active_words(),
186        }
187    }
188}
189
190/// `in_count` counts the elements of `self` missing from `rhs`, and
191/// `out_count` the elements of `rhs` missing from `self`.
192impl Sub<&ShallowBitset> for &SparseBitset {
193    type Output = Difference;
194    fn sub(self, rhs: &ShallowBitset) -> Self::Output {
195        let words = self.active_words();
196        let in_count = words
197            .iter()
198            .zip(&rhs.words)
199            .map(|(&own, &other)| (own & !other).count_ones() as usize)
200            .sum();
201        let out_count = words
202            .iter()
203            .zip(&rhs.words)
204            .map(|(&own, &other)| (other & !own).count_ones() as usize)
205            .sum();
206        Difference {
207            in_count,
208            out_count,
209        }
210    }
211}
212
213impl Sub<ShallowBitset> for SparseBitset {
214    type Output = Difference;
215
216    fn sub(self, rhs: ShallowBitset) -> Self::Output {
217        &self - &rhs
218    }
219}
220impl Sub<&ShallowBitset> for SparseBitset {
221    type Output = Difference;
222
223    fn sub(self, rhs: &ShallowBitset) -> Self::Output {
224        &self - rhs
225    }
226}
227
228impl Sub<ShallowBitset> for &SparseBitset {
229    type Output = Difference;
230
231    fn sub(self, rhs: ShallowBitset) -> Self::Output {
232        self - &rhs
233    }
234}
235
236impl Ord for Difference {
237    fn cmp(&self, other: &Self) -> Ordering {
238        (self.in_count + self.out_count).cmp(&(other.in_count + other.out_count))
239    }
240}
241
242impl PartialOrd for Difference {
243    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
244        Some(self.cmp(other))
245    }
246}
247
248impl PartialEq for Difference {
249    fn eq(&self, other: &Self) -> bool {
250        (self.in_count + self.out_count) == (other.in_count + other.out_count)
251    }
252}
253
254impl Eq for Difference {}
255
256#[cfg(test)]
257mod sparse_test {
258    use crate::bitsets::{BitCollection, Bitset, BitsetInit};
259    use crate::cover::reversible_cover::{ShallowBitset, SparseBitset};
260    use search_trail::UsizeManager;
261
262    #[test]
263    fn create() {
264        let mut cover = SparseBitset::new(10);
265
266        let mut feature = Bitset::new(BitsetInit::Empty(10));
267
268        feature.set(9);
269
270        println!("{:?}", cover.to_vec());
271        assert_eq!(cover.count(), 10);
272
273        cover.intersect_with(&feature, false);
274        println!("{:?}", cover.to_vec());
275        println!("xx {:?}", cover.state_manager.get_usize(cover.nb_non_zero));
276
277        let shallow_cover: ShallowBitset = (&cover).into();
278        println!("Shaloow {:?}", shallow_cover);
279        cover.restore();
280        println!("{:?}", cover.to_vec());
281
282        let _ = &cover - shallow_cover;
283    }
284
285    #[test]
286    fn a_word_emptied_by_an_intersection_counts_as_empty() {
287        // Rows 64..128 only: the intersection empties word 0, which keeps its
288        // old value in storage but is no longer active.
289        let mut feature = Bitset::new(BitsetInit::Empty(128));
290        for row in 64..128 {
291            feature.set(row);
292        }
293        let mut cover = SparseBitset::new(128);
294        let full: ShallowBitset = (&cover).into();
295        cover.intersect_with(&feature, false);
296
297        let difference = &cover - &full;
298        assert_eq!(difference.in_count, 0);
299        assert_eq!(difference.out_count, 64, "rows 0..64 are only in `full`");
300
301        let half: ShallowBitset = (&cover).into();
302        cover.restore();
303        let difference = &cover - &half;
304        assert_eq!(difference.in_count, 64, "rows 0..64 are only in `cover`");
305        assert_eq!(difference.out_count, 0);
306    }
307}