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::{check, Result};
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<()> { unsafe {
131        let r = runtime()?;
132        let cu = r.cuda_launch_kernel()?;
133        let stream_handle: cudaStream_t = self.stream.map_or(core::ptr::null_mut(), |s| s.as_raw());
134        let args_ptr = if self.args.is_empty() {
135            core::ptr::null_mut()
136        } else {
137            self.args.as_mut_ptr()
138        };
139        check(cu(
140            self.kernel.as_launch_ptr(),
141            self.grid.to_sys(),
142            self.block.to_sys(),
143            args_ptr,
144            self.shared_mem_bytes,
145            stream_handle,
146        ))
147    }}
148
149    /// Launch as a cooperative kernel — grid-wide sync via
150    /// `cooperative_groups::this_grid()`. All blocks must fit resident
151    /// on the device simultaneously; use
152    /// [`crate::Kernel::max_active_blocks_per_multiprocessor`] to size
153    /// the grid.
154    ///
155    /// # Safety
156    ///
157    /// Same as [`launch`](Self::launch) plus the kernel must be
158    /// compiled with cooperative-groups support.
159    pub unsafe fn launch_cooperative(mut self) -> Result<()> { unsafe {
160        let r = runtime()?;
161        let cu = r.cuda_launch_cooperative_kernel()?;
162        let stream_handle: cudaStream_t = self.stream.map_or(core::ptr::null_mut(), |s| s.as_raw());
163        let args_ptr = if self.args.is_empty() {
164            core::ptr::null_mut()
165        } else {
166            self.args.as_mut_ptr()
167        };
168        check(cu(
169            self.kernel.as_launch_ptr(),
170            self.grid.to_sys(),
171            self.block.to_sys(),
172            args_ptr,
173            self.shared_mem_bytes,
174            stream_handle,
175        ))
176    }}
177}