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>,
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 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 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 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 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 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 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 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 #[inline]
141 pub fn len(&self) -> usize {
142 self.len
143 }
144
145 #[inline]
147 pub fn byte_size(&self) -> usize {
148 self.len * size_of::<T>()
149 }
150
151 #[inline]
153 pub fn is_empty(&self) -> bool {
154 self.len == 0
155 }
156
157 #[inline]
159 pub fn as_raw(&self) -> *mut c_void {
160 self.ptr
161 }
162
163 #[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 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
194pub 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#[derive(Copy, Clone, Debug, Eq, PartialEq)]
209pub enum PrefetchTarget {
210 Device(i32),
212 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, }
223 }
224}
225
226pub 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
244pub 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
261pub 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 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 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 #[inline]
310 pub fn len(&self) -> usize {
311 self.len
312 }
313
314 #[inline]
316 pub fn is_empty(&self) -> bool {
317 self.len == 0
318 }
319
320 #[inline]
322 pub fn as_ptr(&self) -> *const T {
323 self.ptr
324 }
325
326 #[inline]
330 pub fn as_mut_ptr(&mut self) -> *mut T {
331 self.ptr
332 }
333
334 pub fn as_slice(&self) -> &[T] {
336 unsafe { core::slice::from_raw_parts(self.ptr, self.len) }
339 }
340
341 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
360pub mod pinned_flags {
365 pub use baracuda_cuda_sys::runtime::types::cudaHostAllocFlags::*;
366}
367
368pub 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 pub fn new(len: usize) -> Result<Self> {
391 Self::with_flags(len, 0)
392 }
393
394 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 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 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 #[inline]
431 pub fn len(&self) -> usize {
432 self.len
433 }
434 #[inline]
436 pub fn is_empty(&self) -> bool {
437 self.len == 0
438 }
439 #[inline]
441 pub fn as_ptr(&self) -> *const T {
442 self.ptr
443 }
444 #[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
477pub 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 pub fn register(slice: &'a mut [T]) -> Result<Self> {
499 Self::register_with_flags(slice, 0)
500 }
501
502 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 #[inline]
524 pub fn len(&self) -> usize {
525 self.len
526 }
527 #[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
547impl<T: DeviceRepr> DeviceBuffer<T> {
550 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 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 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 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
620pub 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
645pub 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}