use std::ops::{Deref, DerefMut};
use crate::backend::{
BackendEntry, BackendError, BackendKind, BackendRegistry, BackendResult, ComputeBackend,
CpuBackend, SelectionRequest,
};
#[derive(Debug)]
pub struct SelectedBackend {
kind: BackendKind,
backend: Box<dyn ComputeBackend>,
}
impl SelectedBackend {
#[must_use]
pub fn kind(&self) -> BackendKind {
self.kind
}
#[must_use]
pub fn backend(&self) -> &dyn ComputeBackend {
self.backend.as_ref()
}
pub fn backend_mut(&mut self) -> &mut dyn ComputeBackend {
self.backend.as_mut()
}
#[must_use]
pub fn into_inner(self) -> Box<dyn ComputeBackend> {
self.backend
}
}
impl Deref for SelectedBackend {
type Target = dyn ComputeBackend;
fn deref(&self) -> &Self::Target {
self.backend.as_ref()
}
}
impl DerefMut for SelectedBackend {
fn deref_mut(&mut self) -> &mut Self::Target {
self.backend.as_mut()
}
}
fn cuda_driver_present() -> bool {
crate::Device::count().is_ok_and(|count| count > 0)
}
fn new_cuda() -> BackendResult<Box<dyn ComputeBackend>> {
if cuda_driver_present() {
Ok(Box::new(crate::backend::CudaBackend::new()))
} else {
Err(BackendError::Unsupported(
"no CUDA driver with a usable device on this machine".into(),
))
}
}
#[cfg(feature = "metal")]
fn new_metal() -> BackendResult<Box<dyn ComputeBackend>> {
Ok(Box::new(crate::backend::MetalBackend::new()))
}
#[cfg(not(feature = "metal"))]
fn new_metal() -> BackendResult<Box<dyn ComputeBackend>> {
Err(BackendError::Unsupported(
"not compiled in (enable feature \"metal\")".into(),
))
}
#[cfg(feature = "webgpu")]
fn new_webgpu() -> BackendResult<Box<dyn ComputeBackend>> {
Ok(Box::new(crate::backend::WebGpuBackend::new()))
}
#[cfg(not(feature = "webgpu"))]
fn new_webgpu() -> BackendResult<Box<dyn ComputeBackend>> {
Err(BackendError::Unsupported(
"not compiled in (enable feature \"webgpu\")".into(),
))
}
#[cfg(feature = "vulkan")]
fn new_vulkan() -> BackendResult<Box<dyn ComputeBackend>> {
Ok(Box::new(crate::backend::VulkanBackend::new()))
}
#[cfg(not(feature = "vulkan"))]
fn new_vulkan() -> BackendResult<Box<dyn ComputeBackend>> {
Err(BackendError::Unsupported(
"not compiled in (enable feature \"vulkan\")".into(),
))
}
#[cfg(feature = "rocm")]
fn new_rocm() -> BackendResult<Box<dyn ComputeBackend>> {
Ok(Box::new(crate::backend::RocmBackend::new()))
}
#[cfg(not(feature = "rocm"))]
fn new_rocm() -> BackendResult<Box<dyn ComputeBackend>> {
Err(BackendError::Unsupported(
"not compiled in (enable feature \"rocm\")".into(),
))
}
#[cfg(feature = "level-zero")]
fn new_level_zero() -> BackendResult<Box<dyn ComputeBackend>> {
Ok(Box::new(crate::backend::LevelZeroBackend::new()))
}
#[cfg(not(feature = "level-zero"))]
fn new_level_zero() -> BackendResult<Box<dyn ComputeBackend>> {
Err(BackendError::Unsupported(
"not compiled in (enable feature \"level-zero\")".into(),
))
}
fn instantiate(kind: BackendKind) -> BackendResult<Box<dyn ComputeBackend>> {
match kind {
BackendKind::Cuda => new_cuda(),
BackendKind::Rocm => new_rocm(),
BackendKind::LevelZero => new_level_zero(),
BackendKind::Vulkan => new_vulkan(),
BackendKind::Metal => new_metal(),
BackendKind::WebGpu => new_webgpu(),
BackendKind::Cpu => Ok(Box::new(CpuBackend::new())),
}
}
fn live_entry(kind: BackendKind, backend: &dyn ComputeBackend) -> BackendEntry {
BackendEntry::new(kind, true).with_capabilities(backend.capabilities())
}
#[must_use]
pub fn compiled_in_kinds() -> Vec<BackendKind> {
let mut kinds: Vec<BackendKind> = BackendKind::ALL
.into_iter()
.filter(|kind| match kind {
BackendKind::Cuda | BackendKind::Cpu => true,
BackendKind::Rocm => cfg!(feature = "rocm"),
BackendKind::LevelZero => cfg!(feature = "level-zero"),
BackendKind::Vulkan => cfg!(feature = "vulkan"),
BackendKind::Metal => cfg!(feature = "metal"),
BackendKind::WebGpu => cfg!(feature = "webgpu"),
})
.collect();
kinds.sort_by_key(|kind| std::cmp::Reverse(kind.default_priority()));
kinds
}
#[must_use]
pub fn default_registry() -> BackendRegistry {
let mut registry = BackendRegistry::new();
for kind in BackendKind::ALL {
let entry = match instantiate(kind) {
Ok(mut backend) => match backend.init() {
Ok(()) => live_entry(kind, backend.as_ref()),
Err(_) => BackendEntry::new(kind, false),
},
Err(_) => BackendEntry::new(kind, false),
};
registry.register(entry);
}
registry
}
pub fn select_backend(req: &SelectionRequest) -> BackendResult<SelectedBackend> {
let structural = SelectionRequest {
require_gpu: req.require_gpu,
pin: req.pin,
..SelectionRequest::any()
};
let mut candidates = BackendRegistry::new();
for kind in compiled_in_kinds() {
candidates.register(BackendEntry::new(kind, true));
}
let mut rejected: Vec<String> = Vec::new();
for kind in candidates.fallback_chain(&structural) {
let mut backend = match instantiate(kind) {
Ok(backend) => backend,
Err(e) => {
rejected.push(format!("{kind}: {e}"));
continue;
}
};
if let Err(e) = backend.init() {
rejected.push(format!("{kind}: init failed ({e})"));
continue;
}
let entry = live_entry(kind, backend.as_ref());
if !req.is_satisfied_by(&entry) {
rejected.push(format!("{kind}: [{}] does not satisfy", entry.capabilities));
continue;
}
return Ok(SelectedBackend { kind, backend });
}
Err(BackendError::Unsupported(format!(
"no compute backend satisfies {req:?} — rejected: {}",
if rejected.is_empty() {
"nothing was compiled in".to_string()
} else {
rejected.join("; ")
}
)))
}
pub fn default_backend() -> BackendResult<SelectedBackend> {
select_backend(&SelectionRequest::any())
}
pub fn gpu_backend() -> BackendResult<SelectedBackend> {
select_backend(&SelectionRequest::require_gpu())
}
pub fn backend_for_workload(workload_bytes: usize) -> BackendResult<SelectedBackend> {
let req =
SelectionRequest::any().for_workload(workload_bytes, crate::AUTO_SELECT_THRESHOLD_BYTES);
match select_backend(&req) {
Ok(selected) => Ok(selected),
Err(e) if req != SelectionRequest::any() => {
select_backend(&SelectionRequest::any()).map_err(|_| e)
}
Err(e) => Err(e),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::UnaryOp;
fn to_bytes(values: &[f32]) -> Vec<u8> {
values.iter().flat_map(|v| v.to_ne_bytes()).collect()
}
fn from_bytes(bytes: &[u8]) -> Vec<f32> {
bytes
.chunks_exact(4)
.map(|c| f32::from_ne_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
#[test]
fn compiled_in_kinds_always_offers_the_host_fallback() {
let kinds = compiled_in_kinds();
assert!(
kinds.contains(&BackendKind::Cpu),
"the CPU reference backend must always be constructible"
);
assert_eq!(
kinds.last(),
Some(&BackendKind::Cpu),
"the host backend must rank last: {kinds:?}"
);
}
#[test]
fn compiled_in_kinds_track_the_enabled_features() {
let kinds = compiled_in_kinds();
assert_eq!(kinds.contains(&BackendKind::Metal), cfg!(feature = "metal"));
assert_eq!(
kinds.contains(&BackendKind::WebGpu),
cfg!(feature = "webgpu")
);
}
#[test]
fn compiled_in_kinds_are_ordered_by_priority() {
let kinds = compiled_in_kinds();
for pair in kinds.windows(2) {
assert!(
pair[0].default_priority() >= pair[1].default_priority(),
"selection order must be non-increasing in priority: {kinds:?}"
);
}
}
#[test]
fn default_registry_lists_every_kind_with_cpu_available() {
let registry = default_registry();
assert_eq!(registry.len(), BackendKind::ALL.len());
let cpu = registry
.get(BackendKind::Cpu)
.expect("the CPU entry must be registered");
assert!(cpu.available, "the CPU backend is always available");
assert_eq!(
registry.select_best().expect("selection must succeed"),
default_backend()
.expect("default_backend must succeed")
.kind(),
"the probing registry and the lazy selection must agree"
);
}
#[test]
fn unavailable_backends_are_registered_but_not_selected() {
let registry = default_registry();
for kind in BackendKind::ALL {
let entry = registry.get(kind).expect("every kind must be registered");
assert_eq!(entry.kind, kind);
}
if !cfg!(feature = "vulkan") {
let vulkan = registry
.get(BackendKind::Vulkan)
.expect("Vulkan must still be listed");
assert!(!vulkan.available);
}
}
#[test]
fn cuda_is_never_selected_without_a_driver() {
if cuda_driver_present() {
return; }
let selected = default_backend().expect("a backend must always be available");
assert_ne!(
selected.kind(),
BackendKind::Cuda,
"CudaBackend::init succeeds without a GPU, so it must be filtered by the probe"
);
}
#[cfg(all(target_os = "macos", feature = "metal"))]
#[test]
fn macos_with_metal_feature_selects_metal() {
let metal_opens = {
let mut probe = crate::backend::MetalBackend::new();
probe.init().is_ok()
};
let selected = default_backend().expect("a backend must always be available");
if metal_opens {
assert_eq!(
selected.kind(),
BackendKind::Metal,
"with a working Metal device the facade must select it"
);
} else {
assert_eq!(
selected.kind(),
BackendKind::Cpu,
"without a Metal device the facade must degrade to the host"
);
}
}
fn relu_through(backend: &dyn ComputeBackend, values: &[f32]) -> Vec<f32> {
let bytes = std::mem::size_of_val(values);
let input = backend.alloc(bytes).expect("alloc must succeed");
let output = backend.alloc(bytes).expect("alloc must succeed");
backend
.copy_htod(input, &to_bytes(values))
.expect("host→device copy must succeed");
backend
.unary(UnaryOp::Relu, input, output, values.len())
.expect("relu must succeed");
backend.synchronize().expect("synchronize must succeed");
let mut host = vec![0u8; bytes];
backend
.copy_dtoh(&mut host, output)
.expect("device→host copy must succeed");
backend.free(input).expect("free must succeed");
backend.free(output).expect("free must succeed");
from_bytes(&host)
}
#[cfg(all(target_os = "macos", feature = "metal"))]
#[test]
fn metal_agrees_with_the_cpu_reference_on_negative_inputs() {
let mut metal = crate::backend::MetalBackend::new();
if metal.init().is_err() {
return; }
assert_eq!(metal.name(), "metal");
let mut cpu = CpuBackend::new();
cpu.init().expect("the host backend must initialise");
let values = [-4.0f32, -0.25, 0.0, 0.5, 9.0, -1e-3, 1e3, -7.5];
let on_metal = relu_through(&metal, &values);
let on_cpu = relu_through(&cpu, &values);
assert_eq!(on_metal.len(), values.len());
for (i, (m, c)) in on_metal.iter().zip(on_cpu.iter()).enumerate() {
assert!(
(m - c).abs() < 1e-6,
"relu[{i}]: Metal produced {m}, the CPU reference {c}"
);
}
assert_ne!(
on_metal.as_slice(),
values.as_slice(),
"Metal returned its input unchanged — relu did not run"
);
}
#[cfg(all(target_os = "macos", not(feature = "metal")))]
#[test]
fn macos_without_metal_feature_selects_cpu() {
let selected = default_backend().expect("a backend must always be available");
assert_eq!(
selected.kind(),
BackendKind::Cpu,
"macOS has no CUDA driver, so the host backend must win"
);
}
#[test]
fn gpu_backend_returns_a_gpu_or_an_explanatory_error() {
match gpu_backend() {
Ok(selected) => assert!(
selected.kind().is_gpu(),
"require_gpu must never yield the host backend"
),
Err(e) => {
let msg = e.to_string();
assert!(
msg.contains("rejected") || msg.contains("no compute backend"),
"the error must explain what was tried, got: {msg}"
);
}
}
}
#[test]
fn small_workloads_stay_on_the_host() {
let selected = backend_for_workload(1024).expect("selection must succeed");
assert_eq!(
selected.kind(),
BackendKind::Cpu,
"1 KiB is below AUTO_SELECT_THRESHOLD_BYTES"
);
assert!(selected.is_initialized());
}
#[test]
fn large_workloads_use_the_default_backend() {
let large = backend_for_workload(crate::AUTO_SELECT_THRESHOLD_BYTES * 16)
.expect("selection must succeed");
let default = default_backend().expect("selection must succeed");
assert_eq!(large.kind(), default.kind());
}
#[test]
fn threshold_boundary_is_inclusive_for_the_gpu_side() {
let at = backend_for_workload(crate::AUTO_SELECT_THRESHOLD_BYTES)
.expect("selection must succeed");
let default = default_backend().expect("selection must succeed");
assert_eq!(
at.kind(),
default.kind(),
"exactly at the threshold is not below it"
);
let below = backend_for_workload(crate::AUTO_SELECT_THRESHOLD_BYTES - 1)
.expect("selection must succeed");
assert_eq!(below.kind(), BackendKind::Cpu);
}
#[test]
fn default_backend_runs_a_real_compute_round_trip() {
let backend = default_backend().expect("a backend must always be available");
assert!(
backend.is_initialized(),
"the selected backend must come back initialised"
);
if backend.kind() == BackendKind::Cuda && !cfg!(feature = "ptx") {
eprintln!(
"default_backend_runs_a_real_compute_round_trip: skipping -- Cuda selected \
without the `ptx` feature, which cannot run compute ops by design"
);
return;
}
let input_values = [-2.5f32, -0.5, 0.0, 1.5, 3.25, -7.0, 0.125, 42.0];
let expected = [0.0f32, 0.0, 0.0, 1.5, 3.25, 0.0, 0.125, 42.0];
let got = relu_through(backend.backend(), &input_values);
assert_eq!(got.len(), expected.len());
for (i, (g, e)) in got.iter().zip(expected.iter()).enumerate() {
assert!(
(g - e).abs() < 1e-6,
"relu[{i}] on the {} backend: got {g}, expected {e}",
backend.name()
);
}
}
#[test]
fn selected_backend_can_be_shared_as_an_arc() {
let selected = default_backend().expect("a backend must always be available");
let name = selected.name().to_string();
let shared: std::sync::Arc<dyn ComputeBackend> =
std::sync::Arc::from(selected.into_inner());
assert_eq!(shared.name(), name, "sharing must not change the backend");
assert!(
shared.is_initialized(),
"the shared backend stays initialised"
);
}
}