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