brintos

brintos / llvm-project-archived public Read only

0
0
Text · 8.9 KiB · 72bec11 Raw
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 &registry) const override {177    MlirDialectRegistry cRegistry = wrap(&registry);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