use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use crate::AccelerationMode;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum DeviceKind {
#[default]
Auto,
Cpu,
Gpu,
Npu,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ExecutionProviderSpec {
pub name: String,
#[serde(default)]
pub options: BTreeMap<String, serde_json::Value>,
}
impl ExecutionProviderSpec {
pub fn new(name: impl Into<String>) -> Self {
Self {
name: name.into(),
options: BTreeMap::new(),
}
}
pub fn with_option(
mut self,
key: impl Into<String>,
value: impl Into<serde_json::Value>,
) -> Self {
self.options.insert(key.into(), value.into());
self
}
pub fn cpu() -> Self {
Self::new("cpu")
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct RuntimeOptions {
#[serde(default)]
pub device: DeviceKind,
#[serde(default)]
pub providers: Vec<ExecutionProviderSpec>,
#[serde(default, alias = "max_threads")]
pub max_threads: usize,
#[serde(default = "default_true", alias = "graph_optimization")]
pub graph_optimization: bool,
#[serde(default, flatten)]
pub extra: BTreeMap<String, serde_json::Value>,
}
const fn default_true() -> bool {
true
}
impl Default for RuntimeOptions {
fn default() -> Self {
Self {
device: DeviceKind::Auto,
providers: Vec::new(),
max_threads: 0,
graph_optimization: true,
extra: BTreeMap::new(),
}
}
}
impl RuntimeOptions {
pub fn cpu() -> Self {
Self {
device: DeviceKind::Cpu,
providers: vec![ExecutionProviderSpec::cpu()],
..Self::default()
}
}
pub fn from_acceleration(mode: AccelerationMode) -> Self {
match mode {
AccelerationMode::Auto => Self::default(),
AccelerationMode::Cpu => Self::cpu(),
AccelerationMode::Gpu => Self {
device: DeviceKind::Gpu,
providers: legacy_gpu_provider_order(),
..Self::default()
},
}
}
pub fn legacy_acceleration(&self) -> AccelerationMode {
if self.device == DeviceKind::Cpu
|| (!self.providers.is_empty()
&& self
.providers
.iter()
.all(|provider| provider.name.eq_ignore_ascii_case("cpu")))
{
AccelerationMode::Cpu
} else if self.device == DeviceKind::Gpu
|| self.providers.iter().any(|provider| {
matches!(
provider.name.to_ascii_lowercase().as_str(),
"cuda" | "directml" | "openvino" | "tensorrt" | "coreml" | "qnn"
)
})
{
AccelerationMode::Gpu
} else {
AccelerationMode::Auto
}
}
}
impl From<AccelerationMode> for RuntimeOptions {
fn from(value: AccelerationMode) -> Self {
Self::from_acceleration(value)
}
}
fn legacy_gpu_provider_order() -> Vec<ExecutionProviderSpec> {
vec![
#[cfg(any(
all(target_os = "windows", target_arch = "x86_64"),
all(
target_os = "linux",
any(target_arch = "x86_64", target_arch = "aarch64")
)
))]
ExecutionProviderSpec::new("tensorrt"),
#[cfg(any(target_os = "windows", target_os = "linux"))]
ExecutionProviderSpec::new("cuda"),
#[cfg(target_os = "windows")]
ExecutionProviderSpec::new("directml"),
#[cfg(target_vendor = "apple")]
ExecutionProviderSpec::new("coreml"),
ExecutionProviderSpec::cpu(),
]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn legacy_acceleration_round_trips_without_becoming_a_runtime() {
for mode in [
AccelerationMode::Auto,
AccelerationMode::Cpu,
AccelerationMode::Gpu,
] {
let options = RuntimeOptions::from(mode);
assert_eq!(options.legacy_acceleration(), mode);
}
}
#[cfg(any(
all(target_os = "windows", target_arch = "x86_64"),
all(
target_os = "linux",
any(target_arch = "x86_64", target_arch = "aarch64")
)
))]
#[test]
fn legacy_gpu_prefers_tensorrt_before_cuda() {
let options = RuntimeOptions::from_acceleration(AccelerationMode::Gpu);
let names: Vec<_> = options
.providers
.iter()
.map(|provider| provider.name.as_str())
.collect();
assert_eq!(&names[..2], &["tensorrt", "cuda"]);
assert_eq!(names.last(), Some(&"cpu"));
}
}