[NFC] Refactoring the methods for removing dead allocs and reshape folding into interface patterns. (#12541)

This adds populate* methods to get patterns to remove dead allocations and to fold reshapes into hal.interface.binding.subspan so that they could be used in different passes.
diff --git a/compiler/src/iree/compiler/Codegen/Common/CleanupBufferAllocViewPass.cpp b/compiler/src/iree/compiler/Codegen/Common/CleanupBufferAllocViewPass.cpp
index f34deac..dccf5be 100644
--- a/compiler/src/iree/compiler/Codegen/Common/CleanupBufferAllocViewPass.cpp
+++ b/compiler/src/iree/compiler/Codegen/Common/CleanupBufferAllocViewPass.cpp
@@ -14,6 +14,7 @@
 
 #include "iree/compiler/Codegen/PassDetail.h"
 #include "iree/compiler/Codegen/Passes.h"
+#include "iree/compiler/Codegen/Transforms/Transforms.h"
 #include "iree/compiler/Dialect/Flow/IR/FlowOps.h"
 #include "iree/compiler/Dialect/HAL/IR/HALOps.h"
 #include "mlir/Dialect/Linalg/IR/Linalg.h"
@@ -26,116 +27,15 @@
 namespace mlir {
 namespace iree_compiler {
 
-void populateReshapeToInterfaceTensorPatterns(RewritePatternSet &patterns);
-
 namespace {
 
-/// Folds tensor.expand/collapse_shape into the source
-/// hal.interface.binding.subspan.
-///
-/// For example, this matches the following pattern:
-///
-///   %subspan = hal.interface.binding.subspan ... :
-///       !flow.dispatch.tensor<readonly:tensor<3x3x1x96xf32>>
-///   %tensor = flow.dispatch.tensor.load %subspan :
-///       !flow.dispatch.tensor<readonly:tensor<3x3x1x96xf32>> ->
-///       tensor<3x3x1x96xf32>
-///   %0 = linalg.tensor_reshape %tensor [
-///         affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>
-///       ] : tensor<3x3x1x96xf32> into tensor<864xf32>
-///
-/// And turns it into:
-///
-///   %subspan = hal.interface.binding.subspan ... :
-///       !flow.dispatch.tensor<readonly:tensor<864xf32>>
-///   %0 = flow.dispatch.tensor.load %subspan :
-///       !flow.dispatch.tensor<readonly:tensor<864xf32>> -> tensor<864xf32>
-template <typename TensorReshapeOp>
-struct FoldReshapeIntoInterfaceTensorLoad : OpRewritePattern<TensorReshapeOp> {
-  using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
-
-  LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
-                                PatternRewriter &rewriter) const override {
-    // TODO(antigainst): enable dynamic shape support once they are needed.
-    auto reshapeSrcType =
-        reshapeOp.getSrc().getType().template cast<ShapedType>();
-    auto reshapeDstType = reshapeOp.getType().template cast<ShapedType>();
-    if (!reshapeSrcType.hasStaticShape() || !reshapeDstType.hasStaticShape()) {
-      return failure();
-    }
-
-    auto loadOp =
-        reshapeOp.getSrc()
-            .template getDefiningOp<IREE::Flow::DispatchTensorLoadOp>();
-    if (!loadOp) return failure();
-
-    // Make sure we are loading the full incoming subspan. Otherwise we cannot
-    // simply adjust the subspan's resultant type later.
-    if (!loadOp.offsets().empty() || !loadOp.sizes().empty() ||
-        !loadOp.strides().empty())
-      return failure();
-
-    auto subspanOp =
-        loadOp.getSource()
-            .template getDefiningOp<IREE::HAL::InterfaceBindingSubspanOp>();
-    if (!subspanOp) return failure();
-    assert(subspanOp.getDynamicDims().empty());
-
-    auto tensorAccess = subspanOp.getType()
-                            .template cast<IREE::Flow::DispatchTensorType>()
-                            .getAccess();
-    auto newSubspanType = IREE::Flow::DispatchTensorType::get(
-        tensorAccess, reshapeOp.getResultType());
-
-    Value newSubspanOp = rewriter.create<IREE::HAL::InterfaceBindingSubspanOp>(
-        subspanOp.getLoc(), newSubspanType, subspanOp.getSet(),
-        subspanOp.getBinding(), subspanOp.getDescriptorType(),
-        subspanOp.getByteOffset(), subspanOp.getDynamicDims(),
-        subspanOp.getAlignmentAttr(), subspanOp.getDescriptorFlagsAttr());
-
-    rewriter.replaceOpWithNewOp<IREE::Flow::DispatchTensorLoadOp>(
-        reshapeOp, reshapeOp.getResultType(), newSubspanOp,
-        loadOp.getSourceDims());
-
-    return success();
-  }
-};
-
-// Removes operations with Allocate MemoryEffects but no uses.
-struct RemoveDeadMemAllocs : RewritePattern {
-  RemoveDeadMemAllocs(MLIRContext *context, PatternBenefit benefit = 1)
-      : RewritePattern(MatchAnyOpTypeTag(), benefit, context) {}
-
-  LogicalResult matchAndRewrite(Operation *op,
-                                PatternRewriter &rewriter) const override {
-    auto memEffect = dyn_cast<MemoryEffectOpInterface>(op);
-    if (!memEffect || !memEffect.hasEffect<MemoryEffects::Allocate>()) {
-      return failure();
-    }
-    SmallVector<Operation *> deadUsers;
-    for (OpOperand &use : op->getUses()) {
-      if (auto user = dyn_cast<memref::AssumeAlignmentOp>(use.getOwner())) {
-        deadUsers.push_back(user);
-        continue;
-      }
-      // For any other use, return failure;
-      return failure();
-    }
-    for (auto user : deadUsers) {
-      rewriter.eraseOp(user);
-    }
-    rewriter.eraseOp(op);
-    return success();
-  }
-};
-
 /// Runs canonicalization patterns on interface load/store ops.
 struct CleanupBufferAllocViewPass
     : public CleanupBufferAllocViewBase<CleanupBufferAllocViewPass> {
   void runOnOperation() override {
     RewritePatternSet patterns(&getContext());
     populateReshapeToInterfaceTensorPatterns(patterns);
-    patterns.insert<RemoveDeadMemAllocs>(&getContext());
+    populateRemoveDeadMemAllocPatterns(patterns);
     if (failed(applyPatternsAndFoldGreedily(getOperation(),
                                             std::move(patterns)))) {
       return signalPassFailure();
@@ -145,12 +45,6 @@
 
 }  // namespace
 
-void populateReshapeToInterfaceTensorPatterns(RewritePatternSet &patterns) {
-  patterns.insert<FoldReshapeIntoInterfaceTensorLoad<tensor::CollapseShapeOp>,
-                  FoldReshapeIntoInterfaceTensorLoad<tensor::ExpandShapeOp>>(
-      patterns.getContext());
-}
-
 std::unique_ptr<OperationPass<func::FuncOp>>
 createCleanupBufferAllocViewPass() {
   return std::make_unique<CleanupBufferAllocViewPass>();
diff --git a/compiler/src/iree/compiler/Codegen/Common/Transforms.h b/compiler/src/iree/compiler/Codegen/Common/Transforms.h
index dec9837..bff3822 100644
--- a/compiler/src/iree/compiler/Codegen/Common/Transforms.h
+++ b/compiler/src/iree/compiler/Codegen/Common/Transforms.h
@@ -47,10 +47,6 @@
 void populateTileAndDistributeToWorkgroupsCleanupPatterns(
     RewritePatternSet &patterns, linalg::LinalgTilingOptions options);
 
-/// Populate patterns that fold tensor.expand/collapse_shape into the source
-/// hal.interface.binding.subspan.
-void populateReshapeToInterfaceTensorPatterns(RewritePatternSet &patterns);
-
 }  // namespace iree_compiler
 }  // namespace mlir
 
diff --git a/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp b/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp
index 0fcedcc..213bdc5 100644
--- a/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp
+++ b/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp
@@ -247,5 +247,128 @@
 template void hoistStaticallyBoundAllocationsInFunc<memref::AllocaOp>(
     RewriterBase &rewriter, func::FuncOp funcOp);
 
+//===---------------------------------------------------------------------===//
+// Patterns to fold tensor.expand/collapse_shape into
+// `hal.interface.binding.subspan`
+//===---------------------------------------------------------------------===//
+
+namespace {
+
+/// Folds tensor.expand/collapse_shape into the source
+/// hal.interface.binding.subspan.
+///
+/// For example, this matches the following pattern:
+///
+///   %subspan = hal.interface.binding.subspan ... :
+///       !flow.dispatch.tensor<readonly:tensor<3x3x1x96xf32>>
+///   %tensor = flow.dispatch.tensor.load %subspan :
+///       !flow.dispatch.tensor<readonly:tensor<3x3x1x96xf32>> ->
+///       tensor<3x3x1x96xf32>
+///   %0 = linalg.tensor_reshape %tensor [
+///         affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>
+///       ] : tensor<3x3x1x96xf32> into tensor<864xf32>
+///
+/// And turns it into:
+///
+///   %subspan = hal.interface.binding.subspan ... :
+///       !flow.dispatch.tensor<readonly:tensor<864xf32>>
+///   %0 = flow.dispatch.tensor.load %subspan :
+///       !flow.dispatch.tensor<readonly:tensor<864xf32>> -> tensor<864xf32>
+template <typename TensorReshapeOp>
+struct FoldReshapeIntoInterfaceTensorLoad : OpRewritePattern<TensorReshapeOp> {
+  using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
+
+  LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
+                                PatternRewriter &rewriter) const override {
+    // TODO(antigainst): enable dynamic shape support once they are needed.
+    auto reshapeSrcType =
+        reshapeOp.getSrc().getType().template cast<ShapedType>();
+    auto reshapeDstType = reshapeOp.getType().template cast<ShapedType>();
+    if (!reshapeSrcType.hasStaticShape() || !reshapeDstType.hasStaticShape()) {
+      return failure();
+    }
+
+    auto loadOp =
+        reshapeOp.getSrc()
+            .template getDefiningOp<IREE::Flow::DispatchTensorLoadOp>();
+    if (!loadOp) return failure();
+
+    // Make sure we are loading the full incoming subspan. Otherwise we cannot
+    // simply adjust the subspan's resultant type later.
+    if (!loadOp.offsets().empty() || !loadOp.sizes().empty() ||
+        !loadOp.strides().empty())
+      return failure();
+
+    auto subspanOp =
+        loadOp.getSource()
+            .template getDefiningOp<IREE::HAL::InterfaceBindingSubspanOp>();
+    if (!subspanOp) return failure();
+    assert(subspanOp.getDynamicDims().empty());
+
+    auto tensorAccess = subspanOp.getType()
+                            .template cast<IREE::Flow::DispatchTensorType>()
+                            .getAccess();
+    auto newSubspanType = IREE::Flow::DispatchTensorType::get(
+        tensorAccess, reshapeOp.getResultType());
+
+    Value newSubspanOp = rewriter.create<IREE::HAL::InterfaceBindingSubspanOp>(
+        subspanOp.getLoc(), newSubspanType, subspanOp.getSet(),
+        subspanOp.getBinding(), subspanOp.getDescriptorType(),
+        subspanOp.getByteOffset(), subspanOp.getDynamicDims(),
+        subspanOp.getAlignmentAttr(), subspanOp.getDescriptorFlagsAttr());
+
+    rewriter.replaceOpWithNewOp<IREE::Flow::DispatchTensorLoadOp>(
+        reshapeOp, reshapeOp.getResultType(), newSubspanOp,
+        loadOp.getSourceDims());
+
+    return success();
+  }
+};
+}  // namespace
+
+void populateReshapeToInterfaceTensorPatterns(RewritePatternSet &patterns) {
+  patterns.insert<FoldReshapeIntoInterfaceTensorLoad<tensor::CollapseShapeOp>,
+                  FoldReshapeIntoInterfaceTensorLoad<tensor::ExpandShapeOp>>(
+      patterns.getContext());
+}
+
+//===--------------------------------------------------------------------====//
+// Pattern to remove dead allocations
+//===--------------------------------------------------------------------====//
+
+namespace {
+// Removes operations with Allocate MemoryEffects but no uses.
+struct RemoveDeadMemAllocs : RewritePattern {
+  RemoveDeadMemAllocs(MLIRContext *context, PatternBenefit benefit = 1)
+      : RewritePattern(MatchAnyOpTypeTag(), benefit, context) {}
+
+  LogicalResult matchAndRewrite(Operation *op,
+                                PatternRewriter &rewriter) const override {
+    auto memEffect = dyn_cast<MemoryEffectOpInterface>(op);
+    if (!memEffect || !memEffect.hasEffect<MemoryEffects::Allocate>()) {
+      return failure();
+    }
+    SmallVector<Operation *> deadUsers;
+    for (OpOperand &use : op->getUses()) {
+      if (auto user = dyn_cast<memref::AssumeAlignmentOp>(use.getOwner())) {
+        deadUsers.push_back(user);
+        continue;
+      }
+      // For any other use, return failure;
+      return failure();
+    }
+    for (auto user : deadUsers) {
+      rewriter.eraseOp(user);
+    }
+    rewriter.eraseOp(op);
+    return success();
+  }
+};
+}  // namespace
+
+void populateRemoveDeadMemAllocPatterns(RewritePatternSet &patterns) {
+  patterns.insert<RemoveDeadMemAllocs>(patterns.getContext());
+}
+
 }  // namespace iree_compiler
 }  // namespace mlir
diff --git a/compiler/src/iree/compiler/Codegen/Transforms/Transforms.h b/compiler/src/iree/compiler/Codegen/Transforms/Transforms.h
index ed2ae94..5b7dd53 100644
--- a/compiler/src/iree/compiler/Codegen/Transforms/Transforms.h
+++ b/compiler/src/iree/compiler/Codegen/Transforms/Transforms.h
@@ -88,6 +88,14 @@
 /// |getMinMaxFn| for some know values.
 void populateRemoveSingleIterationLoopPattern(RewritePatternSet &patterns,
                                               GetMinMaxExprFn getMinMaxFn);
+
+/// Populate patterns that fold tensor.expand/collapse_shape into the source
+/// hal.interface.binding.subspan.
+void populateReshapeToInterfaceTensorPatterns(RewritePatternSet &patterns);
+
+/// Populate patterns that remove dead allocations
+void populateRemoveDeadMemAllocPatterns(RewritePatternSet &patterns);
+
 }  // namespace iree_compiler
 }  // namespace mlir
 
diff --git a/compiler/src/iree/compiler/Dialect/VMVX/Transforms/BUILD b/compiler/src/iree/compiler/Dialect/VMVX/Transforms/BUILD
index 26cb330..dd812c2 100644
--- a/compiler/src/iree/compiler/Dialect/VMVX/Transforms/BUILD
+++ b/compiler/src/iree/compiler/Dialect/VMVX/Transforms/BUILD
@@ -54,6 +54,8 @@
         "//compiler/src/iree/compiler/Codegen:PassHeaders",
         "//compiler/src/iree/compiler/Codegen/Common",
         "//compiler/src/iree/compiler/Codegen/LLVMCPU",
+        "//compiler/src/iree/compiler/Codegen/Transforms",
+        "//compiler/src/iree/compiler/Codegen/Utils",
         "//compiler/src/iree/compiler/Dialect/HAL/IR",
         "//compiler/src/iree/compiler/Dialect/HAL/IR:HALDialect",
         "//compiler/src/iree/compiler/Dialect/HAL/Transforms",
diff --git a/compiler/src/iree/compiler/Dialect/VMVX/Transforms/CMakeLists.txt b/compiler/src/iree/compiler/Dialect/VMVX/Transforms/CMakeLists.txt
index 5c15c48..e032390 100644
--- a/compiler/src/iree/compiler/Dialect/VMVX/Transforms/CMakeLists.txt
+++ b/compiler/src/iree/compiler/Dialect/VMVX/Transforms/CMakeLists.txt
@@ -75,6 +75,8 @@
     iree::compiler::Codegen::Common
     iree::compiler::Codegen::LLVMCPU
     iree::compiler::Codegen::PassHeaders
+    iree::compiler::Codegen::Transforms
+    iree::compiler::Codegen::Utils
     iree::compiler::Dialect::HAL::IR
     iree::compiler::Dialect::HAL::IR::HALDialect
     iree::compiler::Dialect::HAL::Transforms