pub fn simd_backend() -> SimdBackendExamples 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
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 }