cubecl_cpp/cuda/
builtin.rs1use 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}