Skip to main content

competitive/algorithm/
bitdp.rs

1use super::{One, Zero};
2use std::ops::{Add, BitAnd, BitOr, BitXor, Div, Not, Shl, Shr, Sub};
3
4pub trait BitDpExt:
5    Sized
6    + Copy
7    + Default
8    + PartialEq
9    + Eq
10    + PartialOrd
11    + Ord
12    + Not<Output = Self>
13    + BitAnd<Output = Self>
14    + BitOr<Output = Self>
15    + BitXor<Output = Self>
16    + Shl<usize, Output = Self>
17    + Shr<usize, Output = Self>
18    + Add<Output = Self>
19    + Sub<Output = Self>
20    + Div<Output = Self>
21    + Zero
22    + One
23{
24    fn contains(self, x: usize) -> bool {
25        self & (Self::one() << x) != Self::zero()
26    }
27    fn insert(self, x: usize) -> Self {
28        self | (Self::one() << x)
29    }
30    fn remove(self, x: usize) -> Self {
31        self & !(Self::one() << x)
32    }
33    fn is_subset(self, elements: Self) -> bool {
34        self & elements == elements
35    }
36    fn is_superset(self, elements: Self) -> bool {
37        elements.is_subset(self)
38    }
39    fn subsets(self) -> Subsets<Self> {
40        Subsets {
41            mask: self,
42            cur: Some(self),
43        }
44    }
45    fn combinations(n: usize, k: usize) -> Combinations<Self> {
46        Combinations {
47            mask: Self::one() << n,
48            cur: Some((Self::one() << k) - Self::one()),
49        }
50    }
51}
52
53impl BitDpExt for u8 {}
54impl BitDpExt for u16 {}
55impl BitDpExt for u32 {}
56impl BitDpExt for u64 {}
57impl BitDpExt for u128 {}
58impl BitDpExt for usize {}
59
60#[derive(Debug, Clone)]
61pub struct Subsets<T> {
62    mask: T,
63    cur: Option<T>,
64}
65
66impl<T> Iterator for Subsets<T>
67where
68    T: BitDpExt,
69{
70    type Item = T;
71    fn next(&mut self) -> Option<Self::Item> {
72        if let Some(cur) = self.cur {
73            self.cur = if cur.is_zero() {
74                None
75            } else {
76                Some((cur - T::one()) & self.mask)
77            };
78            Some(cur)
79        } else {
80            None
81        }
82    }
83}
84
85#[derive(Debug, Clone)]
86pub struct Combinations<T> {
87    mask: T,
88    cur: Option<T>,
89}
90
91impl<T> Iterator for Combinations<T>
92where
93    T: BitDpExt,
94{
95    type Item = T;
96    fn next(&mut self) -> Option<Self::Item> {
97        if let Some(cur) = self.cur {
98            if cur < self.mask {
99                self.cur = if cur == T::zero() {
100                    None
101                } else {
102                    let x = cur & (!cur + T::one());
103                    let y = cur + x;
104                    Some(((cur & !y) / x / (T::one() + T::one())) | y)
105                };
106                Some(cur)
107            } else {
108                None
109            }
110        } else {
111            None
112        }
113    }
114}
115
116#[cfg(test)]
117mod tests {
118    use super::*;
119    use crate::tools::Xorshift;
120    use std::collections::BTreeSet;
121
122    #[test]
123    fn test_bit_operations() {
124        let mut rng = Xorshift::default();
125        let pairs: Vec<_> = (0..256usize)
126            .flat_map(|a| (0..256).map(move |b| (a, b)))
127            .chain(
128                rng.random_iter((0..1usize << 12, 0..1usize << 12))
129                    .take(1000),
130            )
131            .collect();
132        for (a, b) in pairs {
133            let set: BTreeSet<_> = (0..12).filter(|&i| a >> i & 1 != 0).collect();
134            let other: BTreeSet<_> = (0..12).filter(|&i| b >> i & 1 != 0).collect();
135            for i in 0..12 {
136                assert_eq!(a.contains(i), set.contains(&i));
137                let mut inserted = set.clone();
138                inserted.insert(i);
139                assert_eq!(a.insert(i), inserted.iter().map(|i| 1 << i).sum());
140                let mut removed = set.clone();
141                removed.remove(&i);
142                assert_eq!(a.remove(i), removed.iter().map(|i| 1 << i).sum());
143            }
144            assert_eq!(a.is_subset(b), other.is_subset(&set));
145            assert_eq!(a.is_superset(b), other.is_superset(&set));
146        }
147        for a in 0..1usize << 12 {
148            let mut subsets: Vec<_> = a.subsets().collect();
149            subsets.sort();
150            assert_eq!(subsets, (0..=a).filter(|x| x & a == *x).collect::<Vec<_>>());
151        }
152        for n in 0..=12 {
153            for k in 0..=n {
154                let expected: Vec<_> = (0..1usize << n)
155                    .filter(|x| x.count_ones() as usize == k)
156                    .collect();
157                assert_eq!(usize::combinations(n, k).collect::<Vec<_>>(), expected);
158            }
159        }
160    }
161}