[mlir][shape] Add a func to populate ShapeToShape patterns.

Differential Revision: https://reviews.llvm.org/D81933
This commit is contained in:
Alexander Belyaev
2020-06-16 17:52:34 +02:00
parent 1614e35408
commit 7a9258e9bb
2 changed files with 13 additions and 1 deletions
@@ -18,6 +18,8 @@
namespace mlir {
class MLIRContext;
class OwningRewritePatternList;
class Pass;
/// Creates an instance of the ShapeToShapeLowering pass that legalizes Shape
@@ -25,6 +27,9 @@ class Pass;
/// transformed to `shape.reduce`, which can be lowered to SCF and Standard.
std::unique_ptr<Pass> createShapeToShapeLowering();
/// Collects a set of patterns to rewrite ops within the Shape dialect.
void populateShapeRewritePatterns(MLIRContext *context,
OwningRewritePatternList &patterns);
} // end namespace mlir
#endif // MLIR_DIALECT_SHAPE_TRANSFORMS_PASSES_H_
@@ -54,8 +54,10 @@ struct ShapeToShapeLowering
} // namespace
void ShapeToShapeLowering::runOnFunction() {
MLIRContext &ctx = getContext();
OwningRewritePatternList patterns;
patterns.insert<NumElementsOpConverter>(&getContext());
populateShapeRewritePatterns(&ctx, patterns);
ConversionTarget target(getContext());
target.addLegalDialect<ShapeDialect>();
@@ -64,6 +66,11 @@ void ShapeToShapeLowering::runOnFunction() {
signalPassFailure();
}
void mlir::populateShapeRewritePatterns(MLIRContext *context,
OwningRewritePatternList &patterns) {
patterns.insert<NumElementsOpConverter>(context);
}
std::unique_ptr<Pass> mlir::createShapeToShapeLowering() {
return std::make_unique<ShapeToShapeLowering>();
}