baracuda_runtime/
launch.rs1use 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#[derive(Copy, Clone, Debug, Eq, PartialEq)]
14pub struct Dim3 {
15 pub x: u32,
17 pub y: u32,
19 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 #[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#[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 #[inline]
88 pub fn grid(mut self, grid: impl Into<Dim3>) -> Self {
89 self.grid = grid.into();
90 self
91 }
92
93 #[inline]
95 pub fn block(mut self, block: impl Into<Dim3>) -> Self {
96 self.block = block.into();
97 self
98 }
99
100 #[inline]
102 pub fn shared_mem_bytes(mut self, bytes: usize) -> Self {
103 self.shared_mem_bytes = bytes;
104 self
105 }
106
107 #[inline]
109 pub fn stream(mut self, stream: &'k Stream) -> Self {
110 self.stream = Some(stream);
111 self
112 }
113
114 #[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 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 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}