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}