cubecl-cpu 0.9.0

CPU runtime for CubeCL
use cubecl_core::{ir::StorageType, prelude::KernelDefinition};
use cubecl_opt::Optimizer;
use tracel_llvm::mlir_rs::{
    Context, ExecutionEngine,
    ir::{Location, operation::OperationLike},
    pass::{self, PassManager},
};

use super::{passes::shared_memories::SharedMemories, visitor::Visitor};

pub(super) struct Module<'a> {
    module: tracel_llvm::mlir_rs::ir::Module<'a>,
    #[allow(unused)]
    name: String,
    location: Location<'a>,
    context: &'a Context,
}

impl<'a> Module<'a> {
    pub(super) fn new(context: &'a Context, name: String) -> Self {
        let location = Location::unknown(context);
        let module = tracel_llvm::mlir_rs::ir::Module::new(location);
        Self {
            module,
            context,
            name,
            location,
        }
    }

    pub(super) fn visit_kernel(
        &mut self,
        kernel: &KernelDefinition,
        opt: &Optimizer,
        shared_memories: &SharedMemories,
        addr_type: StorageType,
    ) {
        Visitor::visit_kernel(
            self.context,
            self.location,
            kernel,
            &self.module,
            opt,
            shared_memories,
            addr_type,
        )
    }

    pub(super) fn run_pass(&mut self) {
        let pass_manager = PassManager::new(self.context);
        pass_manager.enable_verifier(true);
        #[cfg(feature = "mlir-dump")]
        if let Ok(dir) = std::env::var("CUBECL_DEBUG_MLIR") {
            use std::path::PathBuf;
            use tracel_llvm::mlir_rs::{
                ir::operation::OperationPrintingFlags, pass::PassIrPrintingOptions,
            };

            let dir = dir.to_string() + "/" + &self.name;
            pass_manager.enable_ir_printing(&PassIrPrintingOptions {
                before_all: true,
                after_all: true,
                module_scope: true,
                on_change: true,
                on_failure: true,
                flags: OperationPrintingFlags::new(),
                tree_printing_path: PathBuf::from(dir),
            });
        }
        pass_manager.add_pass(pass::transform::create_canonicalizer());
        pass_manager.add_pass(pass::conversion::create_finalize_mem_ref_to_llvm());
        pass_manager.add_pass(pass::conversion::create_index_to_llvm());
        pass_manager.add_pass(pass::conversion::create_scf_to_control_flow());
        pass_manager.add_pass(pass::conversion::create_control_flow_to_llvm());
        pass_manager.add_pass(pass::conversion::create_math_to_llvm());
        pass_manager.add_pass(pass::conversion::create_math_to_libm());
        pass_manager.add_pass(pass::conversion::create_vector_to_llvm());
        pass_manager.add_pass(pass::conversion::create_arith_to_llvm());
        pass_manager.add_pass(pass::conversion::create_func_to_llvm());
        pass_manager.add_pass(pass::transform::create_inliner());
        pass_manager.add_pass(pass::conversion::create_reconcile_unrealized_casts());
        pass_manager.add_pass(pass::transform::create_sccp());
        pass_manager.add_pass(pass::transform::create_mem_2_reg());
        // pass_manager.add_pass(pass::transform::create_remove_dead_values()); // Needs this to be fixed before https://github.com/llvm/llvm-project/issues/82788
        pass_manager.add_pass(pass::transform::create_control_flow_sink());
        pass_manager.add_pass(pass::transform::create_cse());
        if let Err(err) = pass_manager.run(&mut self.module) {
            panic!("{}", err);
        }
        self.module.as_operation().verify();
    }

    pub(super) fn into_execution_engine(self) -> ExecutionEngine {
        ExecutionEngine::new(&self.module, 0, &[], true)
    }
}