86 lines · cpp
1//===- Canonicalizer.cpp - Canonicalize MLIR operations -------------------===//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// This transformation pass converts operations into their canonical forms by10// folding constants, applying operation identity transformations etc.11//12//===----------------------------------------------------------------------===//13 14#include "mlir/Transforms/Passes.h"15 16#include "mlir/Pass/Pass.h"17#include "mlir/Transforms/GreedyPatternRewriteDriver.h"18 19namespace mlir {20#define GEN_PASS_DEF_CANONICALIZER21#include "mlir/Transforms/Passes.h.inc"22} // namespace mlir23 24using namespace mlir;25 26namespace {27/// Canonicalize operations in nested regions.28struct Canonicalizer : public impl::CanonicalizerBase<Canonicalizer> {29 Canonicalizer() = default;30 Canonicalizer(const GreedyRewriteConfig &config,31 ArrayRef<std::string> disabledPatterns,32 ArrayRef<std::string> enabledPatterns)33 : config(config) {34 this->topDownProcessingEnabled = config.getUseTopDownTraversal();35 this->regionSimplifyLevel = config.getRegionSimplificationLevel();36 this->maxIterations = config.getMaxIterations();37 this->maxNumRewrites = config.getMaxNumRewrites();38 this->disabledPatterns = disabledPatterns;39 this->enabledPatterns = enabledPatterns;40 }41 42 /// Initialize the canonicalizer by building the set of patterns used during43 /// execution.44 LogicalResult initialize(MLIRContext *context) override {45 // Set the config from possible pass options set in the meantime.46 config.setUseTopDownTraversal(topDownProcessingEnabled);47 config.setRegionSimplificationLevel(regionSimplifyLevel);48 config.setMaxIterations(maxIterations);49 config.setMaxNumRewrites(maxNumRewrites);50 51 RewritePatternSet owningPatterns(context);52 for (auto *dialect : context->getLoadedDialects())53 dialect->getCanonicalizationPatterns(owningPatterns);54 for (RegisteredOperationName op : context->getRegisteredOperations())55 op.getCanonicalizationPatterns(owningPatterns, context);56 57 patterns = std::make_shared<FrozenRewritePatternSet>(58 std::move(owningPatterns), disabledPatterns, enabledPatterns);59 return success();60 }61 void runOnOperation() override {62 LogicalResult converged =63 applyPatternsGreedily(getOperation(), *patterns, config);64 // Canonicalization is best-effort. Non-convergence is not a pass failure.65 if (testConvergence && failed(converged))66 signalPassFailure();67 }68 GreedyRewriteConfig config;69 std::shared_ptr<const FrozenRewritePatternSet> patterns;70};71} // namespace72 73/// Create a Canonicalizer pass.74std::unique_ptr<Pass> mlir::createCanonicalizerPass() {75 return std::make_unique<Canonicalizer>();76}77 78/// Creates an instance of the Canonicalizer pass with the specified config.79std::unique_ptr<Pass>80mlir::createCanonicalizerPass(const GreedyRewriteConfig &config,81 ArrayRef<std::string> disabledPatterns,82 ArrayRef<std::string> enabledPatterns) {83 return std::make_unique<Canonicalizer>(config, disabledPatterns,84 enabledPatterns);85}86