Skip to main content

docbert_plaid/
device.rs

1//! Device selection for the candle-backed inner loops.
2//!
3//! Picks CUDA when the crate was built with the `cuda` feature *and* a
4//! GPU is actually present at runtime; falls back to CPU otherwise. The
5//! choice is made once on first call and cached for the rest of the
6//! process — every subsequent tensor allocation lands on the same
7//! device, so there are no surprise host↔device transfers in the
8//! middle of a tight loop.
9//!
10//! Callers shouldn't need to think about devices in normal use:
11//! `default_device()` returns whatever's right for this build.
12
13use std::sync::OnceLock;
14
15use candle_core::Device;
16
17static DEVICE: OnceLock<Device> = OnceLock::new();
18
19/// Return the device docbert-plaid uses for all tensor compute.
20///
21/// Cached on first call. The choice is:
22///
23/// - With the `cuda` feature **and** a usable CUDA device → `Device::Cuda(0)`.
24/// - Otherwise → `Device::Cpu`.
25pub fn default_device() -> &'static Device {
26    DEVICE.get_or_init(select_device)
27}
28
29fn select_device() -> Device {
30    #[cfg(feature = "cuda")]
31    {
32        if let Ok(dev) = Device::new_cuda(0) {
33            return dev;
34        }
35    }
36    Device::Cpu
37}
38
39#[cfg(test)]
40mod tests {
41    use super::*;
42
43    #[test]
44    fn default_device_is_stable_across_calls() {
45        let a = default_device();
46        let b = default_device();
47        // Same backing static, so the references should compare as
48        // pointing at the same value.
49        assert!(std::ptr::eq(a, b));
50    }
51
52    #[test]
53    fn default_device_is_cpu_without_cuda_feature() {
54        // When the cuda feature is off, the choice is unconditional.
55        // When the feature is on but no GPU is present (CI without
56        // CUDA), select_device() also falls back to CPU. So this test
57        // only asserts the no-cuda behaviour to stay portable.
58        #[cfg(not(feature = "cuda"))]
59        {
60            assert!(matches!(default_device(), Device::Cpu));
61        }
62    }
63}