use static_assertions::assert_impl_all;
use super::backend::Backend;
assert_impl_all!(Formula: Send, Sync);
assert_impl_all!(Precision: Send, Sync);
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Formula {
Gemm,
Map,
WindowProduct,
ReduceWindow,
BatchNormTraining,
BatchNormInference,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Precision {
F32,
F64,
}
impl Formula {
pub const ALL: &'static [Formula] = &[
Formula::Gemm,
Formula::Map,
Formula::WindowProduct,
Formula::ReduceWindow,
Formula::BatchNormTraining,
Formula::BatchNormInference,
];
pub const fn chain(self, precision: Precision) -> &'static [Backend] {
match self {
Formula::Gemm => match precision {
Precision::F32 => &[
Backend::Accelerate,
Backend::Metal,
Backend::Cuda,
Backend::Simd,
],
Precision::F64 => &[Backend::Accelerate, Backend::Cuda, Backend::Simd],
},
Formula::Map => match precision {
Precision::F32 => &[Backend::Metal, Backend::Accelerate],
Precision::F64 => &[Backend::Accelerate],
},
Formula::BatchNormTraining => &[Backend::Accelerate],
Formula::WindowProduct | Formula::ReduceWindow | Formula::BatchNormInference => &[],
}
}
}
impl Precision {
pub const ALL: &'static [Precision] = &[Precision::F32, Precision::F64];
}
#[cfg(test)]
#[path = "tests/formula_tests.rs"]
mod tests;