[mlir][sparse] Split SparseTensorRewrite into PreSparsificationRewrite and PostSparsificationRewrite.

Reviewed By: aartbik, wrengr

Differential Revision: https://reviews.llvm.org/D138153
This commit is contained in:
bixia1
2022-11-16 16:28:41 -08:00
parent 662b5f1846
commit f81f0cb75a
14 changed files with 108 additions and 53 deletions

View File

@@ -21,8 +21,9 @@
#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
namespace mlir {
#define GEN_PASS_DEF_SPARSETENSORREWRITE
#define GEN_PASS_DEF_PRESPARSIFICATIONREWRITE
#define GEN_PASS_DEF_SPARSIFICATIONPASS
#define GEN_PASS_DEF_POSTSPARSIFICATIONREWRITE
#define GEN_PASS_DEF_SPARSETENSORCONVERSIONPASS
#define GEN_PASS_DEF_SPARSETENSORCODEGEN
#define GEN_PASS_DEF_SPARSEBUFFERREWRITE
@@ -38,22 +39,17 @@ namespace {
// Passes implementation.
//===----------------------------------------------------------------------===//
struct SparseTensorRewritePass
: public impl::SparseTensorRewriteBase<SparseTensorRewritePass> {
struct PreSparsificationRewritePass
: public impl::PreSparsificationRewriteBase<PreSparsificationRewritePass> {
SparseTensorRewritePass() = default;
SparseTensorRewritePass(const SparseTensorRewritePass &pass) = default;
SparseTensorRewritePass(bool enableRT, bool foreach, bool convert) {
enableRuntimeLibrary = enableRT;
enableForeach = foreach;
enableConvert = convert;
}
PreSparsificationRewritePass() = default;
PreSparsificationRewritePass(const PreSparsificationRewritePass &pass) =
default;
void runOnOperation() override {
auto *ctx = &getContext();
RewritePatternSet patterns(ctx);
populateSparseTensorRewriting(patterns, enableRuntimeLibrary, enableForeach,
enableConvert);
populatePreSparsificationRewriting(patterns);
(void)applyPatternsAndFoldGreedily(getOperation(), std::move(patterns));
}
};
@@ -80,6 +76,28 @@ struct SparsificationPass
}
};
struct PostSparsificationRewritePass
: public impl::PostSparsificationRewriteBase<
PostSparsificationRewritePass> {
PostSparsificationRewritePass() = default;
PostSparsificationRewritePass(const PostSparsificationRewritePass &pass) =
default;
PostSparsificationRewritePass(bool enableRT, bool foreach, bool convert) {
enableRuntimeLibrary = enableRT;
enableForeach = foreach;
enableConvert = convert;
}
void runOnOperation() override {
auto *ctx = &getContext();
RewritePatternSet patterns(ctx);
populatePostSparsificationRewriting(patterns, enableRuntimeLibrary,
enableForeach, enableConvert);
(void)applyPatternsAndFoldGreedily(getOperation(), std::move(patterns));
}
};
struct SparseTensorConversionPass
: public impl::SparseTensorConversionPassBase<SparseTensorConversionPass> {
@@ -254,15 +272,8 @@ mlir::sparseToSparseConversionStrategy(int32_t flag) {
// Pass creation methods.
//===----------------------------------------------------------------------===//
std::unique_ptr<Pass> mlir::createSparseTensorRewritePass() {
return std::make_unique<SparseTensorRewritePass>();
}
std::unique_ptr<Pass> mlir::createSparseTensorRewritePass(bool enableRT,
bool enableForeach,
bool enableConvert) {
return std::make_unique<SparseTensorRewritePass>(enableRT, enableForeach,
enableConvert);
std::unique_ptr<Pass> mlir::createPreSparsificationRewritePass() {
return std::make_unique<PreSparsificationRewritePass>();
}
std::unique_ptr<Pass> mlir::createSparsificationPass() {
@@ -274,6 +285,17 @@ mlir::createSparsificationPass(const SparsificationOptions &options) {
return std::make_unique<SparsificationPass>(options);
}
std::unique_ptr<Pass> mlir::createPostSparsificationRewritePass() {
return std::make_unique<PostSparsificationRewritePass>();
}
std::unique_ptr<Pass>
mlir::createPostSparsificationRewritePass(bool enableRT, bool enableForeach,
bool enableConvert) {
return std::make_unique<PostSparsificationRewritePass>(
enableRT, enableForeach, enableConvert);
}
std::unique_ptr<Pass> mlir::createSparseTensorConversionPass() {
return std::make_unique<SparseTensorConversionPass>();
}