cubecl_core/compute/
launcher.rs1use alloc::{boxed::Box, vec::Vec};
2
3use crate::prelude::{BufferArg, TensorArg, TensorMapArg, TensorMapKind};
4use crate::{InfoBuilder, ScalarArgType};
5#[cfg(feature = "std")]
6use core::cell::RefCell;
7use cubecl_ir::{AddressType, ElemType, Scope, settings::KernelSettings};
8use cubecl_runtime::kernel::BufferIOAttr;
9use cubecl_runtime::server::{BufferBinding, CubeCount, KernelResource, TensorMapBinding};
10use cubecl_runtime::{client::Client, kernel::CubeKernel, server::KernelArguments};
11
12#[cfg(feature = "std")]
13std::thread_local! {
14 static INFO: RefCell<InfoBuilder> = RefCell::new(InfoBuilder::default());
15 static SCOPE: RefCell<Scope> = RefCell::new(Scope::dummy());
17}
18
19pub struct KernelLauncher {
21 resources: Vec<KernelResource>,
22 declared_io: Vec<BufferIOAttr>,
25 declaring: BufferIOAttr,
27 address_type: AddressType,
28 pub settings: KernelSettings,
29 #[cfg(not(feature = "std"))]
30 info: InfoBuilder,
31 #[cfg(not(feature = "std"))]
32 pub scope: Scope,
33}
34
35impl KernelLauncher {
36 #[cfg(feature = "std")]
37 pub fn with_scope<T>(&mut self, fun: impl FnMut(&Scope) -> T) -> T {
38 SCOPE.with_borrow(fun)
39 }
40
41 #[cfg(not(feature = "std"))]
42 pub fn with_scope<T>(&mut self, mut fun: impl FnMut(&Scope) -> T) -> T {
43 fun(&self.scope)
44 }
45
46 #[cfg(feature = "std")]
47 fn with_info<T>(&mut self, fun: impl FnMut(&mut InfoBuilder) -> T) -> T {
48 INFO.with_borrow_mut(fun)
49 }
50
51 #[cfg(not(feature = "std"))]
52 fn with_info<T>(&mut self, mut fun: impl FnMut(&mut InfoBuilder) -> T) -> T {
53 fun(&mut self.info)
54 }
55
56 pub fn register_scalar<C: ScalarArgType>(&mut self, scalar: C) {
58 self.with_info(|info| info.scalars.push(scalar));
59 }
60
61 pub fn register_scalar_raw(&mut self, bytes: &[u8], dtype: ElemType) {
63 self.with_info(|info| info.scalars.push_raw(bytes, dtype));
64 }
65
66 #[track_caller]
68 pub fn launch<K: CubeKernel>(self, cube_count: CubeCount, kernel: K, client: &Client) {
69 let bindings = self.into_bindings();
70 let kernel = Box::new(kernel);
71
72 client.launch(kernel, cube_count, bindings)
73 }
74
75 pub fn discard(self) {
84 let _ = self.into_bindings();
85 }
86
87 fn into_bindings(mut self) -> KernelArguments {
97 let mut bindings = KernelArguments::new();
98 let address_type = self.address_type;
99 let info = self.with_info(|info| info.finish(address_type));
100
101 bindings.resources = self.resources;
102 bindings.declared_io = self.declared_io;
103 bindings.info = info;
104
105 bindings
106 }
107}
108
109impl KernelLauncher {
111 pub fn declare_io(&mut self, io: BufferIOAttr) {
123 self.declaring = io;
124 }
125
126 fn alias_io(&mut self, input_pos: usize) {
135 if self.declaring.is_writable()
136 && let Some(io) = self.declared_io.get_mut(input_pos)
137 {
138 *io = BufferIOAttr::ReadWrite;
139 }
140 }
141
142 fn push_resource(&mut self, resource: KernelResource) {
144 let io = match &resource {
145 KernelResource::TensorMap(_) => BufferIOAttr::ReadWrite,
150 KernelResource::Buffer(_) => self.declaring,
151 };
152 self.declared_io.push(io);
153 self.resources.push(resource);
154 }
155
156 pub fn register_tensor(&mut self, tensor: TensorArg, elem_size: usize) {
158 if let Some(tensor) = self.process_tensor(tensor, elem_size) {
159 self.push_resource(KernelResource::Buffer(tensor));
160 }
161 }
162
163 fn process_tensor(&mut self, tensor: TensorArg, elem_size: usize) -> Option<BufferBinding> {
164 let tensor = match tensor {
165 TensorArg::Handle { handle, .. } => handle,
166 TensorArg::Alias { input_pos, .. } => {
167 self.alias_io(input_pos);
168 return None;
169 }
170 };
171
172 let buffer_len = tensor.handle.size_in_used() / elem_size as u64;
173 let address_type = self.address_type;
174
175 self.with_info(|info| {
176 info.metadata.register_tensor(
177 buffer_len,
178 tensor.shape.clone(),
179 tensor.strides.clone(),
180 address_type,
181 )
182 });
183 Some(tensor.handle)
184 }
185
186 pub fn register_buffer(&mut self, array: BufferArg, elem_size: usize) {
188 if let Some(tensor) = self.process_buffer(array, elem_size) {
189 self.push_resource(KernelResource::Buffer(tensor));
190 }
191 }
192
193 fn process_buffer(&mut self, array: BufferArg, elem_size: usize) -> Option<BufferBinding> {
194 let array = match array {
195 BufferArg::Handle { handle, .. } => handle,
196 BufferArg::Alias { input_pos, .. } => {
197 self.alias_io(input_pos);
198 return None;
199 }
200 };
201
202 let buffer_len = array.handle.size_in_used() / elem_size as u64;
203 let address_type = self.address_type;
204 self.with_info(|info| info.metadata.register_buffer(buffer_len, address_type));
205 Some(array.handle)
206 }
207
208 pub fn register_tensor_map<K: TensorMapKind>(
210 &mut self,
211 map: TensorMapArg<K>,
212 elem_size: usize,
213 ) {
214 let binding = self
215 .process_tensor(map.tensor, elem_size)
216 .expect("Can't use alias for TensorMap");
217
218 let map = map.metadata.clone();
219 self.push_resource(KernelResource::TensorMap(TensorMapBinding { binding, map }));
220 }
221}
222
223impl KernelLauncher {
224 pub fn new(settings: KernelSettings) -> Self {
225 Self {
226 address_type: settings.address_type,
227 settings,
228 resources: Vec::new(),
229 declared_io: Vec::new(),
230 declaring: BufferIOAttr::ReadWrite,
231 #[cfg(not(feature = "std"))]
232 info: InfoBuilder::default(),
233 #[cfg(not(feature = "std"))]
234 scope: Scope::dummy(),
235 }
236 }
237}
238
239#[cfg(test)]
240mod tests {
241 use super::*;
242 use cubecl_ir::settings::{Dim3, ExecutionMode};
243
244 fn settings() -> KernelSettings {
245 KernelSettings::new(Dim3::new_single(), ExecutionMode::Checked, AddressType::U32)
246 }
247
248 fn info_of(launcher: KernelLauncher) -> Vec<u64> {
249 launcher.into_bindings().info.data
250 }
251
252 #[test]
258 fn a_discarded_launcher_leaves_nothing_for_the_next_launch() {
259 let empty = info_of(KernelLauncher::new(settings()));
260
261 let mut registered = KernelLauncher::new(settings());
265 registered.register_scalar(1u32);
266 assert_ne!(info_of(registered), empty);
267
268 let mut dummy = KernelLauncher::new(settings());
269 dummy.register_scalar(1u32);
270 dummy.discard();
271
272 assert_eq!(
273 info_of(KernelLauncher::new(settings())),
274 empty,
275 "a discarded launcher left its scalars behind for the next launch"
276 );
277 }
278}