[Mlir-commits] [mlir] [mlir][bufferization] Fix stale BufferOriginAnalysis in BufferDeallocationSimplification (PR #210105)
Vito Secona
llvmlistbot at llvm.org
Sat Jul 18 11:02:25 PDT 2026
https://github.com/secona updated https://github.com/llvm/llvm-project/pull/210105
>From 7f78ea3f105503b8d107249726710007e6f18b91 Mon Sep 17 00:00:00 2001
From: Vito Secona <secona00 at gmail.com>
Date: Sun, 19 Jul 2026 00:35:39 +0700
Subject: [PATCH] [mlir][bufferization] Fix stale BufferOriginAnalysis in
BufferDeallocationSimplification
---
.../BufferDeallocationSimplification.cpp | 16 ++++++++--
.../buffer-deallocation-simplification.mlir | 29 +++++++++++++++++++
2 files changed, 43 insertions(+), 2 deletions(-)
diff --git a/mlir/lib/Dialect/Bufferization/Transforms/BufferDeallocationSimplification.cpp b/mlir/lib/Dialect/Bufferization/Transforms/BufferDeallocationSimplification.cpp
index a465c957d063e..9660786868498 100644
--- a/mlir/lib/Dialect/Bufferization/Transforms/BufferDeallocationSimplification.cpp
+++ b/mlir/lib/Dialect/Bufferization/Transforms/BufferDeallocationSimplification.cpp
@@ -469,12 +469,24 @@ struct BufferDeallocationSimplificationPass
RetainedMemrefAliasingAlwaysDeallocatedMemref>(&getContext(),
analysis);
- populateDeallocOpCanonicalizationPatterns(patterns, &getContext());
// We don't want that the block structure changes invalidating the
// `BufferOriginAnalysis` so we apply the rewrites with `Normal` level of
- // region simplification
+ // region simplification and disable folding.
if (failed(applyPatternsGreedily(
getOperation(), std::move(patterns),
+ GreedyRewriteConfig()
+ .setRegionSimplificationLevel(GreedySimplifyRegionLevel::Normal)
+ .enableFolding(false))))
+ signalPassFailure();
+
+ // We run canonicalization separately so it can benefit from
+ // folding, which was disabled in the previous pass to avoid invalidating
+ // `BufferOriginAnalysis`.
+ RewritePatternSet canonicalizationPatterns(&getContext());
+ populateDeallocOpCanonicalizationPatterns(canonicalizationPatterns,
+ &getContext());
+ if (failed(applyPatternsGreedily(
+ getOperation(), std::move(canonicalizationPatterns),
GreedyRewriteConfig().setRegionSimplificationLevel(
GreedySimplifyRegionLevel::Normal))))
signalPassFailure();
diff --git a/mlir/test/Dialect/Bufferization/Transforms/buffer-deallocation-simplification.mlir b/mlir/test/Dialect/Bufferization/Transforms/buffer-deallocation-simplification.mlir
index b40a17cf800bf..3acc09d23b59e 100644
--- a/mlir/test/Dialect/Bufferization/Transforms/buffer-deallocation-simplification.mlir
+++ b/mlir/test/Dialect/Bufferization/Transforms/buffer-deallocation-simplification.mlir
@@ -169,3 +169,32 @@ func.func @duplicate_memref(%arg0: memref<5xf32>, %arg1: memref<6xf32>, %c: i1)
// CHECK-LABEL: func @duplicate_memref(
// CHECK: %[[r:.*]] = bufferization.dealloc (%{{.*}} : memref<5xf32>) if (%{{.*}}) retain (%{{.*}} : memref<6xf32>)
// CHECK: return %[[r]]
+
+// -----
+
+module {
+ func.func @memref_cast_folding(%arg0: memref<4x4xf32>) -> memref<4x4xf32, 1> {
+ %cfalse = arith.constant false
+ %ctrue = arith.constant true
+
+ %memspacecast = memref.memory_space_cast %arg0 : memref<4x4xf32> to memref<4x4xf32, 1>
+ %cast = memref.cast %memspacecast : memref<4x4xf32, 1> to memref<?x?xf32, 1>
+ %cast_1 = memref.cast %cast : memref<?x?xf32, 1> to memref<4x4xf32, 1>
+
+ %5 = scf.if %cfalse -> (memref<4x4xf32, 1>) {
+ scf.yield %cast_1 : memref<4x4xf32, 1>
+ } else {
+ %7 = bufferization.clone %cast_1 : memref<4x4xf32, 1> to memref<4x4xf32, 1>
+ scf.yield %7 : memref<4x4xf32, 1>
+ }
+
+ %base_buffer, %offset, %sizes:2, %strides:2 = memref.extract_strided_metadata %arg0 : memref<4x4xf32> -> memref<f32>, index, index, index, index, index
+ %base_buffer_5, %offset_6, %sizes_7:2, %strides_8:2 = memref.extract_strided_metadata %memspacecast : memref<4x4xf32, 1> -> memref<f32, 1>, index, index, index, index, index
+
+ %6 = bufferization.dealloc (%base_buffer, %base_buffer_5 : memref<f32>, memref<f32, 1>) if (%cfalse, %ctrue) retain (%5 : memref<4x4xf32, 1>)
+ return %memspacecast : memref<4x4xf32, 1>
+ }
+}
+
+// CHECK-LABEL: func @memref_cast_folding
+// CHECK: %[[r:.*]] = bufferization.dealloc (%{{.*}} : memref<f32, 1>) if (%{{.*}}) retain (%{{.*}} : memref<4x4xf32, 1>)
More information about the Mlir-commits
mailing list