Skip to main content

avx512_enabled

Function avx512_enabled 

Source
pub fn avx512_enabled() -> bool
Examples found in repository?
crates/competitive/src/tools/avx_helper.rs (line 43)
40pub fn simd_backend() -> SimdBackend {
41    #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
42    {
43        if avx512_enabled() && avx512_supported() {
44            return SimdBackend::Avx512;
45        }
46        if is_x86_feature_detected!("avx2") {
47            return SimdBackend::Avx2;
48        }
49    }
50    SimdBackend::Scalar
51}
More examples
Hide additional examples
crates/competitive/src/data_structure/static_search.rs (line 10)
7fn static_search_backend(bits: u32) -> SimdBackend {
8    #[cfg(target_arch = "x86_64")]
9    {
10        if avx512_enabled()
11            && is_x86_feature_detected!("avx512f")
12            && (bits != 16 || is_x86_feature_detected!("avx512bw"))
13        {
14            return SimdBackend::Avx512;
15        }
16        if is_x86_feature_detected!("avx2") {
17            return SimdBackend::Avx2;
18        }
19    }
20    let _ = bits;
21    SimdBackend::Scalar
22}
crates/competitive/src/num/mint/montgomery_dot_product.rs (line 47)
43    fn add_scaled_assign(x: &mut [MInt<Self>], y: &[MInt<Self>], a: &MInt<Self>) {
44        assert_eq!(x.len(), y.len());
45        #[cfg(target_arch = "x86_64")]
46        if x.len() >= 16 {
47            if x.len() >= 64 && avx512_enabled() && avx512_supported() {
48                unsafe { simd::add_scaled_avx512::<Self>(x, y, a) };
49                return;
50            }
51            if is_x86_feature_detected!("avx2") {
52                unsafe { simd::add_scaled_avx2::<Self>(x, y, a) };
53                return;
54            }
55        }
56        for (x, y) in x.iter_mut().zip(y) {
57            *x += *a * *y;
58        }
59    }
crates/competitive/src/data_structure/bitset.rs (line 50)
45    pub fn from_binary(s: &str) -> Option<Self> {
46        let bytes = s.as_bytes();
47        let mut bits = Self::new(bytes.len());
48        let end = bytes.len() / 64 * 64;
49        #[cfg(target_arch = "x86_64")]
50        let parsed = if avx512_enabled() && is_x86_feature_detected!("avx512bw") {
51            // SAFETY: AVX-512BW is available; each complete chunk contains 64 bytes.
52            unsafe { simd::parse_binary_avx512(&bytes[..end], bits.words_mut()) }
53        } else if is_x86_feature_detected!("avx2") {
54            // SAFETY: AVX2 is available; each complete chunk contains 64 bytes.
55            unsafe { simd::parse_binary_avx2(&bytes[..end], bits.words_mut()) }
56        } else {
57            Self::parse_binary_scalar(&bytes[..end], bits.words_mut())
58        };
59        #[cfg(not(target_arch = "x86_64"))]
60        let parsed = Self::parse_binary_scalar(&bytes[..end], bits.words_mut());
61        if !parsed {
62            return None;
63        }
64        if end != bytes.len() {
65            let mut word = 0;
66            for (i, &b) in bytes[end..].iter().enumerate() {
67                if b != b'0' && b != b'1' {
68                    return None;
69                }
70                word |= u64::from(b & 1) << i;
71            }
72            bits.words_mut()[end / 64] = word;
73        }
74        Some(bits)
75    }
76
77    fn parse_binary_scalar(bytes: &[u8], words: &mut [u64]) -> bool {
78        for (chunk, word) in bytes.as_chunks::<64>().0.iter().zip(words) {
79            for (i, byte) in chunk.as_chunks::<8>().0.iter().enumerate() {
80                let x = u64::from_le_bytes(*byte);
81                if x & 0xfefe_fefe_fefe_fefe != 0x3030_3030_3030_3030 {
82                    return false;
83                }
84                *word |= ((x & 0x0101_0101_0101_0101).wrapping_mul(0x0102_0408_1020_4080) >> 56)
85                    << (i * 8);
86            }
87        }
88        true
89    }
90
91    /// Returns ASCII `0` and `1` in increasing bit-index order.
92    pub fn to_binary(&self) -> String {
93        const TABLE: [[u8; 8]; 256] = {
94            let mut table = [[b'0'; 8]; 256];
95            let mut i = 0;
96            while i < 256 {
97                let mut j = 0;
98                while j < 8 {
99                    table[i][j] |= ((i >> j) & 1) as u8;
100                    j += 1;
101                }
102                i += 1;
103            }
104            table
105        };
106        let mut bytes = vec![b'0'; self.size.div_ceil(8) * 8];
107        #[cfg(target_arch = "x86_64")]
108        let end = if self.size >= 64 && avx512_enabled() && is_x86_feature_detected!("avx512bw") {
109            let end = self.size / 64 * 64;
110            // SAFETY: AVX-512BW is available; each output chunk holds one 64-bit word.
111            unsafe {
112                simd::write_binary_avx512(&mut bytes[..end], self.words());
113            }
114            end
115        } else if self.size >= 64 && is_x86_feature_detected!("avx2") {
116            let end = self.size / 64 * 64;
117            // SAFETY: AVX2 is available; each output chunk holds one complete 64-bit word.
118            unsafe {
119                simd::write_binary_avx2(&mut bytes[..end], self.words());
120            }
121            end
122        } else {
123            0
124        };
125        #[cfg(not(target_arch = "x86_64"))]
126        let end = 0;
127        for (chunk, &word) in bytes[end..].chunks_mut(64).zip(&self.words()[end / 64..]) {
128            for (i, byte) in chunk.as_chunks_mut::<8>().0.iter_mut().enumerate() {
129                byte.copy_from_slice(&TABLE[(word >> (i * 8) & 255) as usize]);
130            }
131        }
132        bytes.truncate(self.size);
133        // SAFETY: every output byte is ASCII `0` or `1`.
134        unsafe { String::from_utf8_unchecked(bytes) }
135    }
136
137    pub fn ones(size: usize) -> Self {
138        let mut self_ = Self {
139            size,
140            bits: vec![Block([u64::MAX; 8]); size.div_ceil(512)],
141        };
142        self_.trim();
143        self_
144    }
145
146    pub fn get(&self, i: usize) -> bool {
147        self.bits[i >> 9].0[i >> 6 & 7] & (1 << (i & 63)) != 0
148    }
149
150    pub fn set(&mut self, i: usize, b: bool) {
151        let word = &mut self.bits[i >> 9].0[i >> 6 & 7];
152        if b {
153            *word |= 1 << (i & 63);
154        } else {
155            *word &= !(1 << (i & 63));
156        }
157    }
158
159    /// Clears all bits.
160    pub fn reset(&mut self) {
161        self.bits.fill(Block::default());
162    }
163
164    /// Sets all bits to `value`.
165    pub fn fill(&mut self, value: bool) {
166        self.bits.fill(Block([if value { u64::MAX } else { 0 }; 8]));
167        self.trim();
168    }
169
170    /// Tests whether any bit is set.
171    #[inline]
172    pub fn any(&self) -> bool {
173        !self.none()
174    }
175
176    /// Tests whether all bits are unset.
177    #[inline]
178    pub fn none(&self) -> bool {
179        #[cfg(target_arch = "x86_64")]
180        if self.bits.len() >= SIMD_MIN_BLOCKS {
181            if self.bits[0].0[0] != 0 {
182                return false;
183            }
184            if avx512_enabled() && is_x86_feature_detected!("avx512f") {
185                // SAFETY: blocks are 64-byte aligned and feature detection checked AVX-512F.
186                return unsafe { simd::none_avx512(&self.bits) };
187            }
188            if is_x86_feature_detected!("avx2") {
189                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
190                return unsafe { simd::none_avx2(&self.bits) };
191            }
192        }
193        self.words().iter().all(|&word| word == 0)
194    }
195
196    /// Tests whether all bits are set.
197    #[inline(always)]
198    pub fn all(&self) -> bool {
199        let words = self.words();
200        let full_words = self.size >> 6;
201        #[cfg(target_arch = "x86_64")]
202        if self.size >> 9 >= SIMD_MIN_BLOCKS {
203            let full_blocks = self.size >> 9;
204            if words[0] != u64::MAX {
205                return false;
206            }
207            let full_blocks_are_set = if avx512_enabled() && is_x86_feature_detected!("avx512f") {
208                // SAFETY: blocks are 64-byte aligned and feature detection checked AVX-512F.
209                Some(unsafe { simd::all_avx512(&self.bits[..full_blocks]) })
210            } else if is_x86_feature_detected!("avx2") {
211                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
212                Some(unsafe { simd::all_avx2(&self.bits[..full_blocks]) })
213            } else {
214                None
215            };
216            if let Some(full_blocks_are_set) = full_blocks_are_set {
217                return full_blocks_are_set
218                    && words[full_blocks * 8..full_words]
219                        .iter()
220                        .all(|&word| word == u64::MAX)
221                    && (self.size & 63 == 0
222                        || words[full_words] == u64::MAX >> (64 - (self.size & 63)));
223            }
224        }
225        if self.size & 63 == 0 {
226            return words.iter().all(|&word| word == u64::MAX);
227        }
228        words[..full_words].iter().all(|&word| word == u64::MAX)
229            && words[full_words] == u64::MAX >> (64 - (self.size & 63))
230    }
231
232    /// Iterates over set-bit indices in ascending order.
233    pub fn iter_ones(&self) -> impl Iterator<Item = usize> + '_ {
234        self.words()
235            .iter()
236            .copied()
237            .enumerate()
238            .flat_map(|(word_index, mut word)| {
239                std::iter::from_fn(move || {
240                    if word == 0 {
241                        None
242                    } else {
243                        let bit = word.trailing_zeros() as usize;
244                        word &= word - 1;
245                        Some((word_index << 6) | bit)
246                    }
247                })
248            })
249    }
250
251    /// Counts set bits.
252    #[inline]
253    pub fn count_ones(&self) -> u64 {
254        let words = self.words();
255        #[cfg(target_arch = "x86_64")]
256        if words.len() >= 8 {
257            if avx512_enabled()
258                && is_x86_feature_detected!("avx512f")
259                && is_x86_feature_detected!("avx512vpopcntdq")
260            {
261                // SAFETY: blocks are aligned and feature detection checked both requirements.
262                return unsafe { simd::count_ones_avx512(&self.bits) };
263            }
264            if is_x86_feature_detected!("avx2") {
265                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
266                return unsafe { simd::count_ones_avx2(&self.bits) };
267            }
268        }
269        words.iter().map(|word| word.count_ones() as u64).sum()
270    }
271
272    /// Counts unset bits.
273    #[inline]
274    pub fn count_zeros(&self) -> u64 {
275        self.size as u64 - self.count_ones()
276    }
277
278    pub fn push(&mut self, b: bool) {
279        if self.size & 511 == 0 {
280            self.bits.push(Block::default());
281        }
282        if b {
283            self.bits[self.size >> 9].0[self.size >> 6 & 7] |= 1 << (self.size & 63);
284        }
285        self.size += 1;
286    }
287
288    pub fn resize(&mut self, new_size: usize) {
289        match self.size.cmp(&new_size) {
290            Ordering::Less => self.bits.resize(new_size.div_ceil(512), Block::default()),
291            Ordering::Equal => {}
292            Ordering::Greater => self.bits.truncate(new_size.div_ceil(512)),
293        }
294        self.size = new_size;
295        self.trim();
296    }
297
298    /// Assigns `self | (self << rhs)` to `self`.
299    #[inline]
300    pub fn shl_bitor_assign(&mut self, rhs: usize) {
301        self.shift_left::<true>(rhs);
302    }
303
304    /// Assigns `self | (self >> rhs)` to `self`.
305    #[inline]
306    pub fn shr_bitor_assign(&mut self, rhs: usize) {
307        self.shift_right::<true>(rhs);
308    }
309
310    /// Returns words in increasing bit-index order; unused high bits in the last word are zero.
311    pub fn words(&self) -> &[u64] {
312        // SAFETY: `Block` is exactly eight contiguous `u64`s, with no trailing padding.
313        unsafe { std::slice::from_raw_parts(self.bits.as_ptr().cast(), self.size.div_ceil(64)) }
314    }
315
316    /// Returns mutable words in increasing bit-index order.
317    /// Callers must keep unused high bits in the last word zero.
318    pub fn words_mut(&mut self) -> &mut [u64] {
319        let len = self.size.div_ceil(64);
320        // SAFETY: `Block` is exactly eight contiguous `u64`s, with no trailing padding.
321        unsafe { std::slice::from_raw_parts_mut(self.bits.as_mut_ptr().cast(), len) }
322    }
323
324    fn trim(&mut self) {
325        let used_words = self.size.div_ceil(64) & 7;
326        if let Some(last) = self.bits.last_mut() {
327            if self.size & 63 != 0 {
328                last.0[(self.size >> 6) & 7] &= u64::MAX >> (64 - (self.size & 63));
329            }
330            if used_words != 0 {
331                last.0[used_words..].fill(0);
332            }
333        }
334    }
335
336    #[inline]
337    fn bitop_assign<const OP: u8>(&mut self, rhs: &Self) {
338        assert_eq!(
339            self.size, rhs.size,
340            "bitwise operations require equal lengths"
341        );
342        #[cfg(target_arch = "x86_64")]
343        if self.bits.len() >= SIMD_MIN_BLOCKS {
344            if avx512_enabled() && is_x86_feature_detected!("avx512f") {
345                // SAFETY: blocks are 64-byte aligned and feature detection checked AVX-512F.
346                unsafe { simd::bitop_avx512::<OP>(&mut self.bits, &rhs.bits) };
347                return;
348            }
349            if is_x86_feature_detected!("avx2") {
350                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
351                unsafe { simd::bitop_avx2::<OP>(&mut self.bits, &rhs.bits) };
352                return;
353            }
354        }
355        for (lhs, &rhs) in self.words_mut().iter_mut().zip(rhs.words()) {
356            *lhs = match OP {
357                BIT_AND => *lhs & rhs,
358                BIT_OR => *lhs | rhs,
359                BIT_XOR => *lhs ^ rhs,
360                _ => unreachable!(),
361            };
362        }
363    }
364
365    #[inline]
366    fn shift_left<const OR_ASSIGN: bool>(&mut self, rhs: usize) {
367        if rhs == 0 {
368            return;
369        }
370        if rhs >= self.size {
371            if !OR_ASSIGN {
372                self.reset();
373            }
374            return;
375        }
376        #[cfg(target_arch = "x86_64")]
377        if self.bits.len() >= SIMD_MIN_BLOCKS {
378            if avx512_enabled()
379                && is_x86_feature_detected!("avx512f")
380                && is_x86_feature_detected!("avx512vbmi2")
381            {
382                // SAFETY: blocks are aligned and feature detection checked AVX-512F and VBMI2.
383                unsafe { simd::shift_left_avx512::<OR_ASSIGN>(&mut self.bits, rhs) };
384                if self.size & 511 != 0 {
385                    self.trim();
386                }
387                return;
388            }
389            if is_x86_feature_detected!("avx2") {
390                // SAFETY: feature detection checked AVX2 support.
391                unsafe { simd::shift_left_avx2::<OR_ASSIGN>(self.words_mut(), rhs) };
392                if self.size & 63 != 0 {
393                    self.trim();
394                }
395                return;
396            }
397        }
398
399        let bits = self.words_mut();
400        let word_shift = rhs >> 6;
401        let bit_shift = rhs & 63;
402        if bit_shift == 0 {
403            for i in (0..bits.len() - word_shift).rev() {
404                if OR_ASSIGN {
405                    bits[i + word_shift] |= bits[i];
406                } else {
407                    bits[i + word_shift] = bits[i];
408                }
409            }
410        } else {
411            for i in (1..bits.len() - word_shift).rev() {
412                let value = (bits[i] << bit_shift) | (bits[i - 1] >> (64 - bit_shift));
413                if OR_ASSIGN {
414                    bits[i + word_shift] |= value;
415                } else {
416                    bits[i + word_shift] = value;
417                }
418            }
419            if OR_ASSIGN {
420                bits[word_shift] |= bits[0] << bit_shift;
421            } else {
422                bits[word_shift] = bits[0] << bit_shift;
423            }
424        }
425        if !OR_ASSIGN {
426            bits[..word_shift].fill(0);
427        }
428        if self.size & 63 != 0 {
429            self.trim();
430        }
431    }
432
433    #[inline]
434    fn shift_right<const OR_ASSIGN: bool>(&mut self, rhs: usize) {
435        if rhs == 0 {
436            return;
437        }
438        if rhs >= self.size {
439            if !OR_ASSIGN {
440                self.reset();
441            }
442            return;
443        }
444        #[cfg(target_arch = "x86_64")]
445        if self.bits.len() >= SIMD_MIN_BLOCKS {
446            if avx512_enabled()
447                && is_x86_feature_detected!("avx512f")
448                && is_x86_feature_detected!("avx512vbmi2")
449            {
450                // SAFETY: blocks are aligned and feature detection checked AVX-512F and VBMI2.
451                unsafe { simd::shift_right_avx512::<OR_ASSIGN>(&mut self.bits, rhs) };
452                return;
453            }
454            if is_x86_feature_detected!("avx2") {
455                // SAFETY: feature detection checked AVX2 support.
456                unsafe { simd::shift_right_avx2::<OR_ASSIGN>(self.words_mut(), rhs) };
457                return;
458            }
459        }
460
461        let bits = self.words_mut();
462        let word_shift = rhs >> 6;
463        let bit_shift = rhs & 63;
464        if bit_shift == 0 {
465            for i in word_shift..bits.len() {
466                if OR_ASSIGN {
467                    bits[i - word_shift] |= bits[i];
468                } else {
469                    bits[i - word_shift] = bits[i];
470                }
471            }
472        } else {
473            for i in word_shift..bits.len() - 1 {
474                let value = (bits[i] >> bit_shift) | (bits[i + 1] << (64 - bit_shift));
475                if OR_ASSIGN {
476                    bits[i - word_shift] |= value;
477                } else {
478                    bits[i - word_shift] = value;
479                }
480            }
481            if OR_ASSIGN {
482                bits[bits.len() - word_shift - 1] |= bits[bits.len() - 1] >> bit_shift;
483            } else {
484                bits[bits.len() - word_shift - 1] = bits[bits.len() - 1] >> bit_shift;
485            }
486        }
487        if !OR_ASSIGN {
488            let end = bits.len() - word_shift;
489            bits[end..].fill(0);
490        }
491    }
crates/competitive/src/num/mint/simd_matrix.rs (line 276)
211    pub unsafe fn matrix_product_avx2(
212        a: &[Vec<Self>],
213        b: &[Vec<Self>],
214        scale: u32,
215    ) -> Vec<Vec<Self>> {
216        let (n, m, p) = (a.len(), b.len(), b.first().map_or(0, Vec::len));
217        assert!(a.iter().all(|row| row.len() == m));
218        assert!(b.iter().all(|row| row.len() == p));
219        let modulus = M::get_mod();
220        let alignment = if n.min(m).min(p) <= 64 { 8 } else { 32 };
221        let (nn, mm, pp) = (
222            n.div_ceil(alignment) * alignment,
223            m.div_ceil(alignment) * alignment,
224            p.div_ceil(alignment) * alignment,
225        );
226        let mut depth = 0;
227        let (mut x, mut y, mut z) = (nn, mm, pp);
228        while x.min(y).min(z) > 64 && x % 16 == 0 && y % 16 == 0 && z % 16 == 0 {
229            depth += 1;
230            x /= 2;
231            y /= 2;
232            z /= 2;
233        }
234        let entries = nn * mm + mm * pp + nn * pp;
235        // A recursive level uses one quarter of its parent's storage; siblings reuse it.
236        let mut data = if entries + entries / 3 >= 1 << 20 {
237            let mut data = Vec::with_capacity(entries + entries / 3);
238            advise_huge_pages(&mut data);
239            data.resize(entries + entries / 3, 0u32);
240            data
241        } else {
242            vec![0u32; entries + entries / 3]
243        };
244        let quotient = (((scale as u64) << 32) / modulus as u64) as u32;
245        let mut inverse = 1u32;
246        for _ in 0..5 {
247            inverse = inverse.wrapping_mul(2u32.wrapping_sub(modulus.wrapping_mul(inverse)));
248        }
249        let inverse = inverse.wrapping_neg();
250        for (offset, row, col) in blocks(nn, mm, depth) {
251            let (nr, nc) = (nn >> depth, mm >> depth);
252            for i in row..(row + nr).min(n) {
253                // SAFETY: MInt is transparent over u32. Reading raw words preserves Montgomery encoding.
254                let values: &[u32] = unsafe { std::slice::from_raw_parts(a[i].as_ptr().cast(), m) };
255                for j in col..(col + nc).min(m) {
256                    let x = values[j];
257                    let q = ((x as u64 * quotient as u64) >> 32) as u32;
258                    let x = x.wrapping_mul(scale).wrapping_sub(q.wrapping_mul(modulus));
259                    data[offset + (i - row) * nc + j - col] = x.min(x.wrapping_sub(modulus));
260                }
261            }
262        }
263        for (offset, row, col) in blocks(mm, pp, depth) {
264            let (nr, nc) = (mm >> depth, pp >> depth);
265            for i in row..(row + nr).min(m) {
266                // SAFETY: MInt is transparent over u32 and row lengths were checked above.
267                let values: &[u32] = unsafe { std::slice::from_raw_parts(b[i].as_ptr().cast(), p) };
268                for j in col..(col + nc).min(p) {
269                    data[nn * mm + offset + (i - row) * nc + j - col] = values[j];
270                }
271            }
272        }
273        let kernel = Kernel {
274            modulus,
275            inverse,
276            avx512: avx512_enabled() && is_x86_feature_detected!("avx512f"),
277        };
278        // SAFETY: padding keeps every leaf dimension divisible by eight. The three matrices
279        // and the geometric scratch space are disjoint parts of the allocated buffer.
280        unsafe {
281            let ptr = data.as_mut_ptr();
282            multiply(
283                ptr,
284                ptr.add(nn * mm),
285                ptr.add(nn * mm + mm * pp),
286                (nn, mm, pp),
287                ptr.add(entries),
288                &kernel,
289            );
290        }
291        let mut result = vec![vec![MInt::new_unchecked(M::mod_zero()); p]; n];
292        for (offset, row, col) in blocks(nn, pp, depth) {
293            let (nr, nc) = (nn >> depth, pp >> depth);
294            for i in row..(row + nr).min(n) {
295                for j in col..(col + nc).min(p) {
296                    result[i][j] = MInt::new_unchecked(
297                        data[nn * mm + mm * pp + offset + (i - row) * nc + j - col],
298                    );
299                }
300            }
301        }
302        result
303    }