openvm_cuda_common/memory_manager/
mod.rs1use 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
45struct 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
58unsafe impl Send for MemoryManager {}
62unsafe impl Sync for MemoryManager {}
63
64impl MemoryManager {
65 pub fn new() -> Self {
66 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 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
162pub 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 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}