use static_assertions::assert_impl_all;
use thiserror::Error;
use super::accelerate::Accelerate;
use super::coverage::{Coverage, Dispatch};
use super::cuda::Cuda;
use super::formula::{Formula, Precision};
use super::fused::Fused;
use super::manifest::Manifest;
use super::metal::Metal;
use super::simd::Simd;
use super::stablehlo::StableHlo;
assert_impl_all!(Backend: Send, Sync);
assert_impl_all!(BackendUnavailable: Send, Sync);
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Backend {
Accelerate,
Metal,
Cuda,
Simd,
Fused,
StableHlo,
}
impl Backend {
pub const ALL: &'static [Backend] = &[
Backend::Accelerate,
Backend::Metal,
Backend::Cuda,
Backend::Simd,
Backend::Fused,
Backend::StableHlo,
];
pub fn coverage(self, formula: Formula) -> Coverage {
match self {
Backend::Accelerate => Accelerate::coverage(formula),
Backend::Metal => Metal::coverage(formula),
Backend::Cuda => Cuda::coverage(formula),
Backend::Simd => Simd::coverage(formula),
Backend::Fused => Fused::coverage(formula),
Backend::StableHlo => StableHlo::coverage(formula),
}
}
pub fn serves(self, formula: Formula, precision: Precision) -> bool {
self.coverage(formula).admits(precision)
}
pub fn dispatch(self) -> Dispatch {
match self {
Backend::Accelerate => Accelerate::DISPATCH,
Backend::Metal => Metal::DISPATCH,
Backend::Cuda => Cuda::DISPATCH,
Backend::Simd => Simd::DISPATCH,
Backend::Fused => Fused::DISPATCH,
Backend::StableHlo => StableHlo::DISPATCH,
}
}
pub fn compiled(self) -> bool {
match self {
Backend::Accelerate => Accelerate::compiled(),
Backend::Metal => Metal::compiled(),
Backend::Cuda => Cuda::compiled(),
Backend::Simd => Simd::compiled(),
Backend::Fused => Fused::compiled(),
Backend::StableHlo => StableHlo::compiled(),
}
}
pub fn status(self) -> Result<(), BackendUnavailable> {
match self {
Backend::Accelerate => Accelerate::status(),
Backend::Metal => Metal::status(),
Backend::Cuda => Cuda::status(),
Backend::Simd => Simd::status(),
Backend::Fused => Fused::status(),
Backend::StableHlo => StableHlo::status(),
}
}
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum BackendUnavailable {
#[error("the backend's cargo feature is off in this build")]
NotCompiled,
#[error("this platform has no such backend")]
PlatformUnsupported,
#[error("backend setup failed: {0}")]
Initialization(String),
#[error("backend disabled after a runtime error: {0}")]
Poisoned(String),
}
#[cfg(test)]
#[path = "tests/backend_tests.rs"]
mod tests;