Skip to main content

cubecl_cpp/shared/
base.rs

1use crate::{
2    cuda::{mma::CudaCmmaCompiler, packed_ops::PackOpsPass},
3    error::EmissionErrors,
4    hip::{arch::AmdWmma, mma::HipCmmaCompiler},
5    shared::{
6        OpExtCPP,
7        builtin::{LowerBuiltins, LowerBuiltinsPass},
8        convert::PromoteUnsupportedTypesPass,
9        lowering::{LowerOpsAfterUnrollCppPass, LowerOpsCppPass},
10        metadata::LowerInfoPass,
11        signature::{
12            CollectIncludesPass, DeclareComplexHelpersOp, DeclareInfoTypeOp,
13            DeclareVectorTypesPass, buffer_io, buffers, shared_memory_size,
14        },
15        unroll::CppUnrollPass,
16    },
17    target::{CppTarget, Shared, Target},
18};
19use cubecl_runtime::kernel::BufferIOAttr;
20
21use super::ComputeKernel;
22use core::marker::PhantomData;
23use cubecl_core::{
24    ir::{
25        AddressType, ContextExt, DeviceProperties, ElemType, FloatKind, IntKind, Type, UIntKind,
26        features::{AtomicUsage, EnumSet, TypeUsage},
27        interfaces::TypedExt,
28        metadata::Info,
29        rewrite::{SimplifyOpsPass, visit_all_values},
30        settings::Dim3,
31        types::scalar::{Complex32Type, Complex64Type},
32    },
33    post_processing::{
34        bitwise::PromoteBitwisePass,
35        checked_io::{CheckedIo, CheckedIoPass},
36        minifloat::{Fp8Container, LowerMinifloatCast, LowerMinifloatCastPass},
37        saturating::LowerSaturatingArithmeticPass,
38    },
39    prelude::KernelDefinition,
40};
41use cubecl_environment::backtrace::BackTrace;
42use cubecl_opt::passes::{
43    alloc_shared_memory::AllocateSharedMemoryBlockPass,
44    annotate_buffer_visibility::AnnotateGlobalVisibilityPass, inst_combine::InstCombinePass,
45    sccp::SCCPPass, simple_cse::SimpleCSEPass, sroa::SROAPass,
46};
47use cubecl_runtime::compiler::{CompilationError, Compiler};
48use pliron::{
49    builtin::ops::{FuncOp, ModuleOp},
50    context::Context,
51    irbuild::match_rewrite::MatchRewrite,
52    op::Op,
53    operation::verify_operation,
54    opts::{dce::DCEPass, mem2reg::Mem2RegPass},
55    pass::{AnalysisManager, NestedOpsPass, OpPass, PMConfig, Pass, Passes},
56};
57use std::fmt::Debug;
58
59pub(crate) fn closure_inference_hack<T, R>(
60    val: &T,
61    ctx: &Context,
62    func: impl FnOnce(&T, &Context) -> R,
63) -> R {
64    func(val, ctx)
65}
66
67macro_rules! scoped_block {
68    ($($lines: expr)*) => {{
69        let mut out = String::from("[&]{\n");
70        $(
71            out.push_str(&$lines);
72            out.push_str("\n");
73        )*
74        out.push_str("}()");
75        out
76    }};
77}
78pub(crate) use scoped_block;
79
80#[derive(Clone, Copy, Debug)]
81pub struct CompilationOptions {
82    pub warp_size: usize,
83    pub supports_features: CppSupportedFeatures,
84    /// AMD only, and `None` on hardware without WMMA.
85    pub amd_wmma: Option<AmdWmma>,
86}
87
88pub struct CompilationState {
89    pub cube_dim: Dim3,
90    pub cluster_dim: Dim3,
91    pub info: Info,
92}
93
94#[derive(Clone, Copy, Debug, Default)]
95pub struct CppSupportedFeatures {
96    pub grid_constants: bool,
97    pub clusters: bool,
98    pub fast_math: bool,
99    pub fast_tanh: bool,
100    pub elect_sync: bool,
101    pub dp4a: bool,
102}
103
104impl Default for CompilationOptions {
105    fn default() -> Self {
106        Self {
107            warp_size: 32,
108            supports_features: Default::default(),
109            amd_wmma: None,
110        }
111    }
112}
113
114#[allow(clippy::too_many_arguments)]
115#[derive(Clone, Copy, Debug, Default)]
116pub struct CppCompiler<T: CppTarget> {
117    _target: PhantomData<T>,
118}
119
120impl<T: CppTarget> Compiler for CppCompiler<T>
121where
122    LowerBuiltins<T>: MatchRewrite,
123{
124    type Representation = ComputeKernel;
125    type CompilationOptions = CompilationOptions;
126
127    fn buffer_io(repr: &Self::Representation) -> Option<Vec<BufferIOAttr>> {
128        Some(repr.io.clone())
129    }
130
131    fn compile(
132        &mut self,
133        kernel: KernelDefinition,
134        compilation_options: &Self::CompilationOptions,
135    ) -> Result<Self::Representation, CompilationError> {
136        let errors = kernel.body.pop_errors();
137        if !errors.is_empty() {
138            let mut reason = "Can't compile cpp kernel\nCaused by:\n  ".to_string();
139            for error in errors {
140                reason += error.as_str();
141                reason += "\n";
142            }
143
144            return Err(CompilationError::Validation {
145                reason,
146                backtrace: BackTrace::capture(),
147            });
148        }
149
150        self.compile_ir(kernel, *compilation_options)
151    }
152
153    fn extension(&self) -> &'static str {
154        "cpp"
155    }
156
157    fn lang_tag(&self) -> &'static str {
158        match T::target() {
159            Target::Cuda => "cuda",
160            Target::Hip => "hip",
161            Target::Metal => "msl",
162        }
163    }
164}
165
166impl<T: CppTarget> CppCompiler<T>
167where
168    LowerBuiltins<T>: MatchRewrite,
169{
170    fn compile_ir(
171        self,
172        kernel: KernelDefinition,
173        compilation_options: CompilationOptions,
174    ) -> Result<ComputeKernel, CompilationError> {
175        let module = kernel.body.state().module;
176        let module_op = module.get_operation();
177        let entry_func = kernel.body.state().entry_func;
178        let mut ctx = kernel.body.into_context().expect("Should be owned scope");
179
180        let state = CompilationState {
181            cube_dim: kernel.settings.cube_dim,
182            cluster_dim: kernel.settings.cluster_dim.unwrap_or(Dim3::new_single()),
183            info: kernel.info,
184        };
185
186        ctx.set_aux_ty(compilation_options);
187        ctx.set_aux_ty(state);
188        ctx.set_aux_ty(T::target());
189
190        ctx.set_aux_ty(CudaCmmaCompiler::Cpp);
191        ctx.set_aux_ty(HipCmmaCompiler::RocWmma);
192
193        verify_operation(module.get_operation(), &ctx)?;
194
195        // This is an op so it can be inserted after the includes, which is important for scalars
196        // that need includes. I wish C++ didn't have ordering dependent declarations...
197        let decl_types = DeclareInfoTypeOp::new(&mut ctx);
198        decl_types
199            .get_operation()
200            .insert_before(&ctx, entry_func.get_operation());
201
202        let mut has_complex = false;
203        visit_all_values(
204            &ctx,
205            &mut has_complex,
206            module_op,
207            |ctx, has_complex, value| {
208                if let Some(ty) = value.try_get_scalar_elem_ty(ctx) {
209                    let ty = ty.deref(ctx);
210                    *has_complex |= ty.is::<Complex32Type>() || ty.is::<Complex64Type>();
211                }
212            },
213        );
214        if has_complex && T::target() == Target::Cuda {
215            DeclareComplexHelpersOp::new(&mut ctx)
216                .get_operation()
217                .insert_before(&ctx, entry_func.get_operation());
218        }
219
220        #[cfg(feature = "pliron-dump")]
221        let dump_dir = kernel_dir_name(&kernel.settings.kernel_name);
222
223        let config = PMConfig {
224            #[cfg(feature = "pliron-dump")]
225            ir_printing_dir: dump_dir.clone(),
226            print_after_all: cfg!(feature = "pliron-dump"),
227            ..Default::default()
228        };
229
230        let mut analyses = AnalysisManager::default();
231        analyses.set_config(config);
232
233        let mut passes = OpPass::<ModuleOp, Passes>::default();
234        let mut func_passes = OpPass::<FuncOp, Passes>::default();
235
236        func_passes.add_pass(LowerInfoPass);
237        func_passes.add_pass(SROAPass);
238        func_passes.add_pass(CheckedIoPass::new(CheckedIo::new(
239            kernel.settings.execution_mode,
240            kernel.settings.kernel_name,
241        )));
242        func_passes.add_pass(AllocateSharedMemoryBlockPass);
243
244        // CUDA converts fp8 with cuda_fp8.h, which carries its own software path below sm_89.
245        let native_fp8 = match T::target() {
246            Target::Cuda => EnumSet::all(),
247            Target::Hip | Target::Metal => EnumSet::empty(),
248        };
249        func_passes.add_pass(LowerMinifloatCastPass::new(LowerMinifloatCast::new(
250            native_fp8,
251            Fp8Container::Bytes,
252        )));
253
254        // Shared lowerings can create ops that need target-specific lowerings, but target-specific
255        // lowerings should take priority. So we just run the target-specific lowerings twice.
256        func_passes.add_pass(LowerOpsCppPass::<T>::default());
257        func_passes.add_pass(LowerOpsCppPass::<Shared>::default());
258        func_passes.add_pass(LowerOpsCppPass::<T>::default());
259
260        if T::target() != Target::Metal {
261            func_passes.add_pass(LowerSaturatingArithmeticPass::default());
262        }
263
264        if T::target() == Target::Cuda {
265            func_passes.add_pass(PackOpsPass::default());
266        }
267
268        func_passes.add_pass(CppUnrollPass::default());
269        func_passes.add_pass(LowerBuiltinsPass::<T>::default());
270        func_passes.add_pass(LowerOpsAfterUnrollCppPass::<T>::default());
271
272        func_passes.add_pass(SCCPPass);
273        func_passes.add_pass(InstCombinePass::default());
274        func_passes.add_pass(SimpleCSEPass::without_memory());
275        func_passes.add_pass(SimplifyOpsPass::default());
276        func_passes.add_pass(DCEPass);
277        func_passes.add_pass(SROAPass);
278
279        // SCCP/DCE may unlock more mem2reg opportunities, and vice versa. So we do a sandwich.
280        func_passes.add_pass(Mem2RegPass);
281
282        func_passes.add_pass(SROAPass);
283        func_passes.add_pass(SCCPPass);
284        func_passes.add_pass(SimpleCSEPass::with_memory());
285        func_passes.add_pass(SimplifyOpsPass::default());
286        func_passes.add_pass(DCEPass);
287
288        func_passes.add_pass(PromoteBitwisePass);
289        func_passes.add_pass(PromoteUnsupportedTypesPass::default());
290
291        passes.add_pass(NestedOpsPass::new(func_passes));
292        passes.add_pass(AnnotateGlobalVisibilityPass);
293        passes.add_pass(DeclareVectorTypesPass);
294        passes.add_pass(CollectIncludesPass::<T>::default());
295
296        passes.run(module_op, &mut ctx, &mut analyses)?;
297
298        #[cfg(feature = "metal")]
299        if T::target() == Target::Metal {
300            crate::metal::builtin::append_msl_builtins(&mut ctx, entry_func);
301        }
302
303        verify_operation(module.get_operation(), &ctx)?;
304
305        let shared_memory_size = shared_memory_size(&ctx, module_op);
306        let buffers = buffers(&ctx, entry_func);
307        let io = buffer_io(&ctx, entry_func);
308
309        // Emit here rather than lazily from `Display`, so an op that survives lowering with no
310        // `OpToCPP` impl fails the compilation instead of panicking on the compiler thread.
311        ctx.set_aux_ty(EmissionErrors::default());
312        let source = module.get_operation().to_cpp(&ctx);
313        let mut errors = ctx.aux_ty::<EmissionErrors>().take();
314        let source = match source {
315            Ok(source) => source,
316            Err(error) => {
317                errors.push(error);
318                String::new()
319            }
320        };
321        if !errors.is_empty() {
322            let mut reason = "Can't emit cpp kernel\nCaused by:\n".to_string();
323            for error in errors {
324                reason += "  ";
325                reason += &error.to_string();
326                reason += "\n";
327            }
328            return Err(CompilationError::Validation {
329                reason,
330                backtrace: BackTrace::capture(),
331            });
332        }
333
334        let compute_kernel = ComputeKernel {
335            shared_memory_size,
336            buffers,
337            io,
338            source,
339        };
340
341        #[cfg(feature = "pliron-dump")]
342        dump_cpp(&compute_kernel, dump_dir);
343
344        Ok(compute_kernel)
345    }
346}
347
348#[cfg(feature = "pliron-dump")]
349fn dump_cpp(kernel: &ComputeKernel, dir: Option<std::path::PathBuf>) {
350    let Some(dir) = dir else {
351        return;
352    };
353
354    let source = kernel.to_string();
355    let source = crate::formatter::format_cpp(&source).unwrap_or(source);
356    std::fs::write(dir.join("module.cpp"), source).unwrap();
357}
358
359pub fn register_supported_types(props: &mut DeviceProperties) {
360    props.register_address_type(AddressType::U32);
361    props.register_address_type(AddressType::U64);
362
363    let supported_types = [
364        ElemType::Index,
365        ElemType::UInt(UIntKind::U8),
366        ElemType::UInt(UIntKind::U16),
367        ElemType::UInt(UIntKind::U32),
368        ElemType::UInt(UIntKind::U64),
369        ElemType::Int(IntKind::I8),
370        ElemType::Int(IntKind::I16),
371        ElemType::Int(IntKind::I32),
372        ElemType::Int(IntKind::I64),
373        ElemType::Float(FloatKind::BF16),
374        ElemType::Float(FloatKind::F16),
375        ElemType::Float(FloatKind::F32),
376        ElemType::Float(FloatKind::Flex32),
377        ElemType::Float(FloatKind::F64),
378        ElemType::Bool,
379    ];
380
381    let supported_atomic_types = [
382        ElemType::Int(IntKind::I32),
383        ElemType::Int(IntKind::I64),
384        ElemType::UInt(UIntKind::U32),
385        ElemType::UInt(UIntKind::U64),
386        ElemType::Float(FloatKind::F32),
387    ];
388
389    for ty in supported_types {
390        props.register_type_usage(ty, TypeUsage::all());
391    }
392
393    for ty in [FloatKind::E4M3, FloatKind::E5M2] {
394        props.register_type_usage(
395            ElemType::Float(ty),
396            TypeUsage::Conversion | TypeUsage::Buffer,
397        );
398    }
399
400    for ty in supported_atomic_types {
401        // Restricted to 32-bit integers because not every min/max/bitwise/CAS overload
402        // exists for 64-bit and float atomics across the C++ dialects (CUDA, HIP, Metal).
403        let usage = match ty {
404            ElemType::Int(IntKind::I32) | ElemType::UInt(UIntKind::U32) => AtomicUsage::all(),
405            _ => AtomicUsage::Add | AtomicUsage::LoadStore | AtomicUsage::Exchange,
406        };
407        props.register_atomic_type_usage(Type::atomic(ty), usage);
408    }
409}
410
411#[cfg(feature = "pliron-dump")]
412pub fn kernel_dir_name(name: &str) -> Option<std::path::PathBuf> {
413    if let Ok(dir) = std::env::var("CUBECL_DEBUG_PLIRON") {
414        let path = sanitize_filename::sanitize_with_options(
415            name,
416            sanitize_filename::Options {
417                replacement: "_",
418                ..Default::default()
419            },
420        );
421        let dir = std::path::PathBuf::from(dir).join(&path);
422        std::fs::create_dir_all(&dir).unwrap();
423        Some(dir)
424    } else {
425        None
426    }
427}