112 lines · cpp
1//===- TestIRDLToCppDialect.cpp - MLIR Test Dialect Types ---------------*-===//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 file includes TestIRDLToCpp dialect.10//11//===----------------------------------------------------------------------===//12 13// #include "mlir/IR/Dialect.h"14#include "mlir/IR/Region.h"15 16#include "mlir/Dialect/SCF/IR/SCF.h"17#include "mlir/IR/BuiltinTypes.h"18#include "mlir/IR/DialectImplementation.h"19#include "mlir/Interfaces/InferTypeOpInterface.h"20#include "mlir/Pass/Pass.h"21#include "mlir/Target/LLVMIR/Dialect/Builtin/BuiltinToLLVMIRTranslation.h"22#include "mlir/Target/LLVMIR/Dialect/LLVMIR/LLVMToLLVMIRTranslation.h"23#include "mlir/Target/LLVMIR/LLVMTranslationInterface.h"24#include "mlir/Target/LLVMIR/ModuleTranslation.h"25#include "mlir/Tools/mlir-translate/Translation.h"26#include "mlir/Transforms/DialectConversion.h"27#include "mlir/Transforms/GreedyPatternRewriteDriver.h"28#include "llvm/ADT/DenseSet.h"29#include "llvm/ADT/TypeSwitch.h"30 31#include "TestIRDLToCppDialect.h"32 33#define GEN_DIALECT_DEF34#include "test_irdl_to_cpp.irdl.mlir.cpp.inc"35 36namespace test {37using namespace mlir;38struct TestOpConversion : public OpConversionPattern<test_irdl_to_cpp::BeefOp> {39 using OpConversionPattern::OpConversionPattern;40 41 LogicalResult42 matchAndRewrite(mlir::test_irdl_to_cpp::BeefOp op, OpAdaptor adaptor,43 ConversionPatternRewriter &rewriter) const override {44 assert(adaptor.getStructuredOperands(0).size() == 1);45 assert(adaptor.getStructuredOperands(1).size() == 1);46 47 auto bar = rewriter.replaceOpWithNewOp<test_irdl_to_cpp::BarOp>(48 op, op->getResultTypes().front());49 rewriter.setInsertionPointAfter(bar);50 51 test_irdl_to_cpp::HashOp::create(rewriter, bar.getLoc(),52 rewriter.getIntegerType(32),53 adaptor.getLhs(), adaptor.getRhs());54 return success();55 }56};57 58struct TestRegionConversion59 : public OpConversionPattern<test_irdl_to_cpp::ConditionalOp> {60 using OpConversionPattern::OpConversionPattern;61 62 LogicalResult63 matchAndRewrite(mlir::test_irdl_to_cpp::ConditionalOp op, OpAdaptor adaptor,64 ConversionPatternRewriter &rewriter) const override {65 // Just exercising the C++ API even though these are not enforced in the66 // dialect definition67 assert(op.getThen().getBlocks().size() == 1);68 assert(adaptor.getElse().getBlocks().size() == 1);69 auto ifOp = scf::IfOp::create(rewriter, op.getLoc(), op.getInput());70 rewriter.replaceOp(op, ifOp);71 return success();72 }73};74 75struct ConvertTestDialectToSomethingPass76 : PassWrapper<ConvertTestDialectToSomethingPass, OperationPass<ModuleOp>> {77 void runOnOperation() override {78 MLIRContext *ctx = &getContext();79 RewritePatternSet patterns(ctx);80 patterns.add<TestOpConversion, TestRegionConversion>(ctx);81 ConversionTarget target(getContext());82 target.addIllegalOp<test_irdl_to_cpp::BeefOp,83 test_irdl_to_cpp::ConditionalOp>();84 target.addLegalOp<test_irdl_to_cpp::BarOp, test_irdl_to_cpp::HashOp,85 scf::IfOp, scf::YieldOp>();86 if (failed(applyPartialConversion(getOperation(), target,87 std::move(patterns))))88 signalPassFailure();89 }90 91 StringRef getArgument() const final { return "test-irdl-conversion-check"; }92 StringRef getDescription() const final {93 return "Checks the convertability of an irdl dialect";94 }95 96 void getDependentDialects(DialectRegistry ®istry) const override {97 registry.insert<scf::SCFDialect>();98 }99};100 101void registerIrdlTestDialect(mlir::DialectRegistry ®istry) {102 registry.insert<mlir::test_irdl_to_cpp::TestIrdlToCppDialect>();103}104 105} // namespace test106 107namespace mlir::test {108void registerTestIrdlTestDialectConversionPass() {109 PassRegistration<::test::ConvertTestDialectToSomethingPass>();110}111} // namespace mlir::test112