Skip to main content

competitive/algorithm/
binary_search.rs

1use std::cmp::Ordering;
2
3/// binary search helper
4pub trait Bisect: Clone {
5    /// Return between two elements if search is not end.
6    fn bisect_middle_point(&self, other: &Self) -> Option<Self>;
7}
8
9macro_rules! impl_bisect_unsigned {
10    ($($t:ty)*) => {
11        $(impl Bisect for $t {
12            fn bisect_middle_point(&self, other: &Self) -> Option<Self> {
13                if self.abs_diff(*other) > 1 { Some(self.midpoint(*other)) } else { None }
14            }
15        })*
16    };
17}
18macro_rules! impl_bisect_signed {
19    ($($t:ty)*) => {
20        $(impl Bisect for $t {
21            fn bisect_middle_point(&self, other: &Self) -> Option<Self> {
22                if self.signum() != other.signum() {
23                    if match self.cmp(other) {
24                        Ordering::Less => self + 1 < *other,
25                        Ordering::Equal => false,
26                        Ordering::Greater => other + 1 < *self,
27                    } {
28                        Some((*self).midpoint(*other))
29                    } else {
30                        None
31                    }
32                } else {
33                    if self.abs_diff(*other) > 1 { Some(self.midpoint(*other)) } else { None }
34                }
35            }
36        })*
37    };
38}
39macro_rules! impl_bisect_float {
40    ($({$t:ident $u:ident $i:ident $e:expr})*) => {
41        $(impl Bisect for $t {
42            fn bisect_middle_point(&self, other: &Self) -> Option<Self> {
43                fn to_float_ord(x: $t) -> $i {
44                    let a = x.to_bits() as $i;
45                    a ^ (((a >> $e) as $u) >> 1) as $i
46                }
47                fn from_float_ord(a: $i) -> $t {
48                    $t::from_bits((a ^ (((a >> $e) as $u) >> 1) as $i) as _)
49                }
50                <$i as Bisect>::bisect_middle_point(&to_float_ord(*self), &to_float_ord(*other)).map(from_float_ord)
51            }
52        })*
53    };
54}
55impl_bisect_unsigned!(u8 u16 u32 u64 u128 usize);
56impl_bisect_signed!(i8 i16 i32 i64 i128 isize);
57impl_bisect_float!({f32 u32 i32 31} {f64 u64 i64 63});
58
59/// binary search for monotone segment
60///
61/// if `ok < err` then search [ok, err) where t(`ok`), t, t, .... t, t(`ret`), f,  ... f, f, f, `err`
62///
63/// if `err < ok` then search (err, ok] where `err`, f, f, f, ... f, t(`ret`), ... t, t, t(`ok`)
64pub fn binary_search<T, F>(mut f: F, mut ok: T, mut err: T) -> T
65where
66    T: Bisect,
67    F: FnMut(&T) -> bool,
68{
69    while let Some(m) = ok.bisect_middle_point(&err) {
70        if f(&m) {
71            ok = m;
72        } else {
73            err = m;
74        }
75    }
76    ok
77}
78
79/// binary search for slice
80pub trait SliceBisectExt<T> {
81    /// Returns the first element that satisfies a predicate.
82    fn find_bisect(&self, f: impl FnMut(&T) -> bool) -> Option<&T>;
83    /// Returns the last element that satisfies a predicate.
84    fn rfind_bisect(&self, f: impl FnMut(&T) -> bool) -> Option<&T>;
85    /// Returns the first index that satisfies a predicate.
86    /// if not found, returns `len()`.
87    fn position_bisect(&self, f: impl FnMut(&T) -> bool) -> usize;
88    /// Returns the last index+1 that satisfies a predicate.
89    /// if not found, returns `0`.
90    fn rposition_bisect(&self, f: impl FnMut(&T) -> bool) -> usize;
91}
92impl<T> SliceBisectExt<T> for [T] {
93    fn find_bisect(&self, f: impl FnMut(&T) -> bool) -> Option<&T> {
94        self.get(self.position_bisect(f))
95    }
96    fn rfind_bisect(&self, f: impl FnMut(&T) -> bool) -> Option<&T> {
97        let pos = self.rposition_bisect(f);
98        if pos == 0 { None } else { self.get(pos - 1) }
99    }
100    fn position_bisect(&self, mut f: impl FnMut(&T) -> bool) -> usize {
101        binary_search(|i| f(&self[*i as usize]), self.len() as i64, -1) as usize
102    }
103    fn rposition_bisect(&self, mut f: impl FnMut(&T) -> bool) -> usize {
104        binary_search(|i| f(&self[i - 1]), 0, self.len() + 1)
105    }
106}
107
108pub fn parallel_binary_search<T, F, G>(mut f: F, q: usize, ok: T, err: T) -> Vec<T>
109where
110    T: Bisect,
111    F: FnMut(&[Option<T>]) -> G,
112    G: Fn(usize) -> bool,
113{
114    let mut ok = vec![ok; q];
115    let mut err = vec![err; q];
116    loop {
117        let m: Vec<_> = ok
118            .iter()
119            .zip(&err)
120            .map(|(ok, err)| ok.bisect_middle_point(err))
121            .collect();
122        if m.iter().all(|m| m.is_none()) {
123            break;
124        }
125        let g = f(&m);
126        for (i, m) in m.into_iter().enumerate() {
127            if let Some(m) = m {
128                if g(i) {
129                    ok[i] = m;
130                } else {
131                    err[i] = m;
132                }
133            }
134        }
135    }
136    ok
137}
138
139#[cfg(test)]
140mod tests {
141    use super::*;
142    use crate::algorithm::SliceCombinationsExt;
143    use crate::tools::{
144        Xorshift,
145        testutil::{integer_boundary_values, sample_usize, structured_sequences},
146    };
147
148    #[test]
149    fn test_slice_bisect() {
150        let mut rng = Xorshift::default();
151        let mut cases = Vec::new();
152        // Every sorted sequence over {-1, 0, 1}, without enumerating its permutations.
153        for n in 0..=32 {
154            [-1, 0, 1].for_each_combinations_with_replacement(n, |xs| cases.push(xs.to_vec()));
155        }
156        let lengths = sample_usize(&mut rng, 32, 0..=1024, 1000);
157        cases.extend(
158            structured_sequences(&mut rng, -10..=10, lengths).map(|mut values| {
159                values.sort_unstable();
160                values
161            }),
162        );
163        cases.sort_unstable();
164        cases.dedup();
165        for values in cases {
166            let n = values.len();
167            let mut first = 0;
168            let mut end = 0;
169            for key in -11..=11 {
170                while first < n && values[first] < key {
171                    first += 1;
172                }
173                while end < n && values[end] <= key {
174                    end += 1;
175                }
176                assert_eq!(
177                    values.position_bisect(|&x| x >= key),
178                    first,
179                    "values={values:?}, key={key}"
180                );
181                assert_eq!(
182                    values.find_bisect(|&x| x >= key),
183                    values.get(first),
184                    "values={values:?}, key={key}"
185                );
186                assert_eq!(
187                    values.rposition_bisect(|&x| x <= key),
188                    end,
189                    "values={values:?}, key={key}"
190                );
191                assert_eq!(
192                    values.rfind_bisect(|&x| x <= key),
193                    values[..end].last(),
194                    "values={values:?}, key={key}"
195                );
196                assert_eq!(
197                    binary_search(|&i: &isize| values[i as usize] >= key, n as isize, -1),
198                    first as isize,
199                    "values={values:?}, key={key}"
200                );
201                assert_eq!(
202                    binary_search(|&i: &isize| values[i as usize] <= key, -1, n as isize),
203                    end as isize - 1,
204                    "values={values:?}, key={key}"
205                );
206            }
207        }
208    }
209
210    #[test]
211    fn test_integer_bisect() {
212        macro_rules! check {
213            ($($ty:ty),*) => {$(
214                let mut rng = Xorshift::default();
215                for boundary in integer_boundary_values!($ty).into_iter()
216                    .chain((0..=u8::MAX).map(|x| x as $ty))
217                    .chain(rng.random_iter(..).take(10_000))
218                {
219                    if boundary < <$ty>::MAX {
220                        assert_eq!(binary_search(|&x| x <= boundary, <$ty>::MIN, <$ty>::MAX), boundary);
221                        assert_eq!(binary_search(|&x| x > boundary, <$ty>::MAX, <$ty>::MIN), boundary + 1);
222                    }
223                    if boundary > <$ty>::MIN {
224                        assert_eq!(binary_search(|&x| x >= boundary, <$ty>::MAX, <$ty>::MIN), boundary);
225                        assert_eq!(binary_search(|&x| x < boundary, <$ty>::MIN, <$ty>::MAX), boundary - 1);
226                    }
227                }
228            )*};
229        }
230        check!(
231            u8, u16, u32, u64, u128, usize, i8, i16, i32, i64, i128, isize
232        );
233    }
234
235    #[test]
236    fn test_float_bisect() {
237        macro_rules! check {
238            ($ty:ty, $bits:ty, $fraction:expr, $max_exponent:expr) => {
239                let mut rng = Xorshift::default();
240                // Powers of two and their adjacent representations cross every
241                // normal/subnormal exponent boundary, in both directions.
242                let mut bits: Vec<$bits> = (1..=$max_exponent)
243                    .flat_map(|exponent| {
244                        let power = exponent << $fraction;
245                        [power - 1, power, power + 1]
246                    })
247                    .collect();
248                bits.extend(integer_boundary_values!($bits));
249                bits.extend((0..10_000).map(|_| {
250                    let bits: $bits = rng.random(..);
251                    bits
252                }));
253                for bits in bits {
254                    for x in [<$ty>::from_bits(bits), -<$ty>::from_bits(bits)] {
255                        if x.is_finite() && x != 0.0 {
256                            assert_eq!(
257                                binary_search(|&y| y <= x, <$ty>::NEG_INFINITY, <$ty>::INFINITY),
258                                x
259                            );
260                            assert_eq!(
261                                binary_search(|&y| y >= x, <$ty>::INFINITY, <$ty>::NEG_INFINITY),
262                                x
263                            );
264                        }
265                    }
266                }
267                for x in 0..=10_000 {
268                    let x = x as $ty;
269                    let actual = binary_search(|&y| y * y <= x, 0.0, x + 1.0);
270                    assert!((actual - x.sqrt()).abs() <= <$ty>::EPSILON * x.sqrt().max(1.0));
271                }
272            };
273        }
274        check!(f32, u32, 23, 254);
275        check!(f64, u64, 52, 2046);
276    }
277}