Skip to main content

competitive/math/min_plus_convolution/
selector.rs

1use super::{
2    Signed, bounded_requirements_from_extrema, concave, convex, min_plus_convolution_bounded_ntt,
3    min_plus_convolution_naive, monotone, output_len, piecewise_linear, sparse,
4};
5
6#[derive(Clone, Copy, Debug, Eq, PartialEq)]
7enum Algorithm {
8    Naive,
9    Sparse,
10    BoundedNtt,
11    ConvexDivideAndConquerLeft,
12    ConvexDivideAndConquerRight,
13    ConvexMerge,
14    ConcaveEnvelopeLeft,
15    ConcaveEnvelopeRight,
16    ConcaveBoth,
17    MonotoneRunsIncreasing,
18    MonotoneRunsDecreasing,
19    LinearLeft,
20    LinearRight,
21    PiecewiseLinearLeft,
22    PiecewiseLinearRight,
23}
24
25struct InputCharacteristics<T> {
26    finite_count: usize,
27    finite_prefix_len: usize,
28    // Finite entries after the first infinity.
29    finite_entries: Vec<(usize, T)>,
30    // Present for finite inputs while the run count fits the selector cache.
31    run_entries: Option<Vec<(usize, T)>>,
32    extrema: Option<(T, T)>,
33    is_convex: bool,
34    is_concave: bool,
35    is_nondecreasing: bool,
36    is_nonincreasing: bool,
37    run_count: usize,
38    piece_count: usize,
39}
40
41fn analyze<T>(values: &[T]) -> InputCharacteristics<T>
42where
43    T: Signed,
44{
45    let Some(&first) = values.first() else {
46        return InputCharacteristics {
47            finite_count: 0,
48            finite_prefix_len: 0,
49            finite_entries: Vec::new(),
50            run_entries: Some(Vec::new()),
51            extrema: None,
52            is_convex: true,
53            is_concave: true,
54            is_nondecreasing: true,
55            is_nonincreasing: true,
56            run_count: 0,
57            piece_count: 0,
58        };
59    };
60    if first.is_maximum() {
61        return analyze_with_infinity(values, 0, None);
62    }
63
64    let mut minimum = first;
65    let mut maximum = first;
66    let mut is_convex = true;
67    let mut is_concave = true;
68    let mut is_nondecreasing = true;
69    let mut is_nonincreasing = true;
70    let mut run_count = 1;
71    let mut run_entries = Some(vec![(0, first)]);
72    let mut piece_count = 1;
73    let mut previous_value = first;
74    let mut previous_slope = None;
75
76    for (index, &value) in values.iter().enumerate().skip(1) {
77        if value.is_maximum() {
78            return analyze_with_infinity(values, index, Some((minimum, maximum)));
79        }
80        minimum = minimum.min(value);
81        maximum = maximum.max(value);
82        is_nondecreasing &= previous_value <= value;
83        is_nonincreasing &= previous_value >= value;
84        if previous_value != value {
85            run_count += 1;
86            if let Some(entries) = &mut run_entries {
87                if entries.len() == MAX_CACHED_RUNS {
88                    run_entries = None;
89                } else {
90                    entries.push((index, value));
91                }
92            }
93        }
94        let slope = value - previous_value;
95        if let Some(previous_slope) = previous_slope {
96            is_convex &= previous_slope <= slope;
97            is_concave &= previous_slope >= slope;
98            piece_count += usize::from(previous_slope != slope);
99        }
100        previous_slope = Some(slope);
101        previous_value = value;
102    }
103    InputCharacteristics {
104        finite_count: values.len(),
105        finite_prefix_len: values.len(),
106        finite_entries: Vec::new(),
107        run_entries,
108        extrema: Some((minimum, maximum)),
109        is_convex,
110        is_concave,
111        is_nondecreasing,
112        is_nonincreasing,
113        run_count,
114        piece_count,
115    }
116}
117
118fn analyze_with_infinity<T>(
119    values: &[T],
120    first_infinity: usize,
121    mut extrema: Option<(T, T)>,
122) -> InputCharacteristics<T>
123where
124    T: Signed,
125{
126    let mut finite_entries = Vec::new();
127    for (offset, &value) in values[first_infinity + 1..].iter().enumerate() {
128        if !value.is_maximum() {
129            finite_entries.push((first_infinity + offset + 1, value));
130            extrema = Some(extrema.map_or((value, value), |(minimum, maximum)| {
131                (minimum.min(value), maximum.max(value))
132            }));
133        }
134    }
135    InputCharacteristics {
136        finite_count: first_infinity + finite_entries.len(),
137        finite_prefix_len: first_infinity,
138        finite_entries,
139        run_entries: None,
140        extrema,
141        is_convex: false,
142        is_concave: false,
143        is_nondecreasing: false,
144        is_nonincreasing: false,
145        run_count: 0,
146        piece_count: 0,
147    }
148}
149
150fn scaled_work(factor: u128, work: u128) -> u128 {
151    factor.saturating_mul(work)
152}
153
154// min_plus/dense: input inspection is above the 15% budget at n=256 and
155// below it at n=1024, so keep the conservative power-of-two boundary.
156const SMALL_PAIR_COUNT: u128 = 1 << 18;
157
158// min_plus_long/structured measures cached run counts through 4096.
159const MAX_CACHED_RUNS: usize = 4096;
160
161fn select_algorithm<T>(
162    a_len: usize,
163    b_len: usize,
164    a_characteristics: &InputCharacteristics<T>,
165    b_characteristics: &InputCharacteristics<T>,
166) -> Algorithm
167where
168    T: Signed + TryFrom<usize>,
169    T::Unsigned: TryInto<usize>,
170{
171    let output = a_len.saturating_add(b_len).saturating_sub(1) as u128;
172    let mut selected = ((a_len as u128) * (b_len as u128), Algorithm::Naive);
173    let mut consider = |work: u128, algorithm| {
174        if work < selected.0 {
175            selected = (work, algorithm);
176        }
177    };
178
179    consider(
180        // min_plus/sparse: 50% finite inputs are already over 10% faster than
181        // the INF-skipping naive scan, while dense pair enumeration loses.
182        scaled_work(
183            3,
184            (a_characteristics.finite_count as u128) * (b_characteristics.finite_count as u128),
185        ),
186        Algorithm::Sparse,
187    );
188    if let (Some(a_extrema), Some(b_extrema)) =
189        (a_characteristics.extrema, b_characteristics.extrema)
190        && let Some(requirements) =
191            bounded_requirements_from_extrema(a_len, b_len, a_extrema, b_extrema)
192    {
193        // min_plus/bounded: small transforms need a wider margin for encoding
194        // overhead; at 2^20 and above the NTT wins from a lower work ratio.
195        consider(
196            scaled_work(
197                if requirements.transform_len < 1 << 20 {
198                    8
199                } else {
200                    6
201                },
202                (requirements.transform_len as u128) * (requirements.transform_len.ilog2() as u128),
203            ),
204            Algorithm::BoundedNtt,
205        );
206    }
207
208    if a_characteristics.is_convex && b_characteristics.is_convex {
209        consider(output, Algorithm::ConvexMerge);
210    } else if a_characteristics.is_convex {
211        consider(
212            scaled_work(2, output),
213            Algorithm::ConvexDivideAndConquerLeft,
214        );
215    } else if b_characteristics.is_convex {
216        consider(
217            scaled_work(2, output),
218            Algorithm::ConvexDivideAndConquerRight,
219        );
220    }
221    if a_characteristics.is_concave && b_characteristics.is_concave {
222        consider(output, Algorithm::ConcaveBoth);
223    } else if a_characteristics.is_concave || b_characteristics.is_concave {
224        consider(
225            scaled_work(8, output.saturating_mul(output.max(1).ilog2() as u128 + 1)),
226            if a_characteristics.is_concave {
227                Algorithm::ConcaveEnvelopeLeft
228            } else {
229                Algorithm::ConcaveEnvelopeRight
230            },
231        );
232    }
233    let increasing = a_characteristics.is_nondecreasing && b_characteristics.is_nondecreasing;
234    let decreasing = a_characteristics.is_nonincreasing && b_characteristics.is_nonincreasing;
235    if increasing || decreasing {
236        consider(
237            scaled_work(
238                4,
239                a_characteristics
240                    .run_count
241                    .saturating_mul(b_characteristics.run_count) as u128,
242            ) + output,
243            if decreasing {
244                Algorithm::MonotoneRunsDecreasing
245            } else {
246                Algorithm::MonotoneRunsIncreasing
247            },
248        );
249    }
250    let piecewise = match (
251        a_characteristics.finite_count == a_len,
252        b_characteristics.finite_count == b_len,
253    ) {
254        (true, true) if a_characteristics.piece_count < b_characteristics.piece_count => {
255            Some((a_characteristics.piece_count, true))
256        }
257        (true, true) => Some((b_characteristics.piece_count, false)),
258        (true, false) => Some((a_characteristics.piece_count, true)),
259        (false, true) => Some((b_characteristics.piece_count, false)),
260        (false, false) => None,
261    };
262    if let Some((pieces, structured_is_left)) = piecewise {
263        if pieces == 1 {
264            consider(
265                output,
266                if structured_is_left {
267                    Algorithm::LinearLeft
268                } else {
269                    Algorithm::LinearRight
270                },
271            );
272        } else {
273            consider(
274                scaled_work(4, (pieces as u128).saturating_mul(output)),
275                if structured_is_left {
276                    Algorithm::PiecewiseLinearLeft
277                } else {
278                    Algorithm::PiecewiseLinearRight
279                },
280            );
281        }
282    }
283    selected.1
284}
285
286/// Computes min-plus convolution after selecting a deterministic exact method
287/// from the observed input structure.
288///
289/// # Panics
290///
291/// Panics if the output length does not fit [`usize`]. Arithmetic overflow is
292/// the caller's responsibility.
293pub fn min_plus_convolution<T>(a: &[T], b: &[T]) -> Vec<T>
294where
295    T: Signed + TryFrom<usize>,
296    T::Unsigned: TryInto<usize>,
297{
298    if (a.len() as u128) * (b.len() as u128) <= SMALL_PAIR_COUNT {
299        return min_plus_convolution_naive(a, b);
300    }
301    let a_characteristics = analyze(a);
302    let distinct_b_characteristics = (!std::ptr::eq(a, b)).then(|| analyze(b));
303    let b_characteristics = distinct_b_characteristics
304        .as_ref()
305        .unwrap_or(&a_characteristics);
306    match select_algorithm(a.len(), b.len(), &a_characteristics, b_characteristics) {
307        Algorithm::Naive => min_plus_convolution_naive(a, b),
308        Algorithm::Sparse => {
309            let a_entries = a[..a_characteristics.finite_prefix_len]
310                .iter()
311                .copied()
312                .enumerate()
313                .chain(a_characteristics.finite_entries.iter().copied());
314            let b_entries = b[..b_characteristics.finite_prefix_len]
315                .iter()
316                .copied()
317                .enumerate()
318                .chain(b_characteristics.finite_entries.iter().copied());
319            sparse(a_entries, b_entries, output_len(a.len(), b.len()))
320        }
321        Algorithm::BoundedNtt => min_plus_convolution_bounded_ntt(a, b),
322        Algorithm::ConvexDivideAndConquerLeft => convex::convex_divide_and_conquer(b, a),
323        Algorithm::ConvexDivideAndConquerRight => convex::convex_divide_and_conquer(a, b),
324        Algorithm::ConvexMerge => convex::convex_merge(a, b),
325        Algorithm::ConcaveEnvelopeLeft => concave::concave_envelope(b, a),
326        Algorithm::ConcaveEnvelopeRight => concave::concave_envelope(a, b),
327        Algorithm::ConcaveBoth => concave::concave_both(a, b),
328        algorithm @ (Algorithm::MonotoneRunsIncreasing | Algorithm::MonotoneRunsDecreasing) => {
329            let increasing = algorithm == Algorithm::MonotoneRunsIncreasing;
330            if let (Some(a_runs), Some(b_runs)) = (
331                &a_characteristics.run_entries,
332                &b_characteristics.run_entries,
333            ) {
334                monotone::monotone_runs_from_entries(a_runs, b_runs, a.len(), b.len(), increasing)
335            } else {
336                monotone::monotone_runs(a, b, increasing)
337            }
338        }
339        Algorithm::LinearLeft => piecewise_linear::linear(b, a),
340        Algorithm::LinearRight => piecewise_linear::linear(a, b),
341        Algorithm::PiecewiseLinearLeft => piecewise_linear::piecewise_linear(b, a),
342        Algorithm::PiecewiseLinearRight => piecewise_linear::piecewise_linear(a, b),
343    }
344}
345
346#[cfg(test)]
347mod tests {
348    use super::*;
349    use crate::tools::Xorshift;
350
351    #[test]
352    fn test_selector() {
353        const LEN: usize = 1024;
354
355        let mut rng = Xorshift::default();
356        let selected =
357            |a: &[i64], b: &[i64]| select_algorithm(a.len(), b.len(), &analyze(a), &analyze(b));
358        let mut sparse_a = vec![i64::MAX; LEN];
359        let mut sparse_b = vec![i64::MAX; LEN];
360        for _ in 0..8 {
361            let i = rng.random(0..LEN);
362            let j = rng.random(0..LEN);
363            sparse_a[i] = rng.random(-1_000_i64..=1_000);
364            sparse_b[j] = rng.random(-1_000_i64..=1_000);
365        }
366        assert_eq!(selected(&sparse_a, &sparse_b), Algorithm::Sparse);
367
368        let mut slopes: Vec<_> = rng.random_iter(-100_i64..=100).take(LEN - 1).collect();
369        slopes.sort_unstable();
370        let mut convex = Vec::with_capacity(LEN);
371        convex.push(rng.random(-1_000_i64..=1_000));
372        for slope in slopes {
373            convex.push(convex[convex.len() - 1] + slope);
374        }
375        assert_eq!(selected(&convex, &convex), Algorithm::ConvexMerge);
376        let arbitrary: Vec<_> = rng
377            .random_iter(-1_000_000_000_i64..=1_000_000_000)
378            .take(LEN)
379            .collect();
380        assert_eq!(
381            selected(&convex, &arbitrary),
382            Algorithm::ConvexDivideAndConquerLeft
383        );
384
385        let concave: Vec<_> = convex.iter().map(|&value| -value).collect();
386        assert_eq!(selected(&concave, &concave), Algorithm::ConcaveBoth);
387        assert_eq!(
388            selected(&concave, &arbitrary),
389            Algorithm::ConcaveEnvelopeLeft
390        );
391
392        let bounded_a: Vec<_> = rng.random_iter(0_i64..=1).take(4096).collect();
393        let bounded_b: Vec<_> = rng.random_iter(0_i64..=1).take(4096).collect();
394        assert_eq!(selected(&bounded_a, &bounded_b), Algorithm::BoundedNtt);
395
396        let mut run_values = Vec::with_capacity(8);
397        let mut value = rng.random(-1_000_i64..=1_000);
398        for _ in 0..8 {
399            run_values.push(value);
400            value += rng.random(1_i64..=10);
401        }
402        let monotone: Vec<_> = (0..LEN).map(|i| run_values[i * 8 / LEN]).collect();
403        let low = rng.random(1_i64..=5);
404        let high = rng.random(6_i64..=10);
405        let mut irregular = Vec::with_capacity(LEN);
406        value = rng.random(-1_000_i64..=1_000);
407        for i in 0..LEN {
408            irregular.push(value);
409            value += if i % 2 == 0 { low } else { high };
410        }
411        assert_eq!(
412            selected(&monotone, &irregular),
413            Algorithm::MonotoneRunsIncreasing
414        );
415
416        let start = rng.random(-1_000_i64..=1_000);
417        let slope = rng.random(-20_i64..=20);
418        let linear: Vec<_> = (0..LEN).map(|i| start + slope * i as i64).collect();
419        assert_eq!(selected(&arbitrary, &linear), Algorithm::LinearRight);
420
421        let piece_slopes = [
422            rng.random(-20_i64..=-11),
423            rng.random(11_i64..=20),
424            rng.random(-10_i64..=-1),
425            rng.random(1_i64..=10),
426        ];
427        let mut piecewise = Vec::with_capacity(LEN);
428        value = rng.random(-1_000_i64..=1_000);
429        for i in 0..LEN {
430            piecewise.push(value);
431            value += piece_slopes[i * piece_slopes.len() / LEN];
432        }
433        assert_eq!(
434            selected(&arbitrary, &piecewise),
435            Algorithm::PiecewiseLinearRight
436        );
437
438        assert_eq!(selected(&irregular, &irregular), Algorithm::Naive);
439    }
440}