1use 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
13pub struct DeviceBuffer<T: DeviceRepr> {
15 ptr: *mut c_void,
16 len: usize,
17 stream: Option<Stream>,
24 _marker: PhantomData<T>,
25}
26
27unsafe impl<T: DeviceRepr + Send> Send for DeviceBuffer<T> {}
28
29impl<T: DeviceRepr> core::fmt::Debug for DeviceBuffer<T> {
30 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
31 f.debug_struct("DeviceBuffer")
32 .field("ptr", &self.ptr)
33 .field("len", &self.len)
34 .field("type", &core::any::type_name::<T>())
35 .finish()
36 }
37}
38
39impl<T: DeviceRepr> DeviceBuffer<T> {
40 pub fn new(len: usize) -> Result<Self> {
42 let r = runtime()?;
43 let cu = r.cuda_malloc()?;
44 let bytes = len
45 .checked_mul(size_of::<T>())
46 .expect("overflow computing allocation size");
47 let mut ptr: *mut c_void = core::ptr::null_mut();
48 check(unsafe { cu(&mut ptr, bytes) })?;
49 Ok(Self {
50 ptr,
51 len,
52 stream: None,
53 _marker: PhantomData,
54 })
55 }
56
57 pub fn zeros(len: usize) -> Result<Self> {
59 let buf = Self::new(len)?;
60 let r = runtime()?;
61 let cu = r.cuda_memset()?;
62 let bytes = len * size_of::<T>();
63 check(unsafe { cu(buf.ptr, 0, bytes) })?;
64 Ok(buf)
65 }
66
67 pub fn from_slice(src: &[T]) -> Result<Self> {
69 let buf = Self::new(src.len())?;
70 buf.copy_from_host(src)?;
71 Ok(buf)
72 }
73
74 pub fn copy_from_host(&self, src: &[T]) -> Result<()> {
76 assert_eq!(src.len(), self.len);
77 let r = runtime()?;
78 let cu = r.cuda_memcpy()?;
79 let bytes = self.len * size_of::<T>();
80 check(unsafe {
81 cu(
82 self.ptr,
83 src.as_ptr() as *const c_void,
84 bytes,
85 cudaMemcpyKind::HostToDevice,
86 )
87 })
88 }
89
90 pub fn copy_to_host(&self, dst: &mut [T]) -> Result<()> {
92 assert_eq!(dst.len(), self.len);
93 let r = runtime()?;
94 let cu = r.cuda_memcpy()?;
95 let bytes = self.len * size_of::<T>();
96 check(unsafe {
97 cu(
98 dst.as_mut_ptr() as *mut c_void,
99 self.ptr,
100 bytes,
101 cudaMemcpyKind::DeviceToHost,
102 )
103 })
104 }
105
106 pub fn copy_from_host_async(&self, src: &[T], stream: &Stream) -> Result<()> {
108 assert_eq!(src.len(), self.len);
109 let r = runtime()?;
110 let cu = r.cuda_memcpy_async()?;
111 let bytes = self.len * size_of::<T>();
112 check(unsafe {
113 cu(
114 self.ptr,
115 src.as_ptr() as *const c_void,
116 bytes,
117 cudaMemcpyKind::HostToDevice,
118 stream.as_raw(),
119 )
120 })
121 }
122
123 pub fn copy_to_host_async(&self, dst: &mut [T], stream: &Stream) -> Result<()> {
125 assert_eq!(dst.len(), self.len);
126 let r = runtime()?;
127 let cu = r.cuda_memcpy_async()?;
128 let bytes = self.len * size_of::<T>();
129 check(unsafe {
130 cu(
131 dst.as_mut_ptr() as *mut c_void,
132 self.ptr,
133 bytes,
134 cudaMemcpyKind::DeviceToHost,
135 stream.as_raw(),
136 )
137 })
138 }
139
140 #[inline]
142 pub fn len(&self) -> usize {
143 self.len
144 }
145
146 #[inline]
148 pub fn byte_size(&self) -> usize {
149 self.len * size_of::<T>()
150 }
151
152 #[inline]
154 pub fn is_empty(&self) -> bool {
155 self.len == 0
156 }
157
158 #[inline]
160 pub fn as_raw(&self) -> *mut c_void {
161 self.ptr
162 }
163
164 #[inline]
167 pub fn as_device_ptr(&self) -> u64 {
168 self.ptr as u64
169 }
170}
171
172impl<T: DeviceRepr> Drop for DeviceBuffer<T> {
173 fn drop(&mut self) {
174 if self.ptr.is_null() {
175 return;
176 }
177 let Ok(r) = runtime() else { return };
178 if let Some(stream) = &self.stream {
191 if let Ok(cu) = r.cuda_free_async() {
192 let _ = check(unsafe { cu(self.ptr, stream.as_raw()) });
193 return;
194 }
195 }
198 if let Ok(cu) = r.cuda_free() {
201 let _ = unsafe { cu(self.ptr) };
202 }
203 }
204}
205
206pub fn mem_get_info() -> Result<(u64, u64)> {
210 let r = runtime()?;
211 let cu = r.cuda_mem_get_info()?;
212 let mut free: usize = 0;
213 let mut total: usize = 0;
214 check(unsafe { cu(&mut free, &mut total) })?;
215 Ok((free as u64, total as u64))
216}
217
218#[derive(Copy, Clone, Debug, Eq, PartialEq)]
221pub enum PrefetchTarget {
222 Device(i32),
224 Host,
226}
227
228impl PrefetchTarget {
229 #[inline]
230 fn as_raw(self) -> i32 {
231 match self {
232 PrefetchTarget::Device(i) => i,
233 PrefetchTarget::Host => -1, }
235 }
236}
237
238pub unsafe fn mem_prefetch_async(
246 dev_ptr: *const core::ffi::c_void,
247 count: usize,
248 target: PrefetchTarget,
249 stream: &Stream,
250) -> Result<()> { unsafe {
251 let r = runtime()?;
252 let cu = r.cuda_mem_prefetch_async()?;
253 check(cu(dev_ptr, count, target.as_raw(), stream.as_raw()))
254}}
255
256pub unsafe fn mem_advise(
263 dev_ptr: *const core::ffi::c_void,
264 count: usize,
265 advice: i32,
266 target: PrefetchTarget,
267) -> Result<()> { unsafe {
268 let r = runtime()?;
269 let cu = r.cuda_mem_advise()?;
270 check(cu(dev_ptr, count, advice, target.as_raw()))
271}}
272
273pub struct ManagedBuffer<T: DeviceRepr> {
278 ptr: *mut T,
279 len: usize,
280 _marker: PhantomData<T>,
281}
282
283unsafe impl<T: DeviceRepr + Send> Send for ManagedBuffer<T> {}
284unsafe impl<T: DeviceRepr + Sync> Sync for ManagedBuffer<T> {}
285
286impl<T: DeviceRepr> core::fmt::Debug for ManagedBuffer<T> {
287 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
288 f.debug_struct("ManagedBuffer")
289 .field("ptr", &self.ptr)
290 .field("len", &self.len)
291 .field("type", &core::any::type_name::<T>())
292 .finish()
293 }
294}
295
296impl<T: DeviceRepr> ManagedBuffer<T> {
297 pub fn new(len: usize) -> Result<Self> {
299 use baracuda_cuda_sys::runtime::types::cudaMemAttach;
300 Self::with_flags(len, cudaMemAttach::GLOBAL)
301 }
302
303 pub fn with_flags(len: usize, flags: u32) -> Result<Self> {
306 let r = runtime()?;
307 let cu = r.cuda_malloc_managed()?;
308 let bytes = len
309 .checked_mul(size_of::<T>())
310 .expect("overflow computing allocation size");
311 let mut ptr: *mut c_void = core::ptr::null_mut();
312 check(unsafe { cu(&mut ptr, bytes, flags) })?;
313 Ok(Self {
314 ptr: ptr as *mut T,
315 len,
316 _marker: PhantomData,
317 })
318 }
319
320 #[inline]
322 pub fn len(&self) -> usize {
323 self.len
324 }
325
326 #[inline]
328 pub fn is_empty(&self) -> bool {
329 self.len == 0
330 }
331
332 #[inline]
334 pub fn as_ptr(&self) -> *const T {
335 self.ptr
336 }
337
338 #[inline]
342 pub fn as_mut_ptr(&mut self) -> *mut T {
343 self.ptr
344 }
345
346 pub fn as_slice(&self) -> &[T] {
348 unsafe { core::slice::from_raw_parts(self.ptr, self.len) }
351 }
352
353 pub fn as_mut_slice(&mut self) -> &mut [T] {
355 unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) }
356 }
357}
358
359impl<T: DeviceRepr> Drop for ManagedBuffer<T> {
360 fn drop(&mut self) {
361 if self.ptr.is_null() {
362 return;
363 }
364 if let Ok(r) = runtime() {
365 if let Ok(cu) = r.cuda_free() {
366 let _ = unsafe { cu(self.ptr as *mut c_void) };
367 }
368 }
369 }
370}
371
372pub mod pinned_flags {
377 pub use baracuda_cuda_sys::runtime::types::cudaHostAllocFlags::*;
378}
379
380pub struct PinnedHostBuffer<T: DeviceRepr> {
383 ptr: *mut T,
384 len: usize,
385 _marker: PhantomData<T>,
386}
387
388unsafe impl<T: DeviceRepr + Send> Send for PinnedHostBuffer<T> {}
389unsafe impl<T: DeviceRepr + Sync> Sync for PinnedHostBuffer<T> {}
390
391impl<T: DeviceRepr> core::fmt::Debug for PinnedHostBuffer<T> {
392 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
393 f.debug_struct("PinnedHostBuffer")
394 .field("ptr", &self.ptr)
395 .field("len", &self.len)
396 .finish()
397 }
398}
399
400impl<T: DeviceRepr> PinnedHostBuffer<T> {
401 pub fn new(len: usize) -> Result<Self> {
403 Self::with_flags(len, 0)
404 }
405
406 pub fn with_flags(len: usize, flags: u32) -> Result<Self> {
408 let r = runtime()?;
409 let cu = r.cuda_host_alloc()?;
410 let bytes = len
411 .checked_mul(size_of::<T>())
412 .expect("overflow computing allocation size");
413 let mut ptr: *mut c_void = core::ptr::null_mut();
414 check(unsafe { cu(&mut ptr, bytes, flags) })?;
415 Ok(Self {
416 ptr: ptr as *mut T,
417 len,
418 _marker: PhantomData,
419 })
420 }
421
422 pub fn device_ptr(&self) -> Result<*mut c_void> {
425 let r = runtime()?;
426 let cu = r.cuda_host_get_device_pointer()?;
427 let mut dev: *mut c_void = core::ptr::null_mut();
428 check(unsafe { cu(&mut dev, self.ptr as *mut c_void, 0) })?;
429 Ok(dev)
430 }
431
432 pub fn flags(&self) -> Result<u32> {
434 let r = runtime()?;
435 let cu = r.cuda_host_get_flags()?;
436 let mut f: core::ffi::c_uint = 0;
437 check(unsafe { cu(&mut f, self.ptr as *mut c_void) })?;
438 Ok(f)
439 }
440
441 #[inline]
443 pub fn len(&self) -> usize {
444 self.len
445 }
446 #[inline]
448 pub fn is_empty(&self) -> bool {
449 self.len == 0
450 }
451 #[inline]
453 pub fn as_ptr(&self) -> *const T {
454 self.ptr
455 }
456 #[inline]
458 pub fn as_mut_ptr(&mut self) -> *mut T {
459 self.ptr
460 }
461}
462
463impl<T: DeviceRepr> core::ops::Deref for PinnedHostBuffer<T> {
464 type Target = [T];
465 fn deref(&self) -> &[T] {
466 unsafe { core::slice::from_raw_parts(self.ptr, self.len) }
467 }
468}
469
470impl<T: DeviceRepr> core::ops::DerefMut for PinnedHostBuffer<T> {
471 fn deref_mut(&mut self) -> &mut [T] {
472 unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) }
473 }
474}
475
476impl<T: DeviceRepr> Drop for PinnedHostBuffer<T> {
477 fn drop(&mut self) {
478 if self.ptr.is_null() {
479 return;
480 }
481 if let Ok(r) = runtime() {
482 if let Ok(cu) = r.cuda_free_host() {
483 let _ = unsafe { cu(self.ptr as *mut c_void) };
484 }
485 }
486 }
487}
488
489pub struct PinnedRegistration<'a, T: DeviceRepr> {
492 ptr: *mut T,
493 len: usize,
494 _borrow: PhantomData<&'a mut [T]>,
495}
496
497unsafe impl<T: DeviceRepr + Send> Send for PinnedRegistration<'_, T> {}
498
499impl<T: DeviceRepr> core::fmt::Debug for PinnedRegistration<'_, T> {
500 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
501 f.debug_struct("PinnedRegistration")
502 .field("ptr", &self.ptr)
503 .field("len", &self.len)
504 .finish()
505 }
506}
507
508impl<'a, T: DeviceRepr> PinnedRegistration<'a, T> {
509 pub fn register(slice: &'a mut [T]) -> Result<Self> {
511 Self::register_with_flags(slice, 0)
512 }
513
514 pub fn register_with_flags(slice: &'a mut [T], flags: u32) -> Result<Self> {
518 let r = runtime()?;
519 let cu = r.cuda_host_register()?;
520 check(unsafe {
521 cu(
522 slice.as_mut_ptr() as *mut c_void,
523 core::mem::size_of_val(slice),
524 flags,
525 )
526 })?;
527 Ok(Self {
528 ptr: slice.as_mut_ptr(),
529 len: slice.len(),
530 _borrow: PhantomData,
531 })
532 }
533
534 #[inline]
536 pub fn len(&self) -> usize {
537 self.len
538 }
539 #[inline]
541 pub fn is_empty(&self) -> bool {
542 self.len == 0
543 }
544}
545
546impl<T: DeviceRepr> Drop for PinnedRegistration<'_, T> {
547 fn drop(&mut self) {
548 if self.ptr.is_null() {
549 return;
550 }
551 if let Ok(r) = runtime() {
552 if let Ok(cu) = r.cuda_host_unregister() {
553 let _ = unsafe { cu(self.ptr as *mut c_void) };
554 }
555 }
556 }
557}
558
559impl<T: DeviceRepr> DeviceBuffer<T> {
562 pub fn new_async(len: usize, stream: &Stream) -> Result<Self> {
587 let r = runtime()?;
588 let cu = r.cuda_malloc_async()?;
589 let bytes = len
590 .checked_mul(size_of::<T>())
591 .expect("overflow computing allocation size");
592 let mut ptr: *mut c_void = core::ptr::null_mut();
593 check(unsafe { cu(&mut ptr, bytes, stream.as_raw()) })?;
594 Ok(Self {
595 ptr,
596 len,
597 stream: Some(stream.clone()),
598 _marker: PhantomData,
599 })
600 }
601
602 pub fn zeros_async(len: usize, stream: &Stream) -> Result<Self> {
615 let buf = Self::new_async(len, stream)?;
616 buf.memset_async(0, stream)?;
617 Ok(buf)
618 }
619
620 pub fn free_async(mut self, stream: &Stream) -> Result<()> {
623 let ptr = core::mem::replace(&mut self.ptr, core::ptr::null_mut());
624 if ptr.is_null() {
625 return Ok(());
626 }
627 let r = runtime()?;
628 let cu = r.cuda_free_async()?;
629 check(unsafe { cu(ptr, stream.as_raw()) })
630 }
631
632 pub fn memset_async(&self, value: u8, stream: &Stream) -> Result<()> {
634 let r = runtime()?;
635 let cu = r.cuda_memset_async()?;
636 let bytes = self.len * size_of::<T>();
637 check(unsafe { cu(self.ptr, value as core::ffi::c_int, bytes, stream.as_raw()) })
638 }
639}
640
641pub fn memcpy_peer<T: DeviceRepr>(
646 dst: &DeviceBuffer<T>,
647 dst_device: &crate::Device,
648 src: &DeviceBuffer<T>,
649 src_device: &crate::Device,
650) -> Result<()> {
651 assert_eq!(dst.len(), src.len());
652 let r = runtime()?;
653 let cu = r.cuda_memcpy_peer()?;
654 let bytes = src.len() * size_of::<T>();
655 check(unsafe {
656 cu(
657 dst.as_raw(),
658 dst_device.ordinal(),
659 src.as_raw(),
660 src_device.ordinal(),
661 bytes,
662 )
663 })
664}
665
666pub fn memcpy_peer_async<T: DeviceRepr>(
668 dst: &DeviceBuffer<T>,
669 dst_device: &crate::Device,
670 src: &DeviceBuffer<T>,
671 src_device: &crate::Device,
672 stream: &Stream,
673) -> Result<()> {
674 assert_eq!(dst.len(), src.len());
675 let r = runtime()?;
676 let cu = r.cuda_memcpy_peer_async()?;
677 let bytes = src.len() * size_of::<T>();
678 check(unsafe {
679 cu(
680 dst.as_raw(),
681 dst_device.ordinal(),
682 src.as_raw(),
683 src_device.ordinal(),
684 bytes,
685 stream.as_raw(),
686 )
687 })
688}