dtrees_rs/cover/
reversible_cover.rs1use crate::bitsets::Bitset;
4use search_trail::{
5 ReversibleU64, ReversibleUsize, SaveAndRestore, StateManager, U64Manager, UsizeManager,
6};
7use std::cmp::Ordering;
8use std::ops::Sub;
9
10pub 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#[derive(Debug)]
25pub struct ShallowBitset {
26 words: Vec<u64>,
27}
28
29#[derive(Default)]
32pub struct Difference {
33 pub(crate) in_count: usize,
34 pub(crate) out_count: usize,
35}
36
37impl SparseBitset {
38 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 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 pub fn is_empty(&self) -> bool {
84 self.state_manager.get_usize(self.nb_non_zero) == 0
85 }
86
87 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 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 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 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 #[inline]
163 pub fn restore(&mut self) {
164 self.state_manager.restore_state();
165 }
166
167 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
190impl 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 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 - ½
304 assert_eq!(difference.in_count, 64, "rows 0..64 are only in `cover`");
305 assert_eq!(difference.out_count, 0);
306 }
307}