Skip to main content

cubecl_cpp/cuda/
builtin.rs

1use cubecl_core::ir::{Builtin, Scope, dialect::general::ReadBuiltinOp, prelude::*};
2use pliron::context::Context;
3
4use crate::{
5    cuda::cuda_op_with_out,
6    shared::{
7        CompilationOptions, CompilationState,
8        builtin::{LowerBuiltins, SharedBuiltin},
9        signature::RequiresIncludesOp,
10    },
11    target::Cuda,
12};
13
14cuda_op_with_out!(ReadBuiltinOp, |op, ctx| {
15    op.builtin(ctx).0.display_cuda(ctx)
16});
17
18#[op_interface_impl]
19impl RequiresIncludesOp<Cuda> for ReadBuiltinOp {
20    fn includes(&self, ctx: &Context) -> Vec<String> {
21        match self.builtin(ctx).0 {
22            Builtin::CubePosCluster
23            | Builtin::CubePosClusterX
24            | Builtin::CubePosClusterY
25            | Builtin::CubePosClusterZ
26            | Builtin::CubeClusterDim
27            | Builtin::CubeClusterDimX
28            | Builtin::CubeClusterDimY
29            | Builtin::CubeClusterDimZ => vec!["cooperative_groups.h".into()],
30            _ => vec![],
31        }
32    }
33}
34
35impl MatchRewrite for LowerBuiltins<Cuda> {
36    fn r#match(&mut self, ctx: &Context, op: Ptr<Operation>) -> bool {
37        op.is_op::<ReadBuiltinOp>(ctx)
38    }
39
40    fn rewrite(
41        &mut self,
42        ctx: &mut Context,
43        rewriter: &mut MatchRewriter,
44        op: Ptr<Operation>,
45    ) -> Result<()> {
46        let builtin = op.as_op::<ReadBuiltinOp>(ctx).unwrap().builtin(ctx).0;
47        let scope = Scope::from_context_and_inserter(ctx, rewriter);
48        if let Some(new_value) = builtin.maybe_lower_shared(&scope) {
49            rewriter.replace_operation_with_values(ctx, op, vec![new_value]);
50        }
51        Ok(())
52    }
53}
54
55pub(crate) trait CudaBuiltin {
56    fn display_cuda(&self, ctx: &Context) -> String;
57}
58
59impl CudaBuiltin for Builtin {
60    fn display_cuda(&self, ctx: &Context) -> String {
61        let clusters = ctx
62            .aux_ty::<CompilationOptions>()
63            .supports_features
64            .clusters;
65        let cluster_dim = ctx.aux_ty::<CompilationState>().cluster_dim;
66        match self {
67            Builtin::CubePosCluster if clusters => {
68                "cooperative_groups::this_cluster().block_rank()".into()
69            }
70            Builtin::CubePosClusterX if clusters => {
71                "cooperative_groups::this_cluster().block_index().x".into()
72            }
73            Builtin::CubePosClusterY if clusters => {
74                "cooperative_groups::this_cluster().block_index().y".into()
75            }
76            Builtin::CubePosClusterZ if clusters => {
77                "cooperative_groups::this_cluster().block_index().z".into()
78            }
79            Builtin::CubeClusterDim if clusters => format!("{}", cluster_dim.num_elems()),
80            Builtin::CubeClusterDimX if clusters => format!("{}", cluster_dim.x),
81            Builtin::CubeClusterDimY if clusters => format!("{}", cluster_dim.y),
82            Builtin::CubeClusterDimZ if clusters => format!("{}", cluster_dim.z),
83            _ => SharedBuiltin::display(self).into(),
84        }
85    }
86}