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};
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/// Wgsl Compiler.
46#[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        // After the unroll so that an fp8 vector is at most one word, see `types.rs`.
135        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        // SCCP/DCE may unlock more mem2reg opportunities, and vice versa. So we do a sandwich.
154        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}