[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