123 lines · cpp
1//===- ExecutionEngine.cpp - C API for MLIR JIT ---------------------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir-c/ExecutionEngine.h"10#include "mlir/CAPI/ExecutionEngine.h"11#include "mlir/CAPI/IR.h"12#include "mlir/CAPI/Support.h"13#include "mlir/ExecutionEngine/OptUtils.h"14#include "mlir/Target/LLVMIR/Dialect/Builtin/BuiltinToLLVMIRTranslation.h"15#include "mlir/Target/LLVMIR/Dialect/LLVMIR/LLVMToLLVMIRTranslation.h"16#include "mlir/Target/LLVMIR/Dialect/OpenMP/OpenMPToLLVMIRTranslation.h"17#include "llvm/ExecutionEngine/Orc/Mangling.h"18#include "llvm/Support/TargetSelect.h"19 20using namespace mlir;21 22extern "C" MlirExecutionEngine23mlirExecutionEngineCreate(MlirModule op, int optLevel, int numPaths,24 const MlirStringRef *sharedLibPaths,25 bool enableObjectDump) {26 static bool initOnce = [] {27 llvm::InitializeNativeTarget();28 llvm::InitializeNativeTargetAsmParser(); // needed for inline_asm29 llvm::InitializeNativeTargetAsmPrinter();30 return true;31 }();32 (void)initOnce;33 34 auto &ctx = *unwrap(op)->getContext();35 mlir::registerBuiltinDialectTranslation(ctx);36 mlir::registerLLVMDialectTranslation(ctx);37 mlir::registerOpenMPDialectTranslation(ctx);38 39 auto tmBuilderOrError = llvm::orc::JITTargetMachineBuilder::detectHost();40 if (!tmBuilderOrError) {41 llvm::errs() << "Failed to create a JITTargetMachineBuilder for the host\n";42 return MlirExecutionEngine{nullptr};43 }44 auto tmOrError = tmBuilderOrError->createTargetMachine();45 if (!tmOrError) {46 llvm::errs() << "Failed to create a TargetMachine for the host\n";47 return MlirExecutionEngine{nullptr};48 }49 50 SmallVector<StringRef> libPaths;51 for (unsigned i = 0; i < static_cast<unsigned>(numPaths); ++i)52 libPaths.push_back(sharedLibPaths[i].data);53 54 // Create a transformer to run all LLVM optimization passes at the55 // specified optimization level.56 auto transformer = mlir::makeOptimizingTransformer(57 optLevel, /*sizeLevel=*/0, /*targetMachine=*/tmOrError->get());58 ExecutionEngineOptions jitOptions;59 jitOptions.transformer = transformer;60 jitOptions.jitCodeGenOptLevel = static_cast<llvm::CodeGenOptLevel>(optLevel);61 jitOptions.sharedLibPaths = libPaths;62 jitOptions.enableObjectDump = enableObjectDump;63 auto jitOrError = ExecutionEngine::create(unwrap(op), jitOptions);64 if (!jitOrError) {65 consumeError(jitOrError.takeError());66 return MlirExecutionEngine{nullptr};67 }68 return wrap(jitOrError->release());69}70 71extern "C" void mlirExecutionEngineInitialize(MlirExecutionEngine jit) {72 unwrap(jit)->initialize();73}74 75extern "C" void mlirExecutionEngineDestroy(MlirExecutionEngine jit) {76 delete (unwrap(jit));77}78 79extern "C" MlirLogicalResult80mlirExecutionEngineInvokePacked(MlirExecutionEngine jit, MlirStringRef name,81 void **arguments) {82 const std::string ifaceName = ("_mlir_ciface_" + unwrap(name)).str();83 llvm::Error error = unwrap(jit)->invokePacked(84 ifaceName, MutableArrayRef<void *>{arguments, (size_t)0});85 if (error)86 return wrap(failure());87 return wrap(success());88}89 90extern "C" void *mlirExecutionEngineLookupPacked(MlirExecutionEngine jit,91 MlirStringRef name) {92 auto optionalFPtr =93 llvm::expectedToOptional(unwrap(jit)->lookupPacked(unwrap(name)));94 if (!optionalFPtr)95 return nullptr;96 return reinterpret_cast<void *>(*optionalFPtr);97}98 99extern "C" void *mlirExecutionEngineLookup(MlirExecutionEngine jit,100 MlirStringRef name) {101 auto optionalFPtr =102 llvm::expectedToOptional(unwrap(jit)->lookup(unwrap(name)));103 if (!optionalFPtr)104 return nullptr;105 return *optionalFPtr;106}107 108extern "C" void mlirExecutionEngineRegisterSymbol(MlirExecutionEngine jit,109 MlirStringRef name,110 void *sym) {111 unwrap(jit)->registerSymbols([&](llvm::orc::MangleAndInterner interner) {112 llvm::orc::SymbolMap symbolMap;113 symbolMap[interner(unwrap(name))] = {llvm::orc::ExecutorAddr::fromPtr(sym),114 llvm::JITSymbolFlags::Exported};115 return symbolMap;116 });117}118 119extern "C" void mlirExecutionEngineDumpToObjectFile(MlirExecutionEngine jit,120 MlirStringRef name) {121 unwrap(jit)->dumpToObjectFile(unwrap(name));122}123