use std::sync::OnceLock;
use candle_core::Device;
static DEVICE: OnceLock<Device> = OnceLock::new();
pub fn default_device() -> &'static Device {
DEVICE.get_or_init(select_device)
}
fn select_device() -> Device {
#[cfg(feature = "cuda")]
{
if let Ok(dev) = Device::new_cuda(0) {
return dev;
}
}
Device::Cpu
}
pub fn device_memory_info() -> Option<(usize, usize)> {
#[cfg(feature = "cuda")]
{
if let Device::Cuda(_) = default_device() {
use candle_core::cuda_backend::cudarc::driver::result;
if let Ok((free, total)) = result::mem_get_info() {
return Some((free, total));
}
}
}
None
}
pub fn release_cached_device_memory() -> Result<(), candle_core::Error> {
#[cfg(feature = "cuda")]
{
if let Device::Cuda(_) = default_device() {
use candle_core::cuda_backend::cudarc::driver::result;
unsafe {
let dev = result::device::get(0).map_err(|e| {
candle_core::Error::Msg(format!(
"release_cached_device_memory: cuDeviceGet(0) failed: {e}"
))
})?;
let pool =
result::device::get_default_mem_pool(dev).map_err(|e| {
candle_core::Error::Msg(format!(
"release_cached_device_memory: cuDeviceGetDefaultMemPool failed: {e}"
))
})?;
result::mem_pool::trim_to(pool, 0).map_err(|e| {
candle_core::Error::Msg(format!(
"release_cached_device_memory: cuMemPoolTrimTo failed: {e}"
))
})?;
}
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_device_is_stable_across_calls() {
let a = default_device();
let b = default_device();
assert!(std::ptr::eq(a, b));
}
#[test]
fn default_device_is_cpu_without_cuda_feature() {
#[cfg(not(feature = "cuda"))]
{
assert!(matches!(default_device(), Device::Cpu));
}
}
}