brintos

brintos / llvm-project-archived public Read only

0
0
Text · 3.3 KiB · 7a99fe8 Raw
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