competitive/math/
bitwisexor_convolve.rs1use super::{ConvolveSteps, Field, Group, Invertible, bitwise_transform};
2use std::{fmt::Debug, marker::PhantomData};
3
4trait FromLength<const EXACT_DIVISION: bool> {
5 fn from_length(len: usize) -> Self;
6}
7
8impl<T> FromLength<false> for T
9where
10 T: From<usize>,
11{
12 fn from_length(len: usize) -> Self {
13 T::from(len)
14 }
15}
16
17impl<T> FromLength<true> for T
18where
19 T: TryFrom<usize>,
20 T::Error: Debug,
21{
22 fn from_length(len: usize) -> Self {
23 T::try_from(len).unwrap()
24 }
25}
26
27pub struct BitwisexorConvolve<M, const EXACT_DIVISION: bool = false> {
29 _marker: PhantomData<fn() -> M>,
30}
31
32impl<G, const EXACT_DIVISION: bool> BitwisexorConvolve<G, EXACT_DIVISION>
33where
34 G: Group,
35{
36 pub fn hadamard_transform(f: &mut [G::T]) {
37 bitwise_transform(f, |x, y| {
38 let t = G::operate(x, y);
39 *y = G::rinv_operate(x, y);
40 *x = t;
41 });
42 }
43}
44
45impl<R, const EXACT_DIVISION: bool> ConvolveSteps for BitwisexorConvolve<R, EXACT_DIVISION>
46where
47 R: Field<
48 T: PartialEq + FromLength<EXACT_DIVISION>,
49 Additive: Invertible,
50 Multiplicative: Invertible,
51 >,
52{
53 type T = Vec<R::T>;
54 type F = Vec<R::T>;
55
56 fn length(t: &Self::T) -> usize {
57 t.len()
58 }
59
60 fn transform(mut t: Self::T, _len: usize) -> Self::F {
61 BitwisexorConvolve::<R::Additive, EXACT_DIVISION>::hadamard_transform(&mut t);
62 t
63 }
64
65 fn inverse_transform(mut f: Self::F, len: usize) -> Self::T {
66 BitwisexorConvolve::<R::Additive, EXACT_DIVISION>::hadamard_transform(&mut f);
67 let len = R::T::from_length(len);
68 if EXACT_DIVISION {
69 for f in &mut f {
70 *f = R::div(f, &len);
71 }
72 } else if !f.is_empty() {
73 let inv_len = R::inv(&len);
74 for f in &mut f {
75 *f = R::mul(f, &inv_len);
76 }
77 }
78 f
79 }
80
81 fn multiply(f: &mut Self::F, g: &Self::F) {
82 for (f, g) in f.iter_mut().zip(g) {
83 *f = R::mul(f, g);
84 }
85 }
86
87 fn convolve(a: Self::T, b: Self::T) -> Self::T {
88 assert_eq!(a.len(), b.len());
89 let len = a.len();
90 let same = a == b;
91 let mut a = Self::transform(a, len);
92 if same {
93 for a in a.iter_mut() {
94 *a = R::mul(a, a);
95 }
96 } else {
97 let b = Self::transform(b, len);
98 Self::multiply(&mut a, &b);
99 }
100 Self::inverse_transform(a, len)
101 }
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use crate::{algebra::AddMulOperation, rand, tools::Xorshift};
108
109 const A: i64 = 100_000;
110
111 #[test]
112 fn test_bitwisexor_convolve() {
113 let mut rng = Xorshift::default();
114
115 for k in 0..12 {
116 let n = 1 << k;
117 rand!(rng, f: [-A..A; n], g: [-A..A; n]);
118 let mut h = vec![0i64; n];
119 for i in 0..n {
120 for j in 0..n {
121 h[i ^ j] += f[i] * g[j];
122 }
123 }
124 let i = BitwisexorConvolve::<AddMulOperation<i64>, true>::convolve(f, g);
125 assert_eq!(h, i);
126 }
127 }
128}