[flang-commits] [flang] db96dfa - [flang][FIRToMemRef] Treat heap-pointer allocas as static, not as dynamic arrays (#223821)

via flang-commits flang-commits at lists.llvm.org
Wed Sep 16 09:47:15 PDT 2026


Author: Susan Tan (ス-ザン タン)
Date: 2026-09-16T16:47:09Z
New Revision: db96dfae2a4fd74cb1dba7b4f9dec2dba4c83864

URL: https://github.com/llvm/llvm-project/commit/db96dfae2a4fd74cb1dba7b4f9dec2dba4c83864
DIFF: https://github.com/llvm/llvm-project/commit/db96dfae2a4fd74cb1dba7b4f9dec2dba4c83864.diff

LOG: [flang][FIRToMemRef] Treat heap-pointer allocas as static, not as dynamic arrays (#223821)

Whole-array assignment of an allocatable inside !$acc kernels with
-Mstack_arrays aborted on !fir.ref<!fir.heap<!fir.array<?xf32>>>. A
fir.alloca !fir.heap<array> is a local heap pointer, but it was
unwrapped and handled as a dynamic array. The memref converter then
peeled only the outer ref and asserted because !fir.heap is not a memref
element.

Keep the pointer as a static allocation. Make convertibleMemrefType use
the same one-pointer peel as convertMemrefType so a nested pointer is
not treated as an f32 array.

Added: 
    

Modified: 
    flang/include/flang/Optimizer/Transforms/FIRToMemRefTypeConverter.h
    flang/unittests/Optimizer/OpenACC/FIROpenACCPointerLikeTypeInterfaceTest.cpp

Removed: 
    


################################################################################
diff  --git a/flang/include/flang/Optimizer/Transforms/FIRToMemRefTypeConverter.h b/flang/include/flang/Optimizer/Transforms/FIRToMemRefTypeConverter.h
index fd434b1f09c9b..9fb509b3d0afa 100644
--- a/flang/include/flang/Optimizer/Transforms/FIRToMemRefTypeConverter.h
+++ b/flang/include/flang/Optimizer/Transforms/FIRToMemRefTypeConverter.h
@@ -32,6 +32,36 @@ class FIRToMemRefTypeConverter : public mlir::TypeConverter {
   bool convertComplexTypes = false;
   bool convertScalarTypesOnly = false;
 
+  mlir::MemRefType convertMemrefBaseType(mlir::Type baseTy) const {
+    if (auto charTy = mlir::dyn_cast<fir::CharacterType>(baseTy)) {
+      unsigned kind = charTy.getFKind();
+      unsigned bitWidth = kindMapping.getCharacterBitsize(kind);
+      mlir::Type elTy = mlir::IntegerType::get(charTy.getContext(), bitWidth);
+
+      if (charTy.hasConstantLen() && charTy.getLen() == 1)
+        return mlir::MemRefType::get({}, elTy);
+      if (charTy.hasConstantLen())
+        return mlir::MemRefType::get({charTy.getLen()}, elTy);
+      return mlir::MemRefType::get({mlir::ShapedType::kDynamic}, elTy);
+    }
+
+    if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(baseTy)) {
+      mlir::Type ty = convertType(seqTy.getElementType());
+      llvm::ArrayRef<int64_t> firShape = seqTy.getShape();
+      llvm::SmallVector<int64_t> shape;
+      for (auto it = firShape.rbegin(); it != firShape.rend(); ++it)
+        shape.push_back(*it);
+      assert(mlir::BaseMemRefType::isValidElementType(ty) &&
+             "got invalid memref element type from array fir type");
+      return mlir::MemRefType::get(shape, ty);
+    }
+
+    mlir::Type ty = convertType(baseTy);
+    assert(mlir::BaseMemRefType::isValidElementType(ty) &&
+           "got invalid memref element type from scalar fir type");
+    return mlir::MemRefType::get({}, ty);
+  }
+
 public:
   explicit FIRToMemRefTypeConverter(mlir::ModuleOp mod)
       : kindMapping(fir::getKindMapping(mod)) {
@@ -64,19 +94,26 @@ class FIRToMemRefTypeConverter : public mlir::TypeConverter {
   void setConvertScalarTypesOnly(bool value) { convertScalarTypesOnly = value; }
 
   /// Return true if the given FIR type can be converted to a MemRef-typed
-  /// descriptor (i.e. is a supported base element for MemRef converting).
+  /// descriptor. Uses the same `dyn_cast_ptrEleTy` / `box` recursion as
+  /// `convertMemrefType`. Nested pointers such as
+  /// `!fir.ref<!fir.heap<!fir.array<?xf32>>>` hold a heap address, not the
+  /// array, and are not convertible.
   bool convertibleMemrefType(mlir::Type ty) {
-    if (auto refTy = mlir::dyn_cast<fir::ReferenceType>(ty))
-      return convertibleMemrefType(refTy.getElementType());
-    else if (auto pointerTy = mlir::dyn_cast<fir::PointerType>(ty))
-      return convertibleMemrefType(pointerTy.getElementType());
-    else if (auto heapTy = mlir::dyn_cast<fir::HeapType>(ty))
-      return convertibleMemrefType(heapTy.getElementType());
-    else if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(ty))
-      return convertibleMemrefType(seqTy.getElementType());
+    if (mlir::Type pointee = fir::dyn_cast_ptrEleTy(ty))
+      ty = pointee;
     else if (auto boxTy = mlir::dyn_cast<fir::BoxType>(ty))
       return convertibleMemrefType(boxTy.getElementType());
 
+    // convertMemrefType peels only one pointer wrapper. A remaining pointer
+    // or box is not a valid memref element.
+    if (fir::isa_ref_type(ty) || fir::isa_box_type(ty))
+      return false;
+
+    if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(ty))
+      ty = seqTy.getElementType();
+    if (fir::isa_ref_type(ty) || fir::isa_box_type(ty))
+      return false;
+
     setConvertScalarTypesOnly(true);
     bool result = convertibleType(ty);
     setConvertScalarTypesOnly(false);
@@ -140,64 +177,19 @@ class FIRToMemRefTypeConverter : public mlir::TypeConverter {
 
   /// Convert a FIR element / aggregate type to a MemRef descriptor type.
   mlir::MemRefType convertMemrefType(mlir::Type firTy) const {
-    auto convertBaseType = [&](mlir::Type firTy) -> mlir::MemRefType {
-      if (auto charTy = mlir::dyn_cast<fir::CharacterType>(firTy)) {
-        unsigned kind = charTy.getFKind();
-        unsigned bitWidth = kindMapping.getCharacterBitsize(kind);
-        mlir::Type elTy = mlir::IntegerType::get(charTy.getContext(), bitWidth);
-
-        if (charTy.hasConstantLen() && charTy.getLen() == 1) {
-          return mlir::MemRefType::get({}, elTy);
-        } else if (charTy.hasConstantLen()) {
-          int64_t len = charTy.getLen();
-          return mlir::MemRefType::get({len}, elTy);
-        } else {
-          return mlir::MemRefType::get({mlir::ShapedType::kDynamic}, elTy);
-        }
-      }
-
-      if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(firTy)) {
-        auto elTy = seqTy.getElementType();
-        mlir::Type ty = convertType(elTy);
-
-        llvm::ArrayRef<int64_t> firShape = seqTy.getShape();
-        llvm::SmallVector<int64_t> shape;
-        for (auto it = firShape.rbegin(); it != firShape.rend(); ++it)
-          shape.push_back(*it);
-
-        assert(mlir::BaseMemRefType::isValidElementType(ty) &&
-               "got invalid memref element type from array fir type");
-        return mlir::MemRefType::get(shape, ty);
-      }
-
-      mlir::Type ty = convertType(firTy);
-      assert(mlir::BaseMemRefType::isValidElementType(ty) &&
-             "got invalid memref element type from scalar fir type");
-      return mlir::MemRefType::get({}, ty);
-    };
-
-    if (auto refTy = mlir::dyn_cast<fir::ReferenceType>(firTy))
-      return convertBaseType(refTy.getElementType());
-
-    if (auto pointerTy = mlir::dyn_cast<fir::PointerType>(firTy))
-      return convertBaseType(pointerTy.getElementType());
-
-    if (auto heapTy = mlir::dyn_cast<fir::HeapType>(firTy))
-      return convertBaseType(heapTy.getElementType());
+    if (mlir::Type pointee = fir::dyn_cast_ptrEleTy(firTy))
+      return convertMemrefBaseType(pointee);
 
     if (auto boxTy = mlir::dyn_cast<fir::BoxType>(firTy)) {
-      auto elTy = boxTy.getElementType();
-
-      auto memRefTy = convertMemrefType(elTy);
-      mlir::MemRefType dynTy = mlir::MemRefType::Builder(memRefTy).setLayout(
+      mlir::MemRefType memRefTy = convertMemrefType(boxTy.getElementType());
+      return mlir::MemRefType::Builder(memRefTy).setLayout(
           mlir::StridedLayoutAttr::get(
               memRefTy.getContext(), mlir::ShapedType::kDynamic,
               llvm::SmallVector<int64_t>(memRefTy.getRank(),
                                          mlir::ShapedType::kDynamic)));
-      return dynTy;
     }
 
-    return convertBaseType(firTy);
+    return convertMemrefBaseType(firTy);
   }
 };
 

diff  --git a/flang/unittests/Optimizer/OpenACC/FIROpenACCPointerLikeTypeInterfaceTest.cpp b/flang/unittests/Optimizer/OpenACC/FIROpenACCPointerLikeTypeInterfaceTest.cpp
index 4000ed6f68091..9e86baf6df78c 100644
--- a/flang/unittests/Optimizer/OpenACC/FIROpenACCPointerLikeTypeInterfaceTest.cpp
+++ b/flang/unittests/Optimizer/OpenACC/FIROpenACCPointerLikeTypeInterfaceTest.cpp
@@ -14,6 +14,7 @@
 #include "flang/Optimizer/Dialect/Support/KindMapping.h"
 #include "flang/Optimizer/OpenACC/Support/RegisterOpenACCExtensions.h"
 #include "flang/Optimizer/Support/InitFIR.h"
+#include "flang/Optimizer/Transforms/FIRToMemRefTypeConverter.h"
 
 using namespace mlir;
 
@@ -121,4 +122,47 @@ TEST_F(FIROpenACCPointerLikeTypeInterfaceTest,
   EXPECT_EQ(elTy.getWidth(), kindMap->getLogicalBitsize(logicalKind));
 }
 
+TEST_F(FIROpenACCPointerLikeTypeInterfaceTest,
+    NestedPointerTypesAreNotConvertible) {
+  Type f32 = Float32Type::get(&context);
+  Type dyn1 = fir::SequenceType::get({ShapedType::kDynamic}, f32);
+  Type dyn2 =
+      fir::SequenceType::get({ShapedType::kDynamic, ShapedType::kDynamic}, f32);
+  Type stat = fir::SequenceType::get({16}, f32);
+
+  fir::FIRToMemRefTypeConverter converter(module);
+  converter.setConvertComplexTypes(true);
+  auto conv = [&](Type t) { return converter.convertibleMemrefType(t); };
+
+  // One pointer/box to the array or scalar is convertible.
+  EXPECT_TRUE(conv(fir::HeapType::get(dyn1)));
+  EXPECT_TRUE(conv(fir::PointerType::get(dyn1)));
+  EXPECT_TRUE(conv(fir::ReferenceType::get(dyn1)));
+  EXPECT_TRUE(conv(fir::HeapType::get(dyn2)));
+  EXPECT_TRUE(conv(fir::HeapType::get(stat)));
+  EXPECT_TRUE(conv(fir::HeapType::get(f32)));
+  EXPECT_TRUE(conv(fir::BoxType::get(fir::HeapType::get(dyn1))));
+  EXPECT_TRUE(conv(fir::BoxType::get(fir::PointerType::get(dyn1))));
+  EXPECT_TRUE(conv(fir::BoxType::get(dyn1)));
+
+  // A remaining pointer after one peel holds an address, not the array.
+  // `!fir.ptr`/`!fir.heap` cannot wrap another pointer type.
+  EXPECT_FALSE(conv(fir::ReferenceType::get(fir::HeapType::get(dyn1))));
+  EXPECT_FALSE(conv(fir::ReferenceType::get(fir::PointerType::get(dyn1))));
+  EXPECT_FALSE(conv(fir::ReferenceType::get(fir::HeapType::get(dyn2))));
+  EXPECT_FALSE(conv(fir::ReferenceType::get(fir::HeapType::get(stat))));
+  EXPECT_FALSE(conv(fir::ReferenceType::get(fir::HeapType::get(f32))));
+  EXPECT_FALSE(conv(
+      fir::ReferenceType::get(fir::BoxType::get(fir::HeapType::get(dyn1)))));
+  EXPECT_FALSE(conv(
+      fir::BoxType::get(fir::ReferenceType::get(fir::HeapType::get(dyn1)))));
+
+  auto asMemRef = [&](Type t) {
+    return cast<acc::PointerLikeType>(t).getAsMemRefType(module);
+  };
+  EXPECT_FALSE(asMemRef(fir::ReferenceType::get(fir::HeapType::get(dyn1))));
+  EXPECT_FALSE(asMemRef(fir::ReferenceType::get(fir::PointerType::get(dyn1))));
+  EXPECT_FALSE(asMemRef(fir::ReferenceType::get(fir::HeapType::get(f32))));
+}
+
 } // namespace


        


More information about the flang-commits mailing list