1use crate::shared::{Builtin, Component, Value};
2
3use super::{Body, Dialect, Elem, Flags, INFO_NAME, Item};
4use cubecl_core::{CubeDim, ir::Id, prelude::Visibility};
5
6use std::{collections::HashSet, fmt::Display};
7
8#[derive(Debug, PartialEq, Eq, Clone)]
9pub struct KernelArg<D: Dialect> {
10 pub id: Id,
11 pub value: Value<D>,
12 pub vis: Visibility,
13}
14
15#[derive(Debug, PartialEq, Eq, Clone)]
16pub struct SharedMemory<D: Dialect> {
17 pub ptr: Value<D>,
18 pub value_ty: Item<D>,
19 pub align: usize,
20 pub offset: usize,
21}
22
23impl<D: Dialect> SharedMemory<D> {
24 pub fn size(&self) -> usize {
25 self.value_ty.size()
26 }
27}
28
29#[derive(Debug, Clone)]
30pub struct ComputeKernel<D: Dialect> {
31 pub tensor_maps: Vec<KernelArg<D>>,
32 pub buffers: Vec<KernelArg<D>>,
33 pub scalars: Vec<(Elem<D>, usize)>,
34 pub info: cubecl_core::Info,
35 pub meta_static_len: usize,
36 pub body: Body<D>,
37 pub cube_dim: CubeDim,
38 pub cluster_dim: Option<CubeDim>,
39 pub extensions: Vec<D::Extension>,
40 pub flags: Flags<D>,
41 pub items: HashSet<super::Item<D>>,
42 pub kernel_name: String,
43}
44
45impl<D: Dialect> ComputeKernel<D> {
46 pub fn shared_memory_size(&self) -> usize {
47 let smems = self.body.shared_memories.iter();
48 let ends = smems.map(|it| it.offset + it.size());
49 ends.max().unwrap_or_default()
50 }
51}
52
53impl<D: Dialect> Display for ComputeKernel<D> {
54 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
55 let mut flags = self.flags.clone();
56 if !self.tensor_maps.is_empty() {
57 flags.inst_tma = true;
58 }
59
60 D::compile_includes(f, &flags)?;
62 D::compile_type_definitions(f, &self.items, &self.scalars, &self.info, &flags)?;
63 D::compile_polyfills(f, &flags)?;
64 D::compile_extensions(f, &self.extensions)?;
65
66 D::compile_kernel_signature(
68 f,
69 &self.kernel_name,
70 &self.tensor_maps,
71 &self.buffers,
72 &self.flags,
73 )?;
74
75 f.write_str(" {\n")?;
77 compile_cube_builtin_bindings_decl::<D>(f, &self.flags)?;
78 write!(f, "{}", self.body)?;
79 f.write_str("\n}")?;
80
81 Ok(())
82 }
83}
84
85pub fn type_definitions<D: Dialect>(f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
86 writeln!(f, "typedef unsigned int uint;")?;
87 writeln!(f, "typedef unsigned char uint8;")?;
88 writeln!(f, "typedef unsigned short uint16;")?;
89 writeln!(f, "typedef unsigned int uint32;")?;
90 writeln!(f, "typedef unsigned long long int uint64;")?;
91
92 writeln!(f, "typedef signed char int8;")?;
93 writeln!(f, "typedef signed short int16;")?;
94 writeln!(f, "typedef signed int int32;")?;
95 writeln!(f, "typedef signed long long int int64;")?;
96
97 define_array_polyfill(f, "__device__")?;
98
99 Ok(())
100}
101
102pub fn define_array_polyfill(
104 f: &mut core::fmt::Formatter<'_>,
105 function_class: &str,
106) -> core::fmt::Result {
107 writeln!(
108 f,
109 "
110template <typename T, size_t N>
111struct array {{
112 T data[N];
113 {function_class} T& operator[](size_t i) {{ return data[i]; }}
114 {function_class} const T& operator[](size_t i) const {{ return data[i]; }}
115}};"
116 )
117}
118
119pub fn type_vectorized_definitions<D: Dialect>(
120 f: &mut std::fmt::Formatter<'_>,
121 items: &HashSet<Item<D>>,
122) -> std::fmt::Result {
123 for item in items.iter().filter(|it| it.vectorization() > 1) {
124 let elem = item.elem();
125 let size = item.vectorization();
126 let alignment = elem.size() * size;
127 if size > 1 {
128 write!(
129 f,
130 "
131struct __align__({alignment}) {item} {{"
132 )?;
133
134 for i in 0..size {
135 write!(
136 f,
137 "
138 {elem} i_{i};"
139 )?;
140 }
141
142 f.write_str("\n};")?;
143 }
144 }
145 Ok(())
146}
147
148pub fn type_info_definition_sized<D: Dialect>(
149 f: &mut std::fmt::Formatter<'_>,
150 info: &cubecl_core::Info,
151 scalars: &[(Elem<D>, usize)],
152 address_type: Item<D>,
153) -> std::fmt::Result {
154 let scalars = info
155 .scalars
156 .iter()
157 .zip(scalars)
158 .map(|(field, (ty, _))| format!("{ty} scalars_{ty}[{}];", field.padded_size()))
159 .collect::<Vec<_>>()
160 .join("\n");
161 let static_meta = info
162 .sized_meta
163 .as_ref()
164 .map(|field| format!("{address_type} static_meta[{}];", field.padded_size()))
165 .unwrap_or_default();
166 write!(
167 f,
168 "
169struct info_st {{
170 {scalars}{static_meta}
171}};
172"
173 )
174}
175
176pub fn compile_bindings<D: Dialect>(
177 f: &mut core::fmt::Formatter<'_>,
178 tensor_maps: &[KernelArg<D>],
179 buffers: &[KernelArg<D>],
180 trailing_comma: bool,
181) -> core::fmt::Result {
182 write!(f, " ")?;
183
184 let mut args = Vec::new();
185
186 args.extend(tensor_maps.iter().map(|binding| {
187 format!(
188 "const __grid_constant__ {} {}",
189 binding.value.item(),
190 binding.value
191 )
192 }));
193 args.extend(buffers.iter().map(|binding| {
194 let ty = binding.value.item();
195 match binding.vis {
196 Visibility::Read | Visibility::Uniform => {
197 format!("const {ty} __restrict__ {}", binding.value)
198 }
199 _ => {
200 format!("{ty} __restrict__ {}", binding.value)
201 }
202 }
203 }));
204
205 write!(f, "{}", args.join(", "))?;
206 if trailing_comma {
207 f.write_str(", ")?;
208 }
209 Ok(())
210}
211
212pub fn compile_info_dynamic<D: Dialect>(
213 f: &mut std::fmt::Formatter<'_>,
214 flags: &Flags<D>,
215) -> core::fmt::Result {
216 if flags.has_info {
217 write!(f, "const info_st* __restrict__ {INFO_NAME}_ptr")
218 } else {
219 Ok(())
220 }
221}
222
223pub fn compile_info_static<D: Dialect>(
224 f: &mut std::fmt::Formatter<'_>,
225 flags: &Flags<D>,
226) -> core::fmt::Result {
227 let mut inputs = Vec::new();
228
229 if flags.has_dynamic_meta {
230 inputs.push(format!(
231 "const {}* __restrict__ dynamic_meta",
232 flags.address_type
233 ))
234 }
235
236 if flags.has_info {
237 inputs.push(format!("const __grid_constant__ info_st {INFO_NAME}"));
238 }
239
240 write!(f, "{}", inputs.join(", "))
241}
242
243fn compile_cube_builtin_bindings_decl<D: Dialect>(
244 f: &mut core::fmt::Formatter<'_>,
245 settings: &Flags<D>,
246) -> core::fmt::Result {
247 if settings.indexes.absolute_pos_tuple {
248 D::compile_absolute_pos_tuple_computation(f)?;
249 }
250
251 if settings.indexes.unit_pos {
252 D::compile_unit_pos_computation(f)?;
253 }
254
255 if settings.indexes.absolute_pos {
256 let value = Builtin::<D>::AbsolutePos(*settings.address_type.elem());
257 let ty = value.item();
258 let absolute_pos_x = Builtin::<D>::AbsolutePosX.fmt_cast_to(ty);
259 let absolute_pos_y = Builtin::<D>::AbsolutePosY.fmt_cast_to(ty);
260 let absolute_pos_z = Builtin::<D>::AbsolutePosZ.fmt_cast_to(ty);
261 let cube_count_x = Builtin::<D>::CubeCountX.fmt_cast_to(ty);
262 let cube_count_y = Builtin::<D>::CubeCountY.fmt_cast_to(ty);
263 let cube_dim_x = Builtin::<D>::CubeDimX.fmt_cast_to(ty);
264 let cube_dim_y = Builtin::<D>::CubeDimY.fmt_cast_to(ty);
265 writeln!(
266 f,
267 "{ty} {value} = (
268 {absolute_pos_z} * {cube_count_x} * {cube_dim_x} * {cube_count_y} * {cube_dim_y})
269 + ({absolute_pos_y} * {cube_count_x} * {cube_dim_x})
270 + {absolute_pos_x};"
271 )?;
272 }
273
274 if settings.indexes.cube_dim {
275 let value = Builtin::<D>::CubeDim;
276 let ty = value.item();
277 let cube_dim_x = Builtin::<D>::CubeDimX;
278 let cube_dim_y = Builtin::<D>::CubeDimY;
279 let cube_dim_z = Builtin::<D>::CubeDimZ;
280 writeln!(
281 f,
282 "{ty} {value} = {cube_dim_x} * {cube_dim_y} * {cube_dim_z};"
283 )?;
284 }
285
286 if settings.indexes.cube_count {
287 let value = Builtin::<D>::CubeCount(*settings.address_type.elem());
288 let ty = value.item();
289 let cube_count_x = Builtin::<D>::CubeCountX.fmt_cast_to(ty);
290 let cube_count_y = Builtin::<D>::CubeCountY.fmt_cast_to(ty);
291 let cube_count_z = Builtin::<D>::CubeCountZ.fmt_cast_to(ty);
292 writeln!(
293 f,
294 "{ty} {value} = {cube_count_x} * {cube_count_y} * {cube_count_z};"
295 )?;
296 }
297
298 if settings.indexes.cube_pos {
299 let value = Builtin::<D>::CubePos(*settings.address_type.elem());
300 let ty = value.item();
301 let cube_pos_x = Builtin::<D>::CubePosX.fmt_cast_to(ty);
302 let cube_pos_y = Builtin::<D>::CubePosY.fmt_cast_to(ty);
303 let cube_pos_z = Builtin::<D>::CubePosZ.fmt_cast_to(ty);
304 let cube_count_x = Builtin::<D>::CubeCountX.fmt_cast_to(ty);
305 let cube_count_y = Builtin::<D>::CubeCountY.fmt_cast_to(ty);
306 writeln!(
307 f,
308 "{ty} {value} = ({cube_pos_z} * {cube_count_y} * {cube_count_x}) + ({cube_pos_y} * {cube_count_x}) + {cube_pos_x};"
309 )?;
310 }
311
312 if settings.indexes.plane_dim_checked {
313 let plane_dim = Builtin::<D>::PlaneDim;
314 let value = Builtin::<D>::PlaneDimChecked;
315 let ty = value.item();
316 let cube_dim_x = Builtin::<D>::CubeDimX;
317 let cube_dim_y = Builtin::<D>::CubeDimY;
318 let cube_dim_z = Builtin::<D>::CubeDimZ;
319 writeln!(
320 f,
321 "{ty} {value} = min({plane_dim}, {cube_dim_x} * {cube_dim_y} * {cube_dim_z});"
322 )?;
323 }
324
325 if settings.thread_block {
326 f.write_str(
327 "
328cooperative_groups::thread_block thread_block = cooperative_groups::this_thread_block();
329",
330 )?;
331 }
332
333 if settings.indexes.cluster_pos {
334 f.write_str(
335 "
336cooperative_groups::cluster_group cluster = cooperative_groups::this_cluster();
337",
338 )?;
339 }
340
341 Ok(())
342}