[mlir][sparse] move sparse tensor rewriting into its own pass

Makes individual testing and debugging easier.

Reviewed By: bixia

Differential Revision: https://reviews.llvm.org/D135319
This commit is contained in:
Aart Bik
2022-10-05 13:38:51 -07:00
parent 617ca92bf1
commit 779dcd2ecc
8 changed files with 59 additions and 8 deletions

View File

@@ -21,6 +21,7 @@
#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
namespace mlir {
#define GEN_PASS_DEF_SPARSETENSORREWRITE
#define GEN_PASS_DEF_SPARSIFICATIONPASS
#define GEN_PASS_DEF_SPARSETENSORCONVERSIONPASS
#define GEN_PASS_DEF_SPARSETENSORCODEGEN
@@ -37,6 +38,23 @@ namespace {
// Passes implementation.
//===----------------------------------------------------------------------===//
struct SparseTensorRewritePass
: public impl::SparseTensorRewriteBase<SparseTensorRewritePass> {
SparseTensorRewritePass() = default;
SparseTensorRewritePass(const SparseTensorRewritePass &pass) = default;
SparseTensorRewritePass(const SparsificationOptions &options) {
enableRuntimeLibrary = options.enableRuntimeLibrary;
}
void runOnOperation() override {
auto *ctx = &getContext();
RewritePatternSet patterns(ctx);
populateSparseTensorRewriting(patterns, enableRuntimeLibrary);
(void)applyPatternsAndFoldGreedily(getOperation(), std::move(patterns));
}
};
struct SparsificationPass
: public impl::SparsificationPassBase<SparsificationPass> {
@@ -53,14 +71,10 @@ struct SparsificationPass
void runOnOperation() override {
auto *ctx = &getContext();
RewritePatternSet prePatterns(ctx);
// Translate strategy flags to strategy options.
SparsificationOptions options(parallelization, vectorization, vectorLength,
enableSIMDIndex32, enableVLAVectorization,
enableRuntimeLibrary);
// Apply pre-rewriting.
populateSparseTensorRewriting(prePatterns, options.enableRuntimeLibrary);
(void)applyPatternsAndFoldGreedily(getOperation(), std::move(prePatterns));
// Apply sparsification and vector cleanup rewriting.
RewritePatternSet patterns(ctx);
populateSparsificationPatterns(patterns, options);
@@ -236,6 +250,15 @@ 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(const SparsificationOptions &options) {
return std::make_unique<SparseTensorRewritePass>(options);
}
std::unique_ptr<Pass> mlir::createSparsificationPass() {
return std::make_unique<SparsificationPass>();
}