237 lines · cpp
1//===- Pass.cpp - C Interface for General Pass Management APIs ------------===//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/Pass.h"10 11#include "mlir/CAPI/IR.h"12#include "mlir/CAPI/Pass.h"13#include "mlir/CAPI/Support.h"14#include "mlir/CAPI/Utils.h"15#include "mlir/Pass/PassManager.h"16#include "llvm/Support/ErrorHandling.h"17#include <optional>18 19using namespace mlir;20 21//===----------------------------------------------------------------------===//22// PassManager/OpPassManager APIs.23//===----------------------------------------------------------------------===//24 25MlirPassManager mlirPassManagerCreate(MlirContext ctx) {26 return wrap(new PassManager(unwrap(ctx)));27}28 29MlirPassManager mlirPassManagerCreateOnOperation(MlirContext ctx,30 MlirStringRef anchorOp) {31 return wrap(new PassManager(unwrap(ctx), unwrap(anchorOp)));32}33 34void mlirPassManagerDestroy(MlirPassManager passManager) {35 delete unwrap(passManager);36}37 38MlirOpPassManager39mlirPassManagerGetAsOpPassManager(MlirPassManager passManager) {40 return wrap(static_cast<OpPassManager *>(unwrap(passManager)));41}42 43MlirLogicalResult mlirPassManagerRunOnOp(MlirPassManager passManager,44 MlirOperation op) {45 return wrap(unwrap(passManager)->run(unwrap(op)));46}47 48void mlirPassManagerEnableIRPrinting(MlirPassManager passManager,49 bool printBeforeAll, bool printAfterAll,50 bool printModuleScope,51 bool printAfterOnlyOnChange,52 bool printAfterOnlyOnFailure,53 MlirOpPrintingFlags flags,54 MlirStringRef treePrintingPath) {55 auto shouldPrintBeforePass = [printBeforeAll](Pass *, Operation *) {56 return printBeforeAll;57 };58 auto shouldPrintAfterPass = [printAfterAll](Pass *, Operation *) {59 return printAfterAll;60 };61 if (unwrap(treePrintingPath).empty())62 return unwrap(passManager)63 ->enableIRPrinting(shouldPrintBeforePass, shouldPrintAfterPass,64 printModuleScope, printAfterOnlyOnChange,65 printAfterOnlyOnFailure, /*out=*/llvm::errs(),66 *unwrap(flags));67 68 unwrap(passManager)69 ->enableIRPrintingToFileTree(shouldPrintBeforePass, shouldPrintAfterPass,70 printModuleScope, printAfterOnlyOnChange,71 printAfterOnlyOnFailure,72 unwrap(treePrintingPath), *unwrap(flags));73}74 75void mlirPassManagerEnableVerifier(MlirPassManager passManager, bool enable) {76 unwrap(passManager)->enableVerifier(enable);77}78 79void mlirPassManagerEnableTiming(MlirPassManager passManager) {80 unwrap(passManager)->enableTiming();81}82 83void mlirPassManagerEnableStatistics(MlirPassManager passManager,84 MlirPassDisplayMode displayMode) {85 PassDisplayMode mode;86 switch (displayMode) {87 case MLIR_PASS_DISPLAY_MODE_LIST:88 mode = PassDisplayMode::List;89 break;90 case MLIR_PASS_DISPLAY_MODE_PIPELINE:91 mode = PassDisplayMode::Pipeline;92 break;93 }94 unwrap(passManager)->enableStatistics(mode);95}96 97MlirOpPassManager mlirPassManagerGetNestedUnder(MlirPassManager passManager,98 MlirStringRef operationName) {99 return wrap(&unwrap(passManager)->nest(unwrap(operationName)));100}101 102MlirOpPassManager mlirOpPassManagerGetNestedUnder(MlirOpPassManager passManager,103 MlirStringRef operationName) {104 return wrap(&unwrap(passManager)->nest(unwrap(operationName)));105}106 107void mlirPassManagerAddOwnedPass(MlirPassManager passManager, MlirPass pass) {108 unwrap(passManager)->addPass(std::unique_ptr<Pass>(unwrap(pass)));109}110 111void mlirOpPassManagerAddOwnedPass(MlirOpPassManager passManager,112 MlirPass pass) {113 unwrap(passManager)->addPass(std::unique_ptr<Pass>(unwrap(pass)));114}115 116MlirLogicalResult mlirOpPassManagerAddPipeline(MlirOpPassManager passManager,117 MlirStringRef pipelineElements,118 MlirStringCallback callback,119 void *userData) {120 detail::CallbackOstream stream(callback, userData);121 return wrap(parsePassPipeline(unwrap(pipelineElements), *unwrap(passManager),122 stream));123}124 125void mlirPrintPassPipeline(MlirOpPassManager passManager,126 MlirStringCallback callback, void *userData) {127 detail::CallbackOstream stream(callback, userData);128 unwrap(passManager)->printAsTextualPipeline(stream);129}130 131MlirLogicalResult mlirParsePassPipeline(MlirOpPassManager passManager,132 MlirStringRef pipeline,133 MlirStringCallback callback,134 void *userData) {135 detail::CallbackOstream stream(callback, userData);136 FailureOr<OpPassManager> pm = parsePassPipeline(unwrap(pipeline), stream);137 if (succeeded(pm))138 *unwrap(passManager) = std::move(*pm);139 return wrap(pm);140}141 142//===----------------------------------------------------------------------===//143// External Pass API.144//===----------------------------------------------------------------------===//145 146namespace mlir {147class ExternalPass;148} // namespace mlir149DEFINE_C_API_PTR_METHODS(MlirExternalPass, mlir::ExternalPass)150 151namespace mlir {152/// This pass class wraps external passes defined in other languages using the153/// MLIR C-interface154class ExternalPass : public Pass {155public:156 ExternalPass(TypeID passID, StringRef name, StringRef argument,157 StringRef description, std::optional<StringRef> opName,158 ArrayRef<MlirDialectHandle> dependentDialects,159 MlirExternalPassCallbacks callbacks, void *userData)160 : Pass(passID, opName), id(passID), name(name), argument(argument),161 description(description), dependentDialects(dependentDialects),162 callbacks(callbacks), userData(userData) {163 if (callbacks.construct)164 callbacks.construct(userData);165 }166 167 ~ExternalPass() override {168 if (callbacks.destruct)169 callbacks.destruct(userData);170 }171 172 StringRef getName() const override { return name; }173 StringRef getArgument() const override { return argument; }174 StringRef getDescription() const override { return description; }175 176 void getDependentDialects(DialectRegistry ®istry) const override {177 MlirDialectRegistry cRegistry = wrap(®istry);178 for (MlirDialectHandle dialect : dependentDialects)179 mlirDialectHandleInsertDialect(dialect, cRegistry);180 }181 182 void signalPassFailure() { Pass::signalPassFailure(); }183 184protected:185 LogicalResult initialize(MLIRContext *ctx) override {186 if (callbacks.initialize)187 return unwrap(callbacks.initialize(wrap(ctx), userData));188 return success();189 }190 191 bool canScheduleOn(RegisteredOperationName opName) const override {192 if (std::optional<StringRef> specifiedOpName = getOpName())193 return opName.getStringRef() == specifiedOpName;194 return true;195 }196 197 void runOnOperation() override {198 callbacks.run(wrap(getOperation()), wrap(this), userData);199 }200 201 std::unique_ptr<Pass> clonePass() const override {202 void *clonedUserData = callbacks.clone(userData);203 return std::make_unique<ExternalPass>(id, name, argument, description,204 getOpName(), dependentDialects,205 callbacks, clonedUserData);206 }207 208private:209 TypeID id;210 std::string name;211 std::string argument;212 std::string description;213 std::vector<MlirDialectHandle> dependentDialects;214 MlirExternalPassCallbacks callbacks;215 void *userData;216};217} // namespace mlir218 219MlirPass mlirCreateExternalPass(MlirTypeID passID, MlirStringRef name,220 MlirStringRef argument,221 MlirStringRef description, MlirStringRef opName,222 intptr_t nDependentDialects,223 MlirDialectHandle *dependentDialects,224 MlirExternalPassCallbacks callbacks,225 void *userData) {226 return wrap(static_cast<mlir::Pass *>(new mlir::ExternalPass(227 unwrap(passID), unwrap(name), unwrap(argument), unwrap(description),228 opName.length > 0 ? std::optional<StringRef>(unwrap(opName))229 : std::nullopt,230 {dependentDialects, static_cast<size_t>(nDependentDialects)}, callbacks,231 userData)));232}233 234void mlirExternalPassSignalFailure(MlirExternalPass pass) {235 unwrap(pass)->signalPassFailure();236}237