[Mlir-commits] [mlir] [mlir] Preserve dropped conversion mappings (PR #209982)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Wed Jul 15 23:46:04 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir

@llvm/pr-subscribers-mlir-core

Author: lianjinfeng2003 (mygitljf)

<details>
<summary>Changes</summary>

I separated missing mappings from values that were explicitly dropped during dialect conversion. This keeps folded replacement chains from reusing values whose defining operations are removed.
I added coverage for folded live uses and partial 1:N replacements. The original reproducer now reports a normal conversion failure instead of crashing.
Fixes #<!-- -->207353 

---
Full diff: https://github.com/llvm/llvm-project/pull/209982.diff


3 Files Affected:

- (modified) mlir/lib/Transforms/Utils/DialectConversion.cpp (+46-31) 
- (modified) mlir/test/Transforms/test-legalize-erased-op-with-uses.mlir (+25-1) 
- (modified) mlir/test/lib/Dialect/Test/TestPatterns.cpp (+18-1) 


``````````diff
diff --git a/mlir/lib/Transforms/Utils/DialectConversion.cpp b/mlir/lib/Transforms/Utils/DialectConversion.cpp
index c76e3808d3b37..da6bd32a96693 100644
--- a/mlir/lib/Transforms/Utils/DialectConversion.cpp
+++ b/mlir/lib/Transforms/Utils/DialectConversion.cpp
@@ -138,8 +138,10 @@ struct ConversionValueMapping {
   /// false positives.
   bool isMappedTo(Value value) const { return mappedTo.contains(value); }
 
-  /// Lookup a value in the mapping.
-  ValueVector lookup(const ValueVector &from) const;
+  /// Lookup a value in the mapping. A missing optional means that there is no
+  /// mapping, an empty vector is an explicit mapping to no values, and a
+  /// non-empty vector contains the mapped values.
+  std::optional<ValueVector> lookup(const ValueVector &from) const;
 
   template <typename T>
   struct IsValueVector : std::is_same<std::decay_t<T>, ValueVector> {};
@@ -220,11 +222,12 @@ static bool isPureTypeConversion(const ValueVector &values) {
   return op && op->hasAttr(kPureTypeConversionMarker);
 }
 
-ValueVector ConversionValueMapping::lookup(const ValueVector &from) const {
+std::optional<ValueVector>
+ConversionValueMapping::lookup(const ValueVector &from) const {
   auto it = mapping.find(from);
   if (it == mapping.end()) {
     // No mapping found: The lookup stops here.
-    return {};
+    return std::nullopt;
   }
   return it->second;
 }
@@ -945,12 +948,16 @@ struct ConversionPatternRewriterImpl : public RewriterBase::Listener {
   ///
   /// If `skipPureTypeConversions` is "true", materializations that are pure
   /// type conversions are not considered.
+  ///
+  /// If `followDroppedValues` is "false", explicit mappings to no values are
+  /// treated as leaves. Otherwise, they are followed as dropped values.
   ValueVector lookupOrDefault(Value from, TypeRange desiredTypes = {},
-                              bool skipPureTypeConversions = false) const;
+                              bool skipPureTypeConversions = false,
+                              bool followDroppedValues = false) const;
 
   /// Lookup the given value within the map, or return an empty vector if the
-  /// value is not mapped. If it is mapped, this follows the same behavior
-  /// as `lookupOrDefault`.
+  /// value is not mapped. Unlike the default `lookupOrDefault` behavior, this
+  /// follows explicit mappings to no values.
   ValueVector lookupOrNull(Value from, TypeRange desiredTypes = {}) const;
 
   //===--------------------------------------------------------------------===//
@@ -1362,9 +1369,10 @@ void ConversionPatternRewriterImpl::applyRewrites() {
 //===----------------------------------------------------------------------===//
 
 ValueVector ConversionPatternRewriterImpl::lookupOrDefault(
-    Value from, TypeRange desiredTypes, bool skipPureTypeConversions) const {
+    Value from, TypeRange desiredTypes, bool skipPureTypeConversions,
+    bool followDroppedValues) const {
   // Helper function that looks up a single value.
-  auto lookup = [&](const ValueVector &values) -> ValueVector {
+  auto lookup = [&](const ValueVector &values) -> std::optional<ValueVector> {
     assert(!values.empty() && "expected non-empty value vector");
 
     // If the pattern rollback is enabled, use the mapping to look up the
@@ -1376,32 +1384,37 @@ ValueVector ConversionPatternRewriterImpl::lookupOrDefault(
     // already been materialized in IR.
     Operation *op = getCommonDefiningOp(values);
     if (!op)
-      return {};
+      return std::nullopt;
     auto castOp = dyn_cast<UnrealizedConversionCastOp>(op);
     if (!castOp)
-      return {};
+      return std::nullopt;
     if (!this->unresolvedMaterializations.contains(castOp))
-      return {};
+      return std::nullopt;
     if (castOp.getOutputs() != values)
-      return {};
-    return castOp.getInputs();
+      return std::nullopt;
+    if (castOp.getInputs().empty())
+      return std::nullopt;
+    return ValueVector(castOp.getInputs());
   };
 
   // Helper function that looks up each value in `values` individually and then
-  // composes the results. If that fails, it tries to look up the entire vector
-  // at once.
-  auto composedLookup = [&](const ValueVector &values) -> ValueVector {
+  // composes the results. If no singleton mapping can be followed, it tries to
+  // look up the entire vector at once.
+  auto composedLookup =
+      [&](const ValueVector &values) -> std::optional<ValueVector> {
     // If possible, replace each value with (one or multiple) mapped values.
     ValueVector next;
+    bool foundMapping = false;
     for (Value v : values) {
-      ValueVector r = lookup({v});
-      if (!r.empty()) {
-        llvm::append_range(next, r);
-      } else {
+      std::optional<ValueVector> r = lookup({v});
+      if (!r || (r->empty() && !followDroppedValues)) {
         next.push_back(v);
+        continue;
       }
+      foundMapping = true;
+      llvm::append_range(next, *r);
     }
-    if (next != values) {
+    if (foundMapping) {
       // At least one value was replaced.
       return next;
     }
@@ -1414,11 +1427,9 @@ ValueVector ConversionPatternRewriterImpl::lookupOrDefault(
     // be stored (and looked up) in the mapping. But for performance reasons,
     // we choose to reuse existing IR (when possible) instead of creating it
     // multiple times.
-    ValueVector r = lookup(values);
-    if (r.empty()) {
-      // No mapping found: The lookup stops here.
-      return {};
-    }
+    std::optional<ValueVector> r = lookup(values);
+    if (r && r->empty() && !followDroppedValues)
+      return std::nullopt;
     return r;
   };
 
@@ -1443,10 +1454,12 @@ ValueVector ConversionPatternRewriterImpl::lookupOrDefault(
       desiredValue = current;
 
     // Lookup next value in the mapping.
-    ValueVector next = composedLookup(current);
-    if (next.empty())
+    std::optional<ValueVector> next = composedLookup(current);
+    if (!next)
       break;
-    current = std::move(next);
+    if (next->empty())
+      return {};
+    current = std::move(*next);
   } while (true);
 
   // If the desired values were found use them, otherwise default to the leaf
@@ -1461,7 +1474,9 @@ ValueVector ConversionPatternRewriterImpl::lookupOrDefault(
 ValueVector
 ConversionPatternRewriterImpl::lookupOrNull(Value from,
                                             TypeRange desiredTypes) const {
-  ValueVector result = lookupOrDefault(from, desiredTypes);
+  ValueVector result = lookupOrDefault(from, desiredTypes,
+                                       /*skipPureTypeConversions=*/false,
+                                       /*followDroppedValues=*/true);
   if (result == ValueVector{from} ||
       (!desiredTypes.empty() && TypeRange(ValueRange(result)) != desiredTypes))
     return {};
diff --git a/mlir/test/Transforms/test-legalize-erased-op-with-uses.mlir b/mlir/test/Transforms/test-legalize-erased-op-with-uses.mlir
index 031442b0ee2da..dea5e2eb6f7e1 100644
--- a/mlir/test/Transforms/test-legalize-erased-op-with-uses.mlir
+++ b/mlir/test/Transforms/test-legalize-erased-op-with-uses.mlir
@@ -1,4 +1,4 @@
-// RUN: mlir-opt %s -test-legalize-unknown-root-patterns -verify-diagnostics
+// RUN: mlir-opt %s -test-legalize-unknown-root-patterns -split-input-file -verify-diagnostics | FileCheck %s
 
 // Test that an error is emitted when an operation is marked as "erased", but
 // has users that live across the conversion.
@@ -8,3 +8,27 @@ func.func @remove_all_ops(%arg0: i32) -> i32 {
   // expected-note at below {{see existing live user here}}
   return %0 : i32
 }
+
+// -----
+
+// Test that folding an identity cast does not hide a live use of an explicitly
+// erased result.
+func.func @remove_op_through_folded_identity_cast() -> i32 {
+  %0 = "test.illegal_op_a"() : () -> i32
+  // expected-error at below {{failed to legalize unresolved materialization from () to ('i32') that remained live after conversion}}
+  %1 = builtin.unrealized_conversion_cast %0 : i32 to i32
+  // expected-note at below {{see existing live user here}}
+  return %1 : i32
+}
+
+// -----
+
+// CHECK-LABEL: func.func @compose_partial_1_to_n_erasure
+// CHECK-SAME:  (%[[ARG:.*]]: i32) -> i32 {
+// CHECK-NEXT:    return %[[ARG]] : i32
+// CHECK-NEXT:  }
+func.func @compose_partial_1_to_n_erasure(%arg0: i32) -> i32 {
+  %0 = "test.illegal_op_b"() : () -> i32
+  %1 = "test.cast"(%0, %arg0) : (i32, i32) -> i32
+  return %1 : i32
+}
diff --git a/mlir/test/lib/Dialect/Test/TestPatterns.cpp b/mlir/test/lib/Dialect/Test/TestPatterns.cpp
index 552a1a473c9fd..451fe299adf13 100644
--- a/mlir/test/lib/Dialect/Test/TestPatterns.cpp
+++ b/mlir/test/lib/Dialect/Test/TestPatterns.cpp
@@ -1881,6 +1881,22 @@ struct RemoveTestDialectOps : public RewritePattern {
   }
 };
 
+/// This pattern replaces the single result of test.cast with its two operands
+/// to test composing a partial 1:N replacement with an erasure.
+struct ReplaceTestDialectOpWithOperands : public ConversionPattern {
+  ReplaceTestDialectOpWithOperands(MLIRContext *context)
+      : ConversionPattern("test.cast", /*benefit=*/2, context) {}
+
+  LogicalResult
+  matchAndRewrite(Operation *op, ArrayRef<ValueRange> operands,
+                  ConversionPatternRewriter &rewriter) const override {
+    if (op->getNumOperands() != 2 || op->getNumResults() != 1)
+      return failure();
+    rewriter.replaceOpWithMultiple(op, {op->getOperands()});
+    return success();
+  }
+};
+
 struct TestUnknownRootOpDriver
     : public mlir::PassWrapper<TestUnknownRootOpDriver, OperationPass<>> {
   MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(TestUnknownRootOpDriver)
@@ -1893,7 +1909,8 @@ struct TestUnknownRootOpDriver
   }
   void runOnOperation() override {
     mlir::RewritePatternSet patterns(&getContext());
-    patterns.add<RemoveTestDialectOps>(&getContext());
+    patterns.add<RemoveTestDialectOps, ReplaceTestDialectOpWithOperands>(
+        &getContext());
 
     mlir::ConversionTarget target(getContext());
     target.addIllegalDialect<TestDialect>();

``````````

</details>


https://github.com/llvm/llvm-project/pull/209982


More information about the Mlir-commits mailing list