Skip to main content

competitive/tools/
avx_helper.rs

1use std::sync::atomic::{AtomicBool, Ordering};
2
3#[derive(Copy, Clone, Debug, Eq, PartialEq)]
4pub enum SimdBackend {
5    Scalar,
6    Avx2,
7    Avx512,
8}
9
10static AVX512_ENABLED: AtomicBool = AtomicBool::new(true);
11
12#[inline]
13pub fn disable_avx512() {
14    AVX512_ENABLED.store(false, Ordering::Relaxed);
15}
16
17#[inline]
18pub fn enable_avx512() {
19    AVX512_ENABLED.store(true, Ordering::Relaxed);
20}
21
22#[inline]
23pub fn avx512_enabled() -> bool {
24    AVX512_ENABLED.load(Ordering::Relaxed)
25}
26
27#[inline]
28pub fn avx512_supported() -> bool {
29    #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
30    return is_x86_feature_detected!("avx512f")
31        && is_x86_feature_detected!("avx512dq")
32        && is_x86_feature_detected!("avx512cd")
33        && is_x86_feature_detected!("avx512bw")
34        && is_x86_feature_detected!("avx512vl");
35    #[cfg(not(any(target_arch = "x86", target_arch = "x86_64")))]
36    false
37}
38
39#[inline]
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}
52
53#[macro_export]
54macro_rules! avx_helper {
55    (@dispatch $backend:path, $kind:ident; $avx512:expr, $avx2:expr, $scalar:expr) => {{
56        #[cfg(target_arch = "x86_64")]
57        {
58            match $backend() {
59                $kind::Avx512 => $avx512,
60                $kind::Avx2 => $avx2,
61                $kind::Scalar => $scalar,
62            }
63        }
64        #[cfg(not(target_arch = "x86_64"))]
65        $scalar
66    }};
67    (@dispatch_avx2_fma $avx2:expr, $scalar:expr) => {{
68        #[cfg(target_arch = "x86_64")]
69        {
70            if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
71                $avx2
72            } else {
73                $scalar
74            }
75        }
76        #[cfg(not(target_arch = "x86_64"))]
77        $scalar
78    }};
79    (@avx512 $(#[$meta:meta])* $vis:vis fn $name:ident$(<$($T:ident),+>)?($($i:ident: $t:ty),*) -> $ret:ty where [$($clauses:tt)*] $body:block) => {
80        $(#[$meta])*
81        $vis fn $name$(<$($T)*>)?($($i: $t),*) -> $ret
82        where
83            $($clauses)*
84        {
85            if $crate::avx512_supported() {
86                $crate::avx_helper!(@def_avx512 fn avx512$(<$($T)*>)?($($i: $t),*) -> $ret where [$($clauses)*] $body);
87                unsafe { avx512$(::<$($T),*>)?($($i),*) }
88            } else if is_x86_feature_detected!("avx2") {
89                $crate::avx_helper!(@def_avx2 fn avx2$(<$($T)*>)?($($i: $t),*) -> $ret where [$($clauses)*] $body);
90                unsafe { avx2$(::<$($T),*>)?($($i),*) }
91            } else {
92                $body
93            }
94        }
95    };
96    (@avx2 $(#[$meta:meta])* $vis:vis fn $name:ident$(<$($T:ident),+>)?($($i:ident: $t:ty),*) -> $ret:ty where [$($clauses:tt)*] $body:block) => {
97        $(#[$meta])*
98        $vis fn $name$(<$($T)*>)?($($i: $t),*) -> $ret
99        where
100            $($clauses)*
101        {
102            if is_x86_feature_detected!("avx2") {
103                $crate::avx_helper!(@def_avx2 fn avx2$(<$($T)*>)?($($i: $t),*) -> $ret where [$($clauses)*] $body);
104                unsafe { avx2$(::<$($T),*>)?($($i),*) }
105            } else {
106                $body
107            }
108        }
109    };
110    (@def_avx512 fn $name:ident$(<$($T:ident),+>)?($($args:tt)*) -> $ret:ty where [$($clauses:tt)*] $body:block) => {
111        #[target_feature(enable = "avx512f,avx512dq,avx512cd,avx512bw,avx512vl")]
112        unsafe fn $name$(<$($T)*>)?($($args)*) -> $ret
113        where
114            $($clauses)*
115        $body
116    };
117    (@def_avx2 fn $name:ident$(<$($T:ident),+>)?($($args:tt)*) -> $ret:ty where [$($clauses:tt)*] $body:block) => {
118        #[target_feature(enable = "avx2")]
119        unsafe fn $name$(<$($T)*>)?($($args)*) -> $ret
120        where
121            $($clauses)*
122        $body
123    };
124    (@$tag:ident $(#[$meta:meta])* $vis:vis fn $name:ident$(<$($T:ident),+>)?($($args:tt)*) -> $ret:ty $body:block) => {
125        $crate::avx_helper!(@$tag $(#[$meta])* $vis fn $name$(<$($T)*>)?($($args)*) -> $ret where [] $body);
126    };
127    (@$tag:ident $(#[$meta:meta])* $vis:vis fn $name:ident$(<$($T:ident),+>)?($($args:tt)*) $($t:tt)*) => {
128        $crate::avx_helper!(@$tag $(#[$meta])* $vis fn $name$(<$($T)*>)?($($args)*) -> () $($t)*);
129    };
130    ($($t:tt)*) => {
131        ::std::compile_error!($($t)*);
132    }
133}