Skip to main content

competitive/math/min_plus_convolution/
mod.rs

1//! Exact min-plus convolution algorithms for structured integer sequences.
2//!
3//! `T::maximum()` represents positive infinity. Callers must choose a signed
4//! integer type that can represent every finite result and every intermediate
5//! arithmetic expression used by the selected algorithm.
6
7use super::{Convolve998244353, ConvolveSteps, Signed, montgomery::MInt998244353};
8
9pub use self::concave::{min_plus_convolution_concave_both, min_plus_convolution_concave_envelope};
10pub use self::convex::{
11    min_plus_convolution_convex_divide_and_conquer, min_plus_convolution_convex_merge,
12    min_plus_convolution_convex_smawk,
13};
14pub use self::monotone::min_plus_convolution_monotone_runs;
15pub use self::near_convex::min_plus_convolution_near_convex_scan;
16pub use self::piecewise_linear::{
17    min_plus_convolution_linear, min_plus_convolution_piecewise_linear,
18};
19pub use self::selector::min_plus_convolution;
20pub use self::squared_distance::min_plus_convolution_with_squared_distance;
21
22mod concave;
23mod convex;
24mod monotone;
25mod near_convex;
26mod piecewise_linear;
27mod selector;
28mod squared_distance;
29
30pub(super) fn output_len(a_len: usize, b_len: usize) -> usize {
31    if a_len == 0 || b_len == 0 {
32        return 0;
33    }
34    a_len
35        .checked_add(b_len)
36        .and_then(|len| len.checked_sub(1))
37        .expect("min-plus convolution output length must fit usize")
38}
39
40pub(super) fn assert_finite<T>(values: &[T])
41where
42    T: Signed,
43{
44    assert!(
45        !values.iter().any(T::is_maximum),
46        "min-plus convolution algorithm requires finite input values"
47    );
48}
49
50/// Computes min-plus convolution by enumerating all input pairs.
51///
52/// # Panics
53///
54/// Panics if the output length does not fit [`usize`]. Arithmetic overflow is
55/// the caller's responsibility.
56pub fn min_plus_convolution_naive<T>(a: &[T], b: &[T]) -> Vec<T>
57where
58    T: Signed,
59{
60    let len = output_len(a.len(), b.len());
61    if len == 0 {
62        return Vec::new();
63    }
64    let (outer, inner) = if a.len() <= b.len() { (a, b) } else { (b, a) };
65    let inner_is_finite = !inner.iter().any(T::is_maximum);
66    let mut result = vec![T::maximum(); len];
67    for (index, &left) in outer.iter().enumerate() {
68        if left.is_maximum() {
69            continue;
70        }
71        let output = &mut result[index..index + inner.len()];
72        if inner_is_finite {
73            for (slot, &right) in output.iter_mut().zip(inner) {
74                *slot = (*slot).min(left + right);
75            }
76        } else {
77            for (slot, &right) in output.iter_mut().zip(inner) {
78                if !right.is_maximum() {
79                    *slot = (*slot).min(left + right);
80                }
81            }
82        }
83    }
84    result
85}
86
87/// Computes min-plus convolution by enumerating finite input pairs only.
88///
89/// If the inputs have `s_a` and `s_b` finite values, the running time is
90/// `O(s_a * s_b + n + m)`.
91///
92/// # Panics
93///
94/// Panics if the output length does not fit [`usize`]. Arithmetic overflow is
95/// the caller's responsibility.
96pub fn min_plus_convolution_sparse<T>(a: &[T], b: &[T]) -> Vec<T>
97where
98    T: Signed,
99{
100    let len = output_len(a.len(), b.len());
101    if len == 0 {
102        return Vec::new();
103    }
104    let a: Vec<_> = a
105        .iter()
106        .copied()
107        .enumerate()
108        .filter(|(_, value)| !value.is_maximum())
109        .collect();
110    let b: Vec<_> = b
111        .iter()
112        .copied()
113        .enumerate()
114        .filter(|(_, value)| !value.is_maximum())
115        .collect();
116    sparse(a.iter().copied(), b.iter().copied(), len)
117}
118
119fn sparse<T>(
120    a: impl IntoIterator<Item = (usize, T)>,
121    b: impl IntoIterator<Item = (usize, T)> + Clone,
122    len: usize,
123) -> Vec<T>
124where
125    T: Signed,
126{
127    let mut result = vec![T::maximum(); len];
128    for (i, left) in a {
129        for (j, right) in b.clone() {
130            result[i + j] = result[i + j].min(left + right);
131        }
132    }
133    result
134}
135
136const MAX_NTT_SIZE: usize = 1 << 23;
137
138#[derive(Clone, Copy, Debug)]
139struct BoundedRequirements<T> {
140    a_min: T,
141    b_min: T,
142    base: usize,
143    transform_len: usize,
144}
145
146fn finite_extrema<T>(values: &[T]) -> Option<(T, T)>
147where
148    T: Signed,
149{
150    let mut finite = values.iter().copied().filter(|value| !value.is_maximum());
151    let first = finite.next()?;
152    Some(finite.fold((first, first), |(minimum, maximum), value| {
153        (minimum.min(value), maximum.max(value))
154    }))
155}
156
157fn bounded_transform_len(a_len: usize, b_len: usize, base: usize) -> Option<usize> {
158    let left_len = a_len.checked_mul(base)?;
159    let right_len = b_len.checked_mul(base)?;
160    let coefficient_len = left_len
161        .checked_add(right_len)
162        .and_then(|len| len.checked_sub(1))?;
163    let transform_len = coefficient_len.checked_next_power_of_two()?;
164    (transform_len <= MAX_NTT_SIZE).then_some(transform_len)
165}
166
167fn bounded_requirements_from_extrema<T>(
168    a_len: usize,
169    b_len: usize,
170    (a_min, a_max): (T, T),
171    (b_min, b_max): (T, T),
172) -> Option<BoundedRequirements<T>>
173where
174    T: Signed,
175    T::Unsigned: TryInto<usize>,
176{
177    let a_span = a_max.abs_diff(a_min).try_into().ok()?;
178    let b_span = b_max.abs_diff(b_min).try_into().ok()?;
179    let base = a_span
180        .checked_add(b_span)
181        .and_then(|span| span.checked_add(1))?;
182    let transform_len = bounded_transform_len(a_len, b_len, base)?;
183    Some(BoundedRequirements {
184        a_min,
185        b_min,
186        base,
187        transform_len,
188    })
189}
190
191/// Computes exact min-plus convolution for small integer value spans using NTT.
192///
193/// # Panics
194///
195/// Panics if the encoded range cannot be represented or requires a transform
196/// longer than `2^23`. Arithmetic overflow is the caller's responsibility.
197pub fn min_plus_convolution_bounded_ntt<T>(a: &[T], b: &[T]) -> Vec<T>
198where
199    T: Signed + TryFrom<usize>,
200    T::Unsigned: TryInto<usize>,
201{
202    let output_len = output_len(a.len(), b.len());
203    if output_len == 0 {
204        return Vec::new();
205    }
206    let (Some(a_extrema), Some(b_extrema)) = (finite_extrema(a), finite_extrema(b)) else {
207        return vec![T::maximum(); output_len];
208    };
209    let requirements = bounded_requirements_from_extrema(a.len(), b.len(), a_extrema, b_extrema)
210        .expect("bounded min-plus convolution encoding must fit the 2^23 NTT limit");
211    let mut left = vec![MInt998244353::from(0_u32); a.len() * requirements.base];
212    let mut right = vec![MInt998244353::from(0_u32); b.len() * requirements.base];
213    for (index, &value) in a.iter().enumerate() {
214        if !value.is_maximum() {
215            let normalized: usize = value
216                .abs_diff(requirements.a_min)
217                .try_into()
218                .ok()
219                .expect("bounded min-plus convolution value span must fit usize");
220            left[index * requirements.base + normalized] = MInt998244353::from(1_u32);
221        }
222    }
223    for (index, &value) in b.iter().enumerate() {
224        if !value.is_maximum() {
225            let normalized: usize = value
226                .abs_diff(requirements.b_min)
227                .try_into()
228                .ok()
229                .expect("bounded min-plus convolution value span must fit usize");
230            right[index * requirements.base + normalized] = MInt998244353::from(1_u32);
231        }
232    }
233    let coefficients = Convolve998244353::convolve(left, right);
234    let encoded_len = output_len
235        .checked_mul(requirements.base)
236        .expect("bounded min-plus convolution encoded length must fit usize");
237    let mut result = Vec::with_capacity(output_len);
238    for chunk in coefficients[..encoded_len].chunks_exact(requirements.base) {
239        let value = if let Some(normalized) = chunk.iter().position(|&value| u32::from(value) != 0)
240        {
241            requirements.a_min
242                + requirements.b_min
243                + T::try_from(normalized)
244                    .ok()
245                    .expect("bounded min-plus convolution value must fit the output type")
246        } else {
247            T::maximum()
248        };
249        result.push(value);
250    }
251    result
252}
253
254#[cfg(test)]
255mod tests {
256    use super::*;
257    use crate::tools::Xorshift;
258
259    #[test]
260    fn test_min_plus_convolution() {
261        let inf = i64::MAX;
262        let mut rng = Xorshift::default();
263        for a_len in 0..=8 {
264            for b_len in 0..=8 {
265                for case in 0..32 {
266                    let mut values = |len, all_infinite| {
267                        let mut values = Vec::with_capacity(len);
268                        for _ in 0..len {
269                            values.push(if all_infinite || rng.random(0_u64..5) == 0 {
270                                inf
271                            } else {
272                                rng.random(-4_i64..=4)
273                            });
274                        }
275                        values
276                    };
277                    let a = values(a_len, case <= 1);
278                    let b = values(b_len, case == 0 || case == 2);
279                    let mut expected = if a.is_empty() || b.is_empty() {
280                        Vec::new()
281                    } else {
282                        vec![inf; a.len() + b.len() - 1]
283                    };
284                    for (i, &left) in a.iter().enumerate() {
285                        if left == inf {
286                            continue;
287                        }
288                        for (j, &right) in b.iter().enumerate() {
289                            if right != inf {
290                                expected[i + j] = expected[i + j].min(left + right);
291                            }
292                        }
293                    }
294                    assert_eq!(min_plus_convolution_naive(&a, &b), expected);
295                    assert_eq!(min_plus_convolution_sparse(&a, &b), expected);
296                    assert_eq!(min_plus_convolution_bounded_ntt(&a, &b), expected);
297                    assert_eq!(min_plus_convolution(&a, &b), expected);
298                }
299            }
300        }
301    }
302
303    #[test]
304    fn test_automatic_selection() {
305        let mut rng = Xorshift::default();
306        for case in 0..30 {
307            let a_len = rng.random(520..=640);
308            let b_len = rng.random(520..=640);
309            let (a, b): (Vec<i64>, Vec<i64>) = match case % 5 {
310                0 => {
311                    let mut a = vec![i64::MAX; a_len];
312                    let mut b = vec![i64::MAX; b_len];
313                    let a_prefix = rng.random(1..=8);
314                    let b_prefix = rng.random(1..=8);
315                    for value in &mut a[..a_prefix] {
316                        *value = rng.random(-1_000_i64..=1_000);
317                    }
318                    for value in &mut b[..b_prefix] {
319                        *value = rng.random(-1_000_i64..=1_000);
320                    }
321                    a[a_len - 1] = rng.random(-1_000_i64..=1_000);
322                    b[b_len - 1] = rng.random(-1_000_i64..=1_000);
323                    for _ in 0..8 {
324                        let i = rng.random(0..a_len);
325                        let j = rng.random(0..b_len);
326                        a[i] = rng.random(-1_000_i64..=1_000);
327                        b[j] = rng.random(-1_000_i64..=1_000);
328                    }
329                    (a, b)
330                }
331                1 => (
332                    rng.random_iter(-2_i64..=2).take(a_len).collect(),
333                    rng.random_iter(-2_i64..=2).take(b_len).collect(),
334                ),
335                2 | 3 => {
336                    let len = if case % 5 == 2 { a_len } else { b_len };
337                    let mut slope = rng.random(-20_i64..=20);
338                    let mut value = rng.random(-1_000_i64..=1_000);
339                    let mut structured = Vec::with_capacity(len);
340                    for _ in 0..len {
341                        structured.push(value);
342                        value += slope;
343                        slope += if case % 5 == 2 {
344                            rng.random(0_i64..=3)
345                        } else {
346                            -rng.random(0_i64..=3)
347                        };
348                    }
349                    if case % 5 == 2 {
350                        (
351                            structured,
352                            rng.random_iter(-1_000_i64..=1_000).take(b_len).collect(),
353                        )
354                    } else {
355                        (
356                            rng.random_iter(-1_000_i64..=1_000).take(a_len).collect(),
357                            structured,
358                        )
359                    }
360                }
361                _ => {
362                    let mut a = Vec::with_capacity(a_len);
363                    let mut b = Vec::with_capacity(b_len);
364                    let mut left = rng.random(-1_000_i64..=1_000);
365                    let mut right = rng.random(-1_000_i64..=1_000);
366                    for _ in 0..a_len {
367                        a.push(left);
368                        left += rng.random(0_i64..=3);
369                    }
370                    for _ in 0..b_len {
371                        b.push(right);
372                        right += rng.random(0_i64..=3);
373                    }
374                    (a, b)
375                }
376            };
377            let expected = min_plus_convolution_naive(&a, &b);
378            assert_eq!(min_plus_convolution(&a, &b), expected);
379            assert_eq!(min_plus_convolution(&b, &a), expected);
380        }
381    }
382}