Skip to main content

openvm_cuda_common/memory_manager/
mod.rs

1use std::{
2    collections::HashMap,
3    ffi::c_void,
4    ptr::NonNull,
5    sync::{Mutex, OnceLock},
6    time::{SystemTime, UNIX_EPOCH},
7};
8
9use bytesize::ByteSize;
10
11use crate::{
12    error::{check, MemoryError},
13    stream::{cudaStream_t, device_synchronize, StreamGuard},
14};
15
16mod cuda;
17mod vm_pool;
18use vm_pool::VirtualMemoryPool;
19
20#[cfg(test)]
21mod tests;
22
23#[link(name = "cudart")]
24extern "C" {
25    fn cudaMallocAsync(dev_ptr: *mut *mut c_void, size: usize, stream: cudaStream_t) -> i32;
26    fn cudaFreeAsync(dev_ptr: *mut c_void, stream: cudaStream_t) -> i32;
27    fn cudaMemGetInfo(free: *mut usize, total: *mut usize) -> i32;
28}
29
30static MEMORY_MANAGER: OnceLock<Mutex<MemoryManager>> = OnceLock::new();
31
32pub fn device_memory_used() -> usize {
33    let mut free = 0usize;
34    let mut total = 0usize;
35    unsafe { cudaMemGetInfo(&mut free, &mut total) };
36    total - free
37}
38
39#[ctor::ctor]
40fn init() {
41    let _ = MEMORY_MANAGER.set(Mutex::new(MemoryManager::new()));
42    tracing::info!("Memory manager initialized at program start");
43}
44
45/// Allocation record for the small-allocation path (`cudaMallocAsync`).
46struct AllocRecord {
47    size: usize,
48    stream: StreamGuard,
49}
50
51pub struct MemoryManager {
52    pool: VirtualMemoryPool,
53    allocated_ptrs: HashMap<NonNull<c_void>, AllocRecord>,
54    current_size: usize,
55    max_used_size: usize,
56}
57
58/// # Safety
59/// `MemoryManager` is not internally synchronized. These impls are safe because
60/// the singleton instance is wrapped in `Mutex` via `MEMORY_MANAGER`.
61unsafe impl Send for MemoryManager {}
62unsafe impl Sync for MemoryManager {}
63
64impl MemoryManager {
65    pub fn new() -> Self {
66        // Create virtual memory pool
67        let pool = VirtualMemoryPool::default();
68
69        Self {
70            pool,
71            allocated_ptrs: HashMap::new(),
72            current_size: 0,
73            max_used_size: 0,
74        }
75    }
76
77    fn d_malloc_on(
78        &mut self,
79        size: usize,
80        stream: &StreamGuard,
81    ) -> Result<*mut c_void, MemoryError> {
82        assert!(size != 0, "Requested size must be non-zero");
83
84        let mut tracked_size = size;
85        let ptr = if size < self.pool.page_size {
86            let mut ptr: *mut c_void = std::ptr::null_mut();
87            check(unsafe { cudaMallocAsync(&mut ptr, size, stream.as_raw()) }).map_err(|e| {
88                tracing::error!("cudaMallocAsync failed: size={}: {:?}", size, e);
89                MemoryError::from(e)
90            })?;
91            self.allocated_ptrs.insert(
92                NonNull::new(ptr).expect("BUG: cudaMallocAsync returned null"),
93                AllocRecord {
94                    size,
95                    stream: stream.clone(),
96                },
97            );
98            ptr
99        } else {
100            tracked_size = size.next_multiple_of(self.pool.page_size);
101            self.pool.malloc_internal(tracked_size, stream)?
102        };
103
104        self.current_size += tracked_size;
105        if self.current_size > self.max_used_size {
106            self.max_used_size = self.current_size;
107        }
108        Ok(ptr)
109    }
110
111    /// Two-stage free: first resolves the record under the lock, returning the
112    /// `StreamGuard` that must be dropped AFTER the lock is released.
113    ///
114    /// # Safety
115    /// - The pointer `ptr` must be a valid, previously allocated device pointer.
116    /// - The caller must ensure that `ptr` is not used after this function is called.
117    /// - The caller must hold the `MEMORY_MANAGER` lock before calling this method.
118    unsafe fn d_free_under_lock(&mut self, ptr: *mut c_void) -> Result<StreamGuard, MemoryError> {
119        let nn = NonNull::new(ptr).ok_or(MemoryError::NullPointer)?;
120
121        if let Some(record) = self.allocated_ptrs.remove(&nn) {
122            let size = record.size;
123            self.current_size -= size;
124            check(unsafe { cudaFreeAsync(ptr, record.stream.as_raw()) }).map_err(|e| {
125                tracing::error!("cudaFreeAsync failed: ptr={:p}: {:?}", ptr, e);
126                MemoryError::from(e)
127            })?;
128            Ok(record.stream)
129        } else {
130            let (freed_size, guard) = self.pool.free_internal(ptr)?;
131            self.current_size -= freed_size;
132            Ok(guard)
133        }
134    }
135}
136
137impl Drop for MemoryManager {
138    fn drop(&mut self) {
139        device_synchronize().unwrap();
140        let ptrs: Vec<*mut c_void> = self.allocated_ptrs.keys().map(|nn| nn.as_ptr()).collect();
141        for &ptr in &ptrs {
142            match unsafe { self.d_free_under_lock(ptr) } {
143                Ok(guard) => drop(guard),
144                Err(e) => tracing::error!("MemoryManager drop: failed to free {:p}: {:?}", ptr, e),
145            }
146        }
147    }
148}
149
150impl Default for MemoryManager {
151    fn default() -> Self {
152        Self::new()
153    }
154}
155
156pub fn d_malloc_on(size: usize, stream: &StreamGuard) -> Result<*mut c_void, MemoryError> {
157    let manager = MEMORY_MANAGER.get().unwrap();
158    let mut manager = manager.lock().map_err(|_| MemoryError::LockError)?;
159    manager.d_malloc_on(size, stream)
160}
161
162/// # Safety
163/// The pointer `ptr` must be a valid, previously allocated device pointer.
164/// The caller must ensure that `ptr` is not used after this function is called.
165pub unsafe fn d_free(ptr: *mut c_void) -> Result<(), MemoryError> {
166    let manager = MEMORY_MANAGER.get().unwrap();
167    let mut manager = manager.lock().map_err(|_| MemoryError::LockError)?;
168    let guard = manager.d_free_under_lock(ptr)?;
169    drop(manager);
170    drop(guard);
171    Ok(())
172}
173
174#[derive(Debug, Clone)]
175pub struct MemTracker {
176    current: usize,
177    label: &'static str,
178}
179
180impl MemTracker {
181    pub fn start(label: &'static str) -> Self {
182        let current = MEMORY_MANAGER
183            .get()
184            .and_then(|m| m.lock().ok())
185            .map(|m| m.current_size)
186            .unwrap_or_default();
187
188        Self { current, label }
189    }
190
191    pub fn start_and_reset_peak(label: &'static str) -> Self {
192        let mut mem = Self::start(label);
193        mem.reset_peak();
194        mem
195    }
196
197    pub fn emit_metrics(&self) {
198        self.emit_metrics_with_label(self.label);
199    }
200
201    pub fn emit_metrics_with_label(&self, label: &'static str) {
202        let Some(manager) = MEMORY_MANAGER.get().and_then(|m| m.lock().ok()) else {
203            return;
204        };
205
206        let ts = SystemTime::now()
207            .duration_since(UNIX_EPOCH)
208            .unwrap()
209            .as_secs_f64()
210            * 1000.0;
211        let current = manager.current_size;
212        // local_peak is local maximum memory size, as observed by the manager, since the last
213        // reset_peak call
214        let local_peak = manager.max_used_size;
215        let reserved = manager.pool.memory_usage();
216        metrics::gauge!("gpu_mem.timestamp_ms", "module" => label).set(ts);
217        metrics::gauge!("gpu_mem.current_bytes", "module" => label).set(current as f64);
218        metrics::gauge!("gpu_mem.local_peak_bytes", "module" => label).set(local_peak as f64);
219        metrics::gauge!("gpu_mem.reserved_bytes", "module" => label).set(reserved as f64);
220    }
221
222    #[inline]
223    pub fn tracing_info(&self, msg: impl Into<Option<&'static str>>) {
224        let Some(manager) = MEMORY_MANAGER.get().and_then(|m| m.lock().ok()) else {
225            tracing::error!("Memory manager not available");
226            return;
227        };
228        let current = manager.current_size;
229        let peak = manager.max_used_size;
230        let used = current as isize - self.current as isize;
231        let sign = if used >= 0 { "+" } else { "-" };
232        let pool_usage = manager.pool.memory_usage();
233        tracing::info!(
234            "GPU mem: used={}{}, current={}, peak={}, in pool={} ({})",
235            sign,
236            ByteSize::b(used.unsigned_abs() as u64),
237            ByteSize::b(current as u64),
238            ByteSize::b(peak as u64),
239            ByteSize::b(pool_usage as u64),
240            msg.into()
241                .map_or(self.label.to_string(), |m| format!("{}:{}", self.label, m))
242        );
243    }
244
245    pub fn reset_peak(&mut self) {
246        if let Some(mut manager) = MEMORY_MANAGER.get().and_then(|m| m.lock().ok()) {
247            manager.max_used_size = manager.current_size;
248        }
249    }
250}
251
252impl Drop for MemTracker {
253    fn drop(&mut self) {
254        self.tracing_info(None);
255    }
256}