baracuda_runtime/
mempool.rs1use 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#[derive(Copy, Clone, Debug, Eq, PartialEq)]
23pub enum AccessFlags {
24 None,
26 Read,
28 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#[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 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 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 #[inline]
125 pub fn as_raw(&self) -> cudaMemPool_t {
126 self.inner.handle
127 }
128
129 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 pub fn release_threshold(&self) -> Result<u64> {
147 self.get_u64_attr(cudaMemPoolAttr::RELEASE_THRESHOLD)
148 }
149
150 pub fn used_bytes(&self) -> Result<u64> {
152 self.get_u64_attr(cudaMemPoolAttr::USED_MEM_CURRENT)
153 }
154
155 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 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 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 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 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 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 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 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
278pub 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 Ok(unsafe { MemoryPool::from_borrowed(handle) })
286}
287
288pub 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
297pub 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}