Skip to main content

competitive/math/
bitwisexor_convolve.rs

1use 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
27/// `EXACT_DIVISION` normalizes with division instead of a multiplicative inverse.
28pub 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}