use objc2_core_ml::MLComputeUnits;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display)]
#[display("{}", self.as_str())]
#[non_exhaustive]
pub enum ComputeUnits {
CpuOnly,
CpuAndGpu,
CpuAndNeuralEngine,
#[default]
All,
}
impl ComputeUnits {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::CpuOnly => "cpu_only",
Self::CpuAndGpu => "cpu_and_gpu",
Self::CpuAndNeuralEngine => "cpu_and_neural_engine",
Self::All => "all",
}
}
#[inline(always)]
pub(crate) const fn to_raw(self) -> MLComputeUnits {
match self {
Self::CpuOnly => MLComputeUnits::CPUOnly,
Self::CpuAndGpu => MLComputeUnits::CPUAndGPU,
Self::CpuAndNeuralEngine => MLComputeUnits::CPUAndNeuralEngine,
Self::All => MLComputeUnits::All,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("unknown compute units name")]
pub struct ParseComputeUnitsError(());
impl core::str::FromStr for ComputeUnits {
type Err = ParseComputeUnitsError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(match s {
"cpu_only" => Self::CpuOnly,
"cpu_and_gpu" => Self::CpuAndGpu,
"cpu_and_neural_engine" => Self::CpuAndNeuralEngine,
"all" => Self::All,
_ => return Err(ParseComputeUnitsError(())),
})
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for ComputeUnits {
#[inline]
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(self.as_str())
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for ComputeUnits {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let name = <String as serde::Deserialize>::deserialize(deserializer)?;
name.parse::<Self>().map_err(serde::de::Error::custom)
}
}
#[cfg(test)]
mod tests;