[Mlir-commits] [mlir] [mlir][dataflow] Fix SparseConstantPropagation corrupting IR via incomplete fold rollback (PR #213312)

Sanshrav Arora llvmlistbot at llvm.org
Fri Jul 31 10:23:00 PDT 2026


https://github.com/sanshrav1311 updated https://github.com/llvm/llvm-project/pull/213312

>From ebec4287cd614805ae88fc3d5905e83c4b423465 Mon Sep 17 00:00:00 2001
From: Sanshrav Arora <sanshrav1311 at gmail.com>
Date: Fri, 31 Jul 2026 22:44:47 +0530
Subject: [PATCH] [mlir][dataflow] Fix SparseConstantPropagation corrupting IR
 via fold rollback

---
 .../DataFlow/ConstantPropagationAnalysis.cpp  | 21 +++++++++++++---
 .../Arith/unsigned-when-equivalent.mlir       | 25 +++++++++++++++++++
 2 files changed, 42 insertions(+), 4 deletions(-)

diff --git a/mlir/lib/Analysis/DataFlow/ConstantPropagationAnalysis.cpp b/mlir/lib/Analysis/DataFlow/ConstantPropagationAnalysis.cpp
index dbf68ac575dbe..593dd77e0241d 100644
--- a/mlir/lib/Analysis/DataFlow/ConstantPropagationAnalysis.cpp
+++ b/mlir/lib/Analysis/DataFlow/ConstantPropagationAnalysis.cpp
@@ -81,16 +81,29 @@ LogicalResult SparseConstantPropagation::visitOperation(
     return success();
   }
 
-  // If the folding was in-place, mark the results as overdefined and reset
-  // the operation. We don't allow in-place folds as the desire here is for
-  // simulated execution, and not general folding.
-  if (foldResults.empty()) {
+  // If the fold mutated the operation, reset it and mark the results as
+  // overdefined. This analysis passes speculative lattice-inferred values to
+  // op->fold(), which may not reflect actual runtime constants. Some fold
+  // implementations (e.g. vector.extract) mutate the op in-place as an
+  // intermediate step before returning a constant result, leaving non-empty
+  // foldResults -- so checking foldResults.empty() alone is not sufficient to
+  // detect all mutations. Instead, compare the live op state against the
+  // snapshot taken before folding.
+  if (op->getOperands() != ArrayRef<Value>(originalOperands) ||
+      op->getAttrDictionary() != originalAttrs) {
     op->setOperands(originalOperands);
     op->setAttrs(originalAttrs);
     setAllToEntryStates(results);
     return success();
   }
 
+  // If the folding was in-place (signalled by empty foldResults) but the op
+  // was not mutated, mark results as overdefined without needing a restore.
+  if (foldResults.empty()) {
+    setAllToEntryStates(results);
+    return success();
+  }
+
   // Merge the fold results into the lattice for this operation.
   assert(foldResults.size() == op->getNumResults() && "invalid result size");
   for (const auto it : llvm::zip(results, foldResults)) {
diff --git a/mlir/test/Dialect/Arith/unsigned-when-equivalent.mlir b/mlir/test/Dialect/Arith/unsigned-when-equivalent.mlir
index 0ea69de8b8f9a..76f906df69af3 100644
--- a/mlir/test/Dialect/Arith/unsigned-when-equivalent.mlir
+++ b/mlir/test/Dialect/Arith/unsigned-when-equivalent.mlir
@@ -114,3 +114,28 @@ func.func @gpu_func(%arg0: memref<2x32xf32>, %arg1: memref<2x32xf32>, %arg2: mem
   } 
   return %arg1 : memref<2x32xf32> 
 }
+
+// CHECK-LABEL: func @vector_extract_loop_dynamic_index
+// CHECK: vector.extract %{{.*}}[%{{.*}}]
+// CHECK-NOT: vector.extract %{{.*}}[0]
+func.func @vector_extract_loop_dynamic_index() {
+  %c0 = arith.constant 0 : index
+  %c1 = arith.constant 1 : index
+  %c13 = arith.constant 13 : index
+  %c0_i32 = arith.constant 0 : i32
+  %cst = arith.constant dense<[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13]> : vector<13xi32>
+  cf.br ^bb1
+^bb1:
+  cf.br ^bb2(%c0, %c0_i32 : index, i32)
+^bb2(%iv: index, %acc: i32):
+  %cond = arith.cmpi slt, %iv, %c13 : index
+  cf.cond_br %cond, ^bb3, ^bb4
+^bb3:
+  %val = vector.extract %cst[%iv] : i32 from vector<13xi32>
+  %acc_next = arith.addi %acc, %val : i32
+  %iv_next = arith.addi %iv, %c1 : index
+  cf.br ^bb2(%iv_next, %acc_next : index, i32)
+^bb4:
+  vector.print %acc : i32
+  return
+}



More information about the Mlir-commits mailing list