Skip to main content

baracuda_runtime/
launch.rs

1//! Kernel launch builder for the Runtime API.
2
3use core::ffi::c_void;
4
5use baracuda_cuda_sys::runtime::{cudaStream_t, runtime, types::dim3};
6use baracuda_types::KernelArg;
7
8use crate::error::{Result, check};
9use crate::module::Kernel;
10use crate::stream::Stream;
11
12/// Grid / block size triple, matching [`baracuda_driver::Dim3`].
13#[derive(Copy, Clone, Debug, Eq, PartialEq)]
14pub struct Dim3 {
15    /// Extent in the X dimension. Must be `>= 1`.
16    pub x: u32,
17    /// Extent in the Y dimension. Use `1` for 1-D launches.
18    pub y: u32,
19    /// Extent in the Z dimension. Use `1` for 1-D and 2-D launches.
20    pub z: u32,
21}
22
23impl Dim3 {
24    #[inline]
25    fn to_sys(self) -> dim3 {
26        dim3::new(self.x, self.y, self.z)
27    }
28}
29
30impl From<u32> for Dim3 {
31    fn from(x: u32) -> Self {
32        Self { x, y: 1, z: 1 }
33    }
34}
35
36impl From<(u32, u32)> for Dim3 {
37    fn from((x, y): (u32, u32)) -> Self {
38        Self { x, y, z: 1 }
39    }
40}
41
42impl From<(u32, u32, u32)> for Dim3 {
43    fn from((x, y, z): (u32, u32, u32)) -> Self {
44        Self { x, y, z }
45    }
46}
47
48impl Kernel {
49    /// Start a kernel-launch builder for this kernel.
50    #[inline]
51    pub fn launch(&self) -> LaunchBuilder<'_> {
52        LaunchBuilder {
53            kernel: self,
54            grid: Dim3 { x: 1, y: 1, z: 1 },
55            block: Dim3 { x: 1, y: 1, z: 1 },
56            shared_mem_bytes: 0,
57            stream: None,
58            args: Vec::new(),
59        }
60    }
61}
62
63/// Builder produced by [`Kernel::launch`].
64#[must_use = "the launch builder does nothing until `.launch()` is called"]
65pub struct LaunchBuilder<'k> {
66    kernel: &'k Kernel,
67    grid: Dim3,
68    block: Dim3,
69    shared_mem_bytes: usize,
70    stream: Option<&'k Stream>,
71    args: Vec<*mut c_void>,
72}
73
74impl core::fmt::Debug for LaunchBuilder<'_> {
75    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
76        f.debug_struct("LaunchBuilder")
77            .field("grid", &self.grid)
78            .field("block", &self.block)
79            .field("shared_mem_bytes", &self.shared_mem_bytes)
80            .field("arg_count", &self.args.len())
81            .finish_non_exhaustive()
82    }
83}
84
85impl<'k> LaunchBuilder<'k> {
86    /// Set the grid dimensions (number of thread blocks).
87    #[inline]
88    pub fn grid(mut self, grid: impl Into<Dim3>) -> Self {
89        self.grid = grid.into();
90        self
91    }
92
93    /// Set the block dimensions (threads per block).
94    #[inline]
95    pub fn block(mut self, block: impl Into<Dim3>) -> Self {
96        self.block = block.into();
97        self
98    }
99
100    /// Reserve `bytes` of dynamic shared memory per block.
101    #[inline]
102    pub fn shared_mem_bytes(mut self, bytes: usize) -> Self {
103        self.shared_mem_bytes = bytes;
104        self
105    }
106
107    /// Enqueue on `stream` instead of the default stream.
108    #[inline]
109    pub fn stream(mut self, stream: &'k Stream) -> Self {
110        self.stream = Some(stream);
111        self
112    }
113
114    /// Append `arg` to the kernel argument list. Arguments are passed
115    /// positionally in the order they are added.
116    #[inline]
117    pub fn arg<K: KernelArg>(mut self, arg: K) -> Self {
118        self.args.push(arg.as_kernel_arg_ptr());
119        self
120    }
121
122    /// Enqueue the kernel.
123    ///
124    /// # Safety
125    ///
126    /// Same rules as [`baracuda_driver::LaunchBuilder::launch`]: argument
127    /// types and order must match the kernel's C signature, referenced
128    /// device memory must stay valid for the duration of device execution,
129    /// and grid/block dims must be within device limits.
130    pub unsafe fn launch(mut self) -> Result<()> {
131        unsafe {
132            let r = runtime()?;
133            let cu = r.cuda_launch_kernel()?;
134            let stream_handle: cudaStream_t =
135                self.stream.map_or(core::ptr::null_mut(), |s| s.as_raw());
136            let args_ptr = if self.args.is_empty() {
137                core::ptr::null_mut()
138            } else {
139                self.args.as_mut_ptr()
140            };
141            check(cu(
142                self.kernel.as_launch_ptr(),
143                self.grid.to_sys(),
144                self.block.to_sys(),
145                args_ptr,
146                self.shared_mem_bytes,
147                stream_handle,
148            ))
149        }
150    }
151
152    /// Launch as a cooperative kernel — grid-wide sync via
153    /// `cooperative_groups::this_grid()`. All blocks must fit resident
154    /// on the device simultaneously; use
155    /// [`crate::Kernel::max_active_blocks_per_multiprocessor`] to size
156    /// the grid.
157    ///
158    /// # Safety
159    ///
160    /// Same as [`launch`](Self::launch) plus the kernel must be
161    /// compiled with cooperative-groups support.
162    pub unsafe fn launch_cooperative(mut self) -> Result<()> {
163        unsafe {
164            let r = runtime()?;
165            let cu = r.cuda_launch_cooperative_kernel()?;
166            let stream_handle: cudaStream_t =
167                self.stream.map_or(core::ptr::null_mut(), |s| s.as_raw());
168            let args_ptr = if self.args.is_empty() {
169                core::ptr::null_mut()
170            } else {
171                self.args.as_mut_ptr()
172            };
173            check(cu(
174                self.kernel.as_launch_ptr(),
175                self.grid.to_sys(),
176                self.block.to_sys(),
177                args_ptr,
178                self.shared_mem_bytes,
179                stream_handle,
180            ))
181        }
182    }
183}