competitive/math/min_plus_convolution/
convex.rs1use 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
29pub 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
88pub 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
254pub 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 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}