Skip to main content

baracuda_runtime/
mempool.rs

1//! Stream-ordered memory pools (Runtime API, CUDA 11.2+).
2//!
3//! Mirrors [`baracuda_driver::mempool`] — a pool is a device-backed
4//! allocator with a configurable release threshold, accessed via
5//! `cudaMallocFromPoolAsync` and returned with `cudaFreeAsync`.
6//! Each device exposes a default pool via [`default_pool`].
7
8use std::sync::Arc;
9
10use baracuda_cuda_sys::runtime::runtime;
11use baracuda_cuda_sys::runtime::types::{
12    cudaMemAccessDesc, cudaMemAllocationHandleType, cudaMemAllocationType, cudaMemLocation,
13    cudaMemLocationType, cudaMemPool_t, cudaMemPoolAttr, cudaMemPoolProps,
14    cudaMemPoolPtrExportData,
15};
16
17use crate::device::Device;
18use crate::error::{Result, check};
19use crate::stream::Stream;
20
21/// Access rights granted to a device for a pool's allocations.
22#[derive(Copy, Clone, Debug, Eq, PartialEq)]
23pub enum AccessFlags {
24    /// No access (`cudaMemAccessFlagsProtNone`).
25    None,
26    /// Read-only access (`cudaMemAccessFlagsProtRead`).
27    Read,
28    /// Full read/write access (`cudaMemAccessFlagsProtReadWrite`).
29    ReadWrite,
30}
31
32impl AccessFlags {
33    #[inline]
34    fn raw(self) -> core::ffi::c_int {
35        use baracuda_cuda_sys::runtime::types::cudaMemAccessFlags;
36        match self {
37            AccessFlags::None => cudaMemAccessFlags::NONE,
38            AccessFlags::Read => cudaMemAccessFlags::READ,
39            AccessFlags::ReadWrite => cudaMemAccessFlags::READ_WRITE,
40        }
41    }
42
43    #[inline]
44    fn from_raw(raw: core::ffi::c_int) -> Self {
45        use baracuda_cuda_sys::runtime::types::cudaMemAccessFlags;
46        match raw {
47            x if x == cudaMemAccessFlags::READ => AccessFlags::Read,
48            x if x == cudaMemAccessFlags::READ_WRITE => AccessFlags::ReadWrite,
49            _ => AccessFlags::None,
50        }
51    }
52}
53
54/// A memory pool. Owned pools are destroyed on last-clone drop; borrowed
55/// pools (returned by [`default_pool`] / [`current_pool`]) are not.
56#[derive(Clone)]
57pub struct MemoryPool {
58    inner: Arc<MemoryPoolInner>,
59}
60
61struct MemoryPoolInner {
62    handle: cudaMemPool_t,
63    owned: bool,
64}
65
66unsafe impl Send for MemoryPoolInner {}
67unsafe impl Sync for MemoryPoolInner {}
68
69impl core::fmt::Debug for MemoryPoolInner {
70    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
71        f.debug_struct("MemoryPool")
72            .field("handle", &self.handle)
73            .field("owned", &self.owned)
74            .finish()
75    }
76}
77
78impl core::fmt::Debug for MemoryPool {
79    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
80        self.inner.fmt(f)
81    }
82}
83
84impl MemoryPool {
85    /// Create a fresh pool backed on `device`.
86    pub fn new(device: &Device) -> Result<Self> {
87        let r = runtime()?;
88        let cu = r.cuda_mem_pool_create()?;
89        let props = cudaMemPoolProps {
90            alloc_type: cudaMemAllocationType::PINNED,
91            handle_types: cudaMemAllocationHandleType::NONE,
92            location: cudaMemLocation {
93                type_: cudaMemLocationType::DEVICE,
94                id: device.ordinal(),
95            },
96            ..Default::default()
97        };
98        let mut handle: cudaMemPool_t = core::ptr::null_mut();
99        check(unsafe { cu(&mut handle, &props) })?;
100        Ok(Self {
101            inner: Arc::new(MemoryPoolInner {
102                handle,
103                owned: true,
104            }),
105        })
106    }
107
108    /// Wrap a raw pool handle without taking ownership.
109    ///
110    /// # Safety
111    ///
112    /// `handle` must outlive this wrapper.
113    pub unsafe fn from_borrowed(handle: cudaMemPool_t) -> Self {
114        Self {
115            inner: Arc::new(MemoryPoolInner {
116                handle,
117                owned: false,
118            }),
119        }
120    }
121
122    /// Raw `cudaMemPool_t` handle. Use with care — owned pools are
123    /// destroyed when the last clone of `self` drops.
124    #[inline]
125    pub fn as_raw(&self) -> cudaMemPool_t {
126        self.inner.handle
127    }
128
129    /// Set the release threshold (bytes retained before the pool starts
130    /// returning memory to the OS). Default is 0.
131    pub fn set_release_threshold(&self, bytes: u64) -> Result<()> {
132        let r = runtime()?;
133        let cu = r.cuda_mem_pool_set_attribute()?;
134        let mut v = bytes;
135        check(unsafe {
136            cu(
137                self.inner.handle,
138                cudaMemPoolAttr::RELEASE_THRESHOLD,
139                &mut v as *mut u64 as *mut core::ffi::c_void,
140            )
141        })
142    }
143
144    /// Read the current release threshold (in bytes) via
145    /// `cudaMemPoolGetAttribute` with `cudaMemPoolAttrReleaseThreshold`.
146    pub fn release_threshold(&self) -> Result<u64> {
147        self.get_u64_attr(cudaMemPoolAttr::RELEASE_THRESHOLD)
148    }
149
150    /// Current bytes handed out to allocations.
151    pub fn used_bytes(&self) -> Result<u64> {
152        self.get_u64_attr(cudaMemPoolAttr::USED_MEM_CURRENT)
153    }
154
155    /// Current bytes reserved for the pool (used + kept-free).
156    pub fn reserved_bytes(&self) -> Result<u64> {
157        self.get_u64_attr(cudaMemPoolAttr::RESERVED_MEM_CURRENT)
158    }
159
160    fn get_u64_attr(&self, attr: i32) -> Result<u64> {
161        let r = runtime()?;
162        let cu = r.cuda_mem_pool_get_attribute()?;
163        let mut v: u64 = 0;
164        check(unsafe {
165            cu(
166                self.inner.handle,
167                attr,
168                &mut v as *mut u64 as *mut core::ffi::c_void,
169            )
170        })?;
171        Ok(v)
172    }
173
174    /// Release memory down to `min_bytes_to_keep`.
175    pub fn trim_to(&self, min_bytes_to_keep: usize) -> Result<()> {
176        let r = runtime()?;
177        let cu = r.cuda_mem_pool_trim_to()?;
178        check(unsafe { cu(self.inner.handle, min_bytes_to_keep) })
179    }
180
181    /// Grant `device` the specified access to allocations from this pool.
182    pub fn set_access(&self, device: &Device, flags: AccessFlags) -> Result<()> {
183        let r = runtime()?;
184        let cu = r.cuda_mem_pool_set_access()?;
185        let desc = cudaMemAccessDesc {
186            location: cudaMemLocation {
187                type_: cudaMemLocationType::DEVICE,
188                id: device.ordinal(),
189            },
190            flags: flags.raw(),
191        };
192        check(unsafe { cu(self.inner.handle, &desc, 1) })
193    }
194
195    /// Query `device`'s access flags for this pool.
196    pub fn access(&self, device: &Device) -> Result<AccessFlags> {
197        let r = runtime()?;
198        let cu = r.cuda_mem_pool_get_access()?;
199        let mut loc = cudaMemLocation {
200            type_: cudaMemLocationType::DEVICE,
201            id: device.ordinal(),
202        };
203        let mut flags: core::ffi::c_int = 0;
204        check(unsafe { cu(&mut flags, self.inner.handle, &mut loc) })?;
205        Ok(AccessFlags::from_raw(flags))
206    }
207
208    /// Allocate `bytes` bytes of device memory from this pool, ordered on
209    /// `stream`. Returns a raw device pointer — free via
210    /// [`crate::DeviceBuffer::free_async`] or by calling
211    /// [`Self::free_async`] on the raw pointer.
212    pub fn alloc_async(&self, bytes: usize, stream: &Stream) -> Result<*mut core::ffi::c_void> {
213        let r = runtime()?;
214        let cu = r.cuda_malloc_from_pool_async()?;
215        let mut ptr: *mut core::ffi::c_void = core::ptr::null_mut();
216        check(unsafe { cu(&mut ptr, bytes, self.inner.handle, stream.as_raw()) })?;
217        Ok(ptr)
218    }
219
220    /// Free a device pointer previously returned by
221    /// [`Self::alloc_async`] (routes through `cudaFreeAsync`).
222    ///
223    /// # Safety
224    ///
225    /// `ptr` must be a live allocation from this (or another) pool.
226    pub unsafe fn free_async(&self, ptr: *mut core::ffi::c_void, stream: &Stream) -> Result<()> {
227        unsafe {
228            let r = runtime()?;
229            let cu = r.cuda_free_async()?;
230            check(cu(ptr, stream.as_raw()))
231        }
232    }
233
234    /// Export a pointer in this pool for sharing with a peer process.
235    ///
236    /// # Safety
237    ///
238    /// `ptr` must be a live allocation from this pool.
239    pub unsafe fn export_pointer(
240        &self,
241        ptr: *mut core::ffi::c_void,
242    ) -> Result<cudaMemPoolPtrExportData> {
243        unsafe {
244            let r = runtime()?;
245            let cu = r.cuda_mem_pool_export_pointer()?;
246            let mut data = cudaMemPoolPtrExportData::default();
247            check(cu(&mut data, ptr))?;
248            Ok(data)
249        }
250    }
251
252    /// Import an exported pointer into this pool.
253    pub fn import_pointer(
254        &self,
255        mut data: cudaMemPoolPtrExportData,
256    ) -> Result<*mut core::ffi::c_void> {
257        let r = runtime()?;
258        let cu = r.cuda_mem_pool_import_pointer()?;
259        let mut ptr: *mut core::ffi::c_void = core::ptr::null_mut();
260        check(unsafe { cu(&mut ptr, self.inner.handle, &mut data) })?;
261        Ok(ptr)
262    }
263}
264
265impl Drop for MemoryPoolInner {
266    fn drop(&mut self) {
267        if !self.owned || self.handle.is_null() {
268            return;
269        }
270        if let Ok(r) = runtime() {
271            if let Ok(cu) = r.cuda_mem_pool_destroy() {
272                let _ = unsafe { cu(self.handle) };
273            }
274        }
275    }
276}
277
278/// Return the device's default memory pool (borrowed — not destroyed on drop).
279pub fn default_pool(device: &Device) -> Result<MemoryPool> {
280    let r = runtime()?;
281    let cu = r.cuda_device_get_default_mem_pool()?;
282    let mut handle: cudaMemPool_t = core::ptr::null_mut();
283    check(unsafe { cu(&mut handle, device.ordinal()) })?;
284    // SAFETY: the runtime owns the default pool; we wrap non-owning.
285    Ok(unsafe { MemoryPool::from_borrowed(handle) })
286}
287
288/// Return the pool currently used by `cudaMallocAsync` on `device`.
289pub fn current_pool(device: &Device) -> Result<MemoryPool> {
290    let r = runtime()?;
291    let cu = r.cuda_device_get_mem_pool()?;
292    let mut handle: cudaMemPool_t = core::ptr::null_mut();
293    check(unsafe { cu(&mut handle, device.ordinal()) })?;
294    Ok(unsafe { MemoryPool::from_borrowed(handle) })
295}
296
297/// Replace the pool used by `cudaMallocAsync` on `device`.
298pub fn set_current_pool(device: &Device, pool: &MemoryPool) -> Result<()> {
299    let r = runtime()?;
300    let cu = r.cuda_device_set_mem_pool()?;
301    check(unsafe { cu(device.ordinal(), pool.as_raw()) })
302}