[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