Skip to main content

competitive/data_structure/
bit_vector.rs

1use std::iter::FromIterator;
2
3/// rank_i(select_i(k)) = k
4/// rank_i(select_i(k) + 1) = k + 1
5pub trait RankSelectDictionaries {
6    fn bit_length(&self) -> usize;
7    /// get k-th bit
8    fn access(&self, k: usize) -> bool;
9    /// Returns the k-th bit and the number of ones before it.
10    fn access_rank1(&self, k: usize) -> (bool, usize) {
11        (self.access(k), self.rank1(k))
12    }
13    /// the number of 1 in [0, k)
14    fn rank1(&self, k: usize) -> usize {
15        (0..k).filter(|&i| self.access(i)).count()
16    }
17    /// the number of 0 in [0, k)
18    fn rank0(&self, k: usize) -> usize {
19        k - self.rank1(k)
20    }
21    /// index of k-th 1
22    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    /// index of k-th 0
39    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)] // BMI2 is confined to a feature-gated function.
122mod 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    /// Builds a bit vector from low-bit-first words. `words.len()` must equal
151    /// `len.div_ceil(64)`. Unused high bits in the final word are ignored.
152    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    /// Words paired with the number of ones preceding each word. The last block
225    /// is partial, or an empty sentinel when the bit length is a multiple of 64.
226    pub fn blocks(&self) -> &[BitVectorBlock] {
227        &self.blocks
228    }
229
230    /// Returns the position of the zero-based occurrence `rank` in `bits`.
231    /// `rank` must be less than the number of set bits.
232    #[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            // SAFETY: BMI2 is available and the caller checked the occurrence count.
237            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}