Skip to content
New issue

Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.

By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.

Already on GitHub? Sign in to your account

Allow partial TorchToTcp conversions using a whitelist #16

Merged
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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);
}
128 changes: 54 additions & 74 deletions lib/Conversion/TorchToTcp/Elementwise.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -650,78 +650,58 @@ 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(AtenOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet<Convert##AtenOp, AtenOp>( \
typeConverter, patterns, target, convertTorchOpsSet)
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(AtenToDtypeOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(AtenClampOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(AtenReluOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(AtenBatchNormOp);
INSERT_ATEN_ELEMENTWISE_OP_PATTERN(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
}
49 changes: 23 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,25 @@ 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(AtenOp) \
torch_to_tcp::addPatternIfOpInConvertTorchOpsSet<Convert##AtenOp, AtenOp>( \
typeConverter, patterns, target, convertTorchOpsSet)
INSERT_ATEN_MISC_OP_PATTERN(AtenBroadcastToOp);
INSERT_ATEN_MISC_OP_PATTERN(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
36 changes: 28 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,14 @@ class ConvertTorchToTcp : public ConvertTorchToTcpBase<ConvertTorchToTcp> {

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

// Usually the default constructor is called which means `convertTorchOps`
// is usually unset. Doing this here allows the initialization of
// `convertTorchOpsSet` to be be delayed to when `runOnOperation` is called.
convertTorchOpsSet.clear();
convertTorchOpsSet.insert(convertTorchOps.begin(), convertTorchOps.end());

ConversionTarget target(*context);
target.addLegalDialect<tcp::TcpDialect, tensor::TensorDialect,
arith::ArithDialect>();
Expand All @@ -57,14 +76,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 +94,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
24 changes: 24 additions & 0 deletions lib/Conversion/TorchToTcp/Utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,10 @@

#include "mlir/Transforms/DialectConversion.h"

#include "torch-mlir/Dialect/Torch/Transforms/Passes.h"

#include "llvm/ADT/StringSet.h"

namespace mlir {
namespace torch_to_tcp {

Expand Down Expand Up @@ -71,6 +75,26 @@ std::optional<Value> getConstTensor(PatternRewriter &rewriter, Operation *op,
bool getConstTensorWithType(ConversionPatternRewriter &rewriter, Operation *op,
Value &constOp, Type resultType, int fillVal);

// Utility function to selectively add a torch->tcp pattern if whitelist op is
// provided
template <typename TorchToTcpPattern, typename AtenOp>
inline void addPatternIfOpInConvertTorchOpsSet(
TypeConverter &typeConverter, RewritePatternSet &patterns,
ConversionTarget &target, const llvm::StringSet<> &convertTorchOpsSet) {
MLIRContext *context = patterns.getContext();
std::optional<OperationName> opName =
TorchToTcpPattern(context).getRootKind();
assert(opName && "All TorchToTcp patterns must target a single op");
// When no ops are specified, convert all.
// When ops are specified, convert those ops only.
if (convertTorchOpsSet.empty() ||
convertTorchOpsSet.contains(
opName->getStringRef().ltrim(torch::Torch::kTorchOpPrefix))) {
target.addIllegalOp<AtenOp>();
patterns.add<TorchToTcpPattern>(typeConverter, context);
}
}

namespace impl {
template <typename T>
std::optional<Value>
Expand Down
Loading