//! Execution provider configuration for ONNX Runtime //! //! Supports CPU, CUDA, CoreML, TensorRT, and DirectML execution providers. use serde::{Deserialize, Serialize}; /// Execution provider type with configuration options #[derive(Debug, Clone, Serialize, Deserialize)] pub enum ExecutionProviderType { /// CPU execution (always available) CPU(CpuOptions), /// NVIDIA CUDA execution #[cfg(feature = "cuda")] CUDA(CudaOptions), /// Apple CoreML execution (macOS/iOS) #[cfg(feature = "coreml")] CoreML(CoreMLOptions), /// NVIDIA TensorRT execution #[cfg(feature = "tensorrt")] TensorRT(TensorRTOptions), /// DirectML execution (Windows) #[cfg(feature = "directml")] DirectML(DirectMLOptions), } impl Default for ExecutionProviderType { fn default() -> Self { ExecutionProviderType::CPU(CpuOptions::default()) } } /// CPU execution provider options #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct CpuOptions { /// Number of intra-op threads (0 = auto) pub intra_op_threads: usize, /// Number of inter-op threads (0 = auto) pub inter_op_threads: usize, /// Enable memory arena pub enable_cpu_mem_arena: bool, } /// CUDA execution provider options #[cfg(feature = "cuda")] #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CudaOptions { /// CUDA device ID pub device_id: i32, /// GPU memory limit in bytes (0 = no limit) pub gpu_mem_limit: usize, /// Arena extend strategy pub arena_extend_strategy: ArenaExtendStrategy, /// Enable CUDA graphs for reduced kernel launch overhead pub enable_cuda_graph: bool, /// cuDNN convolution algorithm search mode pub cudnn_conv_algo_search: CudnnConvAlgoSearch, } #[cfg(feature = "cuda")] impl Default for CudaOptions { fn default() -> Self { Self { device_id: 0, gpu_mem_limit: 0, arena_extend_strategy: ArenaExtendStrategy::NextPowerOfTwo, enable_cuda_graph: false, cudnn_conv_algo_search: CudnnConvAlgoSearch::Exhaustive, } } } /// Arena memory extend strategy #[cfg(feature = "cuda")] #[derive(Debug, Clone, Copy, Serialize, Deserialize)] pub enum ArenaExtendStrategy { /// Extend to next power of two NextPowerOfTwo, /// Extend by exact amount needed SameAsRequested, } /// cuDNN convolution algorithm search mode #[cfg(feature = "cuda")] #[derive(Debug, Clone, Copy, Serialize, Deserialize)] pub enum CudnnConvAlgoSearch { /// Exhaustive search (slowest startup, best performance) Exhaustive, /// Heuristic search (balanced) Heuristic, /// Default algorithm (fastest startup) Default, } /// CoreML execution provider options #[cfg(feature = "coreml")] #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct CoreMLOptions { /// Use CPU only (disable Neural Engine and GPU) pub use_cpu_only: bool, /// Enable on subgraphs pub enable_on_subgraph: bool, /// Only enable on devices with Apple Neural Engine pub only_enable_device_with_ane: bool, } /// TensorRT execution provider options #[cfg(feature = "tensorrt")] #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TensorRTOptions { /// Device ID pub device_id: i32, /// Maximum workspace size in bytes pub max_workspace_size: usize, /// Enable FP16 precision pub fp16_enable: bool, /// Enable INT8 precision pub int8_enable: bool, /// Engine cache path pub engine_cache_path: Option, } #[cfg(feature = "tensorrt")] impl Default for TensorRTOptions { fn default() -> Self { Self { device_id: 0, max_workspace_size: 1 << 30, // 1GB fp16_enable: true, int8_enable: false, engine_cache_path: None, } } } /// DirectML execution provider options #[cfg(feature = "directml")] #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct DirectMLOptions { /// Device ID pub device_id: i32, } /// Detect the best available execution provider based on hardware pub fn detect_best_provider() -> ExecutionProviderType { // Check for CUDA availability #[cfg(feature = "cuda")] { if is_cuda_available() { tracing::info!("CUDA execution provider detected"); return ExecutionProviderType::CUDA(CudaOptions::default()); } } // Check for CoreML on macOS #[cfg(all(feature = "coreml", target_os = "macos"))] { if is_apple_silicon() { tracing::info!("CoreML execution provider detected (Apple Silicon)"); return ExecutionProviderType::CoreML(CoreMLOptions::default()); } } // Check for TensorRT #[cfg(feature = "tensorrt")] { if is_tensorrt_available() { tracing::info!("TensorRT execution provider detected"); return ExecutionProviderType::TensorRT(TensorRTOptions::default()); } } // Check for DirectML on Windows #[cfg(all(feature = "directml", target_os = "windows"))] { tracing::info!("DirectML execution provider detected"); return ExecutionProviderType::DirectML(DirectMLOptions::default()); } // Fall back to CPU tracing::info!("Using CPU execution provider"); ExecutionProviderType::CPU(CpuOptions::default()) } /// Check if CUDA is available #[cfg(feature = "cuda")] fn is_cuda_available() -> bool { // Simple check - in production this would query CUDA runtime std::env::var("CUDA_VISIBLE_DEVICES").is_ok() || std::path::Path::new("/usr/local/cuda").exists() } /// Check if running on Apple Silicon #[cfg(all(feature = "coreml", target_os = "macos"))] fn is_apple_silicon() -> bool { #[cfg(target_arch = "aarch64")] { true } #[cfg(not(target_arch = "aarch64"))] { false } } /// Check if TensorRT is available #[cfg(feature = "tensorrt")] fn is_tensorrt_available() -> bool { std::path::Path::new("/usr/lib/x86_64-linux-gnu/libnvinfer.so").exists() || std::env::var("TENSORRT_ROOT").is_ok() } impl ExecutionProviderType { /// Get a human-readable name for the execution provider pub fn name(&self) -> &'static str { match self { ExecutionProviderType::CPU(_) => "CPU", #[cfg(feature = "cuda")] ExecutionProviderType::CUDA(_) => "CUDA", #[cfg(feature = "coreml")] ExecutionProviderType::CoreML(_) => "CoreML", #[cfg(feature = "tensorrt")] ExecutionProviderType::TensorRT(_) => "TensorRT", #[cfg(feature = "directml")] ExecutionProviderType::DirectML(_) => "DirectML", } } }