competitive/data_structure/
bit_vector.rs1use std::iter::FromIterator;
2
3pub trait RankSelectDictionaries {
6 fn bit_length(&self) -> usize;
7 fn access(&self, k: usize) -> bool;
9 fn access_rank1(&self, k: usize) -> (bool, usize) {
11 (self.access(k), self.rank1(k))
12 }
13 fn rank1(&self, k: usize) -> usize {
15 (0..k).filter(|&i| self.access(i)).count()
16 }
17 fn rank0(&self, k: usize) -> usize {
19 k - self.rank1(k)
20 }
21 fn select1(&self, k: usize) -> Option<usize> {
23 let n = self.bit_length();
24 if self.rank1(n) <= k {
25 return None;
26 }
27 let (mut l, mut r) = (0, n);
28 while r - l > 1 {
29 let m = l.midpoint(r);
30 if self.rank1(m) <= k {
31 l = m;
32 } else {
33 r = m;
34 }
35 }
36 Some(l)
37 }
38 fn select0(&self, k: usize) -> Option<usize> {
40 let n = self.bit_length();
41 if self.rank0(n) <= k {
42 return None;
43 }
44 let (mut l, mut r) = (0, n);
45 while r - l > 1 {
46 let m = l.midpoint(r);
47 if self.rank0(m) <= k {
48 l = m;
49 } else {
50 r = m;
51 }
52 }
53 Some(l)
54 }
55}
56
57macro_rules! impl_rank_select_for_bits {
58 ($($t:ty)*) => {$(
59 impl RankSelectDictionaries for $t {
60 fn bit_length(&self) -> usize {
61 const WORD_SIZE: usize = (0 as $t).count_zeros() as usize;
62 WORD_SIZE
63 }
64 fn access(&self, k: usize) -> bool {
65 const WORD_SIZE: usize = (0 as $t).count_zeros() as usize;
66 if k < WORD_SIZE {
67 self & (1 as $t) << k != 0
68 } else {
69 false
70 }
71 }
72 fn rank1(&self, k: usize) -> usize {
73 const WORD_SIZE: usize = (0 as $t).count_zeros() as usize;
74 if k < WORD_SIZE {
75 (self & !(!(0 as $t) << k)).count_ones() as usize
76 } else {
77 self.count_ones() as usize
78 }
79 }
80 }
81 )*};
82}
83
84impl_rank_select_for_bits!(u8 u16 u32 u64 usize i8 i16 i32 i64 isize u128 i128);
85
86#[inline]
87fn select_word_scalar(mut word: u64, mut rank: usize) -> usize {
88 let count = word.count_ones() as usize;
89 debug_assert!(rank < count);
90 if rank < 4 {
91 for _ in 0..rank {
92 word &= word - 1;
93 }
94 return word.trailing_zeros() as usize;
95 }
96 if count - rank <= 4 {
97 for _ in 0..count - rank - 1 {
98 word &= !(1 << (u64::BITS as usize - 1 - word.leading_zeros() as usize));
99 }
100 return u64::BITS as usize - 1 - word.leading_zeros() as usize;
101 }
102
103 let mut offset = 0;
104 let mut width = u64::BITS as usize / 2;
105 while width != 0 {
106 let mask = u64::MAX >> (u64::BITS as usize - width);
107 let count = (word & mask).count_ones() as usize;
108 if rank < count {
109 word &= mask;
110 } else {
111 word >>= width;
112 rank -= count;
113 offset += width;
114 }
115 width /= 2;
116 }
117 offset
118}
119
120#[cfg(target_arch = "x86_64")]
121#[allow(unsafe_op_in_unsafe_fn)] mod simd {
123 use std::arch::x86_64::_pdep_u64;
124
125 #[target_feature(enable = "bmi2")]
126 #[inline]
127 pub unsafe fn select_word(word: u64, rank: usize) -> usize {
128 _pdep_u64(1 << rank, word).trailing_zeros() as usize
129 }
130}
131
132#[derive(Debug, Clone)]
133#[repr(C)]
134pub struct BitVectorBlock {
135 pub bits: u64,
136 pub rank: usize,
137}
138
139#[derive(Debug, Clone)]
140pub struct BitVector {
141 blocks: Vec<BitVectorBlock>,
142 len: usize,
143 sum: usize,
144 select_samples: [Vec<usize>; 2],
145}
146
147impl BitVector {
148 const WORD_SIZE: usize = u64::BITS as usize;
149
150 pub fn from_words(words: &[u64], len: usize) -> Self {
153 assert_eq!(words.len(), len.div_ceil(Self::WORD_SIZE));
154 let mut sum = 0;
155 let mut blocks = Vec::with_capacity(len / Self::WORD_SIZE + 1);
156 for (i, &bits) in words.iter().enumerate() {
157 let count = (len - i * Self::WORD_SIZE).min(Self::WORD_SIZE);
158 let bits = if count == Self::WORD_SIZE {
159 bits
160 } else {
161 bits & ((1u64 << count) - 1)
162 };
163 blocks.push(BitVectorBlock { bits, rank: sum });
164 sum += bits.count_ones() as usize;
165 }
166 if len.is_multiple_of(Self::WORD_SIZE) {
167 blocks.push(BitVectorBlock { bits: 0, rank: sum });
168 }
169 Self::from_blocks(blocks, len, sum)
170 }
171
172 fn from_blocks(blocks: Vec<BitVectorBlock>, len: usize, sum: usize) -> Self {
173 let mut select_samples = [Vec::new(), Vec::new()];
174 for (i, block) in blocks
175 .iter()
176 .enumerate()
177 .take(len.div_ceil(Self::WORD_SIZE))
178 {
179 let start = [i * Self::WORD_SIZE - block.rank, block.rank];
180 let end1 = blocks.get(i + 1).map_or(sum, |next| next.rank);
181 let end = [((i + 1) * Self::WORD_SIZE).min(len) - end1, end1];
182 for bit in 0..2 {
183 if start[bit].div_ceil(256) != end[bit].div_ceil(256) {
184 select_samples[bit].push(i);
185 }
186 }
187 }
188 Self {
189 blocks,
190 len,
191 sum,
192 select_samples,
193 }
194 }
195
196 pub fn with_capacity(bits: usize) -> Self {
197 let mut blocks = Vec::with_capacity(bits.div_ceil(Self::WORD_SIZE) + 1);
198 blocks.push(BitVectorBlock { bits: 0, rank: 0 });
199 Self {
200 blocks,
201 len: 0,
202 sum: 0,
203 select_samples: [Vec::new(), Vec::new()],
204 }
205 }
206
207 pub fn push(&mut self, bit: bool) {
208 let word = self.len / Self::WORD_SIZE;
209 let rank = if bit { self.sum } else { self.len - self.sum };
210 if rank.is_multiple_of(256) {
211 self.select_samples[bit as usize].push(word);
212 }
213 self.blocks[word].bits |= (bit as u64) << (self.len % Self::WORD_SIZE);
214 self.sum += bit as usize;
215 self.len += 1;
216 if self.len.is_multiple_of(Self::WORD_SIZE) {
217 self.blocks.push(BitVectorBlock {
218 bits: 0,
219 rank: self.sum,
220 });
221 }
222 }
223
224 pub fn blocks(&self) -> &[BitVectorBlock] {
227 &self.blocks
228 }
229
230 #[inline]
233 pub fn select_word(bits: u64, rank: usize) -> usize {
234 #[cfg(target_arch = "x86_64")]
235 if is_x86_feature_detected!("bmi2") {
236 return unsafe { simd::select_word(bits, rank) };
238 }
239 select_word_scalar(bits, rank)
240 }
241}
242
243impl RankSelectDictionaries for BitVector {
244 fn bit_length(&self) -> usize {
245 self.len
246 }
247
248 #[inline]
249 fn access(&self, k: usize) -> bool {
250 self.blocks[k / Self::WORD_SIZE].bits & (1u64 << (k % Self::WORD_SIZE)) != 0
251 }
252
253 #[inline]
254 fn access_rank1(&self, k: usize) -> (bool, usize) {
255 let block = &self.blocks[k / Self::WORD_SIZE];
256 let offset = k % Self::WORD_SIZE;
257 (
258 block.bits & (1u64 << offset) != 0,
259 block.rank + (block.bits & !(u64::MAX << offset)).count_ones() as usize,
260 )
261 }
262
263 #[inline]
264 fn rank1(&self, k: usize) -> usize {
265 self.access_rank1(k).1
266 }
267
268 fn select1(&self, k: usize) -> Option<usize> {
269 if k >= self.sum {
270 return None;
271 }
272 let sample = k / 256;
273 let start = self.select_samples[1][sample];
274 let end = self.select_samples[1]
275 .get(sample + 1)
276 .map_or(self.blocks.len(), |&word| word + 1);
277 let word = start + self.blocks[start..end].partition_point(|block| block.rank <= k) - 1;
278 let rank = k - self.blocks[word].rank;
279 Some(word * Self::WORD_SIZE + Self::select_word(self.blocks[word].bits, rank))
280 }
281
282 fn select0(&self, k: usize) -> Option<usize> {
283 if k >= self.len - self.sum {
284 return None;
285 }
286 let sample = k / 256;
287 let mut word = self.select_samples[0][sample];
288 let end = self.select_samples[0]
289 .get(sample + 1)
290 .map_or(self.blocks.len(), |&word| word + 1);
291 let mut size = end - word;
292 while size > 1 {
293 let half = size / 2;
294 let middle = word + half;
295 word = if middle * Self::WORD_SIZE - self.blocks[middle].rank <= k {
296 middle
297 } else {
298 word
299 };
300 size -= half;
301 }
302 let rank = k - (word * Self::WORD_SIZE - self.blocks[word].rank);
303 Some(word * Self::WORD_SIZE + Self::select_word(!self.blocks[word].bits, rank))
304 }
305}
306
307impl FromIterator<bool> for BitVector {
308 fn from_iter<I: IntoIterator<Item = bool>>(iter: I) -> Self {
309 let iter = iter.into_iter();
310 let mut blocks = Vec::with_capacity(iter.size_hint().0 / Self::WORD_SIZE + 1);
311 let mut len = 0usize;
312 let mut sum = 0;
313 let mut iter = iter.fuse();
314 while let Some(first) = iter.next() {
315 let mut bits = first as u64;
316 let mut count = 1;
317 for (i, bit) in iter.by_ref().take(Self::WORD_SIZE - 1).enumerate() {
318 bits |= (bit as u64) << (i + 1);
319 count += 1;
320 }
321 blocks.push(BitVectorBlock { bits, rank: sum });
322 len += count;
323 sum += bits.count_ones() as usize;
324 }
325 if len.is_multiple_of(Self::WORD_SIZE) {
326 blocks.push(BitVectorBlock { bits: 0, rank: sum });
327 }
328 Self::from_blocks(blocks, len, sum)
329 }
330}
331
332#[cfg(test)]
333mod tests {
334 use super::*;
335 use crate::tools::Xorshift;
336 use crate::tools::testutil::{exhaustive_sequences, integer_boundary_values};
337 use std::iter;
338
339 const Q: usize = 5_000;
340
341 #[test]
342 fn test_rank_select_word() {
343 const WORD_SIZE: usize = u64::BITS as usize;
344 let mut rng = Xorshift::default();
345 for x in (0..=u16::MAX as u64)
346 .chain(integer_boundary_values!(u64))
347 .chain(rng.random_iter(0u64..).take(Q))
348 {
349 for k in 0..=WORD_SIZE {
350 assert_eq!(x.rank1(k), (0..k).filter(|&i| x.access(i)).count());
351 assert_eq!(x.rank0(k), (0..k).filter(|&i| !x.access(i)).count());
352 if k < x.count_ones() as usize {
353 assert_eq!(select_word_scalar(x, k), x.select1(k).unwrap());
354 }
355 if let Some(i) = x.select1(k) {
356 assert_eq!((0..i).filter(|&j| x.access(j)).count(), k);
357 assert!(x.access(i));
358 } else {
359 assert!(x.rank1(WORD_SIZE) <= k);
360 }
361 if let Some(i) = x.select0(k) {
362 assert_eq!((0..i).filter(|&j| !x.access(j)).count(), k);
363 assert!(!x.access(i));
364 } else {
365 assert!(x.rank0(WORD_SIZE) <= k);
366 }
367 }
368 }
369 }
370
371 #[test]
372 fn test_rank_select_bit_vector() {
373 for events in exhaustive_sequences([None, Some(false), Some(true)], 0..=6) {
374 let expected: Vec<_> = events
375 .iter()
376 .copied()
377 .take_while(Option::is_some)
378 .flatten()
379 .collect();
380 let mut events = events.into_iter();
381 let actual: BitVector = iter::from_fn(|| events.next().flatten()).collect();
382 assert_eq!(actual.bit_length(), expected.len());
383 for (i, &bit) in expected.iter().enumerate() {
384 assert_eq!(actual.access(i), bit);
385 }
386 assert_eq!(
387 actual.rank1(expected.len()),
388 expected.iter().filter(|&&x| x).count()
389 );
390 }
391 let mut rng = Xorshift::default();
392 for len in [
393 0,
394 1,
395 BitVector::WORD_SIZE - 1,
396 BitVector::WORD_SIZE,
397 BitVector::WORD_SIZE + 1,
398 16384 - 1,
399 16384,
400 16384 + 1,
401 65537,
402 ] {
403 for pattern in 0..7 {
404 let bits: Vec<_> = (0..len)
405 .map(|index| match pattern {
406 0 => rng.rand(5) != 0,
407 1 => index.is_multiple_of(BitVector::WORD_SIZE * 3 + 1),
408 2 => !index.is_multiple_of(BitVector::WORD_SIZE * 3 + 1),
409 3 => false,
410 4 => true,
411 5 => index % 8193 == 8192,
412 _ => index % 8193 != 8192,
413 })
414 .collect();
415 let mut words = vec![u64::MAX; len.div_ceil(64)];
416 let mut positions = [Vec::new(), Vec::new()];
417 for (i, &bit) in bits.iter().enumerate() {
418 positions[bit as usize].push(i);
419 if !bit {
420 words[i / 64] &= !(1 << (i % 64));
421 }
422 }
423 for split in [0, len / 2, len.saturating_sub(1), len] {
424 let mut pushed = BitVector::with_capacity(len);
425 for &bit in &bits[..split] {
426 pushed.push(bit);
427 }
428 let collected: BitVector = bits[..split].iter().copied().collect();
429 let packed = BitVector::from_words(&words[..split.div_ceil(64)], split);
430 for mut actual in [pushed, collected, packed] {
431 for &bit in &bits[split..] {
432 actual.push(bit);
433 }
434 let mut rank1 = 0;
435 for (index, &bit) in bits.iter().enumerate() {
436 assert_eq!(actual.access(index), bit);
437 assert_eq!(actual.access_rank1(index), (bit, rank1));
438 rank1 += bit as usize;
439 }
440 for end in [0, len / 3, len / 2, len] {
441 assert_eq!(
442 actual.rank1(end),
443 bits[..end].iter().filter(|&&bit| bit).count()
444 );
445 assert_eq!(
446 actual.rank0(end),
447 bits[..end].iter().filter(|&&bit| !bit).count()
448 );
449 }
450 for rank in 0..=positions[1].len() {
451 assert_eq!(actual.select1(rank), positions[1].get(rank).copied());
452 }
453 for rank in 0..=positions[0].len() {
454 assert_eq!(actual.select0(rank), positions[0].get(rank).copied());
455 }
456 }
457 }
458 }
459 }
460 }
461}