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 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 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 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 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 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 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 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}