[Mlir-commits] [mlir] [mlir][memref] Enforce consistent reinterpret_cast metadata (PR #217338)

ioana ghiban llvmlistbot at llvm.org
Wed Aug 19 08:53:19 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/5] [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/5] 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/5] 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/5] 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/5] 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



More information about the Mlir-commits mailing list