Skip to main content

competitive/algorithm/
ternary_search.rs

1use std::ops::RangeInclusive;
2
3/// fibonacci search helper
4pub trait FibonacciSearch: Sized {
5    fn fibonacci_search<T, F>(self, other: Self, f: F) -> (Self, T)
6    where
7        T: PartialOrd,
8        F: FnMut(Self) -> T;
9}
10macro_rules! impl_fibonacci_search_unsigned {
11    ($($t:ty)*) => {
12        $(impl FibonacciSearch for $t {
13            fn fibonacci_search<T, F>(self, other: Self, mut f: F) -> (Self, T)
14            where
15                T: PartialOrd,
16                F: FnMut(Self) -> T,
17            {
18                let l = self;
19                let r = other;
20                assert!(l <= r);
21                const W: usize = [12, 23, 46, 92, 185][<$t>::BITS.ilog2() as usize - 3];
22                const FIB: [$t; W] = {
23                    let mut fib = [0; W];
24                    fib[0] = 1;
25                    fib[1] = 2;
26                    let mut i = 2;
27                    while i < W {
28                        fib[i] = fib[i - 1] + fib[i - 2];
29                        i += 1;
30                    }
31                    fib
32                };
33                let mut s = l;
34                let mut v0 = None;
35                let mut v1 = None;
36                let mut v2 = None;
37                let mut v3 = None;
38                for w in FIB[..FIB.partition_point(|&f| f < r - l)].windows(2).rev() {
39                    let (w0, w1) = (w[0], w[1]);
40                    if w1 > r - s || v1.get_or_insert_with(|| f(s + w0)) <= v2.get_or_insert_with(|| f(s + w1)) {
41                        v3 = v2;
42                        v2 = v1;
43                        v1 = None;
44                    } else {
45                        v0 = v1;
46                        v1 = v2;
47                        v2 = None;
48                        s += w0;
49                    }
50                }
51                let mut kv = (s, v0.unwrap_or_else(|| f(s)));
52                if s < r {
53                    let v = v1.or(v2).unwrap_or_else(|| f(s + 1));
54                    if v < kv.1 {
55                        kv = (s + 1, v);
56                    }
57                    if s + 1 < r {
58                        let v = v3.unwrap_or_else(|| f(s + 2));
59                        if v < kv.1 {
60                            kv = (s + 2, v);
61                        }
62                    }
63                }
64                kv
65            }
66        })*
67    };
68}
69impl_fibonacci_search_unsigned!(u8 u16 u32 u64 u128 usize);
70
71/// ternary search helper
72pub trait Trisect: Clone {
73    type Key: FibonacciSearch;
74    fn trisect_key(self) -> Self::Key;
75    fn trisect_unkey(key: Self::Key) -> Self;
76}
77
78macro_rules! impl_trisect_unsigned {
79    ($($t:ty)*) => {
80        $(impl Trisect for $t {
81            type Key = $t;
82            fn trisect_key(self) -> Self::Key {
83                self
84            }
85            fn trisect_unkey(key: Self::Key) -> Self {
86                key
87            }
88        })*
89    };
90}
91macro_rules! impl_trisect_signed {
92    ($({$i:ident $u:ident})*) => {
93        $(impl Trisect for $i {
94            type Key = $u;
95            fn trisect_key(self) -> Self::Key {
96                (self as $u) ^ (1 << <$u>::BITS - 1)
97            }
98            fn trisect_unkey(key: Self::Key) -> Self {
99                (key ^ (1 << <$u>::BITS - 1)) as $i
100            }
101        })*
102    };
103}
104macro_rules! impl_trisect_float {
105    ($({$t:ident $u:ident $i:ident})*) => {
106        $(impl Trisect for $t {
107            type Key = $u;
108            fn trisect_key(self) -> Self::Key {
109                let a = self.to_bits() as $i;
110                (a ^ (((a >> <$u>::BITS - 1) as $u) >> 1) as $i) as $u ^ (1 << <$u>::BITS - 1)
111            }
112            fn trisect_unkey(key: Self::Key) -> Self {
113                let key = (key  ^ (1 << <$u>::BITS - 1)) as $i;
114                $t::from_bits((key ^ (((key >> <$u>::BITS - 1) as $u) >> 1) as $i) as _)
115            }
116        })*
117    };
118}
119
120impl_trisect_unsigned!(u8 u16 u32 u64 u128 usize);
121impl_trisect_signed!({i8 u8} {i16 u16} {i32 u32} {i64 u64} {i128 u128} {isize usize});
122impl_trisect_float!({f32 u32 i32} {f64 u64 i64});
123
124/// Returns the element that gives the minimum value from the strictly concave up function and the minimum value.
125pub fn ternary_search<K, V, F>(range: RangeInclusive<K>, mut f: F) -> (K, V)
126where
127    K: Trisect,
128    V: PartialOrd,
129    F: FnMut(K) -> V,
130{
131    let (l, r) = range.into_inner();
132    let (k, v) =
133        <K::Key as FibonacciSearch>::fibonacci_search(l.trisect_key(), r.trisect_key(), |x| {
134            f(Trisect::trisect_unkey(x))
135        });
136    (K::trisect_unkey(k), v)
137}
138
139pub fn piecewise_ternary_search<const N: usize, K, V, F>(piece: [K; N], mut f: F) -> (K, V)
140where
141    K: Trisect,
142    V: PartialOrd,
143    F: FnMut(K) -> V,
144{
145    piece
146        .windows(2)
147        .map(|w| ternary_search(w[0].clone()..=w[1].clone(), &mut f))
148        .min_by(|a, b| a.1.partial_cmp(&b.1).unwrap())
149        .unwrap_or_else(|| (piece[0].clone(), f(piece[0].clone())))
150}
151
152pub fn golden_ternary_search<T>(
153    range: RangeInclusive<f64>,
154    count: usize,
155    mut f: impl FnMut(f64) -> T,
156) -> (f64, T)
157where
158    T: PartialOrd,
159{
160    let (mut l, mut r) = range.into_inner();
161    // FIXME: 1.94.0: std::f64::consts::GOLDEN_RATIO;
162    const GOLDEN_RATIO_INV: f64 = 1f64 / 1.618_033_988_749_895_f64;
163    let mut v0 = None;
164    let mut v1 = None;
165    let mut v2 = None;
166    let mut v3 = None;
167    for _ in 0..count {
168        let w = (r - l) * GOLDEN_RATIO_INV;
169        if v1.get_or_insert_with(|| f(r - w)) <= v2.get_or_insert_with(|| f(l + w)) {
170            v3 = v2;
171            v2 = v1;
172            v1 = None;
173            r = l + w;
174        } else {
175            v0 = v1;
176            v1 = v2;
177            v2 = None;
178            l = r - w;
179        }
180    }
181    let kv = (l, v0.unwrap_or_else(|| f(l)));
182    let kv2 = (r, v3.unwrap_or_else(|| f(r)));
183    if kv2.1 < kv.1 { kv2 } else { kv }
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189    use crate::tools::testutil::integer_boundary_values;
190    use crate::{num::DoubleDouble, tools::Xorshift};
191    use std::array;
192
193    #[test]
194    fn test_trisect() {
195        let mut rng = Xorshift::default();
196        macro_rules! check {
197            ($($ty:ty),*) => {$(
198                let mut values = integer_boundary_values!($ty);
199                if <$ty>::BITS == 8 { values = (<$ty>::MIN..=<$ty>::MAX).collect(); }
200                let pairs: Vec<_> = values.iter().flat_map(|&p| values.iter().map(move |&q| (p, q))).chain(rng.random_iter((.., ..)).take(1000)).collect();
201                for (p, q) in pairs {
202                    assert_eq!(p, <$ty>::trisect_unkey(p.trisect_key()));
203                    assert_eq!(p.cmp(&q), p.trisect_key().cmp(&q.trisect_key()));
204                }
205            )*};
206        }
207        check!(u8, i8, u16, i16, u32, i32, u64, i64, usize, isize);
208        for _ in 0..1000 {
209            let p = (rng.randf() - 0.5) * 2e100;
210            let q = (rng.randf() - 0.5) * 2e100;
211            assert_eq!(p, f64::trisect_unkey(p.trisect_key()));
212            assert_eq!(
213                p.partial_cmp(&q),
214                p.trisect_key().partial_cmp(&q.trisect_key())
215            );
216        }
217    }
218
219    #[test]
220    fn test_ternary_search() {
221        for p in 0..=u8::MAX {
222            for l in 0..=u8::MAX {
223                for r in l..=u8::MAX {
224                    assert_eq!(ternary_search(l..=r, |x| p.abs_diff(x)).0, p.clamp(l, r));
225                }
226            }
227        }
228        let mut rng = Xorshift::default();
229        for _ in 0..1000 {
230            let l = rng.random(-100i64..=100);
231            let r = rng.random(l..=100);
232            let p = rng.random(-100i64..=100);
233            let a = rng.random(1..=100i64);
234            let b = rng.random(-100i64..=100);
235            let f = |x: i64| a * (x - p).pow(2) + b;
236            let actual = ternary_search(l..=r, f);
237            assert_eq!(actual.0, p.clamp(l, r));
238            assert_eq!(actual.1, (l..=r).map(f).min().unwrap());
239            let l = rng.random(0..=u8::MAX);
240            let r = rng.random(l..=u8::MAX);
241            let p: u8 = rng.random(..);
242            assert_eq!(
243                ternary_search(l..=r, |x| p.abs_diff(x)).1,
244                (l..=r).map(|x| p.abs_diff(x)).min().unwrap()
245            );
246            let p = (rng.randf() - 0.5) * 2e5;
247            let f = |x| (DoubleDouble::from(x) - DoubleDouble::from(p)).abs();
248            assert_eq!(ternary_search(f64::MIN..=f64::MAX, f).0, p);
249            let actual = golden_ternary_search(-1e100..=1e100, 1000, f).0;
250            assert!((actual - p).abs() <= 2.0 * f64::EPSILON * p.abs().max(1.0));
251            let mut bounds: [f64; 8] = array::from_fn(|_| (rng.randf() - 0.5) * 2e5);
252            bounds.sort_by(f64::total_cmp);
253            let expected = p.clamp(bounds[0], *bounds.last().unwrap());
254            assert_eq!(piecewise_ternary_search(bounds, f).0, expected);
255        }
256    }
257}