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}