Skip to main content

cubecl_wgpu/compiler/wgsl/
compiler.rs

1use 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/// Wgsl Compiler.
47#[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        // After the unroll so that an fp8 vector is at most one word, see `types.rs`.
144        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        // SCCP/DCE may unlock more mem2reg opportunities, and vice versa. So we do a sandwich.
164        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        // Before `rewrite_args` erases the arguments and widens their
179        // visibility for the shader.
180        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}