Skip to main content

simd_backend

Function simd_backend 

Source
pub fn simd_backend() -> SimdBackend
Examples found in repository?
crates/competitive/src/math/number_theoretic_transform.rs (line 31)
20fn batch_ntt_simd_backend(len: usize, width: usize) -> SimdBackend {
21    // power_projection uses power-of-two widths; SIMD scalar tails regress for some other widths.
22    if width < 8 || !width.is_power_of_two() || len < width * 4 {
23        SimdBackend::Scalar
24    } else if width == 8 {
25        if is_x86_feature_detected!("avx2") {
26            SimdBackend::Avx2
27        } else {
28            SimdBackend::Scalar
29        }
30    } else {
31        simd_backend()
32    }
33}
34
35pub struct Convolve<M>(PhantomData<fn() -> M>);
36pub type Convolve998244353 = Convolve<Modulo998244353>;
37/// Raw transforms require each integer coefficient reconstructed by CRT to be below
38/// the product of the three NTT moduli. `convolve` splits products exceeding this bound.
39pub type MIntConvolve<M> = Convolve<(M, (Modulo167772161, Modulo469762049, Modulo754974721))>;
40/// Convolution modulo 2^64. Multiply only freshly transformed operands; reconstruct
41/// and transform again before multiplying another factor.
42pub type U64Convolve = Convolve<(u64, (Modulo167772161, Modulo469762049, Modulo754974721))>;
43
44macro_rules! impl_ntt_modulus {
45    ($([$name:ident, $g:expr]),*) => {
46        $(
47            impl Montgomery32NttModulus for $name {}
48        )*
49    };
50}
51impl_ntt_modulus!(
52    [Modulo167772161, 3],
53    [Modulo469762049, 3],
54    [Modulo754974721, 11],
55    [Modulo998244353, 3]
56);
57
58const fn reduce(z: u64, p: u32, r: u32) -> u32 {
59    let mut z = ((z + r.wrapping_mul(z as u32) as u64 * p as u64) >> 32) as u32;
60    if z >= p {
61        z -= p;
62    }
63    z
64}
65const fn mod_mul(x: u32, y: u32, p: u32, r: u32) -> u32 {
66    reduce(x as u64 * y as u64, p, r)
67}
68const fn mod_pow(mut x: u32, mut y: u32, p: u32, r: u32, mut z: u32) -> u32 {
69    while y > 0 {
70        if y & 1 == 1 {
71            z = mod_mul(z, x, p, r);
72        }
73        x = mod_mul(x, x, p, r);
74        y >>= 1;
75    }
76    z
77}
78
79pub trait Montgomery32NttModulus: Sized + MontgomeryReduction32 {
80    const PRIMITIVE_ROOT: u32 = {
81        let mut g = 3u32;
82        loop {
83            let mut ok = true;
84            let mut d = 1u32;
85            while d * d < Self::MOD {
86                if (Self::MOD - 1) % d == 0 {
87                    let ds = [d, (Self::MOD - 1) / d];
88                    let mut i = 0;
89                    while i < 2 {
90                        ok &= ds[i] == Self::MOD - 1
91                            || mod_pow(
92                                reduce(g as u64 * Self::N2 as u64, Self::MOD, Self::R),
93                                ds[i],
94                                Self::MOD,
95                                Self::R,
96                                Self::N1,
97                            ) != Self::N1;
98                        i += 1;
99                    }
100                }
101                d += 1;
102            }
103            if ok {
104                break;
105            }
106            g += 2;
107        }
108        g
109    };
110    const RANK: u32 = (Self::MOD - 1).trailing_zeros();
111    const INFO: NttInfo = NttInfo::new::<Self>();
112}
113
114#[derive(Debug, PartialEq)]
115pub struct NttInfo {
116    root: [u32; 32],
117    inv_root: [u32; 32],
118    rate3: [u32; 32],
119    rate3_packed: [[u32; 8]; 32],
120    inv_rate3_packed: [[u32; 8]; 32],
121}
122impl NttInfo {
123    const fn new<M>() -> Self
124    where
125        M: Montgomery32NttModulus,
126    {
127        let mut root = [0; 32];
128        let mut inv_root = [0; 32];
129        let mut rate3_values = [0; 32];
130        let mut rate3_packed = [[0; 8]; 32];
131        let mut inv_rate3_packed = [[0; 8]; 32];
132        let rank = M::RANK as usize;
133
134        let g = reduce(M::PRIMITIVE_ROOT as u64 * M::N2 as u64, M::MOD, M::R);
135        root[rank] = mod_pow(g, (M::MOD - 1) >> rank, M::MOD, M::R, M::N1);
136        inv_root[rank] = mod_pow(root[rank], M::MOD - 2, M::MOD, M::R, M::N1);
137        let mut i = rank - 1;
138        loop {
139            root[i] = mod_mul(root[i + 1], root[i + 1], M::MOD, M::R);
140            inv_root[i] = mod_mul(inv_root[i + 1], inv_root[i + 1], M::MOD, M::R);
141            if i == 0 {
142                break;
143            }
144            i -= 1;
145        }
146
147        let (mut i, mut prod, mut inv_prod) = (0, M::N1, M::N1);
148        while i < rank - 2 {
149            let rate3 = mod_mul(root[i + 3], prod, M::MOD, M::R);
150            rate3_values[i] = rate3;
151            let rate3_2 = mod_mul(rate3, rate3, M::MOD, M::R);
152            let rate3_3 = mod_mul(rate3_2, rate3, M::MOD, M::R);
153            let inv_rate3 = mod_mul(inv_root[i + 3], inv_prod, M::MOD, M::R);
154            let inv_rate3_2 = mod_mul(inv_rate3, inv_rate3, M::MOD, M::R);
155            let inv_rate3_3 = mod_mul(inv_rate3_2, inv_rate3, M::MOD, M::R);
156            rate3_packed[i] = [
157                rate3.wrapping_mul(M::R),
158                rate3,
159                rate3_2.wrapping_mul(M::R),
160                rate3_2,
161                rate3_3.wrapping_mul(M::R),
162                rate3_3,
163                0,
164                0,
165            ];
166            inv_rate3_packed[i] = [
167                inv_rate3.wrapping_mul(M::R),
168                inv_rate3,
169                inv_rate3_2.wrapping_mul(M::R),
170                inv_rate3_2,
171                inv_rate3_3.wrapping_mul(M::R),
172                inv_rate3_3,
173                0,
174                0,
175            ];
176            prod = mod_mul(prod, inv_root[i + 3], M::MOD, M::R);
177            inv_prod = mod_mul(inv_prod, root[i + 3], M::MOD, M::R);
178            i += 1;
179        }
180
181        NttInfo {
182            root,
183            inv_root,
184            rate3: rate3_values,
185            rate3_packed,
186            inv_rate3_packed,
187        }
188    }
189}
190
191const LAZY_THRESHOLD: u32 = 1 << 30;
192
193#[inline]
194fn add_scalar<M>(x: u32, y: u32) -> u32
195where
196    M: Montgomery32NttModulus,
197{
198    let modulus = if M::MOD < LAZY_THRESHOLD {
199        M::MOD * 2
200    } else {
201        M::MOD
202    };
203    let sum = x + y;
204    if sum >= modulus { sum - modulus } else { sum }
205}
206
207#[inline]
208fn sub_scalar<M>(x: u32, y: u32) -> u32
209where
210    M: Montgomery32NttModulus,
211{
212    let modulus = if M::MOD < LAZY_THRESHOLD {
213        M::MOD * 2
214    } else {
215        M::MOD
216    };
217    if x < y { x + modulus - y } else { x - y }
218}
219
220#[inline]
221fn mul_scalar<M>(x: u32, y: u32) -> u32
222where
223    M: Montgomery32NttModulus,
224{
225    if M::MOD < LAZY_THRESHOLD {
226        let z = x as u64 * y as u64;
227        ((z + M::R.wrapping_mul(z as u32) as u64 * M::MOD as u64) >> 32) as u32
228    } else {
229        M::mod_mul(x, y)
230    }
231}
232
233fn ntt_scalar<M>(a: &mut [MInt<M>])
234where
235    M: Montgomery32NttModulus,
236{
237    ntt_batch_scalar(a, 1);
238}
239
240fn ntt_batch<M>(a: &mut [MInt<M>], width: usize)
241where
242    M: Montgomery32NttModulus,
243{
244    #[cfg(target_arch = "x86_64")]
245    {
246        match batch_ntt_simd_backend(a.len(), width) {
247            SimdBackend::Avx512 => {
248                // SAFETY: backend detection checked all required AVX-512 features.
249                unsafe { ntt_simd::ntt_batch_avx512(a, width) };
250                return;
251            }
252            SimdBackend::Avx2 => {
253                // SAFETY: backend detection checked AVX2.
254                unsafe { ntt_simd::ntt_batch_avx2::<_, false>(a, width) };
255                return;
256            }
257            SimdBackend::Scalar => {}
258        }
259    }
260    ntt_batch_scalar(a, width);
261}
262
263fn ntt_batch_scalar<M>(a: &mut [MInt<M>], width: usize)
264where
265    M: Montgomery32NttModulus,
266{
267    let n = a.len() / width;
268    if n <= 1 {
269        return;
270    }
271    let mut v = n / 2;
272    if n.trailing_zeros() & 1 == 1 {
273        let (l, r) = a.split_at_mut(v * width);
274        for (x0, x1) in l.iter_mut().zip(r) {
275            let a0 = *x0;
276            let a1 = *x1;
277            *x0 = a0 + a1;
278            *x1 = a0 - a1;
279        }
280        v >>= 1;
281    }
282    let imag = MInt::<M>::new_unchecked(M::INFO.root[2]);
283    while v > 1 {
284        let mut w1 = MInt::<M>::one();
285        let mut w2 = w1;
286        let mut w3 = w1;
287        for (s, a) in a.chunks_exact_mut((v << 1) * width).enumerate() {
288            let (l, r) = a.split_at_mut(v * width);
289            let (ll, lr) = l.split_at_mut((v >> 1) * width);
290            let (rl, rr) = r.split_at_mut((v >> 1) * width);
291            for (((x0, x1), x2), x3) in ll.iter_mut().zip(lr).zip(rl).zip(rr) {
292                let a0 = *x0;
293                let a1 = *x1 * w1;
294                let a2 = *x2 * w2;
295                let a3 = *x3 * w3;
296                let a0pa2 = a0 + a2;
297                let a0na2 = a0 - a2;
298                let a1pa3 = a1 + a3;
299                let a1na3imag = (a1 - a3) * imag;
300                *x0 = a0pa2 + a1pa3;
301                *x1 = a0pa2 - a1pa3;
302                *x2 = a0na2 + a1na3imag;
303                *x3 = a0na2 - a1na3imag;
304            }
305            let rate = &M::INFO.rate3_packed[s.trailing_ones() as usize];
306            w1 *= MInt::<M>::new_unchecked(rate[1]);
307            w2 *= MInt::<M>::new_unchecked(rate[3]);
308            w3 *= MInt::<M>::new_unchecked(rate[5]);
309        }
310        v >>= 2;
311    }
312}
313
314fn intt_scalar<M>(a: &mut [MInt<M>])
315where
316    M: Montgomery32NttModulus,
317{
318    intt_batch_scalar(a, 1);
319}
320
321fn intt_batch<M>(a: &mut [MInt<M>], width: usize)
322where
323    M: Montgomery32NttModulus,
324{
325    #[cfg(target_arch = "x86_64")]
326    {
327        match batch_ntt_simd_backend(a.len(), width) {
328            SimdBackend::Avx512 => {
329                // SAFETY: backend detection checked all required AVX-512 features.
330                unsafe { ntt_simd::intt_batch_avx512::<_, false>(a, width) };
331                return;
332            }
333            SimdBackend::Avx2 => {
334                // SAFETY: backend detection checked AVX2.
335                unsafe { ntt_simd::intt_batch_avx2::<_, false>(a, width) };
336                return;
337            }
338            SimdBackend::Scalar => {}
339        }
340    }
341    intt_batch_scalar(a, width);
342}
343
344fn intt_batch_scalar<M>(a: &mut [MInt<M>], width: usize)
345where
346    M: Montgomery32NttModulus,
347{
348    let n = a.len() / width;
349    if n <= 1 {
350        return;
351    }
352    // MInt is transparent over u32; lazy residues stay below 2 * MOD and are
353    // normalized before the typed slice is used again.
354    let a = unsafe { std::slice::from_raw_parts_mut(a.as_mut_ptr().cast::<u32>(), a.len()) };
355    let mut v = 1;
356    let limit = if n.trailing_zeros() & 1 == 1 {
357        n / 2
358    } else {
359        n
360    };
361    let iimag = M::INFO.inv_root[2];
362    while v < limit {
363        let mut w1 = M::N1;
364        let mut w2 = w1;
365        let mut w3 = w1;
366        for (s, a) in a.chunks_exact_mut((v << 2) * width).enumerate() {
367            let (l, r) = a.split_at_mut((v << 1) * width);
368            let (ll, lr) = l.split_at_mut(v * width);
369            let (rl, rr) = r.split_at_mut(v * width);
370            for (((x0, x1), x2), x3) in ll.iter_mut().zip(lr).zip(rl).zip(rr) {
371                let a0 = *x0;
372                let a1 = *x1;
373                let a2 = *x2;
374                let a3 = *x3;
375                let a0pa1 = add_scalar::<M>(a0, a1);
376                let a0na1 = sub_scalar::<M>(a0, a1);
377                let a2pa3 = add_scalar::<M>(a2, a3);
378                let a2na3iimag = mul_scalar::<M>(sub_scalar::<M>(a2, a3), iimag);
379                *x0 = add_scalar::<M>(a0pa1, a2pa3);
380                *x1 = mul_scalar::<M>(add_scalar::<M>(a0na1, a2na3iimag), w1);
381                *x2 = mul_scalar::<M>(sub_scalar::<M>(a0pa1, a2pa3), w2);
382                *x3 = mul_scalar::<M>(sub_scalar::<M>(a0na1, a2na3iimag), w3);
383            }
384            let rate = &M::INFO.inv_rate3_packed[s.trailing_ones() as usize];
385            w1 = M::mod_mul(w1, rate[1]);
386            w2 = M::mod_mul(w2, rate[3]);
387            w3 = M::mod_mul(w3, rate[5]);
388        }
389        v <<= 2;
390    }
391    if n.trailing_zeros() & 1 == 1 {
392        let (l, r) = a.split_at_mut(n / 2 * width);
393        for (x0, x1) in l.iter_mut().zip(r) {
394            let a0 = *x0;
395            let a1 = *x1;
396            *x0 = add_scalar::<M>(a0, a1);
397            *x1 = sub_scalar::<M>(a0, a1);
398        }
399    }
400    let inv = M::mod_inv(<M as MIntConvert<u32>>::from(n as u32));
401    for a in a {
402        *a = M::mod_mul(*a, inv);
403    }
404}
405
406fn ntt<M>(a: &mut [MInt<M>])
407where
408    M: Montgomery32NttModulus,
409{
410    #[cfg(target_arch = "x86_64")]
411    match simd_backend() {
412        SimdBackend::Avx512 => unsafe { ntt_simd::ntt_batch_avx512(a, 1) },
413        SimdBackend::Avx2 => unsafe { ntt_simd::ntt_batch_avx2::<_, true>(a, 1) },
414        SimdBackend::Scalar => ntt_scalar(a),
415    }
416    #[cfg(not(target_arch = "x86_64"))]
417    ntt_scalar(a);
418}
419
420fn intt<M>(a: &mut [MInt<M>])
421where
422    M: Montgomery32NttModulus,
423{
424    #[cfg(target_arch = "x86_64")]
425    match simd_backend() {
426        SimdBackend::Avx512 => unsafe { ntt_simd::intt_batch_avx512::<_, true>(a, 1) },
427        SimdBackend::Avx2 => unsafe { ntt_simd::intt_batch_avx2::<_, true>(a, 1) },
428        SimdBackend::Scalar => intt_scalar(a),
429    }
430    #[cfg(not(target_arch = "x86_64"))]
431    intt_scalar(a);
432}
More examples
Hide additional examples
crates/competitive/src/math/fast_fourier_transform.rs (line 628)
623fn fft_dispatch<const INVERSE: bool>(a: &mut [Complex<f64>]) {
624    RotateCache::ensure(a.len() / 2);
625    RotateCache::with(|cache| {
626        #[cfg(target_arch = "x86_64")]
627        if a.len() >= 16 {
628            match simd_backend() {
629                SimdBackend::Avx512 => {
630                    return unsafe { fft_avx512::<INVERSE>(a, cache) };
631                }
632                SimdBackend::Avx2 => return unsafe { fft_avx2::<INVERSE>(a, cache) },
633                SimdBackend::Scalar => {}
634            }
635        }
636        fft_kernel!(a, cache, INVERSE);
637    });
638}
crates/competitive/src/math/bit_matrix.rs (line 157)
155    fn eliminate(&mut self, cols: usize, full: bool, require_full_rank: bool) -> Vec<usize> {
156        #[cfg(target_arch = "x86_64")]
157        match simd_backend() {
158            // SAFETY: the dispatcher checks the required CPU features.
159            SimdBackend::Avx512 => {
160                return unsafe { self.eliminate_avx512(cols, full, require_full_rank) };
161            }
162            // SAFETY: the dispatcher checks AVX2 support.
163            SimdBackend::Avx2 => {
164                return unsafe { self.eliminate_avx2(cols, full, require_full_rank) };
165            }
166            SimdBackend::Scalar => {}
167        }
168        self.eliminate_impl(cols, full, require_full_rank)
169    }
170
171    // Inlined into each target-feature entry point to vectorize the row operations.
172    #[inline(always)]
173    fn eliminate_impl(&mut self, cols: usize, full: bool, require_full_rank: bool) -> Vec<usize> {
174        let n = self.shape.0;
175        let mut pivots = Vec::with_capacity(n.min(cols));
176        if n < 32 {
177            let mut c = 0;
178            while c < cols {
179                let r = pivots.len();
180                if r == n {
181                    break;
182                }
183                let Some(p) = (r..n).find(|&i| self[i].get(c)) else {
184                    if require_full_rank {
185                        return pivots;
186                    }
187                    c = self.next_column(r, c + 1, cols);
188                    continue;
189                };
190                self.data.swap(r, p);
191                let (upper, lower) = self.data.split_at_mut(r);
192                let (pivot, lower) = lower.split_first_mut().unwrap();
193                for row in lower
194                    .iter_mut()
195                    .chain(upper.iter_mut().take(if full { r } else { 0 }))
196                {
197                    if row.get(c) {
198                        xor(&mut row.words_mut()[c / 64..], &pivot.words()[c / 64..]);
199                    }
200                }
201                pivots.push(c);
202                c += 1;
203            }
204            return pivots;
205        }
206        // Eliminating a shared leading coefficient preserves the two-coefficient bound.
207        if self
208            .data
209            .iter()
210            .all(|row| row.iter_ones().take_while(|&c| c < cols).take(3).count() <= 2)
211        {
212            return self.eliminate_sparse(cols, full, require_full_rank);
213        }
214        let block: usize = if n < 512 {
215            4
216        } else if n < 1536 {
217            8
218        } else {
219            32
220        };
221        let mut table = vec![BitSet::new(self.shape.1); (1 << block.min(8)) * block.div_ceil(8)];
222        let mut reduced = vec![0; n];
223        let mut start = 0;
224        while start < cols {
225            let first = pivots.len();
226            let word = start / 64;
227            reduced.fill(first);
228            let end = cols.min(start + block);
229            let mut c = start;
230            while c < end {
231                let r = pivots.len();
232                let mut pivot = None;
233                for (i, reduced) in reduced.iter_mut().enumerate().skip(r) {
234                    let (upper, lower) = self.data.split_at_mut(i);
235                    let row = &mut lower[0];
236                    for p in *reduced..r {
237                        if row.get(pivots[p]) {
238                            xor(&mut row.words_mut()[word..], &upper[p].words()[word..]);
239                        }
240                    }
241                    *reduced = r;
242                    if row.get(c) {
243                        pivot = Some(i);
244                        break;
245                    }
246                }
247                if let Some(p) = pivot {
248                    self.data.swap(r, p);
249                    reduced.swap(r, p);
250                    pivots.push(c);
251                    c += 1;
252                } else if require_full_rank {
253                    return pivots;
254                } else {
255                    c = self.next_column(r, c + 1, end);
256                }
257            }
258            let rank = pivots.len();
259            if first == rank {
260                let next = self.next_column(rank, start + block, cols);
261                if next == cols {
262                    break;
263                }
264                start = next / block * block;
265                continue;
266            }
267            // Make the panel's pivot columns an identity matrix before indexing its table.
268            for r in (first..rank).rev() {
269                let (upper, lower) = self.data.split_at_mut(r);
270                for row in &mut upper[first..] {
271                    if row.get(pivots[r]) {
272                        xor(&mut row.words_mut()[word..], &lower[0].words()[word..]);
273                    }
274                }
275            }
276            let mut indices = [[0usize; 256]; 4];
277            let mut masks = [0usize; 4];
278            for group in 0..block.div_ceil(8) {
279                let mut keys = [0usize; 256];
280                let mut count = 0;
281                for (row, &c) in pivots[first..].iter().enumerate() {
282                    if (c - start) / 8 != group {
283                        continue;
284                    }
285                    let half = 1 << count;
286                    count += 1;
287                    for index in 0..half {
288                        keys[index + half] = keys[index] | (1 << ((c - start) % 8));
289                        let (lower, upper) = table.split_at_mut(group * 256 + index + half);
290                        let target = &mut upper[0].words_mut()[word..];
291                        let source = &lower[group * 256 + index].words()[word..];
292                        let pivot = &self[first + row].words()[word..];
293                        for ((x, y), z) in target.iter_mut().zip(source).zip(pivot) {
294                            *x = y ^ z;
295                        }
296                    }
297                }
298                masks[group] = keys[(1 << count) - 1];
299                for (i, &key) in keys[..1 << count].iter().enumerate() {
300                    indices[group][key] = i;
301                }
302            }
303            for i in (rank..n).chain(0..if full { first } else { 0 }) {
304                let key = (self[i].words()[word] >> (start & 63)) as usize;
305                let x = indices[0][key & masks[0]];
306                if block <= 8 {
307                    if x != 0 {
308                        xor(&mut self[i].words_mut()[word..], &table[x].words()[word..]);
309                    }
310                    continue;
311                }
312                let y = indices[1][(key >> 8) & masks[1]];
313                let z = indices[2][(key >> 16) & masks[2]];
314                let w = indices[3][(key >> 24) & masks[3]];
315                if x | y | z | w != 0 {
316                    let p = &table[x].words()[word..];
317                    let q = &table[256 + y].words()[word..];
318                    let r = &table[512 + z].words()[word..];
319                    let s = &table[768 + w].words()[word..];
320                    for ((((x, y), z), r), s) in self[i].words_mut()[word..]
321                        .iter_mut()
322                        .zip(p)
323                        .zip(q)
324                        .zip(r)
325                        .zip(s)
326                    {
327                        *x ^= y ^ z ^ r ^ s;
328                    }
329                }
330            }
331            if rank == n {
332                break;
333            }
334            start += block;
335        }
336        pivots
337    }
338
339    fn next_column(&self, row: usize, mut start: usize, cols: usize) -> usize {
340        while start < cols {
341            let word = start / 64;
342            let bits = self.data[row..]
343                .iter()
344                .fold(0, |x, row| x | row.words()[word])
345                & (u64::MAX << (start & 63));
346            if bits != 0 {
347                return (word * 64 + bits.trailing_zeros() as usize).min(cols);
348            }
349            start = (word + 1) * 64;
350        }
351        cols
352    }
353
354    #[inline(always)]
355    fn eliminate_sparse(&mut self, cols: usize, full: bool, require_full_rank: bool) -> Vec<usize> {
356        let n = self.shape.0;
357        let mut basis = vec![n; cols];
358        let mut pivots = Vec::new();
359        for i in 0..n {
360            loop {
361                let Some(c) = self[i].iter_ones().next().filter(|&c| c < cols) else {
362                    if require_full_rank {
363                        return pivots;
364                    }
365                    break;
366                };
367                if basis[c] == n {
368                    basis[c] = i;
369                    pivots.push(c);
370                    break;
371                }
372                let (upper, lower) = self.data.split_at_mut(i);
373                xor(
374                    &mut lower[0].words_mut()[c / 64..],
375                    &upper[basis[c]].words()[c / 64..],
376                );
377            }
378        }
379        pivots.sort_unstable();
380        self.data
381            .sort_by_cached_key(|row| row.iter_ones().next().filter(|&c| c < cols).unwrap_or(cols));
382        if full {
383            for (i, &c) in pivots.iter().enumerate() {
384                basis[c] = i;
385            }
386            for i in (0..pivots.len()).rev() {
387                let next = self[i].iter_ones().take_while(|&c| c < cols).nth(1);
388                if let Some(c) = next
389                    && basis[c] != n
390                {
391                    let (upper, lower) = self.data.split_at_mut(basis[c]);
392                    xor(
393                        &mut upper[i].words_mut()[c / 64..],
394                        &lower[0].words()[c / 64..],
395                    );
396                }
397            }
398        }
399        pivots
400    }
401
402    #[inline(always)]
403    fn mul_impl(&self, rhs: &Self) -> Self {
404        let mut result = Self::zeros((self.shape.0, rhs.shape.1));
405        let ones = self.data.iter().map(BitSet::count_ones).sum::<u64>();
406        let size = self.shape.0 as u64 * self.shape.1 as u64;
407        if self.shape.0 < 256 || self.shape.1 < 32 || ones <= size / 8 {
408            for (a, c) in self.data.iter().zip(&mut result.data) {
409                for j in a.iter_ones() {
410                    xor(c.words_mut(), rhs[j].words());
411                }
412            }
413            return result;
414        }
415        if size - ones <= size / 8 {
416            let mut sum = BitSet::new(rhs.shape.1);
417            for row in &rhs.data {
418                sum ^= row;
419            }
420            for (a, c) in self.data.iter().zip(&mut result.data) {
421                c.words_mut().copy_from_slice(sum.words());
422                for j in (!a.clone()).iter_ones() {
423                    xor(c.words_mut(), rhs[j].words());
424                }
425            }
426            return result;
427        }
428        let width = rhs.shape.1.div_ceil(64);
429        if width == 0 {
430            return result;
431        }
432
433        // Separate the table groups by a cache line to avoid mapping them to the same sets.
434        let group = 256 * width + 8;
435        let mut storage = BitSet::new(8 * group * 64);
436        let table = storage.words_mut();
437        for start in (0..self.shape.1).step_by(64) {
438            for (t, table) in table.chunks_exact_mut(group).enumerate() {
439                let col = start + t * 8;
440                for bit in 0..self.shape.1.saturating_sub(col).min(8) {
441                    let row = rhs[col + bit].words();
442                    let half = (1 << bit) * width;
443                    let (lower, upper) = table.split_at_mut(half);
444                    for (source, dest) in
445                        lower.chunks_exact(width).zip(upper.chunks_exact_mut(width))
446                    {
447                        for ((x, y), z) in dest.iter_mut().zip(source).zip(row) {
448                            *x = y ^ z;
449                        }
450                    }
451                }
452            }
453            for (a, c) in self.data.iter().zip(&mut result.data) {
454                let key = a.words()[start / 64];
455                let offset = (key & 255) as usize * width;
456                let p0 = &table[offset..offset + width];
457                let offset = group + (key >> 8 & 255) as usize * width;
458                let p1 = &table[offset..offset + width];
459                let offset = 2 * group + (key >> 16 & 255) as usize * width;
460                let p2 = &table[offset..offset + width];
461                let offset = 3 * group + (key >> 24 & 255) as usize * width;
462                let p3 = &table[offset..offset + width];
463                let offset = 4 * group + (key >> 32 & 255) as usize * width;
464                let p4 = &table[offset..offset + width];
465                let offset = 5 * group + (key >> 40 & 255) as usize * width;
466                let p5 = &table[offset..offset + width];
467                let offset = 6 * group + (key >> 48 & 255) as usize * width;
468                let p6 = &table[offset..offset + width];
469                let offset = 7 * group + (key >> 56 & 255) as usize * width;
470                let p7 = &table[offset..offset + width];
471                for ((((((((x, p0), p1), p2), p3), p4), p5), p6), p7) in c
472                    .words_mut()
473                    .iter_mut()
474                    .zip(p0)
475                    .zip(p1)
476                    .zip(p2)
477                    .zip(p3)
478                    .zip(p4)
479                    .zip(p5)
480                    .zip(p6)
481                    .zip(p7)
482                {
483                    *x ^= p0 ^ p1 ^ p2 ^ p3 ^ p4 ^ p5 ^ p6 ^ p7;
484                }
485            }
486        }
487        result
488    }
489
490    #[cfg(target_arch = "x86_64")]
491    #[target_feature(enable = "avx2")]
492    unsafe fn eliminate_avx2(
493        &mut self,
494        cols: usize,
495        full: bool,
496        require_full_rank: bool,
497    ) -> Vec<usize> {
498        self.eliminate_impl(cols, full, require_full_rank)
499    }
500    #[cfg(target_arch = "x86_64")]
501    #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
502    unsafe fn eliminate_avx512(
503        &mut self,
504        cols: usize,
505        full: bool,
506        require_full_rank: bool,
507    ) -> Vec<usize> {
508        self.eliminate_impl(cols, full, require_full_rank)
509    }
510    #[cfg(target_arch = "x86_64")]
511    #[target_feature(enable = "avx2")]
512    unsafe fn mul_avx2(&self, rhs: &Self) -> Self {
513        self.mul_impl(rhs)
514    }
515    #[cfg(target_arch = "x86_64")]
516    #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
517    unsafe fn mul_avx512(&self, rhs: &Self) -> Self {
518        self.mul_impl(rhs)
519    }
520}
521
522#[inline(always)]
523fn xor(row: &mut [u64], pivot: &[u64]) {
524    for (x, y) in row.iter_mut().zip(pivot) {
525        *x ^= y;
526    }
527}
528
529impl Index<usize> for BitMatrix {
530    type Output = BitSet;
531    fn index(&self, i: usize) -> &Self::Output {
532        &self.data[i]
533    }
534}
535impl IndexMut<usize> for BitMatrix {
536    fn index_mut(&mut self, i: usize) -> &mut Self::Output {
537        &mut self.data[i]
538    }
539}
540impl BitXorAssign<&Self> for BitMatrix {
541    fn bitxor_assign(&mut self, rhs: &Self) {
542        assert_eq!(self.shape, rhs.shape);
543        for (a, b) in self.data.iter_mut().zip(&rhs.data) {
544            *a ^= b;
545        }
546    }
547}
548impl Mul<&BitMatrix> for &BitMatrix {
549    type Output = BitMatrix;
550    fn mul(self, rhs: &BitMatrix) -> BitMatrix {
551        assert_eq!(self.shape.1, rhs.shape.0);
552        #[cfg(target_arch = "x86_64")]
553        match simd_backend() {
554            // SAFETY: the dispatcher checks the required CPU features.
555            SimdBackend::Avx512 => return unsafe { self.mul_avx512(rhs) },
556            // SAFETY: the dispatcher checks AVX2 support.
557            SimdBackend::Avx2 => return unsafe { self.mul_avx2(rhs) },
558            SimdBackend::Scalar => {}
559        }
560        self.mul_impl(rhs)
561    }
crates/competitive/src/data_structure/wavelet_matrix.rs (line 458)
455    pub fn new(v: Vec<T>) -> Self {
456        if v.len() <= u32::MAX as usize {
457            #[cfg(target_arch = "x86_64")]
458            let backend = super::simd_backend();
459            Self::from_values(
460                v,
461                |i| i as u32,
462                |i| i as usize,
463                |indices, d| {
464                    #[cfg(target_arch = "x86_64")]
465                    if backend != super::SimdBackend::Scalar && is_x86_feature_detected!("avx2") {
466                        // SAFETY: AVX2 is available.
467                        return unsafe { simd::pack_words(indices, d) };
468                    }
469                    Self::pack_words(indices, |i| (i >> d) & 1 != 0)
470                },
471                |indices, words, zeros, next| {
472                    #[cfg(target_arch = "x86_64")]
473                    if backend != super::SimdBackend::Scalar && is_x86_feature_detected!("avx2") {
474                        // SAFETY: AVX2 is available, and the partition buffers have equal length.
475                        unsafe { simd::partition_avx2(indices, words, zeros, next) };
476                        return;
477                    }
478                    Self::partition(indices, words, zeros, next);
479                },
480            )
481        } else {
482            Self::from_values(
483                v,
484                |i| i,
485                |i| i,
486                |indices, d| Self::pack_words(indices, |i| (i >> d) & 1 != 0),
487                Self::partition,
488            )
489        }
490    }
491
492    fn pack_words<I: Copy>(indices: &[I], bit: impl Fn(I) -> bool) -> Vec<u64> {
493        indices
494            .chunks(64)
495            .map(|chunk| {
496                chunk
497                    .iter()
498                    .enumerate()
499                    .fold(0, |word, (i, &index)| word | ((bit(index) as u64) << i))
500            })
501            .collect()
502    }
503
504    fn partition<I: Copy>(indices: &[I], words: &[u64], mut one: usize, next: &mut [I]) {
505        let mut zero = 0;
506        for (chunk, &word) in indices.chunks(64).zip(words) {
507            if word == 0 {
508                next[zero..zero + chunk.len()].copy_from_slice(chunk);
509                zero += chunk.len();
510            } else if word == u64::MAX {
511                next[one..one + chunk.len()].copy_from_slice(chunk);
512                one += chunk.len();
513            } else {
514                for (i, &index) in chunk.iter().enumerate() {
515                    let bit = (word >> i) & 1 != 0;
516                    next[if bit { one } else { zero }] = index;
517                    zero += !bit as usize;
518                    one += bit as usize;
519                }
520            }
521        }
522    }
523
524    fn from_values<I: Copy>(
525        v: Vec<T>,
526        code: impl Fn(usize) -> I,
527        index: impl Fn(I) -> usize,
528        pack: impl Fn(&[I], usize) -> Vec<u64>,
529        partition: impl Fn(&[I], &[u64], usize, &mut [I]),
530    ) -> Self {
531        let len = v.len();
532        let mut sorted: Vec<_> = v
533            .into_iter()
534            .enumerate()
535            .map(|(i, value)| (value, code(i)))
536            .collect();
537        sorted.sort_unstable_by(|a, b| a.0.cmp(&b.0));
538        let mut values = Vec::with_capacity(len);
539        let mut indices = vec![code(0); len];
540        for (value, i) in sorted {
541            if values.last().is_none_or(|last| last != &value) {
542                values.push(value);
543            }
544            indices[index(i)] = code(values.len() - 1);
545        }
546        let compress = VecCompress::from_sorted_unique(values);
547        let bit_length = usize::BITS as usize - compress.size().leading_zeros() as usize;
548        let mut bit_vectors = Vec::with_capacity(bit_length);
549        let mut zeros = Vec::with_capacity(bit_length);
550        let quad_bits =
551            usize::BITS as usize - compress.size().saturating_sub(1).leading_zeros() as usize;
552        let mut quad_vectors = Vec::with_capacity(quad_bits.div_ceil(2));
553        let mut next = indices.clone();
554        for d in (0..bit_length).rev() {
555            let words = pack(&indices, d);
556            if len <= u32::MAX as usize && d < quad_bits && (d % 2 == 1 || d + 1 == quad_bits) {
557                if d % 2 == 1 {
558                    let low = pack(&indices, d - 1);
559                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&low, Some(&words), len));
560                } else {
561                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&words, None, len));
562                }
563            }
564            let bits = BitVector::from_words(&words, len);
565            let zero_count = bits.rank0(len);
566            if d == 0 {
567                zeros.push(zero_count);
568                bit_vectors.push(bits);
569                break;
570            }
571            partition(&indices, &words, zero_count, &mut next);
572            zeros.push(zero_count);
573            bit_vectors.push(bits);
574            mem::swap(&mut indices, &mut next);
575        }
576        Self {
577            len,
578            bit_length,
579            zeros,
580            bit_vectors,
581            quad_vectors,
582            compress,
583            #[cfg(target_arch = "x86_64")]
584            backend: match super::simd_backend() {
585                super::SimdBackend::Avx512 if !is_x86_feature_detected!("avx512vpopcntdq") => {
586                    super::SimdBackend::Avx2
587                }
588                backend => backend,
589            },
590        }
591    }