[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:
@@ -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>();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user