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::{Result, check};
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<()> {
251 unsafe {
252 let r = runtime()?;
253 let cu = r.cuda_mem_prefetch_async()?;
254 check(cu(dev_ptr, count, target.as_raw(), stream.as_raw()))
255 }
256}
257
258pub unsafe fn mem_advise(
265 dev_ptr: *const core::ffi::c_void,
266 count: usize,
267 advice: i32,
268 target: PrefetchTarget,
269) -> Result<()> {
270 unsafe {
271 let r = runtime()?;
272 let cu = r.cuda_mem_advise()?;
273 check(cu(dev_ptr, count, advice, target.as_raw()))
274 }
275}
276
277pub struct ManagedBuffer<T: DeviceRepr> {
282 ptr: *mut T,
283 len: usize,
284 _marker: PhantomData<T>,
285}
286
287unsafe impl<T: DeviceRepr + Send> Send for ManagedBuffer<T> {}
288unsafe impl<T: DeviceRepr + Sync> Sync for ManagedBuffer<T> {}
289
290impl<T: DeviceRepr> core::fmt::Debug for ManagedBuffer<T> {
291 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
292 f.debug_struct("ManagedBuffer")
293 .field("ptr", &self.ptr)
294 .field("len", &self.len)
295 .field("type", &core::any::type_name::<T>())
296 .finish()
297 }
298}
299
300impl<T: DeviceRepr> ManagedBuffer<T> {
301 pub fn new(len: usize) -> Result<Self> {
303 use baracuda_cuda_sys::runtime::types::cudaMemAttach;
304 Self::with_flags(len, cudaMemAttach::GLOBAL)
305 }
306
307 pub fn with_flags(len: usize, flags: u32) -> Result<Self> {
310 let r = runtime()?;
311 let cu = r.cuda_malloc_managed()?;
312 let bytes = len
313 .checked_mul(size_of::<T>())
314 .expect("overflow computing allocation size");
315 let mut ptr: *mut c_void = core::ptr::null_mut();
316 check(unsafe { cu(&mut ptr, bytes, flags) })?;
317 Ok(Self {
318 ptr: ptr as *mut T,
319 len,
320 _marker: PhantomData,
321 })
322 }
323
324 #[inline]
326 pub fn len(&self) -> usize {
327 self.len
328 }
329
330 #[inline]
332 pub fn is_empty(&self) -> bool {
333 self.len == 0
334 }
335
336 #[inline]
338 pub fn as_ptr(&self) -> *const T {
339 self.ptr
340 }
341
342 #[inline]
346 pub fn as_mut_ptr(&mut self) -> *mut T {
347 self.ptr
348 }
349
350 pub fn as_slice(&self) -> &[T] {
352 unsafe { core::slice::from_raw_parts(self.ptr, self.len) }
355 }
356
357 pub fn as_mut_slice(&mut self) -> &mut [T] {
359 unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) }
360 }
361}
362
363impl<T: DeviceRepr> Drop for ManagedBuffer<T> {
364 fn drop(&mut self) {
365 if self.ptr.is_null() {
366 return;
367 }
368 if let Ok(r) = runtime() {
369 if let Ok(cu) = r.cuda_free() {
370 let _ = unsafe { cu(self.ptr as *mut c_void) };
371 }
372 }
373 }
374}
375
376pub mod pinned_flags {
381 pub use baracuda_cuda_sys::runtime::types::cudaHostAllocFlags::*;
382}
383
384pub struct PinnedHostBuffer<T: DeviceRepr> {
387 ptr: *mut T,
388 len: usize,
389 _marker: PhantomData<T>,
390}
391
392unsafe impl<T: DeviceRepr + Send> Send for PinnedHostBuffer<T> {}
393unsafe impl<T: DeviceRepr + Sync> Sync for PinnedHostBuffer<T> {}
394
395impl<T: DeviceRepr> core::fmt::Debug for PinnedHostBuffer<T> {
396 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
397 f.debug_struct("PinnedHostBuffer")
398 .field("ptr", &self.ptr)
399 .field("len", &self.len)
400 .finish()
401 }
402}
403
404impl<T: DeviceRepr> PinnedHostBuffer<T> {
405 pub fn new(len: usize) -> Result<Self> {
407 Self::with_flags(len, 0)
408 }
409
410 pub fn with_flags(len: usize, flags: u32) -> Result<Self> {
412 let r = runtime()?;
413 let cu = r.cuda_host_alloc()?;
414 let bytes = len
415 .checked_mul(size_of::<T>())
416 .expect("overflow computing allocation size");
417 let mut ptr: *mut c_void = core::ptr::null_mut();
418 check(unsafe { cu(&mut ptr, bytes, flags) })?;
419 Ok(Self {
420 ptr: ptr as *mut T,
421 len,
422 _marker: PhantomData,
423 })
424 }
425
426 pub fn device_ptr(&self) -> Result<*mut c_void> {
429 let r = runtime()?;
430 let cu = r.cuda_host_get_device_pointer()?;
431 let mut dev: *mut c_void = core::ptr::null_mut();
432 check(unsafe { cu(&mut dev, self.ptr as *mut c_void, 0) })?;
433 Ok(dev)
434 }
435
436 pub fn flags(&self) -> Result<u32> {
438 let r = runtime()?;
439 let cu = r.cuda_host_get_flags()?;
440 let mut f: core::ffi::c_uint = 0;
441 check(unsafe { cu(&mut f, self.ptr as *mut c_void) })?;
442 Ok(f)
443 }
444
445 #[inline]
447 pub fn len(&self) -> usize {
448 self.len
449 }
450 #[inline]
452 pub fn is_empty(&self) -> bool {
453 self.len == 0
454 }
455 #[inline]
457 pub fn as_ptr(&self) -> *const T {
458 self.ptr
459 }
460 #[inline]
462 pub fn as_mut_ptr(&mut self) -> *mut T {
463 self.ptr
464 }
465}
466
467impl<T: DeviceRepr> core::ops::Deref for PinnedHostBuffer<T> {
468 type Target = [T];
469 fn deref(&self) -> &[T] {
470 unsafe { core::slice::from_raw_parts(self.ptr, self.len) }
471 }
472}
473
474impl<T: DeviceRepr> core::ops::DerefMut for PinnedHostBuffer<T> {
475 fn deref_mut(&mut self) -> &mut [T] {
476 unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) }
477 }
478}
479
480impl<T: DeviceRepr> Drop for PinnedHostBuffer<T> {
481 fn drop(&mut self) {
482 if self.ptr.is_null() {
483 return;
484 }
485 if let Ok(r) = runtime() {
486 if let Ok(cu) = r.cuda_free_host() {
487 let _ = unsafe { cu(self.ptr as *mut c_void) };
488 }
489 }
490 }
491}
492
493pub struct PinnedRegistration<'a, T: DeviceRepr> {
496 ptr: *mut T,
497 len: usize,
498 _borrow: PhantomData<&'a mut [T]>,
499}
500
501unsafe impl<T: DeviceRepr + Send> Send for PinnedRegistration<'_, T> {}
502
503impl<T: DeviceRepr> core::fmt::Debug for PinnedRegistration<'_, T> {
504 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
505 f.debug_struct("PinnedRegistration")
506 .field("ptr", &self.ptr)
507 .field("len", &self.len)
508 .finish()
509 }
510}
511
512impl<'a, T: DeviceRepr> PinnedRegistration<'a, T> {
513 pub fn register(slice: &'a mut [T]) -> Result<Self> {
515 Self::register_with_flags(slice, 0)
516 }
517
518 pub fn register_with_flags(slice: &'a mut [T], flags: u32) -> Result<Self> {
522 let r = runtime()?;
523 let cu = r.cuda_host_register()?;
524 check(unsafe {
525 cu(
526 slice.as_mut_ptr() as *mut c_void,
527 core::mem::size_of_val(slice),
528 flags,
529 )
530 })?;
531 Ok(Self {
532 ptr: slice.as_mut_ptr(),
533 len: slice.len(),
534 _borrow: PhantomData,
535 })
536 }
537
538 #[inline]
540 pub fn len(&self) -> usize {
541 self.len
542 }
543 #[inline]
545 pub fn is_empty(&self) -> bool {
546 self.len == 0
547 }
548}
549
550impl<T: DeviceRepr> Drop for PinnedRegistration<'_, T> {
551 fn drop(&mut self) {
552 if self.ptr.is_null() {
553 return;
554 }
555 if let Ok(r) = runtime() {
556 if let Ok(cu) = r.cuda_host_unregister() {
557 let _ = unsafe { cu(self.ptr as *mut c_void) };
558 }
559 }
560 }
561}
562
563impl<T: DeviceRepr> DeviceBuffer<T> {
566 pub fn new_async(len: usize, stream: &Stream) -> Result<Self> {
591 let r = runtime()?;
592 let cu = r.cuda_malloc_async()?;
593 let bytes = len
594 .checked_mul(size_of::<T>())
595 .expect("overflow computing allocation size");
596 let mut ptr: *mut c_void = core::ptr::null_mut();
597 check(unsafe { cu(&mut ptr, bytes, stream.as_raw()) })?;
598 Ok(Self {
599 ptr,
600 len,
601 stream: Some(stream.clone()),
602 _marker: PhantomData,
603 })
604 }
605
606 pub fn zeros_async(len: usize, stream: &Stream) -> Result<Self> {
619 let buf = Self::new_async(len, stream)?;
620 buf.memset_async(0, stream)?;
621 Ok(buf)
622 }
623
624 pub fn free_async(mut self, stream: &Stream) -> Result<()> {
627 let ptr = core::mem::replace(&mut self.ptr, core::ptr::null_mut());
628 if ptr.is_null() {
629 return Ok(());
630 }
631 let r = runtime()?;
632 let cu = r.cuda_free_async()?;
633 check(unsafe { cu(ptr, stream.as_raw()) })
634 }
635
636 pub fn memset_async(&self, value: u8, stream: &Stream) -> Result<()> {
638 let r = runtime()?;
639 let cu = r.cuda_memset_async()?;
640 let bytes = self.len * size_of::<T>();
641 check(unsafe { cu(self.ptr, value as core::ffi::c_int, bytes, stream.as_raw()) })
642 }
643}
644
645pub fn memcpy_peer<T: DeviceRepr>(
650 dst: &DeviceBuffer<T>,
651 dst_device: &crate::Device,
652 src: &DeviceBuffer<T>,
653 src_device: &crate::Device,
654) -> Result<()> {
655 assert_eq!(dst.len(), src.len());
656 let r = runtime()?;
657 let cu = r.cuda_memcpy_peer()?;
658 let bytes = src.len() * size_of::<T>();
659 check(unsafe {
660 cu(
661 dst.as_raw(),
662 dst_device.ordinal(),
663 src.as_raw(),
664 src_device.ordinal(),
665 bytes,
666 )
667 })
668}
669
670pub fn memcpy_peer_async<T: DeviceRepr>(
672 dst: &DeviceBuffer<T>,
673 dst_device: &crate::Device,
674 src: &DeviceBuffer<T>,
675 src_device: &crate::Device,
676 stream: &Stream,
677) -> Result<()> {
678 assert_eq!(dst.len(), src.len());
679 let r = runtime()?;
680 let cu = r.cuda_memcpy_peer_async()?;
681 let bytes = src.len() * size_of::<T>();
682 check(unsafe {
683 cu(
684 dst.as_raw(),
685 dst_device.ordinal(),
686 src.as_raw(),
687 src_device.ordinal(),
688 bytes,
689 stream.as_raw(),
690 )
691 })
692}