Skip to main content

baracuda_runtime/
user_object.rs

1//! Runtime-API graph user objects (CUDA 12.0+).
2//!
3//! Refcounted RAII slot you can attach to a graph via
4//! [`Graph::retain_user_object`]; when the graph releases the last
5//! reference, the destructor runs. Mirrors the Driver-side wrapper.
6
7use core::ffi::c_void;
8
9use baracuda_cuda_sys::runtime::{cudaUserObject_t, runtime};
10
11use crate::error::{Result, check};
12
13/// A refcounted user object.
14pub struct UserObject {
15    handle: cudaUserObject_t,
16}
17
18unsafe impl Send for UserObject {}
19unsafe impl Sync for UserObject {}
20
21impl core::fmt::Debug for UserObject {
22    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
23        f.debug_struct("UserObject")
24            .field("handle", &self.handle)
25            .finish_non_exhaustive()
26    }
27}
28
29unsafe extern "C" fn destroy_trampoline(user_data: *mut c_void) {
30    if user_data.is_null() {
31        return;
32    }
33    let f: Box<Box<dyn FnOnce() + Send>> =
34        unsafe { Box::from_raw(user_data as *mut Box<dyn FnOnce() + Send>) };
35    (*f)();
36}
37
38impl UserObject {
39    /// Create a user object whose destructor is `destroy`.
40    /// `initial_refcount` must be >= 1.
41    pub fn new<F>(destroy: F, initial_refcount: u32) -> Result<Self>
42    where
43        F: FnOnce() + Send + 'static,
44    {
45        let boxed: Box<Box<dyn FnOnce() + Send>> = Box::new(Box::new(destroy));
46        let raw = Box::into_raw(boxed) as *mut c_void;
47        let r = runtime()?;
48        let cu = r.cuda_user_object_create()?;
49        let mut object: cudaUserObject_t = core::ptr::null_mut();
50        // CUDA requires flags == cudaGraphUserObjectMove (1) currently.
51        const CUDA_USER_OBJECT_NO_DESTRUCTOR_SYNC: core::ffi::c_uint = 1;
52        let rc = unsafe {
53            cu(
54                &mut object,
55                raw,
56                Some(destroy_trampoline),
57                initial_refcount,
58                CUDA_USER_OBJECT_NO_DESTRUCTOR_SYNC,
59            )
60        };
61        if rc != baracuda_cuda_sys::runtime::cudaError_t::Success {
62            drop(unsafe { Box::from_raw(raw as *mut Box<dyn FnOnce() + Send>) });
63            return Err(crate::error::Error::Status { status: rc });
64        }
65        Ok(Self { handle: object })
66    }
67
68    /// Safe wrapper for `cudaUserObjectRetain`. Add `count` references to
69    /// this user object.
70    pub fn retain(&self, count: u32) -> Result<()> {
71        let r = runtime()?;
72        let cu = r.cuda_user_object_retain()?;
73        check(unsafe { cu(self.handle, count) })
74    }
75
76    /// Safe wrapper for `cudaUserObjectRelease`. Drop `count` references
77    /// from this user object; when the count reaches zero the destructor
78    /// supplied at creation runs.
79    pub fn release(&self, count: u32) -> Result<()> {
80        let r = runtime()?;
81        let cu = r.cuda_user_object_release()?;
82        check(unsafe { cu(self.handle, count) })
83    }
84
85    /// Raw `cudaUserObject_t` handle. Use with care — owned by `self`.
86    #[inline]
87    pub fn as_raw(&self) -> cudaUserObject_t {
88        self.handle
89    }
90}
91
92impl Drop for UserObject {
93    fn drop(&mut self) {
94        if self.handle.is_null() {
95            return;
96        }
97        if let Ok(r) = runtime() {
98            if let Ok(cu) = r.cuda_user_object_release() {
99                let _ = unsafe { cu(self.handle, 1) };
100            }
101        }
102    }
103}
104
105impl crate::Graph {
106    /// Have this graph retain `count` references to `object`.
107    pub fn retain_user_object(&self, object: &UserObject, count: u32, flags: u32) -> Result<()> {
108        let r = runtime()?;
109        let cu = r.cuda_graph_retain_user_object()?;
110        check(unsafe { cu(self.as_raw(), object.as_raw(), count, flags) })
111    }
112
113    /// Safe wrapper for `cudaGraphReleaseUserObject`. Drop `count`
114    /// references this graph holds on `object`.
115    pub fn release_user_object(&self, object: &UserObject, count: u32) -> Result<()> {
116        let r = runtime()?;
117        let cu = r.cuda_graph_release_user_object()?;
118        check(unsafe { cu(self.as_raw(), object.as_raw(), count) })
119    }
120}