Skip to main content

competitive/math/
subset_convolve.rs

1use super::{ConvolveSteps, Invertible, Ring, advise_huge_pages};
2use std::marker::PhantomData;
3
4pub struct SubsetConvolve<M> {
5    _marker: PhantomData<fn() -> M>,
6}
7
8impl<R> SubsetConvolve<R>
9where
10    R: Ring<T: PartialEq, Additive: Invertible>,
11{
12    fn ranked(t: Vec<R::T>, len: usize) -> (Vec<R::T>, usize) {
13        let width = len.trailing_zeros() as usize + 1;
14        let mut ranked = Vec::with_capacity(len * width);
15        advise_huge_pages(&mut ranked);
16        ranked.resize(len * width, R::zero());
17        for (i, value) in t.into_iter().enumerate() {
18            ranked[i * width + i.count_ones() as usize] = value;
19        }
20        (ranked, width)
21    }
22
23    fn diagonal(ranked: Vec<R::T>, width: usize) -> Vec<R::T> {
24        ranked
25            .chunks_exact(width)
26            .enumerate()
27            .map(|(i, row)| row[i.count_ones() as usize].clone())
28            .collect()
29    }
30
31    #[inline]
32    fn multiply_row(
33        x: &[R::T],
34        y: &[R::T],
35        right: &mut [R::T],
36        output: &mut [R::T],
37        rank: usize,
38    ) -> usize {
39        for (right, y) in right[..=rank].iter_mut().zip(y[..=rank].iter().rev()) {
40            right.clone_from(y);
41        }
42        let end = (rank * 2).min(x.len() - 1);
43        for (degree, output) in output.iter_mut().enumerate().take(end + 1).skip(rank) {
44            let first = degree - rank;
45            *output = R::dot_product(&x[first..=rank], &right[..=rank - first]);
46        }
47        end
48    }
49}
50
51impl<R> ConvolveSteps for SubsetConvolve<R>
52where
53    R: Ring<T: PartialEq, Additive: Invertible>,
54{
55    type T = Vec<R::T>;
56    type F = (Vec<R::T>, usize);
57
58    fn length(t: &Self::T) -> usize {
59        t.len()
60    }
61
62    fn transform(t: Self::T, len: usize) -> Self::F {
63        let (mut f, width) = Self::ranked(t, len);
64        let k = width - 1;
65        for bit in 0..k {
66            let half = 1 << bit;
67            for base in (0..len).step_by(half * 2) {
68                for lower in base..base + half {
69                    let upper = lower + half;
70                    let ranks = lower.count_ones() as usize + 1;
71                    let (lower_rows, upper_rows) = f.split_at_mut(upper * width);
72                    let lower_row = &lower_rows[lower * width..lower * width + ranks];
73                    let upper_row = &mut upper_rows[..ranks];
74                    for (upper, lower) in upper_row.iter_mut().zip(lower_row) {
75                        R::add_assign(upper, lower);
76                    }
77                }
78            }
79        }
80        (f, width)
81    }
82
83    fn inverse_transform((mut f, width): Self::F, len: usize) -> Self::T {
84        let k = width - 1;
85        for bit in 0..k {
86            let half = 1 << bit;
87            for base in (0..len).step_by(half * 2) {
88                for lower in base..base + half {
89                    let upper = lower + half;
90                    let rank = lower.count_ones() as usize;
91                    let (lower_rows, upper_rows) = f.split_at_mut(upper * width);
92                    let lower_row = &lower_rows[lower * width + rank..lower * width + width];
93                    let upper_row = &mut upper_rows[rank..width];
94                    for (upper, lower) in upper_row.iter_mut().zip(lower_row) {
95                        R::sub_assign(upper, lower);
96                    }
97                }
98            }
99        }
100        Self::diagonal(f, width)
101    }
102
103    fn multiply(f: &mut Self::F, g: &Self::F) {
104        let (f, width) = f;
105        let (g, _) = g;
106        let mut right = vec![R::zero(); *width];
107        let mut output = vec![R::zero(); *width];
108        for (i, f) in f.chunks_exact_mut(*width).enumerate() {
109            let rank = i.count_ones() as usize;
110            let g = &g[i * *width..(i + 1) * *width];
111            let end = Self::multiply_row(f, g, &mut right, &mut output, rank);
112            f[rank..=end].clone_from_slice(&output[rank..=end]);
113        }
114    }
115
116    fn convolve(a: Self::T, b: Self::T) -> Self::T {
117        assert_eq!(a.len(), b.len());
118        let len = a.len();
119        let same = a == b;
120        let (mut x, width) = Self::ranked(a, len);
121        let (mut y, _) = if same {
122            (x.clone(), width)
123        } else {
124            Self::ranked(b, len)
125        };
126        let mut right = vec![R::zero(); width];
127        let mut output = vec![R::zero(); width];
128        for i in 0..len {
129            for bit in (0..(i | len).trailing_zeros() as usize).rev() {
130                let half = width << bit;
131                let start = i * width;
132                let (lower, upper) = x[start..start + half * 2].split_at_mut(half);
133                for (upper, lower) in upper.iter_mut().zip(lower) {
134                    R::add_assign(upper, lower);
135                }
136                let (lower, upper) = y[start..start + half * 2].split_at_mut(half);
137                for (upper, lower) in upper.iter_mut().zip(lower) {
138                    R::add_assign(upper, lower);
139                }
140            }
141
142            let rank = i.count_ones() as usize;
143            let start = i * width;
144            let x_row = &x[start..start + width];
145            let y_row = &y[start..start + width];
146            output.fill(R::zero());
147            Self::multiply_row(x_row, y_row, &mut right, &mut output, rank);
148            x[start..start + width].clone_from_slice(&output);
149
150            for bit in 0..i.trailing_ones() as usize {
151                let end = (i + 1) * width;
152                let half = width << bit;
153                let (lower, upper) = x[end - half * 2..end].split_at_mut(half);
154                for (upper, lower) in upper.iter_mut().zip(lower) {
155                    R::sub_assign(upper, lower);
156                }
157            }
158        }
159        Self::diagonal(x, width)
160    }
161}
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166    use crate::{algebra::AddMulOperation, rand, tools::Xorshift};
167
168    const A: i64 = 100_000;
169
170    #[test]
171    fn test_subset_convolve() {
172        let mut rng = Xorshift::default();
173
174        for k in 0..12 {
175            let n = 1 << k;
176            rand!(rng, f: [-A..A; n], g: [-A..A; n]);
177            let mut h = vec![0i64; n];
178            for i in 0..n {
179                for j in 0..n {
180                    if i & j == 0 {
181                        h[i | j] += f[i] * g[j];
182                    }
183                }
184            }
185            let mut transformed = SubsetConvolve::<AddMulOperation<i64>>::transform(f.clone(), n);
186            let other = SubsetConvolve::<AddMulOperation<i64>>::transform(g.clone(), n);
187            SubsetConvolve::<AddMulOperation<i64>>::multiply(&mut transformed, &other);
188            let j = SubsetConvolve::<AddMulOperation<i64>>::inverse_transform(transformed, n);
189            let i = SubsetConvolve::<AddMulOperation<i64>>::convolve(f, g);
190            assert_eq!(h, i);
191            assert_eq!(h, j);
192        }
193    }
194}