#include "bishengir/Dialect/HACC/Utils/Utils.h"
#include "bishengir/Dialect/Linalg/IR/LinalgCanonicalizations.h"
#include "bishengir/Dialect/MemRef/IR/MemRefImpl.h"
#include "bishengir/Transforms/Passes.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
#include "mlir/IR/BuiltinOps.h"
#include "mlir/Rewrite/FrozenRewritePatternSet.h"
#include "mlir/Support/LLVM.h"
#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
#define DEBUG_TYPE "bishengir-canonicalize-ext"
namespace mlir {
#define GEN_PASS_DEF_CANONICALIZER
#include "mlir/Transforms/Passes.h.inc"
}
namespace {
using namespace mlir;
struct FoldTransposeWithTranspose : OpRewritePattern<linalg::TransposeOp> {
using OpRewritePattern<linalg::TransposeOp>::OpRewritePattern;
LogicalResult matchAndRewrite(linalg::TransposeOp transposeOp,
PatternRewriter &rewriter) const override {
auto defTransposeOp =
transposeOp.getInput().getDefiningOp<linalg::TransposeOp>();
if (!defTransposeOp)
return failure();
ArrayRef<int64_t> defPerms = defTransposeOp.getPermutation();
ArrayRef<int64_t> perms = transposeOp.getPermutation();
SmallVector<int64_t> foldedPerms;
foldedPerms.reserve(perms.size());
for (int64_t perm : perms)
foldedPerms.push_back(defPerms[perm]);
rewriter.replaceOpWithNewOp<linalg::TransposeOp>(
transposeOp, defTransposeOp.getInput(), transposeOp.getInit(),
foldedPerms);
return success();
}
};
struct ExtendedCanonicalizer
: public mlir::impl::CanonicalizerBase<ExtendedCanonicalizer> {
using mlir::impl::CanonicalizerBase<ExtendedCanonicalizer>::CanonicalizerBase;
static constexpr ::llvm::StringLiteral getArgumentName() {
return ::llvm::StringLiteral("canonicalize-ext");
}
::llvm::StringRef getArgument() const final { return "canonicalize-ext"; }
::llvm::StringRef getDescription() const final {
return "Canonicalize operations";
}
static constexpr ::llvm::StringLiteral getPassName() {
return ::llvm::StringLiteral("ExtendedCanonicalizer");
}
::llvm::StringRef getName() const final { return "ExtendedCanonicalizer"; }
void runOnOperation() final {
auto *context = getOperation()->getContext();
RewritePatternSet patterns(context);
for (auto *dialect : context->getLoadedDialects())
dialect->getCanonicalizationPatterns(patterns);
for (RegisteredOperationName op : context->getRegisteredOperations())
op.getCanonicalizationPatterns(patterns, context);
mlir::memref::getExtendedCanonicalizationPatterns(patterns);
linalg::getExtendedCanonicalizationPatterns(patterns);
auto moduleOp = dyn_cast<ModuleOp>(getOperation());
if (moduleOp && !hacc::utils::isAscend950(moduleOp))
patterns.add<FoldTransposeWithTranspose>(context);
FrozenRewritePatternSet filteredPatterns{std::move(patterns),
disabledPatterns, enabledPatterns};
GreedyRewriteConfig config;
config.useTopDownTraversal = topDownProcessingEnabled;
config.enableRegionSimplification = enableRegionSimplification;
config.maxIterations = maxIterations;
config.maxNumRewrites = maxNumRewrites;
if (auto converged =
applyPatternsGreedily(getOperation(), filteredPatterns, config);
testConvergence && failed(converged))
signalPassFailure();
}
};
}
std::unique_ptr<mlir::Pass> bishengir::createExtendedCanonicalizerPass(
const CanonicalizerOptions &options) {
return std::make_unique<ExtendedCanonicalizer>(options);
}