cubecl_wgpu/compiler/wgsl/
compiler.rs1use super::shader::ComputeShader;
2use crate::compiler::wgsl::{
3 self, EnableFeaturesPass, builtin::LowerBuiltinsPass, lower::LowerOpsWgslPass,
4 metadata::declare_info, rewrite_args, shared_memory_size, types,
5};
6
7use cubecl_core::{
8 WgpuCompilationOptions,
9 post_processing::{
10 checked_io::{CheckedIo, CheckedIoPass},
11 minifloat::{
12 Fp8Container, LowerMinifloatCast, LowerMinifloatCastPass, LowerMinifloatCompare,
13 LowerMinifloatComparePass,
14 },
15 saturating::LowerSaturatingArithmeticPass,
16 unroll::UnrollPass,
17 },
18};
19use cubecl_environment::backtrace::BackTrace;
20use cubecl_ir::{
21 ContextExt,
22 features::EnumSet,
23 pliron::{
24 builtin::ops::{FuncOp, ModuleOp},
25 operation::verify_operation,
26 opts::{constants::sccp::SCCPPass, dce::DCEPass, mem2reg::Mem2RegPass},
27 },
28 prelude::{AnalysisManager, NestedOpsPass, Op, OpPass, PMConfig, Pass, Passes},
29 rewrite::SimplifyOpsPass,
30 settings::Dim3,
31};
32use cubecl_opt::passes::{
33 annotate_buffer_visibility::AnnotateGlobalVisibilityPass, simple_cse::SimpleCSEPass,
34 sroa::SROAPass,
35};
36use cubecl_runtime::compiler::CompilationError;
37use cubecl_runtime::kernel;
38
39const MAX_VECTOR_SIZE: usize = 4;
40
41pub struct KernelInfo {
42 pub cube_dim: Dim3,
43}
44
45#[derive(Clone, Default)]
47pub struct WgslCompiler;
48
49impl core::fmt::Debug for WgslCompiler {
50 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
51 f.write_str("WgslCompiler")
52 }
53}
54
55impl cubecl_core::Compiler for WgslCompiler {
56 type Representation = ComputeShader;
57 type CompilationOptions = WgpuCompilationOptions;
58
59 fn compile(
60 &mut self,
61 shader: kernel::KernelDefinition,
62 compilation_options: &Self::CompilationOptions,
63 ) -> Result<Self::Representation, CompilationError> {
64 self.compile_shader(shader, compilation_options)
65 }
66
67 fn extension(&self) -> &'static str {
68 "wgsl"
69 }
70}
71
72impl WgslCompiler {
73 fn compile_shader(
74 &mut self,
75 value: kernel::KernelDefinition,
76 compilation_options: &WgpuCompilationOptions,
77 ) -> Result<wgsl::ComputeShader, CompilationError> {
78 let errors = value.body.pop_errors();
79 if !errors.is_empty() {
80 let mut reason = "Can't compile wgsl kernel".to_string();
81 for error in errors {
82 reason += error.as_str();
83 reason += "\n";
84 }
85
86 return Err(CompilationError::Validation {
87 reason,
88 backtrace: BackTrace::capture(),
89 });
90 }
91
92 #[cfg(feature = "pliron-dump")]
93 let ir_printing_dir = kernel_dir_name(&value.settings.kernel_name);
94
95 let module = value.body.state().module;
96 let entry_func = value.body.state().entry_func;
97 let module_op = module.get_operation();
98 let mut ctx = value.body.into_context().expect("Should be unique");
99 ctx.set_aux_ty(value.info);
100 ctx.set_aux_ty(KernelInfo {
101 cube_dim: value.settings.cube_dim,
102 });
103 ctx.set_aux_ty(*compilation_options);
104
105 #[cfg(feature = "pliron-dump")]
106 if let Some(print_dir) = &ir_printing_dir {
107 use pliron::printable::Printable;
108 let str = std::format!("{}", module_op.disp(&ctx));
109 std::fs::write(print_dir.join("initial.plir"), &str).unwrap();
110 }
111
112 verify_operation(module_op, &ctx)?;
113 types::check_fp8_lanes(&ctx, module_op)?;
114
115 let config = PMConfig {
116 #[cfg(feature = "pliron-dump")]
117 ir_printing_dir,
118 print_after_all: cfg!(feature = "pliron-dump"),
119 ..Default::default()
120 };
121
122 let mut analyses = AnalysisManager::default();
123 analyses.set_config(config);
124
125 let mut passes = OpPass::<ModuleOp, Passes>::default();
126 let mut func_passes = OpPass::<FuncOp, Passes>::default();
127
128 func_passes.add_pass(SROAPass);
129 func_passes.add_pass(CheckedIoPass::new(CheckedIo::new(
130 value.settings.execution_mode,
131 value.settings.kernel_name.clone(),
132 )));
133 func_passes.add_pass(UnrollPass::new(MAX_VECTOR_SIZE));
134 func_passes.add_pass(LowerMinifloatCastPass::new(LowerMinifloatCast::new(
136 EnumSet::empty(),
137 Fp8Container::Words,
138 )));
139 func_passes.add_pass(LowerMinifloatComparePass::new(LowerMinifloatCompare::new(
140 Fp8Container::Words,
141 )));
142
143 func_passes.add_pass(LowerOpsWgslPass::default());
144 func_passes.add_pass(LowerSaturatingArithmeticPass::default());
145 func_passes.add_pass(LowerBuiltinsPass);
146
147 func_passes.add_pass(SCCPPass);
148 func_passes.add_pass(SimpleCSEPass);
149 func_passes.add_pass(SimplifyOpsPass::default());
150 func_passes.add_pass(DCEPass);
151 func_passes.add_pass(SROAPass);
152
153 func_passes.add_pass(Mem2RegPass);
155
156 func_passes.add_pass(SROAPass);
157 func_passes.add_pass(SCCPPass);
158 func_passes.add_pass(SimpleCSEPass);
159 func_passes.add_pass(SimplifyOpsPass::default());
160 func_passes.add_pass(DCEPass);
161
162 passes.add_pass(NestedOpsPass::new(func_passes));
163 passes.add_pass(AnnotateGlobalVisibilityPass);
164 passes.add_pass(EnableFeaturesPass);
165
166 passes.run(module_op, &mut ctx, &mut analyses)?;
167
168 let buffers = rewrite_args(&mut ctx, entry_func);
169 declare_info(&mut ctx, entry_func, buffers.len());
170 let shared_memory_size = shared_memory_size(&ctx, module_op);
171
172 verify_operation(module.get_operation(), &ctx)?;
173
174 Ok(ComputeShader {
175 buffers,
176 shared_memory_size,
177 ctx,
178 })
179 }
180}
181
182#[cfg(feature = "pliron-dump")]
183pub fn kernel_dir_name(name: &str) -> Option<std::path::PathBuf> {
184 if let Ok(dir) = std::env::var("CUBECL_DEBUG_PLIRON") {
185 let path = sanitize_filename::sanitize_with_options(
186 name,
187 sanitize_filename::Options {
188 replacement: "_",
189 ..Default::default()
190 },
191 );
192 let dir = std::path::PathBuf::from(dir).join(&path);
193 std::fs::create_dir_all(&dir).unwrap();
194 Some(dir)
195 } else {
196 None
197 }
198}