Skip to main content

cubecl_cpp/shared/
kernel.rs

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        // Program Scope -----------------------------------------------------
61        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        // Kernel signature --------------------------------------------------
67        D::compile_kernel_signature(
68            f,
69            &self.kernel_name,
70            &self.tensor_maps,
71            &self.buffers,
72            &self.flags,
73        )?;
74
75        // Body --------------------------------------------------------------
76        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
102/// Define a minimal version of C++'s `std::array` so we can match Rust semantics on arrays.
103pub 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}