Skip to content

Commit

Permalink
add whitelist pass option to TorchToTcp
Browse files Browse the repository at this point in the history
add private method

move addPatternIfOpInConvertTorchOpsSet to utility function

add back typeConverter

add inline function in Utils.h

update Misc patterns as well

update

elementwise
  • Loading branch information
sjain-stanford committed Nov 2, 2023
1 parent 1808051 commit 7a6d51d
Show file tree
Hide file tree
Showing 10 changed files with 163 additions and 124 deletions.
1 change: 1 addition & 0 deletions BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,7 @@ cc_library(
"@torch-mlir//:TorchMLIRConversionUtils",
"@torch-mlir//:TorchMLIRTorchBackendTypeConversion",
"@torch-mlir//:TorchMLIRTorchConversionDialect",
"@torch-mlir//:TorchMLIRTorchPasses",
],
)

Expand Down
7 changes: 6 additions & 1 deletion include/mlir-tcp/Conversion/Passes.td
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,12 @@ def ConvertTorchToTcp : Pass<"convert-torch-to-tcp", "func::FuncOp"> {
let description = [{
Convert Torch ops to Tcp ops.
}];
let constructor = "mlir::tcp::createConvertTorchToTcpPass()";
let constructor = "mlir::tcp::createConvertTorchToTcpPass(/*convertTorchOps=*/{})";
let options = [
ListOption<"convertTorchOps", "convert-torch-ops", "std::string",
"List of Torch operation names that should be converted to Tcp",
"llvm::cl::ZeroOrMore">
];
}

//===----------------------------------------------------------------------===//
Expand Down
3 changes: 2 additions & 1 deletion include/mlir-tcp/Conversion/TorchToTcp/TorchToTcp.h
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,8 @@ namespace mlir {

namespace tcp {

std::unique_ptr<OperationPass<func::FuncOp>> createConvertTorchToTcpPass();
std::unique_ptr<OperationPass<func::FuncOp>>
createConvertTorchToTcpPass(llvm::ArrayRef<std::string> convertTorchOps);

} // namespace tcp
} // namespace mlir
9 changes: 5 additions & 4 deletions lib/Conversion/TorchToTcp/DataMovement.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
#include "torch-mlir/Dialect/Torch/IR/TorchOps.h"
#include "torch-mlir/Dialect/Torch/Utils/Utils.h"

#include "llvm/ADT/StringSet.h"

using namespace mlir;
using namespace mlir::tcp;
using namespace mlir::torch;
Expand Down Expand Up @@ -60,8 +62,7 @@ class ConvertAtenCatOp : public OpConversionPattern<AtenCatOp> {

void torch_to_tcp::populateDataMovementPatternsAndLegality(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target) {
MLIRContext *context = patterns.getContext();
target.addIllegalOp<AtenCatOp>();
patterns.add<ConvertAtenCatOp>(typeConverter, context);
ConversionTarget &target, const llvm::StringSet<> &convertTorchOpsSet) {
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet<ConvertAtenCatOp, AtenCatOp>(
typeConverter, patterns, target, convertTorchOpsSet);
}
129 changes: 55 additions & 74 deletions lib/Conversion/TorchToTcp/Elementwise.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -650,78 +650,59 @@ class ConvertAtenToDtypeOp : public OpConversionPattern<AtenToDtypeOp> {

void torch_to_tcp::populateElementwisePatternsAndLegality(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target) {
MLIRContext *context = patterns.getContext();

target.addIllegalOp<AtenToDtypeOp>();
patterns.add<ConvertAtenToDtypeOp>(typeConverter, context);

target.addIllegalOp<AtenClampOp>();
patterns.add<ConvertAtenClampOp>(typeConverter, context);
target.addIllegalOp<AtenReluOp>();
patterns.add<ConvertAtenReluOp>(typeConverter, context);

target.addIllegalOp<AtenAddTensorOp>();
target.addIllegalOp<AtenSubTensorOp>();
target.addIllegalOp<AtenAddScalarOp>();
target.addIllegalOp<AtenSubScalarOp>();
patterns.add<ConvertAtenAddSubOp<AtenAddTensorOp, tcp::AddOp>>(typeConverter,
context);
patterns.add<ConvertAtenAddSubOp<AtenSubTensorOp, tcp::SubOp>>(typeConverter,
context);
patterns.add<ConvertAtenAddSubOp<AtenAddScalarOp, tcp::AddOp>>(typeConverter,
context);
patterns.add<ConvertAtenAddSubOp<AtenSubScalarOp, tcp::SubOp>>(typeConverter,
context);

target.addIllegalOp<AtenMulTensorOp>();
target.addIllegalOp<AtenMulScalarOp>();
patterns.add<ConvertAtenMulOp<AtenMulTensorOp>>(typeConverter, context);
patterns.add<ConvertAtenMulOp<AtenMulScalarOp>>(typeConverter, context);

target.addIllegalOp<AtenDivTensorOp>();
target.addIllegalOp<AtenDivScalarOp>();
patterns.add<ConvertAtenDivOp<AtenDivTensorOp>>(typeConverter, context);
patterns.add<ConvertAtenDivOp<AtenDivScalarOp>>(typeConverter, context);

target.addIllegalOp<AtenCeilOp>();
target.addIllegalOp<AtenFloorOp>();
target.addIllegalOp<AtenSigmoidOp>();
target.addIllegalOp<AtenTanhOp>();
target.addIllegalOp<AtenSinOp>();
target.addIllegalOp<AtenCosOp>();
target.addIllegalOp<AtenLogOp>();
target.addIllegalOp<AtenNegOp>();
target.addIllegalOp<AtenAtanOp>();
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenFloorOp, tcp::FloorOp>>(
typeConverter, context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenCeilOp, tcp::CeilOp>>(typeConverter,
context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenSigmoidOp, tcp::SigmoidOp>>(
typeConverter, context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenTanhOp, tcp::TanhOp>>(typeConverter,
context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenSinOp, tcp::SinOp>>(typeConverter,
context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenCosOp, tcp::CosOp>>(typeConverter,
context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenLogOp, tcp::LogOp>>(typeConverter,
context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenNegOp, tcp::NegOp>>(typeConverter,
context);
patterns.add<ConvertAtenUnaryFpOnlyOp<AtenAtanOp, tcp::AtanOp>>(typeConverter,
context);

target.addIllegalOp<AtenAbsOp>();
target.addIllegalOp<AtenSqrtOp>();
patterns.add<ConvertAtenUnaryIntOrFpOp<AtenAbsOp, tcp::AbsOp>>(typeConverter,
context);
patterns.add<ConvertAtenUnaryIntOrFpOp<AtenSqrtOp, tcp::SqrtOp>>(
typeConverter, context);

target.addIllegalOp<AtenBatchNormOp>();
patterns.add<ConvertAtenBatchNormOp>(typeConverter, context);

target.addIllegalOp<AtenAtan2Op>();
patterns.add<ConvertAtenAtan2Op>(typeConverter, context);
ConversionTarget &target, const llvm::StringSet<> &convertTorchOpsSet) {

#define INSERT_ATEN_ELEMENTWISE_OP_PATTERN(ConvertAtenOpPattern, AtenOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet<ConvertAtenOpPattern, \
AtenOp>( \
typeConverter, patterns, target, convertTorchOpsSet)
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(ConvertAtenToDtypeOp, AtenToDtypeOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(ConvertAtenClampOp, AtenClampOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(ConvertAtenReluOp, AtenReluOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(ConvertAtenBatchNormOp, AtenBatchNormOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(ConvertAtenAtan2Op, AtenAtan2Op);
#undef INSERT_ATEN_ELEMENTWISE_OP_PATTERN

#define INSERT_ATEN_ELEMENTWISE_ADD_SUB_PATTERN(AtenOp, TcpOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet< \
ConvertAtenAddSubOp<AtenOp, TcpOp>, AtenOp>(typeConverter, patterns, \
target, convertTorchOpsSet)
INSERT_ATEN_ELEMENTWISE_ADD_SUB_PATTERN(AtenAddTensorOp, tcp::AddOp);
INSERT_ATEN_ELEMENTWISE_ADD_SUB_PATTERN(AtenSubTensorOp, tcp::SubOp);
INSERT_ATEN_ELEMENTWISE_ADD_SUB_PATTERN(AtenAddScalarOp, tcp::AddOp);
INSERT_ATEN_ELEMENTWISE_ADD_SUB_PATTERN(AtenSubScalarOp, tcp::SubOp);
#undef INSERT_ATEN_ELEMENTWISE_ADD_SUB_PATTERN

#define INSERT_ATEN_ELEMENTWISE_MUL_DIV_PATTERN(ConvertAtenOpPattern, AtenOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet< \
ConvertAtenOpPattern<AtenOp>, AtenOp>(typeConverter, patterns, target, \
convertTorchOpsSet)
INSERT_ATEN_ELEMENTWISE_MUL_DIV_PATTERN(ConvertAtenMulOp, AtenMulTensorOp);
INSERT_ATEN_ELEMENTWISE_MUL_DIV_PATTERN(ConvertAtenMulOp, AtenMulScalarOp);
INSERT_ATEN_ELEMENTWISE_MUL_DIV_PATTERN(ConvertAtenDivOp, AtenDivTensorOp);
INSERT_ATEN_ELEMENTWISE_MUL_DIV_PATTERN(ConvertAtenDivOp, AtenDivScalarOp);
#undef INSERT_ATEN_ELEMENTWISE_MUL_DIV_PATTERN

#define INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenOp, TcpOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet< \
ConvertAtenUnaryFpOnlyOp<AtenOp, TcpOp>, AtenOp>( \
typeConverter, patterns, target, convertTorchOpsSet)
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenCeilOp, tcp::CeilOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenFloorOp, tcp::FloorOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenSigmoidOp, tcp::SigmoidOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenTanhOp, tcp::TanhOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenSinOp, tcp::SinOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenCosOp, tcp::CosOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenLogOp, tcp::LogOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenNegOp, tcp::NegOp);
INSERT_ATEN_UNARY_FP_ONLY_PATTERN(AtenAtanOp, tcp::AtanOp);
#undef INSERT_ATEN_UNARY_FP_ONLY_PATTERN

#define INSERT_ATEN_UNARY_INT_OR_FP_PATTERN(AtenOp, TcpOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet< \
ConvertAtenUnaryIntOrFpOp<AtenOp, TcpOp>, AtenOp>( \
typeConverter, patterns, target, convertTorchOpsSet)
INSERT_ATEN_UNARY_INT_OR_FP_PATTERN(AtenAbsOp, tcp::AbsOp);
INSERT_ATEN_UNARY_INT_OR_FP_PATTERN(AtenSqrtOp, tcp::SqrtOp);
#undef INSERT_ATEN_UNARY_INT_OR_FP_PATTERN
}
51 changes: 25 additions & 26 deletions lib/Conversion/TorchToTcp/Misc.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -134,7 +134,7 @@ class ConvertValueTensorLiteralOp
};

template <typename AtenOpT, int fillVal>
class ConvertAtenZerosOnesPatternOp : public OpConversionPattern<AtenOpT> {
class ConvertAtenZerosOnesOp : public OpConversionPattern<AtenOpT> {
public:
using OpConversionPattern<AtenOpT>::OpConversionPattern;
using OpAdaptor = typename AtenOpT::Adaptor;
Expand Down Expand Up @@ -180,7 +180,7 @@ class ConvertAtenZerosOnesPatternOp : public OpConversionPattern<AtenOpT> {
};

template <typename AtenOpT, int fillVal>
class ConvertAtenZerosOnesLikePatternOp : public OpConversionPattern<AtenOpT> {
class ConvertAtenZerosOnesLikeOp : public OpConversionPattern<AtenOpT> {
public:
using OpConversionPattern<AtenOpT>::OpConversionPattern;
using OpAdaptor = typename AtenOpT::Adaptor;
Expand Down Expand Up @@ -220,28 +220,27 @@ class ConvertAtenZerosOnesLikePatternOp : public OpConversionPattern<AtenOpT> {

} // namespace

void torch_to_tcp::populateMiscPatternsAndLegality(TypeConverter &typeConverter,
RewritePatternSet &patterns,
ConversionTarget &target) {
MLIRContext *context = patterns.getContext();

target.addIllegalOp<AtenBroadcastToOp>();
patterns.add<ConvertAtenBroadcastToOp>(typeConverter, context);

target.addIllegalOp<ValueTensorLiteralOp>();
patterns.add<ConvertValueTensorLiteralOp>(typeConverter, context);

target.addIllegalOp<AtenZerosOp>();
patterns.add<ConvertAtenZerosOnesPatternOp<AtenZerosOp, 0>>(typeConverter,
context);
target.addIllegalOp<AtenOnesOp>();
patterns.add<ConvertAtenZerosOnesPatternOp<AtenOnesOp, 1>>(typeConverter,
context);

target.addIllegalOp<AtenZerosLikeOp>();
patterns.add<ConvertAtenZerosOnesLikePatternOp<AtenZerosLikeOp, 0>>(
typeConverter, context);
target.addIllegalOp<AtenOnesLikeOp>();
patterns.add<ConvertAtenZerosOnesLikePatternOp<AtenOnesLikeOp, 1>>(
typeConverter, context);
void torch_to_tcp::populateMiscPatternsAndLegality(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target, const llvm::StringSet<> &convertTorchOpsSet) {

#define INSERT_ATEN_MISC_OP_PATTERN(ConvertAtenOpPattern, AtenOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet<ConvertAtenOpPattern, \
AtenOp>( \
typeConverter, patterns, target, convertTorchOpsSet)
INSERT_ATEN_MISC_OP_PATTERN(ConvertAtenBroadcastToOp, AtenBroadcastToOp);
INSERT_ATEN_MISC_OP_PATTERN(ConvertValueTensorLiteralOp,
ValueTensorLiteralOp);
#undef INSERT_ATEN_MISC_OP_PATTERN

#define INSERT_ATEN_ZEROS_ONES_PATTERN(ConvertAtenOpPattern, AtenOp, Val) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet< \
ConvertAtenOpPattern<AtenOp, Val>, AtenOp>(typeConverter, patterns, \
target, convertTorchOpsSet)
INSERT_ATEN_ZEROS_ONES_PATTERN(ConvertAtenZerosOnesOp, AtenZerosOp, 0);
INSERT_ATEN_ZEROS_ONES_PATTERN(ConvertAtenZerosOnesOp, AtenOnesOp, 1);
INSERT_ATEN_ZEROS_ONES_PATTERN(ConvertAtenZerosOnesLikeOp, AtenZerosLikeOp,
0);
INSERT_ATEN_ZEROS_ONES_PATTERN(ConvertAtenZerosOnesLikeOp, AtenOnesLikeOp, 1);
#undef INSERT_ATEN_ZEROS_ONES_PATTERN
}
21 changes: 12 additions & 9 deletions lib/Conversion/TorchToTcp/PopulatePatterns.h
Original file line number Diff line number Diff line change
Expand Up @@ -9,19 +9,22 @@

#include "mlir/Transforms/DialectConversion.h"

#include "llvm/ADT/StringSet.h"

namespace mlir {
namespace torch_to_tcp {

void populateElementwisePatternsAndLegality(TypeConverter &typeConverter,
RewritePatternSet &patterns,
ConversionTarget &target);
void populateMiscPatternsAndLegality(TypeConverter &typeConverter,
RewritePatternSet &patterns,
ConversionTarget &target);
void populateElementwisePatternsAndLegality(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target, const llvm::StringSet<> &convertTorchOpsSet);

void populateMiscPatternsAndLegality(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target, const llvm::StringSet<> &convertTorchOpsSet);

void populateDataMovementPatternsAndLegality(TypeConverter &typeConverter,
RewritePatternSet &patterns,
ConversionTarget &target);
void populateDataMovementPatternsAndLegality(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target, const llvm::StringSet<> &convertTorchOpsSet);

} // namespace torch_to_tcp
} // namespace mlir
37 changes: 29 additions & 8 deletions lib/Conversion/TorchToTcp/TorchToTcp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,9 @@
#include "torch-mlir/Dialect/TorchConversion/IR/TorchConversionOps.h"
#include "torch-mlir/Dialect/TorchConversion/Transforms/BackendTypeConversion.h"

#include "llvm/ADT/ArrayRef.h"
#include "llvm/ADT/StringSet.h"

using namespace mlir;
using namespace mlir::torch;
using namespace mlir::torch::Torch;
Expand All @@ -40,7 +43,15 @@ namespace tcp {
namespace {

class ConvertTorchToTcp : public ConvertTorchToTcpBase<ConvertTorchToTcp> {
private:
llvm::StringSet<> convertTorchOpsSet;

public:
ConvertTorchToTcp() = default;
ConvertTorchToTcp(ArrayRef<std::string> convertTorchOps) {
this->convertTorchOps = convertTorchOps;
}

void getDependentDialects(DialectRegistry &registry) const override {
registry.insert<tcp::TcpDialect>();
registry.insert<tensor::TensorDialect>();
Expand All @@ -49,6 +60,15 @@ class ConvertTorchToTcp : public ConvertTorchToTcpBase<ConvertTorchToTcp> {

void runOnOperation() override {
MLIRContext *context = &getContext();
RewritePatternSet patterns(context);

// The strings in the `convertTorchOps` ArrayRef don't exist during the call
// to the constructor `ConvertTorchToTcp`, so the creation of the
// `convertTorchOpsSet` must be delayed to when `runOnOperation` gets
// called.
convertTorchOpsSet.clear();
convertTorchOpsSet.insert(convertTorchOps.begin(), convertTorchOps.end());

ConversionTarget target(*context);
target.addLegalDialect<tcp::TcpDialect, tensor::TensorDialect,
arith::ArithDialect>();
Expand All @@ -57,14 +77,14 @@ class ConvertTorchToTcp : public ConvertTorchToTcpBase<ConvertTorchToTcp> {
typeConverter.addConversion([](Type type) { return type; });
TorchConversion::setupBackendTypeConversion(target, typeConverter);

RewritePatternSet patterns(context);
torch_to_tcp::populateElementwisePatternsAndLegality(
typeConverter, patterns, target, convertTorchOpsSet);

torch_to_tcp::populateElementwisePatternsAndLegality(typeConverter,
patterns, target);
torch_to_tcp::populateMiscPatternsAndLegality(typeConverter, patterns,
target);
torch_to_tcp::populateDataMovementPatternsAndLegality(typeConverter,
patterns, target);
target, convertTorchOpsSet);

torch_to_tcp::populateDataMovementPatternsAndLegality(
typeConverter, patterns, target, convertTorchOpsSet);

if (failed(applyPartialConversion(getOperation(), target,
std::move(patterns)))) {
Expand All @@ -75,8 +95,9 @@ class ConvertTorchToTcp : public ConvertTorchToTcpBase<ConvertTorchToTcp> {

} // namespace

std::unique_ptr<OperationPass<func::FuncOp>> createConvertTorchToTcpPass() {
return std::make_unique<ConvertTorchToTcp>();
std::unique_ptr<OperationPass<func::FuncOp>>
createConvertTorchToTcpPass(llvm::ArrayRef<std::string> convertTorchOps) {
return std::make_unique<ConvertTorchToTcp>(convertTorchOps);
}

} // namespace tcp
Expand Down
Loading

0 comments on commit 7a6d51d

Please sign in to comment.