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::{Result, check};
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<()> {
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 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}