[Mlir-commits] [mlir] 359b475 - [mlir][memref] Remove unsafe `getType()` from ReshapeOp (#205105)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jun 26 02:28:32 PDT 2026
Author: Longsheng Mou
Date: 2026-06-26T17:28:28+08:00
New Revision: 359b47535c3f9def01f48d7d08c76a42b41d0645
URL: https://github.com/llvm/llvm-project/commit/359b47535c3f9def01f48d7d08c76a42b41d0645
DIFF: https://github.com/llvm/llvm-project/commit/359b47535c3f9def01f48d7d08c76a42b41d0645.diff
LOG: [mlir][memref] Remove unsafe `getType()` from ReshapeOp (#205105)
Remove the unsafe `getType` method from ReshapeOp. It unconditionally
casts the result to `MemRefType`, but `memref.reshape` may return an
`UnrankedMemRefType`, leading to an assertion failure. The redundant
build method is also removed alongside this change. Fixes #203812.
Added:
Modified:
mlir/include/mlir/Dialect/MemRef/IR/MemRefOps.td
mlir/lib/Dialect/MemRef/Transforms/ExpandOps.cpp
mlir/test/Transforms/test-bubble-down-memory-space-casts.mlir
Removed:
################################################################################
diff --git a/mlir/include/mlir/Dialect/MemRef/IR/MemRefOps.td b/mlir/include/mlir/Dialect/MemRef/IR/MemRefOps.td
index d316fc6c2eef8..deeeda2cb8480 100644
--- a/mlir/include/mlir/Dialect/MemRef/IR/MemRefOps.td
+++ b/mlir/include/mlir/Dialect/MemRef/IR/MemRefOps.td
@@ -1760,15 +1760,7 @@ def MemRef_ReshapeOp: MemRef_Op<"reshape", [
"dynamically-sized shape", [MemRead]>:$shape);
let results = (outs AnyRankedOrUnrankedMemRef:$result);
- let builders = [OpBuilder<
- (ins "MemRefType":$resultType, "Value":$operand, "Value":$shape), [{
- $_state.addOperands(operand);
- $_state.addOperands(shape);
- $_state.addTypes(resultType);
- }]>];
-
let extraClassDeclaration = [{
- MemRefType getType() { return ::llvm::cast<MemRefType>(getResult().getType()); }
Value getViewSource() { return getSource(); }
}];
diff --git a/mlir/lib/Dialect/MemRef/Transforms/ExpandOps.cpp b/mlir/lib/Dialect/MemRef/Transforms/ExpandOps.cpp
index 17a148cc31dc0..de5878407a683 100644
--- a/mlir/lib/Dialect/MemRef/Transforms/ExpandOps.cpp
+++ b/mlir/lib/Dialect/MemRef/Transforms/ExpandOps.cpp
@@ -44,10 +44,11 @@ struct MemRefReshapeOpConverter : public OpRewritePattern<memref::ReshapeOp> {
Location loc = op.getLoc();
Value stride = nullptr;
int64_t staticStride = 1;
+ MemRefType resultType = cast<MemRefType>(op.getType());
for (int i = rank - 1; i >= 0; --i) {
Value size;
// Load dynamic sizes from the shape input, use constants for static dims.
- if (op.getType().isDynamicDim(i)) {
+ if (resultType.isDynamicDim(i)) {
Value index = arith::ConstantIndexOp::create(rewriter, loc, i);
size = memref::LoadOp::create(rewriter, loc, op.getShape(), index);
if (!isa<IndexType>(size.getType()))
@@ -55,7 +56,7 @@ struct MemRefReshapeOpConverter : public OpRewritePattern<memref::ReshapeOp> {
rewriter.getIndexType(), size);
sizes[i] = size;
} else {
- auto sizeAttr = rewriter.getIndexAttr(op.getType().getDimSize(i));
+ auto sizeAttr = rewriter.getIndexAttr(resultType.getDimSize(i));
size = arith::ConstantOp::create(rewriter, loc, sizeAttr);
sizes[i] = sizeAttr;
}
@@ -67,18 +68,18 @@ struct MemRefReshapeOpConverter : public OpRewritePattern<memref::ReshapeOp> {
if (i > 0) {
if (stride) {
stride = arith::MulIOp::create(rewriter, loc, stride, size);
- } else if (op.getType().isDynamicDim(i)) {
+ } else if (resultType.isDynamicDim(i)) {
stride = arith::MulIOp::create(
rewriter, loc,
arith::ConstantIndexOp::create(rewriter, loc, staticStride),
size);
} else {
- staticStride *= op.getType().getDimSize(i);
+ staticStride *= resultType.getDimSize(i);
}
}
}
rewriter.replaceOpWithNewOp<memref::ReinterpretCastOp>(
- op, op.getType(), op.getSource(), /*offset=*/rewriter.getIndexAttr(0),
+ op, resultType, op.getSource(), /*offset=*/rewriter.getIndexAttr(0),
sizes, strides);
return success();
}
diff --git a/mlir/test/Transforms/test-bubble-down-memory-space-casts.mlir b/mlir/test/Transforms/test-bubble-down-memory-space-casts.mlir
index e4fce89cffb45..a03c93b7148e1 100644
--- a/mlir/test/Transforms/test-bubble-down-memory-space-casts.mlir
+++ b/mlir/test/Transforms/test-bubble-down-memory-space-casts.mlir
@@ -107,6 +107,19 @@ func.func @reshape(%arg0: memref<?x?xf32, 1>, %arg1: memref<1xindex>) -> memref<
return %reshape : memref<?xf32>
}
+// CHECK-LABEL: func.func @reshape_unranked(
+// CHECK-SAME: %[[ARG0:.*]]: memref<?x?xf32, 1>,
+// CHECK-SAME: %[[ARG1:.*]]: memref<?xindex>) -> memref<*xf32> {
+// CHECK: %[[RESHAPE_0:.*]] = memref.reshape %[[ARG0]](%[[ARG1]]) : (memref<?x?xf32, 1>, memref<?xindex>) -> memref<*xf32, 1>
+// CHECK: %[[MEMORY_SPACE_CAST_0:.*]] = memref.memory_space_cast %[[RESHAPE_0]] : memref<*xf32, 1> to memref<*xf32>
+// CHECK: return %[[MEMORY_SPACE_CAST_0]] : memref<*xf32>
+// CHECK: }
+func.func @reshape_unranked(%arg0: memref<?x?xf32, 1>, %arg1: memref<?xindex>) -> memref<*xf32> {
+ %memspacecast = memref.memory_space_cast %arg0 : memref<?x?xf32, 1> to memref<?x?xf32>
+ %reshape = memref.reshape %memspacecast(%arg1) : (memref<?x?xf32>, memref<?xindex>) -> memref<*xf32>
+ return %reshape : memref<*xf32>
+}
+
// CHECK-LABEL: func.func @expand_shape(
// CHECK-SAME: %[[ARG0:.*]]: memref<12xf32, 1>) -> memref<3x4xf32> {
// CHECK: %[[VAL_0:.*]] = memref.expand_shape %[[ARG0]] {{\[\[}}0, 1]] output_shape [3, 4] : memref<12xf32, 1> into memref<3x4xf32, 1>
More information about the Mlir-commits
mailing list