From 779dcd2ecce84983fcec0a23dee66f85772cb17a Mon Sep 17 00:00:00 2001 From: Aart Bik Date: Wed, 5 Oct 2022 13:38:51 -0700 Subject: [PATCH] [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 --- .../Dialect/SparseTensor/Transforms/Passes.h | 4 +++ .../Dialect/SparseTensor/Transforms/Passes.td | 23 ++++++++++++++ .../Pipelines/SparseTensorPipelines.cpp | 1 + .../Transforms/SparseTensorPasses.cpp | 31 ++++++++++++++++--- mlir/test/Dialect/SparseTensor/rewriting.mlir | 2 +- .../SparseTensor/sparse_concat_codegen.mlir | 2 +- .../SparseTensor/sparse_fill_zero.mlir | 2 +- .../Dialect/SparseTensor/sparse_sddmm.mlir | 2 +- 8 files changed, 59 insertions(+), 8 deletions(-) diff --git a/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.h b/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.h index e6e65b7b0d68..fd99e4f57af0 100644 --- a/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.h +++ b/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.h @@ -163,6 +163,10 @@ std::unique_ptr createSparseTensorCodegenPass(); void populateSparseTensorRewriting(RewritePatternSet &patterns, bool enableRT); +std::unique_ptr createSparseTensorRewritePass(); +std::unique_ptr +createSparseTensorRewritePass(const SparsificationOptions &options); + std::unique_ptr createDenseBufferizationPass( const bufferization::OneShotBufferizationOptions &options); diff --git a/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.td b/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.td index d97bead8d296..26c78aea50a8 100644 --- a/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.td +++ b/mlir/include/mlir/Dialect/SparseTensor/Transforms/Passes.td @@ -11,6 +11,27 @@ include "mlir/Pass/PassBase.td" +def SparseTensorRewrite : Pass<"sparse-tensor-rewrite", "ModuleOp"> { + let summary = "Applies sparse tensor rewriting rules prior to sparsification"; + let description = [{ + A pass that applies rewriting rules to sparse tensor operations prior + to running the actual sparsification pass. + }]; + let constructor = "mlir::createSparseTensorRewritePass()"; + let dependentDialects = [ + "arith::ArithDialect", + "bufferization::BufferizationDialect", + "linalg::LinalgDialect", + "memref::MemRefDialect", + "scf::SCFDialect", + "sparse_tensor::SparseTensorDialect", + ]; + let options = [ + Option<"enableRuntimeLibrary", "enable-runtime-library", "bool", + "true", "Enable runtime library for manipulating sparse tensors"> + ]; +} + def SparsificationPass : Pass<"sparsification", "ModuleOp"> { let summary = "Automatically generate sparse tensor code from sparse tensor types"; let description = [{ @@ -57,6 +78,7 @@ def SparsificationPass : Pass<"sparsification", "ModuleOp"> { "arith::ArithDialect", "bufferization::BufferizationDialect", "LLVM::LLVMDialect", + "linalg::LinalgDialect", "memref::MemRefDialect", "scf::SCFDialect", "sparse_tensor::SparseTensorDialect", @@ -193,4 +215,5 @@ def SparseBufferRewrite : Pass<"sparse-buffer-rewrite", "ModuleOp"> { "sparse_tensor::SparseTensorDialect", ]; } + #endif // MLIR_DIALECT_SPARSETENSOR_TRANSFORMS_PASSES diff --git a/mlir/lib/Dialect/SparseTensor/Pipelines/SparseTensorPipelines.cpp b/mlir/lib/Dialect/SparseTensor/Pipelines/SparseTensorPipelines.cpp index abecf4679e4e..0cd17998e81b 100644 --- a/mlir/lib/Dialect/SparseTensor/Pipelines/SparseTensorPipelines.cpp +++ b/mlir/lib/Dialect/SparseTensor/Pipelines/SparseTensorPipelines.cpp @@ -58,6 +58,7 @@ void mlir::sparse_tensor::buildSparseCompiler( /*analysisOnly=*/options.testBufferizationAnalysisOnly))); if (options.testBufferizationAnalysisOnly) return; + pm.addPass(createSparseTensorRewritePass(options.sparsificationOptions())); pm.addPass(createSparsificationPass(options.sparsificationOptions())); if (options.enableRuntimeLibrary) pm.addPass(createSparseTensorConversionPass( diff --git a/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp b/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp index b208dfeb5558..5ea55c822634 100644 --- a/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp +++ b/mlir/lib/Dialect/SparseTensor/Transforms/SparseTensorPasses.cpp @@ -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() = 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 { @@ -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 mlir::createSparseTensorRewritePass() { + return std::make_unique(); +} + +std::unique_ptr +mlir::createSparseTensorRewritePass(const SparsificationOptions &options) { + return std::make_unique(options); +} + std::unique_ptr mlir::createSparsificationPass() { return std::make_unique(); } diff --git a/mlir/test/Dialect/SparseTensor/rewriting.mlir b/mlir/test/Dialect/SparseTensor/rewriting.mlir index 000c3560f1e0..f142ecf7ff34 100755 --- a/mlir/test/Dialect/SparseTensor/rewriting.mlir +++ b/mlir/test/Dialect/SparseTensor/rewriting.mlir @@ -1,4 +1,4 @@ -// RUN: mlir-opt %s -sparsification | FileCheck %s +// RUN: mlir-opt %s -sparse-tensor-rewrite | FileCheck %s #SparseVector = #sparse_tensor.encoding<{ dimLevelType = ["compressed"] diff --git a/mlir/test/Dialect/SparseTensor/sparse_concat_codegen.mlir b/mlir/test/Dialect/SparseTensor/sparse_concat_codegen.mlir index 018f39112207..a46da498b387 100644 --- a/mlir/test/Dialect/SparseTensor/sparse_concat_codegen.mlir +++ b/mlir/test/Dialect/SparseTensor/sparse_concat_codegen.mlir @@ -1,4 +1,4 @@ -// RUN: mlir-opt %s --sparsification=enable-runtime-library=false | FileCheck %s +// RUN: mlir-opt %s --sparse-tensor-rewrite=enable-runtime-library=false --sparsification | FileCheck %s #DCSR = #sparse_tensor.encoding<{dimLevelType = ["compressed", "compressed"]}> diff --git a/mlir/test/Dialect/SparseTensor/sparse_fill_zero.mlir b/mlir/test/Dialect/SparseTensor/sparse_fill_zero.mlir index 132566653c97..2b388bafc2cf 100644 --- a/mlir/test/Dialect/SparseTensor/sparse_fill_zero.mlir +++ b/mlir/test/Dialect/SparseTensor/sparse_fill_zero.mlir @@ -1,4 +1,4 @@ -// RUN: mlir-opt %s --linalg-generalize-named-ops --sparsification --sparse-tensor-conversion --canonicalize --cse | FileCheck %s +// RUN: mlir-opt %s --linalg-generalize-named-ops --sparse-tensor-rewrite --sparsification --sparse-tensor-conversion --canonicalize --cse | FileCheck %s #DCSR = #sparse_tensor.encoding<{ dimLevelType = [ "compressed", "compressed" ] }> diff --git a/mlir/test/Dialect/SparseTensor/sparse_sddmm.mlir b/mlir/test/Dialect/SparseTensor/sparse_sddmm.mlir index 1a521800f333..ad1ad1d524be 100755 --- a/mlir/test/Dialect/SparseTensor/sparse_sddmm.mlir +++ b/mlir/test/Dialect/SparseTensor/sparse_sddmm.mlir @@ -1,4 +1,4 @@ -// RUN: mlir-opt %s --tensor-copy-insertion --sparsification --cse | FileCheck %s +// RUN: mlir-opt %s --tensor-copy-insertion --sparse-tensor-rewrite --sparsification --cse | FileCheck %s #SM = #sparse_tensor.encoding<{ dimLevelType = [ "compressed", "compressed" ] }>