use thiserror::Error;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum GpuBackend {
Cuda,
Rocm,
Metal,
Wgpu,
Cpu,
}
impl Default for GpuBackend {
fn default() -> Self {
#[cfg(target_os = "macos")]
return Self::Metal;
#[cfg(not(target_os = "macos"))]
return Self::Cuda;
}
}
#[derive(Debug, Error)]
pub enum BackendError {
#[error("Backend not available: {backend:?}")]
NotAvailable { backend: GpuBackend },
#[error("Backend initialization failed: {reason}")]
InitializationFailed { reason: String },
#[error("Operation not supported by backend: {operation}")]
UnsupportedOperation { operation: String },
#[error("Backend error: {message}")]
BackendSpecific { message: String },
#[error("Device error: {device_id}")]
DeviceError { device_id: u32 },
}
#[derive(Debug, Clone)]
pub struct DeviceCapabilities {
pub name: String,
pub total_memory: usize,
pub available_memory: usize,
pub supports_f16: bool,
pub supports_bf16: bool,
pub supports_tensor_cores: bool,
pub max_threads_per_block: u32,
pub max_shared_memory_per_block: usize,
pub multiprocessor_count: u32,
pub compute_capability: (u32, u32),
}
#[derive(Debug, Clone)]
pub struct LaunchConfig {
pub grid_size: (u32, u32, u32),
pub block_size: (u32, u32, u32),
pub shared_memory_size: usize,
pub stream: Option<u64>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_backend_matches_the_build_platform() {
let backend = GpuBackend::default();
#[cfg(target_os = "macos")]
assert_eq!(backend, GpuBackend::Metal);
#[cfg(not(target_os = "macos"))]
assert_eq!(backend, GpuBackend::Cuda);
}
#[test]
fn device_capabilities_is_a_plain_data_struct() {
let caps = DeviceCapabilities {
name: "test device".to_string(),
total_memory: 1024,
available_memory: 512,
supports_f16: false,
supports_bf16: false,
supports_tensor_cores: false,
max_threads_per_block: 256,
max_shared_memory_per_block: 0,
multiprocessor_count: 1,
compute_capability: (0, 0),
};
assert_eq!(caps.name, "test device");
assert_eq!(caps.total_memory, 1024);
}
}