[Mlir-commits] [mlir] [mlir][memref] Enforce consistent reinterpret_cast metadata (PR #217338)
ioana ghiban
llvmlistbot at llvm.org
Thu Aug 20 02:41:09 PDT 2026
https://github.com/ioghiban updated https://github.com/llvm/llvm-project/pull/217338
>From a55d4077e7c891d6cbdade390f7f7da2c67c7cd4 Mon Sep 17 00:00:00 2001
From: Ioana Ghiban <ioana.ghiban at arm.com>
Date: Wed, 19 Aug 2026 14:20:27 +0200
Subject: [PATCH 1/6] [mlir][memref] Enforce consistent reinterpret_cast
metadata
---
.../GPU/Transforms/DecomposeMemRefs.cpp | 30 ++++++++++++-
mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp | 29 +++++++------
.../Transforms/ExpandStridedMetadata.cpp | 42 +++++++++++++++++--
.../MemRefToSPIRV/memref-to-spirv.mlir | 23 +++++-----
mlir/test/Dialect/GPU/decompose-memrefs.mlir | 3 +-
mlir/test/Dialect/MemRef/canonicalize.mlir | 26 ++++++------
.../MemRef/elide-reinterpret-cast-load.mlir | 5 ++-
.../MemRef/expand-strided-metadata.mlir | 23 ++++++++++
mlir/test/Dialect/MemRef/ops.mlir | 16 +++----
9 files changed, 141 insertions(+), 56 deletions(-)
diff --git a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
index 7b30906abc2fd..c2b7749032b1a 100644
--- a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
+++ b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
@@ -27,6 +27,23 @@ namespace mlir {
using namespace mlir;
+static MemRefType updateTypeFromDescriptor(MemRefType type, OpFoldResult offset,
+ ArrayRef<OpFoldResult> sizes,
+ ArrayRef<OpFoldResult> strides) {
+ SmallVector<OpFoldResult> offsets{offset};
+ SmallVector<int64_t> staticOffsets = decomposeMixedValues(offsets).first;
+ SmallVector<int64_t> staticSizes = decomposeMixedValues(sizes).first;
+ SmallVector<int64_t> staticStrides = decomposeMixedValues(strides).first;
+ auto layout = StridedLayoutAttr::get(type.getContext(), staticOffsets.front(),
+ staticStrides);
+ MemRefType updatedType = MemRefType::get(staticSizes, type.getElementType(),
+ layout, type.getMemorySpace());
+ if (!type.getLayout().isIdentity())
+ return updatedType;
+ MemRefType canonicalType = updatedType.canonicalizeStridedLayout();
+ return canonicalType.getLayout().isIdentity() ? canonicalType : updatedType;
+}
+
static MemRefType inferCastResultType(Value source, OpFoldResult offset) {
auto sourceType = cast<BaseMemRefType>(source.getType());
SmallVector<int64_t> staticOffsets;
@@ -212,8 +229,17 @@ struct FlattenSubview : public OpRewritePattern<memref::SubViewOp> {
finalStrides.push_back(strides[i]);
}
- rewriter.replaceOpWithNewOp<memref::ReinterpretCastOp>(
- op, resultType, base, finalOffset, finalSizes, finalStrides);
+ resultType = updateTypeFromDescriptor(resultType, finalOffset, finalSizes,
+ finalStrides);
+ auto reinterpretCast = memref::ReinterpretCastOp::create(
+ rewriter, op.getLoc(), resultType, base, finalOffset, finalSizes,
+ finalStrides);
+ if (resultType == op.getType()) {
+ rewriter.replaceOp(op, reinterpretCast);
+ return success();
+ }
+ rewriter.replaceOpWithNewOp<memref::CastOp>(op, op.getType(),
+ reinterpretCast);
return success();
}
};
diff --git a/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp b/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
index 0ef57172e380c..251845b83936b 100644
--- a/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
+++ b/mlir/lib/Dialect/MemRef/IR/MemRefOps.cpp
@@ -2093,15 +2093,18 @@ LogicalResult ReinterpretCastOp::verify() {
"result")))
return failure();
+ auto printDynamicOrValue = [](int64_t value) {
+ return ShapedType::isDynamic(value) ? std::string("dynamic")
+ : std::to_string(value);
+ };
+
// Match sizes in result memref type and in static_sizes attribute.
for (auto [idx, resultSize, expectedSize] :
llvm::enumerate(resultType.getShape(), getStaticSizes())) {
- if (ShapedType::isStatic(resultSize) && resultSize != expectedSize)
+ if (resultSize != expectedSize)
return emitError("expected result type with size = ")
- << (ShapedType::isDynamic(expectedSize)
- ? std::string("dynamic")
- : std::to_string(expectedSize))
- << " instead of " << resultSize << " in dim = " << idx;
+ << printDynamicOrValue(expectedSize) << " instead of "
+ << printDynamicOrValue(resultSize) << " in dim = " << idx;
}
// Match offset and strides in static_offset and static_strides attributes. If
@@ -2115,22 +2118,18 @@ LogicalResult ReinterpretCastOp::verify() {
// Match offset in result memref type and in static_offsets attribute.
int64_t expectedOffset = getStaticOffsets().front();
- if (ShapedType::isStatic(resultOffset) && resultOffset != expectedOffset)
+ if (resultOffset != expectedOffset)
return emitError("expected result type with offset = ")
- << (ShapedType::isDynamic(expectedOffset)
- ? std::string("dynamic")
- : std::to_string(expectedOffset))
- << " instead of " << resultOffset;
+ << printDynamicOrValue(expectedOffset) << " instead of "
+ << printDynamicOrValue(resultOffset);
// Match strides in result memref type and in static_strides attribute.
for (auto [idx, resultStride, expectedStride] :
llvm::enumerate(resultStrides, getStaticStrides())) {
- if (ShapedType::isStatic(resultStride) && resultStride != expectedStride)
+ if (resultStride != expectedStride)
return emitError("expected result type with stride = ")
- << (ShapedType::isDynamic(expectedStride)
- ? std::string("dynamic")
- : std::to_string(expectedStride))
- << " instead of " << resultStride << " in dim = " << idx;
+ << printDynamicOrValue(expectedStride) << " instead of "
+ << printDynamicOrValue(resultStride) << " in dim = " << idx;
}
return success();
diff --git a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
index 20d543b7210b1..73255fa813ace 100644
--- a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
+++ b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
@@ -45,6 +45,23 @@ struct StridedMetadata {
SmallVector<OpFoldResult> strides;
};
+static MemRefType updateTypeFromDescriptor(MemRefType type,
+ const StridedMetadata &metadata) {
+ SmallVector<OpFoldResult> offsets{metadata.offset};
+ SmallVector<int64_t> staticOffsets = decomposeMixedValues(offsets).first;
+ SmallVector<int64_t> staticSizes = decomposeMixedValues(metadata.sizes).first;
+ SmallVector<int64_t> staticStrides =
+ decomposeMixedValues(metadata.strides).first;
+ auto layout = StridedLayoutAttr::get(type.getContext(), staticOffsets.front(),
+ staticStrides);
+ MemRefType updatedType = MemRefType::get(staticSizes, type.getElementType(),
+ layout, type.getMemorySpace());
+ if (!type.getLayout().isIdentity())
+ return updatedType;
+ MemRefType canonicalType = updatedType.canonicalizeStridedLayout();
+ return canonicalType.getLayout().isIdentity() ? canonicalType : updatedType;
+}
+
/// From `subview(memref, subOffset, subSizes, subStrides))` compute
///
/// \verbatim
@@ -197,10 +214,18 @@ struct SubviewFolder : public OpRewritePattern<memref::SubViewOp> {
"failed to resolve subview metadata");
}
- rewriter.replaceOpWithNewOp<memref::ReinterpretCastOp>(
- subview, subview.getType(), stridedMetadata->basePtr,
+ MemRefType resultType =
+ updateTypeFromDescriptor(subview.getType(), *stridedMetadata);
+ auto reinterpretCast = memref::ReinterpretCastOp::create(
+ rewriter, subview.getLoc(), resultType, stridedMetadata->basePtr,
stridedMetadata->offset, stridedMetadata->sizes,
stridedMetadata->strides);
+ if (resultType == subview.getType()) {
+ rewriter.replaceOp(subview, reinterpretCast);
+ return success();
+ }
+ rewriter.replaceOpWithNewOp<memref::CastOp>(subview, subview.getType(),
+ reinterpretCast);
return success();
}
};
@@ -600,10 +625,19 @@ struct ReshapeFolder : public OpRewritePattern<ReassociativeReshapeLikeOp> {
"failed to resolve reshape metadata");
}
- rewriter.replaceOpWithNewOp<memref::ReinterpretCastOp>(
- reshape, reshape.getType(), stridedMetadata->basePtr,
+ MemRefType resultType = reshape.getResultType();
+ if (isa<memref::CollapseShapeOp>(reshape.getOperation()))
+ resultType = updateTypeFromDescriptor(resultType, *stridedMetadata);
+ auto reinterpretCast = memref::ReinterpretCastOp::create(
+ rewriter, reshape.getLoc(), resultType, stridedMetadata->basePtr,
stridedMetadata->offset, stridedMetadata->sizes,
stridedMetadata->strides);
+ if (resultType == reshape.getResultType()) {
+ rewriter.replaceOp(reshape, reinterpretCast);
+ return success();
+ }
+ rewriter.replaceOpWithNewOp<memref::CastOp>(
+ reshape, reshape.getResultType(), reinterpretCast);
return success();
}
};
diff --git a/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir b/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
index ebddeadf3a31a..827b25c61ccc9 100644
--- a/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
+++ b/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
@@ -488,30 +488,31 @@ func.func @reinterpret_cast(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgr
// CHECK: %[[RET:.*]] = spirv.InBoundsPtrAccessChain %[[MEM1]][%[[OFF1]]] : !spirv.ptr<f32, CrossWorkgroup>, i32
// CHECK: %[[RET1:.*]] = builtin.unrealized_conversion_cast %[[RET]] : !spirv.ptr<f32, CrossWorkgroup> to memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
// CHECK: return %[[RET1]]
- %ret = memref.reinterpret_cast %arg to offset: [%arg1], sizes: [10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
+ %c10 = arith.constant 10 : index
+ %ret = memref.reinterpret_cast %arg to offset: [%arg1], sizes: [%c10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
return %ret : memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
}
// CHECK-LABEL: func.func @reinterpret_cast_0
// CHECK-SAME: (%[[MEM:.*]]: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>)
-func.func @reinterpret_cast_0(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>) -> memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>> {
-// CHECK-DAG: %[[MEM1:.*]] = builtin.unrealized_conversion_cast %[[MEM]] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to !spirv.ptr<f32, CrossWorkgroup>
-// CHECK-DAG: %[[RET:.*]] = builtin.unrealized_conversion_cast %[[MEM1]] : !spirv.ptr<f32, CrossWorkgroup> to memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
-// CHECK: return %[[RET]]
- %ret = memref.reinterpret_cast %arg to offset: [0], sizes: [10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
- return %ret : memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
+func.func @reinterpret_cast_0(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>) -> memref<?xf32, #spirv.storage_class<CrossWorkgroup>> {
+// CHECK: return %[[MEM]]
+ %c10 = arith.constant 10 : index
+ %ret = memref.reinterpret_cast %arg to offset: [0], sizes: [%c10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, #spirv.storage_class<CrossWorkgroup>>
+ return %ret : memref<?xf32, #spirv.storage_class<CrossWorkgroup>>
}
// CHECK-LABEL: func.func @reinterpret_cast_5
// CHECK-SAME: (%[[MEM:.*]]: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>)
-func.func @reinterpret_cast_5(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>) -> memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>> {
+func.func @reinterpret_cast_5(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>) -> memref<?xf32, strided<[1], offset: 5>, #spirv.storage_class<CrossWorkgroup>> {
// CHECK: %[[MEM1:.*]] = builtin.unrealized_conversion_cast %[[MEM]] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to !spirv.ptr<f32, CrossWorkgroup>
// CHECK: %[[OFF:.*]] = spirv.Constant 5 : i32
// CHECK: %[[RET:.*]] = spirv.InBoundsPtrAccessChain %[[MEM1]][%[[OFF]]] : !spirv.ptr<f32, CrossWorkgroup>, i32
-// CHECK: %[[RET1:.*]] = builtin.unrealized_conversion_cast %[[RET]] : !spirv.ptr<f32, CrossWorkgroup> to memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
+// CHECK: %[[RET1:.*]] = builtin.unrealized_conversion_cast %[[RET]] : !spirv.ptr<f32, CrossWorkgroup> to memref<?xf32, strided<[1], offset: 5>, #spirv.storage_class<CrossWorkgroup>>
// CHECK: return %[[RET1]]
- %ret = memref.reinterpret_cast %arg to offset: [5], sizes: [10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
- return %ret : memref<?xf32, strided<[1], offset: ?>, #spirv.storage_class<CrossWorkgroup>>
+ %c10 = arith.constant 10 : index
+ %ret = memref.reinterpret_cast %arg to offset: [5], sizes: [%c10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, strided<[1], offset: 5>, #spirv.storage_class<CrossWorkgroup>>
+ return %ret : memref<?xf32, strided<[1], offset: 5>, #spirv.storage_class<CrossWorkgroup>>
}
} // end module
diff --git a/mlir/test/Dialect/GPU/decompose-memrefs.mlir b/mlir/test/Dialect/GPU/decompose-memrefs.mlir
index 1a19221948451..ac9578d2b3067 100644
--- a/mlir/test/Dialect/GPU/decompose-memrefs.mlir
+++ b/mlir/test/Dialect/GPU/decompose-memrefs.mlir
@@ -88,7 +88,8 @@ func.func @decompose_load(%arg0 : memref<?x?x?xf32>) {
// CHECK-SAME: threads(%[[TX:.*]], %[[TY:.*]], %[[TZ:.*]]) in
// CHECK: %[[IDX:.*]] = affine.apply #[[MAP]]()[%[[TX]], %[[STRIDES]]#0, %[[TY]], %[[STRIDES]]#1, %[[TZ]]]
// CHECK: %[[PTR:.*]] = memref.reinterpret_cast %[[BASE]] to offset: [%[[IDX]]], sizes: [%{{.*}}, %{{.*}}, %{{.*}}], strides: [%[[STRIDES]]#0, %[[STRIDES]]#1, 1]
-// CHECK: "test.test"(%[[PTR]]) : (memref<?x?x?xf32, strided<[?, ?, ?], offset: ?>>) -> ()
+// CHECK: %[[CAST:.*]] = memref.cast %[[PTR]] : memref<?x?x?xf32, strided<[?, ?, 1], offset: ?>> to memref<?x?x?xf32, strided<[?, ?, ?], offset: ?>>
+// CHECK: "test.test"(%[[CAST]]) : (memref<?x?x?xf32, strided<[?, ?, ?], offset: ?>>) -> ()
func.func @decompose_subview(%arg0 : memref<?x?x?xf32>) {
%c0 = arith.constant 0 : index
%c1 = arith.constant 1 : index
diff --git a/mlir/test/Dialect/MemRef/canonicalize.mlir b/mlir/test/Dialect/MemRef/canonicalize.mlir
index 6c4fd6f8f58d6..9fc9d6a30e29c 100644
--- a/mlir/test/Dialect/MemRef/canonicalize.mlir
+++ b/mlir/test/Dialect/MemRef/canonicalize.mlir
@@ -1235,13 +1235,13 @@ func.func @reinterpret_of_extract_strided_metadata_w_type_mistach(%arg0 : memref
// same constant value, the match is valid.
// CHECK-LABEL: func @reinterpret_of_extract_strided_metadata_w_constants
// CHECK-SAME: (%[[ARG:.*]]: memref<8x2xf32>)
-// CHECK: %[[CAST:.*]] = memref.cast %[[ARG]] : memref<8x2xf32> to memref<?x?xf32,
+// CHECK: %[[CAST:.*]] = memref.cast %[[ARG]] : memref<8x2xf32> to memref<?x2xf32,
// CHECK: return %[[CAST]]
-func.func @reinterpret_of_extract_strided_metadata_w_constants(%arg0 : memref<8x2xf32>) -> memref<?x?xf32, strided<[?, ?], offset: ?>> {
+func.func @reinterpret_of_extract_strided_metadata_w_constants(%arg0 : memref<8x2xf32>) -> memref<?x2xf32, strided<[2, ?], offset: 0>> {
%base, %offset, %sizes:2, %strides:2 = memref.extract_strided_metadata %arg0 : memref<8x2xf32> -> memref<f32>, index, index, index, index, index
%c8 = arith.constant 8: index
- %m2 = memref.reinterpret_cast %base to offset: [0], sizes: [%c8, 2], strides: [2, %strides#1] : memref<f32> to memref<?x?xf32, strided<[?, ?], offset: ?>>
- return %m2 : memref<?x?xf32, strided<[?, ?], offset: ?>>
+ %m2 = memref.reinterpret_cast %base to offset: [0], sizes: [%c8, 2], strides: [2, %strides#1] : memref<f32> to memref<?x2xf32, strided<[2, ?], offset: 0>>
+ return %m2 : memref<?x2xf32, strided<[2, ?], offset: 0>>
}
// -----
@@ -1265,10 +1265,10 @@ func.func @reinterpret_of_extract_strided_metadata_same_type(%arg0 : memref<?x?x
// CHECK: %[[RES:.*]] = memref.reinterpret_cast %[[ARG]] to offset: [0], sizes: [4, 2, 2], strides: [1, 1, 1]
// CHECK: %[[CAST:.*]] = memref.cast %[[RES]]
// CHECK: return %[[CAST]]
-func.func @reinterpret_of_extract_strided_metadata_w_different_stride(%arg0 : memref<8x2xf32>) -> memref<?x?x?xf32, strided<[?, ?, ?], offset: ?>> {
+func.func @reinterpret_of_extract_strided_metadata_w_different_stride(%arg0 : memref<8x2xf32>) -> memref<4x2x2xf32, strided<[1, 1, ?], offset: ?>> {
%base, %offset, %sizes:2, %strides:2 = memref.extract_strided_metadata %arg0 : memref<8x2xf32> -> memref<f32>, index, index, index, index, index
- %m2 = memref.reinterpret_cast %base to offset: [%offset], sizes: [4, 2, 2], strides: [1, 1, %strides#1] : memref<f32> to memref<?x?x?xf32, strided<[?, ?, ?], offset: ?>>
- return %m2 : memref<?x?x?xf32, strided<[?, ?, ?], offset: ?>>
+ %m2 = memref.reinterpret_cast %base to offset: [%offset], sizes: [4, 2, 2], strides: [1, 1, %strides#1] : memref<f32> to memref<4x2x2xf32, strided<[1, 1, ?], offset: ?>>
+ return %m2 : memref<4x2x2xf32, strided<[1, 1, ?], offset: ?>>
}
// -----
@@ -1279,10 +1279,10 @@ func.func @reinterpret_of_extract_strided_metadata_w_different_stride(%arg0 : me
// CHECK: %[[RES:.*]] = memref.reinterpret_cast %[[ARG]] to offset: [1], sizes: [8, 2], strides: [2, 1]
// CHECK: %[[CAST:.*]] = memref.cast %[[RES]]
// CHECK: return %[[CAST]]
-func.func @reinterpret_of_extract_strided_metadata_w_different_offset(%arg0 : memref<8x2xf32>) -> memref<?x?xf32, strided<[?, ?], offset: ?>> {
+func.func @reinterpret_of_extract_strided_metadata_w_different_offset(%arg0 : memref<8x2xf32>) -> memref<?x?xf32, strided<[?, ?], offset: 1>> {
%base, %offset, %sizes:2, %strides:2 = memref.extract_strided_metadata %arg0 : memref<8x2xf32> -> memref<f32>, index, index, index, index, index
- %m2 = memref.reinterpret_cast %base to offset: [1], sizes: [%sizes#0, %sizes#1], strides: [%strides#0, %strides#1] : memref<f32> to memref<?x?xf32, strided<[?, ?], offset: ?>>
- return %m2 : memref<?x?xf32, strided<[?, ?], offset: ?>>
+ %m2 = memref.reinterpret_cast %base to offset: [1], sizes: [%sizes#0, %sizes#1], strides: [%strides#0, %strides#1] : memref<f32> to memref<?x?xf32, strided<[?, ?], offset: 1>>
+ return %m2 : memref<?x?xf32, strided<[?, ?], offset: 1>>
}
// -----
@@ -1348,12 +1348,12 @@ func.func @reinterpret_cast_with_negative_size_and_offset(%arg0: memref<2x3xf32>
// CHECK-SAME: (%[[ARG:.*]]: memref<2x3xf32>)
// CHECK: %[[NEG:.*]] = arith.constant -1 : index
// CHECK: memref.reinterpret_cast %[[ARG]] to offset: [%[[NEG]]], sizes: [%[[NEG]], %[[NEG]]], strides: [2, 1]
-func.func @reinterpret_cast_no_fold_with_all_negative_size_and_offset(%arg0: memref<2x3xf32>) -> memref<?x?xf32, strided<[?, ?], offset: ?>> {
+func.func @reinterpret_cast_no_fold_with_all_negative_size_and_offset(%arg0: memref<2x3xf32>) -> memref<?x?xf32, strided<[2, 1], offset: ?>> {
%neg = arith.constant -1 : index
%output = memref.reinterpret_cast %arg0 to
offset: [%neg], sizes: [%neg, %neg], strides: [2, 1]
- : memref<2x3xf32> to memref<?x?xf32, strided<[?, ?], offset: ?>>
- return %output : memref<?x?xf32, strided<[?, ?], offset: ?>>
+ : memref<2x3xf32> to memref<?x?xf32, strided<[2, 1], offset: ?>>
+ return %output : memref<?x?xf32, strided<[2, 1], offset: ?>>
}
// -----
diff --git a/mlir/test/Dialect/MemRef/elide-reinterpret-cast-load.mlir b/mlir/test/Dialect/MemRef/elide-reinterpret-cast-load.mlir
index ce505485c5715..e262072f543b9 100644
--- a/mlir/test/Dialect/MemRef/elide-reinterpret-cast-load.mlir
+++ b/mlir/test/Dialect/MemRef/elide-reinterpret-cast-load.mlir
@@ -281,9 +281,10 @@ func.func private @negative_dynamic_shape(%dim : index,
// CHECK: %[[RC:.*]] = memref.reinterpret_cast %[[SRC]]
%reinterpret_cast = memref.reinterpret_cast %src
to offset: [0], sizes: [1, %dim], strides: [1, 1]
- : memref<?xf32> to memref<1x?xf32>
+ : memref<?xf32> to memref<1x?xf32, strided<[1, 1]>>
// CHECK: memref.load %[[RC]]
- %0 = memref.load %reinterpret_cast[%idx_1, %idx_2] : memref<1x?xf32>
+ %0 = memref.load %reinterpret_cast[%idx_1, %idx_2] :
+ memref<1x?xf32, strided<[1, 1]>>
return
}
diff --git a/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir b/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir
index 70c5e1aee85dc..ded9a4ae0e111 100644
--- a/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir
+++ b/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir
@@ -72,6 +72,29 @@ func.func @simplify_subview_all_dynamic(
// -----
+// Check that constant folding of the descriptor is reflected in the result
+// type of the reinterpret_cast. The subview sizes remain dynamic because they
+// are SSA operands, while its offset and strides fold to static values.
+// CHECK-LABEL: func @simplify_subview_folded_descriptor
+// CHECK: %[[CAST:.*]] = memref.reinterpret_cast %{{.*}} to offset: [10],
+// CHECK-SAME: sizes: [%{{.*}}, %{{.*}}], strides: [8, 2]
+// CHECK-SAME: memref<f32> to memref<?x?xf32, strided<[8, 2], offset: 10>>
+// CHECK: memref.load %[[CAST]]
+// CHECK-SAME: memref<?x?xf32, strided<[8, 2], offset: 10>>
+func.func @simplify_subview_folded_descriptor(%base: memref<4x8xf32>) -> f32 {
+ %c0 = arith.constant 0 : index
+ %c1 = arith.constant 1 : index
+ %c2 = arith.constant 2 : index
+ %c3 = arith.constant 3 : index
+ %subview = memref.subview %base[%c1, %c2] [%c2, %c3] [%c1, %c2] :
+ memref<4x8xf32> to memref<?x?xf32, strided<[?, ?], offset: ?>>
+ %result = memref.load %subview[%c0, %c0] :
+ memref<?x?xf32, strided<[?, ?], offset: ?>>
+ return %result : f32
+}
+
+// -----
+
// Check that we simplify extract_strided_metadata of subview to
// base_buf, base_offset, base_sizes, base_strides = extract_strided_metadata
// strides = base_stride_i * subview_stride_i
diff --git a/mlir/test/Dialect/MemRef/ops.mlir b/mlir/test/Dialect/MemRef/ops.mlir
index 0434d657822d6..1f13d028af4d9 100644
--- a/mlir/test/Dialect/MemRef/ops.mlir
+++ b/mlir/test/Dialect/MemRef/ops.mlir
@@ -130,21 +130,21 @@ func.func @memref_reinterpret_cast(%in: memref<?xf32>)
}
// CHECK-LABEL: func @memref_reinterpret_cast_static_to_dynamic_sizes
-func.func @memref_reinterpret_cast_static_to_dynamic_sizes(%in: memref<?xf32>)
- -> memref<10x?xf32, strided<[?, 1], offset: ?>> {
+func.func @memref_reinterpret_cast_static_to_dynamic_sizes(%in: memref<?xf32>, %size: index)
+ -> memref<10x?xf32, strided<[1, 1], offset: 1>> {
%out = memref.reinterpret_cast %in to
- offset: [1], sizes: [10, 10], strides: [1, 1]
- : memref<?xf32> to memref<10x?xf32, strided<[?, 1], offset: ?>>
- return %out : memref<10x?xf32, strided<[?, 1], offset: ?>>
+ offset: [1], sizes: [10, %size], strides: [1, 1]
+ : memref<?xf32> to memref<10x?xf32, strided<[1, 1], offset: 1>>
+ return %out : memref<10x?xf32, strided<[1, 1], offset: 1>>
}
// CHECK-LABEL: func @memref_reinterpret_cast_dynamic_offset
func.func @memref_reinterpret_cast_dynamic_offset(%in: memref<?xf32>, %offset: index)
- -> memref<10x?xf32, strided<[?, 1], offset: ?>> {
+ -> memref<10x10xf32, strided<[1, 1], offset: ?>> {
%out = memref.reinterpret_cast %in to
offset: [%offset], sizes: [10, 10], strides: [1, 1]
- : memref<?xf32> to memref<10x?xf32, strided<[?, 1], offset: ?>>
- return %out : memref<10x?xf32, strided<[?, 1], offset: ?>>
+ : memref<?xf32> to memref<10x10xf32, strided<[1, 1], offset: ?>>
+ return %out : memref<10x10xf32, strided<[1, 1], offset: ?>>
}
// CHECK-LABEL: func @memref_reshape(
>From b6cdb855042f7dd11948af5baf401653e7cb9806 Mon Sep 17 00:00:00 2001
From: Ioana Ghiban <ioana.ghiban at arm.com>
Date: Wed, 19 Aug 2026 16:17:25 +0200
Subject: [PATCH 2/6] Add comments
---
mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp | 1 +
mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp | 2 ++
2 files changed, 3 insertions(+)
diff --git a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
index c2b7749032b1a..053b88340ebed 100644
--- a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
+++ b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
@@ -238,6 +238,7 @@ struct FlattenSubview : public OpRewritePattern<memref::SubViewOp> {
rewriter.replaceOp(op, reinterpretCast);
return success();
}
+ // Preserve the original result type expected by existing users.
rewriter.replaceOpWithNewOp<memref::CastOp>(op, op.getType(),
reinterpretCast);
return success();
diff --git a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
index 73255fa813ace..eebdf7cb980cb 100644
--- a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
+++ b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
@@ -224,6 +224,7 @@ struct SubviewFolder : public OpRewritePattern<memref::SubViewOp> {
rewriter.replaceOp(subview, reinterpretCast);
return success();
}
+ // Preserve the original result type expected by existing users.
rewriter.replaceOpWithNewOp<memref::CastOp>(subview, subview.getType(),
reinterpretCast);
return success();
@@ -636,6 +637,7 @@ struct ReshapeFolder : public OpRewritePattern<ReassociativeReshapeLikeOp> {
rewriter.replaceOp(reshape, reinterpretCast);
return success();
}
+ // Preserve the original result type expected by existing users.
rewriter.replaceOpWithNewOp<memref::CastOp>(
reshape, reshape.getResultType(), reinterpretCast);
return success();
>From 93c2f348868f97f4e6d0d98beec51697c8e6bc15 Mon Sep 17 00:00:00 2001
From: Ioana Ghiban <ioana.ghiban at arm.com>
Date: Wed, 19 Aug 2026 16:53:23 +0200
Subject: [PATCH 3/6] More suggestive variable naming
---
mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp | 6 +++---
.../MemRef/Transforms/ExpandStridedMetadata.cpp | 12 ++++++------
2 files changed, 9 insertions(+), 9 deletions(-)
diff --git a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
index 053b88340ebed..479e84a438ae0 100644
--- a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
+++ b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
@@ -231,16 +231,16 @@ struct FlattenSubview : public OpRewritePattern<memref::SubViewOp> {
resultType = updateTypeFromDescriptor(resultType, finalOffset, finalSizes,
finalStrides);
- auto reinterpretCast = memref::ReinterpretCastOp::create(
+ auto flattenedSubview = memref::ReinterpretCastOp::create(
rewriter, op.getLoc(), resultType, base, finalOffset, finalSizes,
finalStrides);
if (resultType == op.getType()) {
- rewriter.replaceOp(op, reinterpretCast);
+ rewriter.replaceOp(op, flattenedSubview);
return success();
}
// Preserve the original result type expected by existing users.
rewriter.replaceOpWithNewOp<memref::CastOp>(op, op.getType(),
- reinterpretCast);
+ flattenedSubview);
return success();
}
};
diff --git a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
index eebdf7cb980cb..5548fd9983d51 100644
--- a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
+++ b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
@@ -216,17 +216,17 @@ struct SubviewFolder : public OpRewritePattern<memref::SubViewOp> {
MemRefType resultType =
updateTypeFromDescriptor(subview.getType(), *stridedMetadata);
- auto reinterpretCast = memref::ReinterpretCastOp::create(
+ auto foldedSubview = memref::ReinterpretCastOp::create(
rewriter, subview.getLoc(), resultType, stridedMetadata->basePtr,
stridedMetadata->offset, stridedMetadata->sizes,
stridedMetadata->strides);
if (resultType == subview.getType()) {
- rewriter.replaceOp(subview, reinterpretCast);
+ rewriter.replaceOp(subview, foldedSubview);
return success();
}
// Preserve the original result type expected by existing users.
rewriter.replaceOpWithNewOp<memref::CastOp>(subview, subview.getType(),
- reinterpretCast);
+ foldedSubview);
return success();
}
};
@@ -629,17 +629,17 @@ struct ReshapeFolder : public OpRewritePattern<ReassociativeReshapeLikeOp> {
MemRefType resultType = reshape.getResultType();
if (isa<memref::CollapseShapeOp>(reshape.getOperation()))
resultType = updateTypeFromDescriptor(resultType, *stridedMetadata);
- auto reinterpretCast = memref::ReinterpretCastOp::create(
+ auto foldedReshape = memref::ReinterpretCastOp::create(
rewriter, reshape.getLoc(), resultType, stridedMetadata->basePtr,
stridedMetadata->offset, stridedMetadata->sizes,
stridedMetadata->strides);
if (resultType == reshape.getResultType()) {
- rewriter.replaceOp(reshape, reinterpretCast);
+ rewriter.replaceOp(reshape, foldedReshape);
return success();
}
// Preserve the original result type expected by existing users.
rewriter.replaceOpWithNewOp<memref::CastOp>(
- reshape, reshape.getResultType(), reinterpretCast);
+ reshape, reshape.getResultType(), foldedReshape);
return success();
}
};
>From 22769e2fff146cb4e53a40d029339dea90c7f0a2 Mon Sep 17 00:00:00 2001
From: Ioana Ghiban <ioana.ghiban at arm.com>
Date: Wed, 19 Aug 2026 17:06:27 +0200
Subject: [PATCH 4/6] Fix memref-to-spirv test
---
.../test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir | 10 ++++++----
1 file changed, 6 insertions(+), 4 deletions(-)
diff --git a/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir b/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
index 827b25c61ccc9..3b6925f538181 100644
--- a/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
+++ b/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
@@ -495,11 +495,13 @@ func.func @reinterpret_cast(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgr
// CHECK-LABEL: func.func @reinterpret_cast_0
// CHECK-SAME: (%[[MEM:.*]]: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>)
-func.func @reinterpret_cast_0(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>) -> memref<?xf32, #spirv.storage_class<CrossWorkgroup>> {
-// CHECK: return %[[MEM]]
+func.func @reinterpret_cast_0(%arg: memref<?xf32, #spirv.storage_class<CrossWorkgroup>>) -> memref<?xf32, strided<[1]>, #spirv.storage_class<CrossWorkgroup>> {
+// CHECK-DAG: %[[MEM1:.*]] = builtin.unrealized_conversion_cast %[[MEM]] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to !spirv.ptr<f32, CrossWorkgroup>
+// CHECK-DAG: %[[RET:.*]] = builtin.unrealized_conversion_cast %[[MEM1]] : !spirv.ptr<f32, CrossWorkgroup> to memref<?xf32, strided<[1]>, #spirv.storage_class<CrossWorkgroup>>
+// CHECK: return %[[RET]]
%c10 = arith.constant 10 : index
- %ret = memref.reinterpret_cast %arg to offset: [0], sizes: [%c10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, #spirv.storage_class<CrossWorkgroup>>
- return %ret : memref<?xf32, #spirv.storage_class<CrossWorkgroup>>
+ %ret = memref.reinterpret_cast %arg to offset: [0], sizes: [%c10], strides: [1] : memref<?xf32, #spirv.storage_class<CrossWorkgroup>> to memref<?xf32, strided<[1]>, #spirv.storage_class<CrossWorkgroup>>
+ return %ret : memref<?xf32, strided<[1]>, #spirv.storage_class<CrossWorkgroup>>
}
// CHECK-LABEL: func.func @reinterpret_cast_5
>From e6fa7629558fad2dea5b0b66965d11be43b8010f Mon Sep 17 00:00:00 2001
From: Ioana Ghiban <ioana.ghiban at arm.com>
Date: Wed, 19 Aug 2026 17:41:05 +0200
Subject: [PATCH 5/6] Move helper function into MemRefUtils
---
.../mlir/Dialect/MemRef/Utils/MemRefUtils.h | 7 +++++
mlir/lib/Dialect/GPU/CMakeLists.txt | 1 +
.../GPU/Transforms/DecomposeMemRefs.cpp | 22 +++------------
.../Transforms/ExpandStridedMetadata.cpp | 27 +++++--------------
mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp | 19 +++++++++++++
5 files changed, 37 insertions(+), 39 deletions(-)
diff --git a/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h b/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h
index c93da6ebada7c..e0b84ce7620b1 100644
--- a/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h
+++ b/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h
@@ -198,6 +198,13 @@ LogicalResult resolveSourceIndicesRankReducingSubview(
/// negative.
bool hasNegativeStaticStride(MemRefType memRefTy);
+/// Returns a memref type matching the descriptor's offset, sizes, and strides.
+/// Static attributes become static type metadata; SSA values remain dynamic.
+/// The caller remains responsible for ensuring the descriptor is semantically
+/// correct.
+MemRefType updateTypeFromDescriptor(MemRefType type, OpFoldResult offset,
+ ArrayRef<OpFoldResult> sizes,
+ ArrayRef<OpFoldResult> strides);
} // namespace memref
} // namespace mlir
diff --git a/mlir/lib/Dialect/GPU/CMakeLists.txt b/mlir/lib/Dialect/GPU/CMakeLists.txt
index 547812da0ab97..b5262d1213e9b 100644
--- a/mlir/lib/Dialect/GPU/CMakeLists.txt
+++ b/mlir/lib/Dialect/GPU/CMakeLists.txt
@@ -70,6 +70,7 @@ add_mlir_dialect_library(MLIRGPUTransforms
MLIRIndexDialect
MLIRLLVMDialect
MLIRMemRefDialect
+ MLIRMemRefUtils
MLIRNVVMTarget
MLIRPass
MLIRROCDLDialect
diff --git a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
index 479e84a438ae0..37bad5fa489f2 100644
--- a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
+++ b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
@@ -14,6 +14,7 @@
#include "mlir/Dialect/GPU/IR/GPUDialect.h"
#include "mlir/Dialect/GPU/Transforms/Passes.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
+#include "mlir/Dialect/MemRef/Utils/MemRefUtils.h"
#include "mlir/Dialect/Utils/IndexingUtils.h"
#include "mlir/IR/AffineExpr.h"
#include "mlir/IR/Builders.h"
@@ -27,23 +28,6 @@ namespace mlir {
using namespace mlir;
-static MemRefType updateTypeFromDescriptor(MemRefType type, OpFoldResult offset,
- ArrayRef<OpFoldResult> sizes,
- ArrayRef<OpFoldResult> strides) {
- SmallVector<OpFoldResult> offsets{offset};
- SmallVector<int64_t> staticOffsets = decomposeMixedValues(offsets).first;
- SmallVector<int64_t> staticSizes = decomposeMixedValues(sizes).first;
- SmallVector<int64_t> staticStrides = decomposeMixedValues(strides).first;
- auto layout = StridedLayoutAttr::get(type.getContext(), staticOffsets.front(),
- staticStrides);
- MemRefType updatedType = MemRefType::get(staticSizes, type.getElementType(),
- layout, type.getMemorySpace());
- if (!type.getLayout().isIdentity())
- return updatedType;
- MemRefType canonicalType = updatedType.canonicalizeStridedLayout();
- return canonicalType.getLayout().isIdentity() ? canonicalType : updatedType;
-}
-
static MemRefType inferCastResultType(Value source, OpFoldResult offset) {
auto sourceType = cast<BaseMemRefType>(source.getType());
SmallVector<int64_t> staticOffsets;
@@ -229,8 +213,8 @@ struct FlattenSubview : public OpRewritePattern<memref::SubViewOp> {
finalStrides.push_back(strides[i]);
}
- resultType = updateTypeFromDescriptor(resultType, finalOffset, finalSizes,
- finalStrides);
+ resultType = memref::updateTypeFromDescriptor(resultType, finalOffset,
+ finalSizes, finalStrides);
auto flattenedSubview = memref::ReinterpretCastOp::create(
rewriter, op.getLoc(), resultType, base, finalOffset, finalSizes,
finalStrides);
diff --git a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
index 5548fd9983d51..53a16b72a4a83 100644
--- a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
+++ b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
@@ -18,6 +18,7 @@
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/MemRef/Transforms/Passes.h"
#include "mlir/Dialect/MemRef/Transforms/Transforms.h"
+#include "mlir/Dialect/MemRef/Utils/MemRefUtils.h"
#include "mlir/Dialect/Utils/IndexingUtils.h"
#include "mlir/IR/AffineMap.h"
#include "mlir/IR/BuiltinTypes.h"
@@ -45,23 +46,6 @@ struct StridedMetadata {
SmallVector<OpFoldResult> strides;
};
-static MemRefType updateTypeFromDescriptor(MemRefType type,
- const StridedMetadata &metadata) {
- SmallVector<OpFoldResult> offsets{metadata.offset};
- SmallVector<int64_t> staticOffsets = decomposeMixedValues(offsets).first;
- SmallVector<int64_t> staticSizes = decomposeMixedValues(metadata.sizes).first;
- SmallVector<int64_t> staticStrides =
- decomposeMixedValues(metadata.strides).first;
- auto layout = StridedLayoutAttr::get(type.getContext(), staticOffsets.front(),
- staticStrides);
- MemRefType updatedType = MemRefType::get(staticSizes, type.getElementType(),
- layout, type.getMemorySpace());
- if (!type.getLayout().isIdentity())
- return updatedType;
- MemRefType canonicalType = updatedType.canonicalizeStridedLayout();
- return canonicalType.getLayout().isIdentity() ? canonicalType : updatedType;
-}
-
/// From `subview(memref, subOffset, subSizes, subStrides))` compute
///
/// \verbatim
@@ -214,8 +198,9 @@ struct SubviewFolder : public OpRewritePattern<memref::SubViewOp> {
"failed to resolve subview metadata");
}
- MemRefType resultType =
- updateTypeFromDescriptor(subview.getType(), *stridedMetadata);
+ MemRefType resultType = memref::updateTypeFromDescriptor(
+ subview.getType(), stridedMetadata->offset, stridedMetadata->sizes,
+ stridedMetadata->strides);
auto foldedSubview = memref::ReinterpretCastOp::create(
rewriter, subview.getLoc(), resultType, stridedMetadata->basePtr,
stridedMetadata->offset, stridedMetadata->sizes,
@@ -628,7 +613,9 @@ struct ReshapeFolder : public OpRewritePattern<ReassociativeReshapeLikeOp> {
MemRefType resultType = reshape.getResultType();
if (isa<memref::CollapseShapeOp>(reshape.getOperation()))
- resultType = updateTypeFromDescriptor(resultType, *stridedMetadata);
+ resultType = memref::updateTypeFromDescriptor(
+ resultType, stridedMetadata->offset, stridedMetadata->sizes,
+ stridedMetadata->strides);
auto foldedReshape = memref::ReinterpretCastOp::create(
rewriter, reshape.getLoc(), resultType, stridedMetadata->basePtr,
stridedMetadata->offset, stridedMetadata->sizes,
diff --git a/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp b/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp
index bc019d601dcd9..d1b52c2b0dee1 100644
--- a/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp
+++ b/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp
@@ -351,5 +351,24 @@ bool hasNegativeStaticStride(MemRefType memRefTy) {
});
}
+MemRefType updateTypeFromDescriptor(MemRefType type, OpFoldResult offset,
+ ArrayRef<OpFoldResult> sizes,
+ ArrayRef<OpFoldResult> strides) {
+ SmallVector<OpFoldResult> offsets{offset};
+ SmallVector<int64_t> staticOffsets = decomposeMixedValues(offsets).first;
+ SmallVector<int64_t> staticSizes = decomposeMixedValues(sizes).first;
+ SmallVector<int64_t> staticStrides = decomposeMixedValues(strides).first;
+ auto layout = StridedLayoutAttr::get(type.getContext(), staticOffsets.front(),
+ staticStrides);
+ // Build a MemRefType using the original element type and memory space,
+ // but derive its shape, offset, and strides from the supplied descriptor.
+ MemRefType updatedType = MemRefType::get(staticSizes, type.getElementType(),
+ layout, type.getMemorySpace());
+ if (!type.getLayout().isIdentity())
+ return updatedType;
+ MemRefType canonicalType = updatedType.canonicalizeStridedLayout();
+ return canonicalType.getLayout().isIdentity() ? canonicalType : updatedType;
+}
+
} // namespace memref
} // namespace mlir
>From e56e398f546281aedccb8dd68c2605f13385ba51 Mon Sep 17 00:00:00 2001
From: Ioana Ghiban <ioana.ghiban at arm.com>
Date: Thu, 20 Aug 2026 11:40:31 +0200
Subject: [PATCH 6/6] Address first round of comments
---
.../mlir/Dialect/MemRef/Utils/MemRefUtils.h | 14 +++++++-------
.../Dialect/GPU/Transforms/DecomposeMemRefs.cpp | 4 ++--
.../MemRef/Transforms/ExpandStridedMetadata.cpp | 4 ++--
mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp | 8 ++++----
.../Dialect/MemRef/expand-strided-metadata.mlir | 16 +++++++++-------
5 files changed, 24 insertions(+), 22 deletions(-)
diff --git a/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h b/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h
index e0b84ce7620b1..471986dd52d2d 100644
--- a/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h
+++ b/mlir/include/mlir/Dialect/MemRef/Utils/MemRefUtils.h
@@ -198,13 +198,13 @@ LogicalResult resolveSourceIndicesRankReducingSubview(
/// negative.
bool hasNegativeStaticStride(MemRefType memRefTy);
-/// Returns a memref type matching the descriptor's offset, sizes, and strides.
-/// Static attributes become static type metadata; SSA values remain dynamic.
-/// The caller remains responsible for ensuring the descriptor is semantically
-/// correct.
-MemRefType updateTypeFromDescriptor(MemRefType type, OpFoldResult offset,
- ArrayRef<OpFoldResult> sizes,
- ArrayRef<OpFoldResult> strides);
+/// Return a memref type matching the provided offset, size, and stride
+/// metadata. Static attributes become static type metadata; SSA values remain
+/// dynamic. The caller remains responsible for ensuring the metadata is
+/// semantically correct.
+MemRefType updateTypeFromMetadata(MemRefType type, OpFoldResult offset,
+ ArrayRef<OpFoldResult> sizes,
+ ArrayRef<OpFoldResult> strides);
} // namespace memref
} // namespace mlir
diff --git a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
index 37bad5fa489f2..bdaffd5563d2e 100644
--- a/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
+++ b/mlir/lib/Dialect/GPU/Transforms/DecomposeMemRefs.cpp
@@ -213,8 +213,8 @@ struct FlattenSubview : public OpRewritePattern<memref::SubViewOp> {
finalStrides.push_back(strides[i]);
}
- resultType = memref::updateTypeFromDescriptor(resultType, finalOffset,
- finalSizes, finalStrides);
+ resultType = memref::updateTypeFromMetadata(resultType, finalOffset,
+ finalSizes, finalStrides);
auto flattenedSubview = memref::ReinterpretCastOp::create(
rewriter, op.getLoc(), resultType, base, finalOffset, finalSizes,
finalStrides);
diff --git a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
index 53a16b72a4a83..dd1434a80e6d0 100644
--- a/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
+++ b/mlir/lib/Dialect/MemRef/Transforms/ExpandStridedMetadata.cpp
@@ -198,7 +198,7 @@ struct SubviewFolder : public OpRewritePattern<memref::SubViewOp> {
"failed to resolve subview metadata");
}
- MemRefType resultType = memref::updateTypeFromDescriptor(
+ MemRefType resultType = memref::updateTypeFromMetadata(
subview.getType(), stridedMetadata->offset, stridedMetadata->sizes,
stridedMetadata->strides);
auto foldedSubview = memref::ReinterpretCastOp::create(
@@ -613,7 +613,7 @@ struct ReshapeFolder : public OpRewritePattern<ReassociativeReshapeLikeOp> {
MemRefType resultType = reshape.getResultType();
if (isa<memref::CollapseShapeOp>(reshape.getOperation()))
- resultType = memref::updateTypeFromDescriptor(
+ resultType = memref::updateTypeFromMetadata(
resultType, stridedMetadata->offset, stridedMetadata->sizes,
stridedMetadata->strides);
auto foldedReshape = memref::ReinterpretCastOp::create(
diff --git a/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp b/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp
index d1b52c2b0dee1..0653313e2f0fe 100644
--- a/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp
+++ b/mlir/lib/Dialect/MemRef/Utils/MemRefUtils.cpp
@@ -351,9 +351,9 @@ bool hasNegativeStaticStride(MemRefType memRefTy) {
});
}
-MemRefType updateTypeFromDescriptor(MemRefType type, OpFoldResult offset,
- ArrayRef<OpFoldResult> sizes,
- ArrayRef<OpFoldResult> strides) {
+MemRefType updateTypeFromMetadata(MemRefType type, OpFoldResult offset,
+ ArrayRef<OpFoldResult> sizes,
+ ArrayRef<OpFoldResult> strides) {
SmallVector<OpFoldResult> offsets{offset};
SmallVector<int64_t> staticOffsets = decomposeMixedValues(offsets).first;
SmallVector<int64_t> staticSizes = decomposeMixedValues(sizes).first;
@@ -361,7 +361,7 @@ MemRefType updateTypeFromDescriptor(MemRefType type, OpFoldResult offset,
auto layout = StridedLayoutAttr::get(type.getContext(), staticOffsets.front(),
staticStrides);
// Build a MemRefType using the original element type and memory space,
- // but derive its shape, offset, and strides from the supplied descriptor.
+ // but derive its shape, offset, and strides from the supplied metadata.
MemRefType updatedType = MemRefType::get(staticSizes, type.getElementType(),
layout, type.getMemorySpace());
if (!type.getLayout().isIdentity())
diff --git a/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir b/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir
index ded9a4ae0e111..85ef1b9b515c6 100644
--- a/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir
+++ b/mlir/test/Dialect/MemRef/expand-strided-metadata.mlir
@@ -72,16 +72,18 @@ func.func @simplify_subview_all_dynamic(
// -----
-// Check that constant folding of the descriptor is reflected in the result
-// type of the reinterpret_cast. The subview sizes remain dynamic because they
-// are SSA operands, while its offset and strides fold to static values.
-// CHECK-LABEL: func @simplify_subview_folded_descriptor
-// CHECK: %[[CAST:.*]] = memref.reinterpret_cast %{{.*}} to offset: [10],
-// CHECK-SAME: sizes: [%{{.*}}, %{{.*}}], strides: [8, 2]
+// Check that folded offsets and strides are reflected in the result type of
+// the reinterpret_cast. Sizes remain dynamic because they are forwarded
+// without folding.
+// CHECK-LABEL: func @simplify_subview_folded_metadata
+// CHECK: %[[CAST:.*]] = memref.reinterpret_cast %{{[^ ]+}} to
+// CHECK-SAME: offset: [10],
+// CHECK-SAME: sizes: [%{{[^ ]+}}, %{{[^ ]+}}],
+// CHECK-SAME: strides: [8, 2]
// CHECK-SAME: memref<f32> to memref<?x?xf32, strided<[8, 2], offset: 10>>
// CHECK: memref.load %[[CAST]]
// CHECK-SAME: memref<?x?xf32, strided<[8, 2], offset: 10>>
-func.func @simplify_subview_folded_descriptor(%base: memref<4x8xf32>) -> f32 {
+func.func @simplify_subview_folded_metadata(%base: memref<4x8xf32>) -> f32 {
%c0 = arith.constant 0 : index
%c1 = arith.constant 1 : index
%c2 = arith.constant 2 : index
More information about the Mlir-commits
mailing list