[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