competitive/math/
subset_convolve.rs1use 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}