[Mlir-commits] [mlir] [mlir][memref] Make memref.cast areCastCompatible return true when meet same types (PR #192029)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Tue Apr 14 02:19:41 PDT 2026
llvmbot wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
@llvm/pr-subscribers-mlir-memref
Author: lonely eagle (linuxlonelyeagle)
<details>
<summary>Changes</summary>
When both the source and destination types of `memref.cast` are unranked, it causes an IR verification failure, which impacts downstream projects. To address this, this PR now allows the operation to return true if the source and destination types are identical.
---
Full diff: https://github.com/llvm/llvm-project/pull/192029.diff
2 Files Affected:
- (modified) mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp (+2)
- (modified) mlir/test/Dialect/MemRef/invalid.mlir (+14-3)
``````````diff
diff --git a/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp b/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
index 27c1649ee4ed3..31e4640499276 100644
--- a/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
+++ b/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
@@ -737,6 +737,8 @@ bool CastOp::canFoldIntoConsumerOp(CastOp castOp) {
bool CastOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
if (inputs.size() != 1 || outputs.size() != 1)
return false;
+ if (inputs == outputs)
+ return true;
Type a = inputs.front(), b = outputs.front();
auto aT = llvm::dyn_cast<MemRefType>(a);
auto bT = llvm::dyn_cast<MemRefType>(b);
diff --git a/mlir/test/Dialect/MemRef/invalid.mlir b/mlir/test/Dialect/MemRef/invalid.mlir
index d3670fde08d81..2f061a1bb773e 100644
--- a/mlir/test/Dialect/MemRef/invalid.mlir
+++ b/mlir/test/Dialect/MemRef/invalid.mlir
@@ -894,12 +894,23 @@ func.func @invalid_memref_cast() {
// -----
-// unranked to unranked
+// unranked incompatible element types
func.func @invalid_memref_cast() {
%0 = memref.alloc() : memref<2x5xf32, 0>
%1 = memref.cast %0 : memref<2x5xf32, 0> to memref<*xf32, 0>
- // expected-error at +1 {{operand type 'memref<*xf32>' and result type 'memref<*xf32>' are cast incompatible}}
- %2 = memref.cast %1 : memref<*xf32, 0> to memref<*xf32, 0>
+ // expected-error at +1 {{operand type 'memref<*xf32>' and result type 'memref<*xi32>' are cast incompatible}}
+ %2 = memref.cast %1 : memref<*xf32, 0> to memref<*xi32, 0>
+ return
+}
+
+// -----
+
+// unranked incompatible memory space
+func.func @invalid_memref_cast() {
+ %0 = memref.alloc() : memref<2x5xf32, 0>
+ %1 = memref.cast %0 : memref<2x5xf32, 0> to memref<*xf32, 0>
+ // expected-error at +1 {{operand type 'memref<*xf32>' and result type 'memref<*xf32, 1>' are cast incompatible}}
+ %2 = memref.cast %1 : memref<*xf32, 0> to memref<*xf32, 1>
return
}
``````````
</details>
https://github.com/llvm/llvm-project/pull/192029
More information about the Mlir-commits
mailing list