Re-add pattern for removing dead interface bindings. (#12560)
Re-adds the pattern for removing dead `hal.interface.binding.subspan` ops where the only users are `memref.assume_alignment` ops that no longer applies after making it a pure op.
diff --git a/compiler/src/iree/compiler/Codegen/Common/test/remove_dead_allocs.mlir b/compiler/src/iree/compiler/Codegen/Common/test/remove_dead_allocs.mlir
index 146d3d0..50040b1 100644
--- a/compiler/src/iree/compiler/Codegen/Common/test/remove_dead_allocs.mlir
+++ b/compiler/src/iree/compiler/Codegen/Common/test/remove_dead_allocs.mlir
@@ -16,3 +16,13 @@
// CHECK-LABEL: func.func @alloc_keep
// CHECK-NEXT: %[[ALLOC:.+]] = memref.alloc
// CHECK-NEXT: return %[[ALLOC]]
+
+// -----
+
+func.func @cleanup_only_assume_alignment_uses() {
+ %0 = hal.interface.binding.subspan set(0) binding(0) type(storage_buffer) : memref<42xf32>
+ memref.assume_alignment %0, 64 : memref<42xf32>
+ return
+}
+// CHECK-LABEL: func.func @cleanup_only_assume_alignment_uses()
+// CHECK-NEXT: return
diff --git a/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp b/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp
index 213bdc5..f8a1221 100644
--- a/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp
+++ b/compiler/src/iree/compiler/Codegen/Transforms/Transforms.cpp
@@ -337,6 +337,26 @@
//===--------------------------------------------------------------------====//
namespace {
+
+// Erases the operation if its only users are memref.assume_alignment ops.
+static LogicalResult eraseAlignmentOnlyDeadOp(PatternRewriter &rewriter,
+ Operation *op) {
+ 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();
+}
+
// Removes operations with Allocate MemoryEffects but no uses.
struct RemoveDeadMemAllocs : RewritePattern {
RemoveDeadMemAllocs(MLIRContext *context, PatternBenefit benefit = 1)
@@ -348,26 +368,27 @@
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();
+ return eraseAlignmentOnlyDeadOp(rewriter, op);
+ }
+};
+
+// Removes hal.interface.binding.subspan ops with only assume_alignment uses.
+struct RemoveDeadInterfaceBindings
+ : OpRewritePattern<IREE::HAL::InterfaceBindingSubspanOp> {
+ RemoveDeadInterfaceBindings(MLIRContext *context, PatternBenefit benefit = 1)
+ : OpRewritePattern<IREE::HAL::InterfaceBindingSubspanOp>(context,
+ benefit) {}
+
+ LogicalResult matchAndRewrite(IREE::HAL::InterfaceBindingSubspanOp op,
+ PatternRewriter &rewriter) const override {
+ return eraseAlignmentOnlyDeadOp(rewriter, op);
}
};
} // namespace
void populateRemoveDeadMemAllocPatterns(RewritePatternSet &patterns) {
patterns.insert<RemoveDeadMemAllocs>(patterns.getContext());
+ patterns.insert<RemoveDeadInterfaceBindings>(patterns.getContext());
}
} // namespace iree_compiler