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/// Return `(free, total)` device-memory bytes for the selected
40/// device, or `None` when we're not on CUDA.
41///
42/// Thin wrapper around `cuMemGetInfo` via cudarc; exposed so the
43/// PLAID builder can annotate its progress output with the headroom
44/// the CUDA mempool has left before a big allocation.
45pub fn device_memory_info() -> Option<(usize, usize)> {
46    #[cfg(feature = "cuda")]
47    {
48        if let Device::Cuda(_) = default_device() {
49            use candle_core::cuda_backend::cudarc::driver::result;
50            if let Ok((free, total)) = result::mem_get_info() {
51                return Some((free, total));
52            }
53        }
54    }
55    None
56}
57
58/// Release cached but currently-unused device memory back to the
59/// driver.
60///
61/// Candle hands every dropped `Tensor` back to CUDA's async memory
62/// pool (via `cuMemFreeAsync`). The pool keeps the bytes around for
63/// fast reuse, which is normally what you want — but when the encoder
64/// model finishes embedding, its ~2 GB of ModernBert per-batch
65/// caches stay committed to the pool even after the callers drop the
66/// model. That's enough to block a subsequent 3.47 GB `Tensor`
67/// allocation on a 12 GB card and surface as `CUDA_ERROR_OUT_OF_MEMORY`.
68///
69/// This function asks CUDA to trim the default mempool to zero
70/// retained bytes, handing the freed blocks back to the driver so
71/// the next large allocation can grow into them. Without the `cuda`
72/// feature (or when running on CPU) it's a no-op.
73///
74/// Safe wrapper around `cuDeviceGetDefaultMemPool` +
75/// `cuMemPoolTrimTo`. Only trims device 0 — docbert currently only
76/// ever uses the default CUDA device.
77pub fn release_cached_device_memory() -> Result<(), candle_core::Error> {
78    #[cfg(feature = "cuda")]
79    {
80        if let Device::Cuda(_) = default_device() {
81            use candle_core::cuda_backend::cudarc::driver::result;
82            // Safety: `result::device::get` is a safe wrapper that
83            // returns a valid device handle; `get_default_mem_pool`
84            // and `trim_to` are unsafe because they take raw CUDA
85            // handles, but we only ever pass handles produced by
86            // cudarc in the same call — the documented preconditions
87            // ("valid device", "valid pool") hold. Trimming the pool
88            // while no outstanding allocations reference it is
89            // always-safe per the CUDA docs.
90            unsafe {
91                let dev = result::device::get(0).map_err(|e| {
92                    candle_core::Error::Msg(format!(
93                        "release_cached_device_memory: cuDeviceGet(0) failed: {e}"
94                    ))
95                })?;
96                let pool =
97                    result::device::get_default_mem_pool(dev).map_err(|e| {
98                        candle_core::Error::Msg(format!(
99                            "release_cached_device_memory: cuDeviceGetDefaultMemPool failed: {e}"
100                        ))
101                    })?;
102                result::mem_pool::trim_to(pool, 0).map_err(|e| {
103                    candle_core::Error::Msg(format!(
104                        "release_cached_device_memory: cuMemPoolTrimTo failed: {e}"
105                    ))
106                })?;
107            }
108        }
109    }
110    Ok(())
111}
112
113#[cfg(test)]
114mod tests {
115    use super::*;
116
117    #[test]
118    fn default_device_is_stable_across_calls() {
119        let a = default_device();
120        let b = default_device();
121        // Same backing static, so the references should compare as
122        // pointing at the same value.
123        assert!(std::ptr::eq(a, b));
124    }
125
126    #[test]
127    fn default_device_is_cpu_without_cuda_feature() {
128        // When the cuda feature is off, the choice is unconditional.
129        // When the feature is on but no GPU is present (CI without
130        // CUDA), select_device() also falls back to CPU. So this test
131        // only asserts the no-cuda behaviour to stay portable.
132        #[cfg(not(feature = "cuda"))]
133        {
134            assert!(matches!(default_device(), Device::Cpu));
135        }
136    }
137}