Skip to main content

competitive/math/min_plus_convolution/
convex.rs

1use super::{Signed, assert_finite, output_len};
2use std::{
3    cmp::Ordering,
4    ops::{Range, RangeInclusive},
5};
6
7pub(super) fn is_convex<T>(values: &[T]) -> bool
8where
9    T: Signed,
10{
11    values
12        .windows(3)
13        .all(|window| window[1] - window[0] <= window[2] - window[1])
14}
15
16fn orient_one_convex<'a, T>(a: &'a [T], b: &'a [T]) -> (&'a [T], &'a [T])
17where
18    T: Signed,
19{
20    if is_convex(b) {
21        (a, b)
22    } else if is_convex(a) {
23        (b, a)
24    } else {
25        panic!("at least one min-plus convolution input must be convex")
26    }
27}
28
29/// Computes convolution of two convex inputs by merging their slope sequences.
30///
31/// The running time is `O(n + m)`.
32///
33/// # Panics
34///
35/// Panics unless both inputs are finite and convex.
36pub fn min_plus_convolution_convex_merge<T>(a: &[T], b: &[T]) -> Vec<T>
37where
38    T: Signed,
39{
40    let len = output_len(a.len(), b.len());
41    if len == 0 {
42        return Vec::new();
43    }
44    assert_finite(a);
45    assert_finite(b);
46    assert!(is_convex(a) && is_convex(b), "both inputs must be convex");
47    convex_merge(a, b)
48}
49
50pub(super) fn convex_merge<T>(a: &[T], b: &[T]) -> Vec<T>
51where
52    T: Signed,
53{
54    let len = output_len(a.len(), b.len());
55    let mut a_slopes = a.windows(2).map(|window| window[1] - window[0]);
56    let mut b_slopes = b.windows(2).map(|window| window[1] - window[0]);
57    let mut next_a = a_slopes.next();
58    let mut next_b = b_slopes.next();
59    let mut current = a[0] + b[0];
60    let mut result = Vec::with_capacity(len);
61    result.push(current);
62    while next_a.is_some() || next_b.is_some() {
63        let slope = match (next_a, next_b) {
64            (Some(left), Some(right)) if left <= right => {
65                next_a = a_slopes.next();
66                left
67            }
68            (Some(_), Some(right)) => {
69                next_b = b_slopes.next();
70                right
71            }
72            (Some(left), None) => {
73                next_a = a_slopes.next();
74                left
75            }
76            (None, Some(right)) => {
77                next_b = b_slopes.next();
78                right
79            }
80            (None, None) => break,
81        };
82        current += slope;
83        result.push(current);
84    }
85    result
86}
87
88/// Computes convolution when one input is convex using monotone divide and conquer.
89///
90/// # Panics
91///
92/// Panics unless both inputs are finite and at least one is convex.
93pub fn min_plus_convolution_convex_divide_and_conquer<T>(a: &[T], b: &[T]) -> Vec<T>
94where
95    T: Signed,
96{
97    let len = output_len(a.len(), b.len());
98    if len == 0 {
99        return Vec::new();
100    }
101    assert_finite(a);
102    assert_finite(b);
103    let (arbitrary, convex) = orient_one_convex(a, b);
104    convex_divide_and_conquer(arbitrary, convex)
105}
106
107pub(super) fn convex_divide_and_conquer<T>(arbitrary: &[T], convex: &[T]) -> Vec<T>
108where
109    T: Signed,
110{
111    let len = output_len(arbitrary.len(), convex.len());
112    let mut result = vec![T::zero(); len];
113
114    fn solve<T>(
115        arbitrary: &[T],
116        convex: &[T],
117        result: &mut [T],
118        rows: Range<usize>,
119        options: RangeInclusive<usize>,
120    ) where
121        T: Signed,
122    {
123        if rows.is_empty() {
124            return;
125        }
126        let row = (rows.start + rows.end) / 2;
127        let first = (*options.start()).max(row.saturating_sub(convex.len() - 1));
128        let last = (*options.end()).min(row).min(arbitrary.len() - 1);
129        let mut best_col = first;
130        let mut best_value = arbitrary[first] + convex[row - first];
131        for col in first + 1..=last {
132            let value = arbitrary[col] + convex[row - col];
133            if value < best_value {
134                best_value = value;
135                best_col = col;
136            }
137        }
138        result[row] = best_value;
139        solve(
140            arbitrary,
141            convex,
142            result,
143            rows.start..row,
144            *options.start()..=best_col,
145        );
146        solve(
147            arbitrary,
148            convex,
149            result,
150            row + 1..rows.end,
151            best_col..=*options.end(),
152        );
153    }
154
155    solve(
156        arbitrary,
157        convex,
158        &mut result,
159        0..len,
160        0..=arbitrary.len() - 1,
161    );
162    result
163}
164
165#[derive(Clone, Copy, Eq, PartialEq)]
166enum MatrixValue<T> {
167    Finite(T),
168    Infinite,
169}
170
171impl<T> Ord for MatrixValue<T>
172where
173    T: Ord,
174{
175    fn cmp(&self, other: &Self) -> Ordering {
176        match (self, other) {
177            (Self::Finite(left), Self::Finite(right)) => left.cmp(right),
178            (Self::Finite(_), Self::Infinite) => Ordering::Less,
179            (Self::Infinite, Self::Finite(_)) => Ordering::Greater,
180            (Self::Infinite, Self::Infinite) => Ordering::Equal,
181        }
182    }
183}
184
185impl<T> PartialOrd for MatrixValue<T>
186where
187    T: Ord,
188{
189    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
190        Some(self.cmp(other))
191    }
192}
193
194fn smawk<T, F>(rows: usize, cols: usize, cost: &F) -> Vec<usize>
195where
196    T: Ord,
197    F: Fn(usize, usize) -> T,
198{
199    fn solve<T, F>(rows: &[usize], cols: &[usize], cost: &F, argmins: &mut [usize])
200    where
201        T: Ord,
202        F: Fn(usize, usize) -> T,
203    {
204        if rows.is_empty() {
205            return;
206        }
207        let mut reduced = Vec::with_capacity(rows.len().min(cols.len()));
208        for &col in cols {
209            while let Some(&previous) = reduced.last() {
210                let row = rows[reduced.len() - 1];
211                if cost(row, col) <= cost(row, previous) {
212                    reduced.pop();
213                } else {
214                    break;
215                }
216            }
217            if reduced.len() < rows.len() {
218                reduced.push(col);
219            }
220        }
221        let odd_rows: Vec<_> = rows.iter().copied().skip(1).step_by(2).collect();
222        solve(&odd_rows, &reduced, cost, argmins);
223        let mut lower = 0;
224        for row_position in (0..rows.len()).step_by(2) {
225            let upper = if row_position + 1 < rows.len() {
226                let target = argmins[rows[row_position + 1]];
227                lower
228                    + reduced[lower..]
229                        .iter()
230                        .position(|&col| col == target)
231                        .expect("SMAWK odd-row minimum must remain in the reduced columns")
232            } else {
233                reduced.len() - 1
234            };
235            let row = rows[row_position];
236            let mut best = lower;
237            for position in lower + 1..=upper {
238                if cost(row, reduced[position]) <= cost(row, reduced[best]) {
239                    best = position;
240                }
241            }
242            argmins[row] = reduced[best];
243            lower = upper;
244        }
245    }
246
247    let row_indices: Vec<_> = (0..rows).collect();
248    let col_indices: Vec<_> = (0..cols).collect();
249    let mut argmins = vec![0; rows];
250    solve(&row_indices, &col_indices, cost, &mut argmins);
251    argmins
252}
253
254/// Computes convolution when one input is convex using SMAWK in `O(n + m)`.
255///
256/// # Panics
257///
258/// Panics unless both inputs are finite and at least one is convex.
259pub fn min_plus_convolution_convex_smawk<T>(a: &[T], b: &[T]) -> Vec<T>
260where
261    T: Signed,
262{
263    let len = output_len(a.len(), b.len());
264    if len == 0 {
265        return Vec::new();
266    }
267    assert_finite(a);
268    assert_finite(b);
269    let (arbitrary, convex) = orient_one_convex(a, b);
270    convex_smawk(arbitrary, convex)
271}
272
273pub(super) fn convex_smawk<T>(arbitrary: &[T], convex: &[T]) -> Vec<T>
274where
275    T: Signed,
276{
277    let len = output_len(arbitrary.len(), convex.len());
278    let cost = |row: usize, col: usize| {
279        row.checked_sub(col)
280            .filter(|&index| index < convex.len())
281            .map_or(MatrixValue::Infinite, |index| {
282                MatrixValue::Finite(arbitrary[col] + convex[index])
283            })
284    };
285    let argmins = smawk(len, arbitrary.len(), &cost);
286    argmins
287        .into_iter()
288        .enumerate()
289        .map(|(row, col)| {
290            row.checked_sub(col)
291                .filter(|&index| index < convex.len())
292                .map(|index| arbitrary[col] + convex[index])
293                .expect("SMAWK minimum must be a valid convolution entry")
294        })
295        .collect()
296}
297
298#[cfg(test)]
299mod tests {
300    use super::*;
301    use crate::{
302        math::min_plus_convolution::min_plus_convolution_naive,
303        tools::{
304            Xorshift,
305            testutil::{exhaustive_sequences, sample_usize},
306        },
307    };
308
309    #[test]
310    fn test_convex_algorithms() {
311        let mut rng = Xorshift::default();
312        let inputs: Vec<_> = exhaustive_sequences([-2i64, 0, 3], 0..=6).collect();
313        let convex: Vec<_> = inputs.iter().filter(|input| is_convex(input)).collect();
314        let exhaustive = inputs
315            .iter()
316            .flat_map(|a| convex.iter().map(move |&b| (a.clone(), b.clone())));
317        let mut random = Vec::new();
318        // Check every pair of small lengths, then cross allocation boundaries.
319        let mut lengths: Vec<_> = (0..=32)
320            .flat_map(|n| (0..=32).map(move |m| (n, m)))
321            .collect();
322        for n in sample_usize(&mut rng, 32, 0..=512, 1000) {
323            lengths.push((n, rng.random(0usize..=512)));
324        }
325        for (n, m) in lengths {
326            let arbitrary: Vec<_> = rng.random_iter(-1000..=1000).take(n).collect();
327            let mut slopes: Vec<_> = rng
328                .random_iter(-100i64..=100)
329                .take(m.saturating_sub(1))
330                .collect();
331            slopes.sort_unstable();
332            let mut structured = Vec::new();
333            if m != 0 {
334                structured.push(rng.random(-1000..=1000));
335            }
336            for slope in slopes {
337                structured.push(structured.last().unwrap() + slope);
338            }
339            random.push((arbitrary, structured.clone()));
340            let mut other: Vec<_> = rng
341                .random_iter(-100i64..=100)
342                .take(n.saturating_sub(1))
343                .collect();
344            other.sort_unstable();
345            let mut convex = Vec::new();
346            if n != 0 {
347                convex.push(rng.random(-1000..=1000));
348            }
349            for slope in other {
350                convex.push(convex.last().unwrap() + slope);
351            }
352            random.push((convex, structured));
353        }
354        for (a, b) in exhaustive.chain(random) {
355            let expected = min_plus_convolution_naive(&a, &b);
356            assert_eq!(
357                min_plus_convolution_convex_divide_and_conquer(&a, &b),
358                expected,
359                "a={a:?}, b={b:?}"
360            );
361            assert_eq!(
362                min_plus_convolution_convex_smawk(&a, &b),
363                expected,
364                "a={a:?}, b={b:?}"
365            );
366            if is_convex(&a) {
367                assert_eq!(
368                    min_plus_convolution_convex_merge(&a, &b),
369                    expected,
370                    "a={a:?}, b={b:?}"
371                );
372            }
373        }
374    }
375}