Skip to main content

baracuda_runtime/
memory.rs

1//! Runtime-API device memory.
2
3use core::ffi::c_void;
4use core::marker::PhantomData;
5use core::mem::size_of;
6
7use baracuda_cuda_sys::runtime::{cudaMemcpyKind, runtime};
8use baracuda_types::DeviceRepr;
9
10use crate::error::{check, Result};
11use crate::stream::Stream;
12
13/// Owned, typed allocation of device memory (Runtime API).
14pub struct DeviceBuffer<T: DeviceRepr> {
15    ptr: *mut c_void,
16    len: usize,
17    /// Origin stream for buffers allocated via a `*_async` constructor.
18    /// `Some` ⇒ [`Drop`] reclaims via `cudaFreeAsync` on this stream
19    /// (stream-ordered free, safe even while a kernel that used the
20    /// buffer is still pending); `None` ⇒ synchronous `cudaFree`.
21    /// Cloning a `Stream` is just an `Arc` bump, so retaining it is cheap.
22    stream: Option<Stream>,
23    _marker: PhantomData<T>,
24}
25
26unsafe impl<T: DeviceRepr + Send> Send for DeviceBuffer<T> {}
27
28impl<T: DeviceRepr> core::fmt::Debug for DeviceBuffer<T> {
29    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
30        f.debug_struct("DeviceBuffer")
31            .field("ptr", &self.ptr)
32            .field("len", &self.len)
33            .field("type", &core::any::type_name::<T>())
34            .finish()
35    }
36}
37
38impl<T: DeviceRepr> DeviceBuffer<T> {
39    /// Allocate an uninitialized buffer of `len` elements on the current device.
40    pub fn new(len: usize) -> Result<Self> {
41        let r = runtime()?;
42        let cu = r.cuda_malloc()?;
43        let bytes = len
44            .checked_mul(size_of::<T>())
45            .expect("overflow computing allocation size");
46        let mut ptr: *mut c_void = core::ptr::null_mut();
47        check(unsafe { cu(&mut ptr, bytes) })?;
48        Ok(Self {
49            ptr,
50            len,
51            stream: None,
52            _marker: PhantomData,
53        })
54    }
55
56    /// Allocate and zero-fill.
57    pub fn zeros(len: usize) -> Result<Self> {
58        let buf = Self::new(len)?;
59        let r = runtime()?;
60        let cu = r.cuda_memset()?;
61        let bytes = len * size_of::<T>();
62        check(unsafe { cu(buf.ptr, 0, bytes) })?;
63        Ok(buf)
64    }
65
66    /// Allocate and synchronously copy `src` from host memory.
67    pub fn from_slice(src: &[T]) -> Result<Self> {
68        let buf = Self::new(src.len())?;
69        buf.copy_from_host(src)?;
70        Ok(buf)
71    }
72
73    /// Synchronous H2D copy.
74    pub fn copy_from_host(&self, src: &[T]) -> Result<()> {
75        assert_eq!(src.len(), self.len);
76        let r = runtime()?;
77        let cu = r.cuda_memcpy()?;
78        let bytes = self.len * size_of::<T>();
79        check(unsafe {
80            cu(
81                self.ptr,
82                src.as_ptr() as *const c_void,
83                bytes,
84                cudaMemcpyKind::HostToDevice,
85            )
86        })
87    }
88
89    /// Synchronous D2H copy.
90    pub fn copy_to_host(&self, dst: &mut [T]) -> Result<()> {
91        assert_eq!(dst.len(), self.len);
92        let r = runtime()?;
93        let cu = r.cuda_memcpy()?;
94        let bytes = self.len * size_of::<T>();
95        check(unsafe {
96            cu(
97                dst.as_mut_ptr() as *mut c_void,
98                self.ptr,
99                bytes,
100                cudaMemcpyKind::DeviceToHost,
101            )
102        })
103    }
104
105    /// Asynchronous H2D copy on `stream`.
106    pub fn copy_from_host_async(&self, src: &[T], stream: &Stream) -> Result<()> {
107        assert_eq!(src.len(), self.len);
108        let r = runtime()?;
109        let cu = r.cuda_memcpy_async()?;
110        let bytes = self.len * size_of::<T>();
111        check(unsafe {
112            cu(
113                self.ptr,
114                src.as_ptr() as *const c_void,
115                bytes,
116                cudaMemcpyKind::HostToDevice,
117                stream.as_raw(),
118            )
119        })
120    }
121
122    /// Asynchronous D2H copy on `stream`.
123    pub fn copy_to_host_async(&self, dst: &mut [T], stream: &Stream) -> Result<()> {
124        assert_eq!(dst.len(), self.len);
125        let r = runtime()?;
126        let cu = r.cuda_memcpy_async()?;
127        let bytes = self.len * size_of::<T>();
128        check(unsafe {
129            cu(
130                dst.as_mut_ptr() as *mut c_void,
131                self.ptr,
132                bytes,
133                cudaMemcpyKind::DeviceToHost,
134                stream.as_raw(),
135            )
136        })
137    }
138
139    /// Number of elements.
140    #[inline]
141    pub fn len(&self) -> usize {
142        self.len
143    }
144
145    /// Size in bytes.
146    #[inline]
147    pub fn byte_size(&self) -> usize {
148        self.len * size_of::<T>()
149    }
150
151    /// `true` if zero elements.
152    #[inline]
153    pub fn is_empty(&self) -> bool {
154        self.len == 0
155    }
156
157    /// Raw device pointer. Use with care.
158    #[inline]
159    pub fn as_raw(&self) -> *mut c_void {
160        self.ptr
161    }
162
163    /// Raw device pointer as the u64 value kernels expect. Convenience
164    /// wrapper around [`as_raw`](Self::as_raw).
165    #[inline]
166    pub fn as_device_ptr(&self) -> u64 {
167        self.ptr as u64
168    }
169}
170
171impl<T: DeviceRepr> Drop for DeviceBuffer<T> {
172    fn drop(&mut self) {
173        if self.ptr.is_null() {
174            return;
175        }
176        let Ok(r) = runtime() else { return };
177        // Stream-ordered free for `*_async`-allocated buffers: enqueue
178        // `cudaFreeAsync` on the origin stream so the reclaim is ordered
179        // after any kernel still using the buffer. If that fails to load
180        // (e.g. pre-11.2 runtime) fall through to the synchronous free.
181        if let Some(stream) = &self.stream {
182            if let Ok(cu) = r.cuda_free_async() {
183                if check(unsafe { cu(self.ptr, stream.as_raw()) }).is_ok() {
184                    return;
185                }
186            }
187        }
188        if let Ok(cu) = r.cuda_free() {
189            let _ = unsafe { cu(self.ptr) };
190        }
191    }
192}
193
194// ---- Mem info / prefetch / advise ----------------------------------------
195
196/// `cudaMemGetInfo` — `(free, total)` bytes on the current device.
197pub fn mem_get_info() -> Result<(u64, u64)> {
198    let r = runtime()?;
199    let cu = r.cuda_mem_get_info()?;
200    let mut free: usize = 0;
201    let mut total: usize = 0;
202    check(unsafe { cu(&mut free, &mut total) })?;
203    Ok((free as u64, total as u64))
204}
205
206/// Target for [`mem_prefetch_async`] / [`mem_advise`]. The CUDA Runtime
207/// API's v1 variants take an ordinal — pass `cudaCpuDeviceId` (-1) for host.
208#[derive(Copy, Clone, Debug, Eq, PartialEq)]
209pub enum PrefetchTarget {
210    /// Prefetch to a specific CUDA device (by ordinal).
211    Device(i32),
212    /// Prefetch to the host CPU.
213    Host,
214}
215
216impl PrefetchTarget {
217    #[inline]
218    fn as_raw(self) -> i32 {
219        match self {
220            PrefetchTarget::Device(i) => i,
221            PrefetchTarget::Host => -1, // cudaCpuDeviceId
222        }
223    }
224}
225
226/// Prefetch `count` bytes of unified memory at `dev_ptr` to `target`,
227/// ordered on `stream`. `dev_ptr` must be a managed-memory allocation
228/// (from [`ManagedBuffer`] or `cudaMallocManaged`).
229///
230/// # Safety
231///
232/// `dev_ptr..dev_ptr+count` must be a live managed allocation.
233pub unsafe fn mem_prefetch_async(
234    dev_ptr: *const core::ffi::c_void,
235    count: usize,
236    target: PrefetchTarget,
237    stream: &Stream,
238) -> Result<()> { unsafe {
239    let r = runtime()?;
240    let cu = r.cuda_mem_prefetch_async()?;
241    check(cu(dev_ptr, count, target.as_raw(), stream.as_raw()))
242}}
243
244/// `cudaMemAdvise` — unified-memory placement hint. `advice` is a
245/// constant from [`baracuda_cuda_sys::runtime::types::cudaMemoryAdvise`].
246///
247/// # Safety
248///
249/// `dev_ptr..dev_ptr+count` must be a live managed allocation.
250pub unsafe fn mem_advise(
251    dev_ptr: *const core::ffi::c_void,
252    count: usize,
253    advice: i32,
254    target: PrefetchTarget,
255) -> Result<()> { unsafe {
256    let r = runtime()?;
257    let cu = r.cuda_mem_advise()?;
258    check(cu(dev_ptr, count, advice, target.as_raw()))
259}}
260
261// ---- Managed memory -------------------------------------------------------
262
263/// Unified managed-memory buffer — allocated via `cudaMallocManaged`.
264/// Accessible from both host and device without explicit copies.
265pub struct ManagedBuffer<T: DeviceRepr> {
266    ptr: *mut T,
267    len: usize,
268    _marker: PhantomData<T>,
269}
270
271unsafe impl<T: DeviceRepr + Send> Send for ManagedBuffer<T> {}
272unsafe impl<T: DeviceRepr + Sync> Sync for ManagedBuffer<T> {}
273
274impl<T: DeviceRepr> core::fmt::Debug for ManagedBuffer<T> {
275    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
276        f.debug_struct("ManagedBuffer")
277            .field("ptr", &self.ptr)
278            .field("len", &self.len)
279            .field("type", &core::any::type_name::<T>())
280            .finish()
281    }
282}
283
284impl<T: DeviceRepr> ManagedBuffer<T> {
285    /// Allocate `len` managed elements with the default attach (`GLOBAL`).
286    pub fn new(len: usize) -> Result<Self> {
287        use baracuda_cuda_sys::runtime::types::cudaMemAttach;
288        Self::with_flags(len, cudaMemAttach::GLOBAL)
289    }
290
291    /// Allocate with explicit attach flags (see
292    /// [`baracuda_cuda_sys::runtime::types::cudaMemAttach`]).
293    pub fn with_flags(len: usize, flags: u32) -> Result<Self> {
294        let r = runtime()?;
295        let cu = r.cuda_malloc_managed()?;
296        let bytes = len
297            .checked_mul(size_of::<T>())
298            .expect("overflow computing allocation size");
299        let mut ptr: *mut c_void = core::ptr::null_mut();
300        check(unsafe { cu(&mut ptr, bytes, flags) })?;
301        Ok(Self {
302            ptr: ptr as *mut T,
303            len,
304            _marker: PhantomData,
305        })
306    }
307
308    /// Number of elements.
309    #[inline]
310    pub fn len(&self) -> usize {
311        self.len
312    }
313
314    /// `true` if the allocation has zero elements.
315    #[inline]
316    pub fn is_empty(&self) -> bool {
317        self.len == 0
318    }
319
320    /// Raw pointer — usable from both host and device code.
321    #[inline]
322    pub fn as_ptr(&self) -> *const T {
323        self.ptr
324    }
325
326    /// Raw mutable pointer to the start of the managed allocation. Use
327    /// with care — managed memory is host-accessible but the driver may
328    /// migrate the backing pages between host and device.
329    #[inline]
330    pub fn as_mut_ptr(&mut self) -> *mut T {
331        self.ptr
332    }
333
334    /// Access as a host slice (synchronizes through device cache on access).
335    pub fn as_slice(&self) -> &[T] {
336        // SAFETY: ptr is live for len elements; managed memory is
337        // host-accessible on supported platforms.
338        unsafe { core::slice::from_raw_parts(self.ptr, self.len) }
339    }
340
341    /// Mutable host slice view of the managed allocation.
342    pub fn as_mut_slice(&mut self) -> &mut [T] {
343        unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) }
344    }
345}
346
347impl<T: DeviceRepr> Drop for ManagedBuffer<T> {
348    fn drop(&mut self) {
349        if self.ptr.is_null() {
350            return;
351        }
352        if let Ok(r) = runtime() {
353            if let Ok(cu) = r.cuda_free() {
354                let _ = unsafe { cu(self.ptr as *mut c_void) };
355            }
356        }
357    }
358}
359
360// ---- Pinned host memory --------------------------------------------------
361
362/// Flags for `cudaHostAlloc`. See
363/// [`baracuda_cuda_sys::runtime::types::cudaHostAllocFlags`] for raw values.
364pub mod pinned_flags {
365    pub use baracuda_cuda_sys::runtime::types::cudaHostAllocFlags::*;
366}
367
368/// Pinned (page-locked) host allocation — CUDA-owned memory that supports
369/// real async H↔D copies without staging.
370pub struct PinnedHostBuffer<T: DeviceRepr> {
371    ptr: *mut T,
372    len: usize,
373    _marker: PhantomData<T>,
374}
375
376unsafe impl<T: DeviceRepr + Send> Send for PinnedHostBuffer<T> {}
377unsafe impl<T: DeviceRepr + Sync> Sync for PinnedHostBuffer<T> {}
378
379impl<T: DeviceRepr> core::fmt::Debug for PinnedHostBuffer<T> {
380    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
381        f.debug_struct("PinnedHostBuffer")
382            .field("ptr", &self.ptr)
383            .field("len", &self.len)
384            .finish()
385    }
386}
387
388impl<T: DeviceRepr> PinnedHostBuffer<T> {
389    /// Allocate `len` pinned elements with default flags.
390    pub fn new(len: usize) -> Result<Self> {
391        Self::with_flags(len, 0)
392    }
393
394    /// Allocate with `cudaHostAllocFlags` bitmask.
395    pub fn with_flags(len: usize, flags: u32) -> Result<Self> {
396        let r = runtime()?;
397        let cu = r.cuda_host_alloc()?;
398        let bytes = len
399            .checked_mul(size_of::<T>())
400            .expect("overflow computing allocation size");
401        let mut ptr: *mut c_void = core::ptr::null_mut();
402        check(unsafe { cu(&mut ptr, bytes, flags) })?;
403        Ok(Self {
404            ptr: ptr as *mut T,
405            len,
406            _marker: PhantomData,
407        })
408    }
409
410    /// Device-side pointer that aliases this pinned region (requires
411    /// `MAPPED` flag at alloc time).
412    pub fn device_ptr(&self) -> Result<*mut c_void> {
413        let r = runtime()?;
414        let cu = r.cuda_host_get_device_pointer()?;
415        let mut dev: *mut c_void = core::ptr::null_mut();
416        check(unsafe { cu(&mut dev, self.ptr as *mut c_void, 0) })?;
417        Ok(dev)
418    }
419
420    /// Query the flags this buffer was created with.
421    pub fn flags(&self) -> Result<u32> {
422        let r = runtime()?;
423        let cu = r.cuda_host_get_flags()?;
424        let mut f: core::ffi::c_uint = 0;
425        check(unsafe { cu(&mut f, self.ptr as *mut c_void) })?;
426        Ok(f)
427    }
428
429    /// Length of the allocation, in elements of `T`.
430    #[inline]
431    pub fn len(&self) -> usize {
432        self.len
433    }
434    /// `true` if the allocation has zero elements.
435    #[inline]
436    pub fn is_empty(&self) -> bool {
437        self.len == 0
438    }
439    /// Raw const host pointer to the first element.
440    #[inline]
441    pub fn as_ptr(&self) -> *const T {
442        self.ptr
443    }
444    /// Raw mutable host pointer to the first element.
445    #[inline]
446    pub fn as_mut_ptr(&mut self) -> *mut T {
447        self.ptr
448    }
449}
450
451impl<T: DeviceRepr> core::ops::Deref for PinnedHostBuffer<T> {
452    type Target = [T];
453    fn deref(&self) -> &[T] {
454        unsafe { core::slice::from_raw_parts(self.ptr, self.len) }
455    }
456}
457
458impl<T: DeviceRepr> core::ops::DerefMut for PinnedHostBuffer<T> {
459    fn deref_mut(&mut self) -> &mut [T] {
460        unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) }
461    }
462}
463
464impl<T: DeviceRepr> Drop for PinnedHostBuffer<T> {
465    fn drop(&mut self) {
466        if self.ptr.is_null() {
467            return;
468        }
469        if let Ok(r) = runtime() {
470            if let Ok(cu) = r.cuda_free_host() {
471                let _ = unsafe { cu(self.ptr as *mut c_void) };
472            }
473        }
474    }
475}
476
477/// RAII guard for `cudaHostRegister` — pins an existing host slice and
478/// unregisters on drop.
479pub struct PinnedRegistration<'a, T: DeviceRepr> {
480    ptr: *mut T,
481    len: usize,
482    _borrow: PhantomData<&'a mut [T]>,
483}
484
485unsafe impl<T: DeviceRepr + Send> Send for PinnedRegistration<'_, T> {}
486
487impl<T: DeviceRepr> core::fmt::Debug for PinnedRegistration<'_, T> {
488    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
489        f.debug_struct("PinnedRegistration")
490            .field("ptr", &self.ptr)
491            .field("len", &self.len)
492            .finish()
493    }
494}
495
496impl<'a, T: DeviceRepr> PinnedRegistration<'a, T> {
497    /// Pin `slice` with `flags = 0` until the guard drops.
498    pub fn register(slice: &'a mut [T]) -> Result<Self> {
499        Self::register_with_flags(slice, 0)
500    }
501
502    /// Safe wrapper for `cudaHostRegister`. Pin `slice` for the lifetime
503    /// of the returned guard, passing `flags` verbatim (see
504    /// [`cudaHostRegisterFlags`](baracuda_cuda_sys::runtime::types::cudaHostRegisterFlags)).
505    pub fn register_with_flags(slice: &'a mut [T], flags: u32) -> Result<Self> {
506        let r = runtime()?;
507        let cu = r.cuda_host_register()?;
508        check(unsafe {
509            cu(
510                slice.as_mut_ptr() as *mut c_void,
511                core::mem::size_of_val(slice),
512                flags,
513            )
514        })?;
515        Ok(Self {
516            ptr: slice.as_mut_ptr(),
517            len: slice.len(),
518            _borrow: PhantomData,
519        })
520    }
521
522    /// Length of the pinned slice, in elements of `T`.
523    #[inline]
524    pub fn len(&self) -> usize {
525        self.len
526    }
527    /// `true` if the pinned slice has zero elements.
528    #[inline]
529    pub fn is_empty(&self) -> bool {
530        self.len == 0
531    }
532}
533
534impl<T: DeviceRepr> Drop for PinnedRegistration<'_, T> {
535    fn drop(&mut self) {
536        if self.ptr.is_null() {
537            return;
538        }
539        if let Ok(r) = runtime() {
540            if let Ok(cu) = r.cuda_host_unregister() {
541                let _ = unsafe { cu(self.ptr as *mut c_void) };
542            }
543        }
544    }
545}
546
547// ---- Async alloc / free --------------------------------------------------
548
549impl<T: DeviceRepr> DeviceBuffer<T> {
550    /// Asynchronously allocate `len` elements on `stream` from the device's
551    /// default memory pool (CUDA 11.2+).
552    ///
553    /// The buffer **retains `stream`** (a cheap `Arc` clone), so it also
554    /// *frees* stream-ordered: [`Drop`] enqueues `cudaFreeAsync` on
555    /// `stream`, ordered by the driver *after* every operation already
556    /// submitted to it — including a kernel still reading the buffer. That
557    /// makes "launch on the stream, then let the buffer drop" safe by
558    /// construction, with no host synchronize and no retention pool: the
559    /// per-op `stream.synchronize()` that today keeps scratch / output
560    /// buffers alive can be dropped. Freed blocks return to the device's
561    /// stream-ordered memory pool for reuse across a chain.
562    ///
563    /// [`free_async`](Self::free_async) is still available if you want to
564    /// free at an explicit point rather than at scope exit.
565    pub fn new_async(len: usize, stream: &Stream) -> Result<Self> {
566        let r = runtime()?;
567        let cu = r.cuda_malloc_async()?;
568        let bytes = len
569            .checked_mul(size_of::<T>())
570            .expect("overflow computing allocation size");
571        let mut ptr: *mut c_void = core::ptr::null_mut();
572        check(unsafe { cu(&mut ptr, bytes, stream.as_raw()) })?;
573        Ok(Self {
574            ptr,
575            len,
576            stream: Some(stream.clone()),
577            _marker: PhantomData,
578        })
579    }
580
581    /// Stream-ordered allocate-and-zero: `cudaMallocAsync` on `stream`
582    /// followed by `cudaMemsetAsync` on the same stream, so the buffer is
583    /// zeroed in stream order before any subsequent op observes it.
584    ///
585    /// This is the async counterpart of [`zeros`](Self::zeros), and the
586    /// constructor to use for **output buffers** that should free
587    /// stream-ordered: like [`new_async`](Self::new_async) the buffer
588    /// retains `stream` and reclaims via `cudaFreeAsync` on [`Drop`], so
589    /// an output written by a pipelined kernel can be evicted without a
590    /// host synchronize or an executor-side lifetime guard.
591    ///
592    /// Requires CUDA 11.2+.
593    pub fn zeros_async(len: usize, stream: &Stream) -> Result<Self> {
594        let buf = Self::new_async(len, stream)?;
595        buf.memset_async(0, stream)?;
596        Ok(buf)
597    }
598
599    /// Free this buffer asynchronously on `stream`. Consumes `self` so
600    /// the sync `Drop` does not also free.
601    pub fn free_async(mut self, stream: &Stream) -> Result<()> {
602        let ptr = core::mem::replace(&mut self.ptr, core::ptr::null_mut());
603        if ptr.is_null() {
604            return Ok(());
605        }
606        let r = runtime()?;
607        let cu = r.cuda_free_async()?;
608        check(unsafe { cu(ptr, stream.as_raw()) })
609    }
610
611    /// Asynchronous memset of `self` to byte value `value` on `stream`.
612    pub fn memset_async(&self, value: u8, stream: &Stream) -> Result<()> {
613        let r = runtime()?;
614        let cu = r.cuda_memset_async()?;
615        let bytes = self.len * size_of::<T>();
616        check(unsafe { cu(self.ptr, value as core::ffi::c_int, bytes, stream.as_raw()) })
617    }
618}
619
620// ---- Peer memcpy ---------------------------------------------------------
621
622/// Peer-to-peer device memory copy. Both buffers must be on enabled-peer
623/// devices (see [`crate::Device::enable_peer_access`]).
624pub fn memcpy_peer<T: DeviceRepr>(
625    dst: &DeviceBuffer<T>,
626    dst_device: &crate::Device,
627    src: &DeviceBuffer<T>,
628    src_device: &crate::Device,
629) -> Result<()> {
630    assert_eq!(dst.len(), src.len());
631    let r = runtime()?;
632    let cu = r.cuda_memcpy_peer()?;
633    let bytes = src.len() * size_of::<T>();
634    check(unsafe {
635        cu(
636            dst.as_raw(),
637            dst_device.ordinal(),
638            src.as_raw(),
639            src_device.ordinal(),
640            bytes,
641        )
642    })
643}
644
645/// Async peer-to-peer memcpy ordered on `stream`.
646pub fn memcpy_peer_async<T: DeviceRepr>(
647    dst: &DeviceBuffer<T>,
648    dst_device: &crate::Device,
649    src: &DeviceBuffer<T>,
650    src_device: &crate::Device,
651    stream: &Stream,
652) -> Result<()> {
653    assert_eq!(dst.len(), src.len());
654    let r = runtime()?;
655    let cu = r.cuda_memcpy_peer_async()?;
656    let bytes = src.len() * size_of::<T>();
657    check(unsafe {
658        cu(
659            dst.as_raw(),
660            dst_device.ordinal(),
661            src.as_raw(),
662            src_device.ordinal(),
663            bytes,
664            stream.as_raw(),
665        )
666    })
667}