#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CpuBackend {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
SimdX86,
#[cfg(feature = "mlas")]
Mlas,
#[cfg(target_os = "android")]
Xnnpack,
#[cfg(any(target_os = "macos", target_os = "ios"))]
Accelerate,
Generic,
}
impl CpuBackend {
pub fn auto_detect() -> Self {
if let Some(backend) = Self::from_env_override(std::env::var("NXRT_CPU_GEMM_BACKEND").ok())
{
return backend;
}
#[cfg(target_os = "android")]
{
Self::Xnnpack
}
#[cfg(any(target_os = "macos", target_os = "ios"))]
{
Self::Accelerate
}
#[cfg(all(
not(target_os = "android"),
not(target_os = "macos"),
not(target_os = "ios")
))]
{
#[cfg(all(feature = "mlas", target_arch = "x86_64"))]
{
Self::Mlas
}
#[cfg(not(all(feature = "mlas", target_arch = "x86_64")))]
{
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if has_simd_x86() {
return Self::SimdX86;
}
}
Self::Generic
}
}
}
fn from_env_override(value: Option<String>) -> Option<Self> {
let value = value?;
if value.eq_ignore_ascii_case("generic") {
return Some(Self::Generic);
}
if value.eq_ignore_ascii_case("simd") {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
return Some(Self::simd_x86_or_generic(has_simd_x86()));
#[cfg(not(any(target_arch = "x86", target_arch = "x86_64")))]
return Some(Self::Generic);
}
if value.eq_ignore_ascii_case("mlas") {
#[cfg(all(feature = "mlas", target_arch = "x86_64"))]
return Some(Self::Mlas);
}
None
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
fn simd_x86_or_generic(supported: bool) -> Self {
if supported {
Self::SimdX86
} else {
Self::Generic
}
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[inline]
pub fn has_simd_x86() -> bool {
#[cfg(test)]
if matches!(
std::env::var("ONNX_RUNTIME_EP_CPU_FORCE_NO_SIMD_X86").as_deref(),
Ok("1")
) {
return false;
}
std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn auto_detect_is_stable() {
assert_eq!(CpuBackend::auto_detect(), CpuBackend::auto_detect());
}
#[cfg(all(
not(target_os = "android"),
not(target_os = "macos"),
not(target_os = "ios")
))]
#[test]
fn auto_detect_tracks_simd_x86_support() {
let expected = {
#[cfg(all(feature = "mlas", target_arch = "x86_64"))]
{
CpuBackend::Mlas
}
#[cfg(all(
not(all(feature = "mlas", target_arch = "x86_64")),
any(target_arch = "x86", target_arch = "x86_64")
))]
{
if has_simd_x86() {
CpuBackend::SimdX86
} else {
CpuBackend::Generic
}
}
#[cfg(not(any(target_arch = "x86", target_arch = "x86_64")))]
{
CpuBackend::Generic
}
};
assert_eq!(CpuBackend::auto_detect(), expected);
}
#[test]
fn backend_env_override_is_case_insensitive() {
assert_eq!(
CpuBackend::from_env_override(Some("GeNeRiC".into())),
Some(CpuBackend::Generic)
);
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
assert_eq!(
CpuBackend::from_env_override(Some("SIMD".into())),
Some(CpuBackend::simd_x86_or_generic(has_simd_x86()))
);
#[cfg(all(feature = "mlas", target_arch = "x86_64"))]
assert_eq!(
CpuBackend::from_env_override(Some("mLaS".into())),
Some(CpuBackend::Mlas)
);
assert_eq!(CpuBackend::from_env_override(Some("unknown".into())), None);
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[test]
fn forced_simd_falls_back_to_generic_without_required_cpu_features() {
assert_eq!(CpuBackend::simd_x86_or_generic(false), CpuBackend::Generic);
}
}