Skip to main content

competitive/math/min_plus_convolution/
squared_distance.rs

1use super::Signed;
2
3fn index_as_value<T>(index: usize) -> T
4where
5    T: Signed + TryFrom<usize>,
6{
7    T::try_from(index)
8        .ok()
9        .expect("squared-distance index must fit the value type")
10}
11
12fn first_position_of_new_source<T>(values: &[T], previous_source: usize, new_source: usize) -> T
13where
14    T: Signed + TryFrom<usize>,
15{
16    let previous_index: T = index_as_value(previous_source);
17    let new_index: T = index_as_value(new_source);
18    let numerator = values[new_source] + new_index * new_index
19        - values[previous_source]
20        - previous_index * previous_index;
21    let denominator = (new_index - previous_index) + (new_index - previous_index);
22    numerator.div_euclid(denominator)
23        + if numerator.rem_euclid(denominator).is_zero() {
24            T::zero()
25        } else {
26            T::one()
27        }
28}
29
30/// Computes min-plus convolution with squared distance in linear time.
31///
32/// The value at `p` is `min_q(values[q] + (p - q)^2)`. `T::maximum()`
33/// represents an unreachable source.
34///
35/// # Panics
36///
37/// Panics if an index cannot be represented by `T`. Arithmetic overflow is
38/// the caller's responsibility.
39pub fn min_plus_convolution_with_squared_distance<T>(values: &[T]) -> Vec<T>
40where
41    T: Signed + TryFrom<usize>,
42{
43    if values.is_empty() {
44        return Vec::new();
45    }
46    let mut sources = Vec::with_capacity(values.len());
47    let mut first_positions = Vec::with_capacity(values.len());
48    for (new_source, &value) in values.iter().enumerate() {
49        if value.is_maximum() {
50            continue;
51        }
52        let mut first_position = T::zero();
53        while let Some(&previous_source) = sources.last() {
54            first_position = first_position_of_new_source(values, previous_source, new_source);
55            if first_position > first_positions[first_positions.len() - 1] {
56                break;
57            }
58            sources.pop();
59            first_positions.pop();
60        }
61        if sources.is_empty() {
62            first_position = T::zero();
63        }
64        if first_position < index_as_value(values.len()) {
65            sources.push(new_source);
66            first_positions.push(first_position.max(T::zero()));
67        }
68    }
69    if sources.is_empty() {
70        return vec![T::maximum(); values.len()];
71    }
72    let mut active_source_index = 0;
73    let mut result = Vec::with_capacity(values.len());
74    for position in 0..values.len() {
75        while active_source_index + 1 < sources.len()
76            && first_positions[active_source_index + 1] <= index_as_value(position)
77        {
78            active_source_index += 1;
79        }
80        let source = sources[active_source_index];
81        let distance: T = index_as_value(source.abs_diff(position));
82        result.push(values[source] + distance * distance);
83    }
84    result
85}
86
87#[cfg(test)]
88mod tests {
89    use super::*;
90    use crate::tools::Xorshift;
91
92    #[test]
93    fn test_min_plus_convolution_with_squared_distance() {
94        let mut rng = Xorshift::default();
95        for len in 0..=32 {
96            for case in 0..32 {
97                let values: Vec<_> = (0..len)
98                    .map(|_| {
99                        if case == 0 || rng.random(0_u64..5) == 0 {
100                            i64::MAX
101                        } else {
102                            rng.random(-50_i64..=50)
103                        }
104                    })
105                    .collect();
106                let expected: Vec<_> = (0..len)
107                    .map(|position| {
108                        values
109                            .iter()
110                            .copied()
111                            .enumerate()
112                            .filter(|(_, value)| *value != i64::MAX)
113                            .map(|(source, value)| {
114                                let distance = source.abs_diff(position) as i64;
115                                value + distance * distance
116                            })
117                            .min()
118                            .unwrap_or(i64::MAX)
119                    })
120                    .collect();
121                assert_eq!(
122                    min_plus_convolution_with_squared_distance(&values),
123                    expected
124                );
125            }
126        }
127    }
128}