[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