summaryrefslogtreecommitdiff
path: root/mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp
diff options
context:
space:
mode:
authorStella Laurenzo <stellaraccident@gmail.com>2024-03-15 22:22:09 -0700
committerGitHub <noreply@github.com>2024-03-15 22:22:09 -0700
commitdbbdee2ea2156170062813fb3d7f2c023d65e02d (patch)
tree19ea0e6699cb51f78ad7261aff9dfbc413d6d529 /mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp
parent426bf0c915aca9e9d78b6192898b95a44d9afcf4 (diff)
[mlir] Make the ml_program dialect allow all of its operations to be inlined. (#85479)
Diffstat (limited to 'mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp')
-rw-r--r--mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp15
1 files changed, 14 insertions, 1 deletions
diff --git a/mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp b/mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp
index 1a8fe208d409..bda1032ed988 100644
--- a/mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp
+++ b/mlir/lib/Dialect/MLProgram/IR/MLProgramDialect.cpp
@@ -8,6 +8,7 @@
#include "mlir/Dialect/MLProgram/IR/MLProgram.h"
#include "mlir/IR/DialectImplementation.h"
+#include "mlir/Transforms/InliningUtils.h"
#include "llvm/ADT/TypeSwitch.h"
using namespace mlir;
@@ -24,6 +25,18 @@ using namespace mlir::ml_program;
#include "mlir/Dialect/MLProgram/IR/MLProgramTypes.cpp.inc"
namespace {
+
+struct MLProgramInlinerInterface : public DialectInlinerInterface {
+ using DialectInlinerInterface::DialectInlinerInterface;
+
+ bool isLegalToInline(Operation *, Region *, bool,
+ IRMapping &) const override {
+ // We have no specific opinion on whether ops defined in this dialect should
+ // be inlined.
+ return true;
+ }
+};
+
struct MLProgramOpAsmDialectInterface : public OpAsmDialectInterface {
using OpAsmDialectInterface::OpAsmDialectInterface;
@@ -53,5 +66,5 @@ void ml_program::MLProgramDialect::initialize() {
#include "mlir/Dialect/MLProgram/IR/MLProgramOps.cpp.inc"
>();
- addInterfaces<MLProgramOpAsmDialectInterface>();
+ addInterfaces<MLProgramInlinerInterface, MLProgramOpAsmDialectInterface>();
}