FLANG
FIRToMemRefTypeConverter.h
1//===---- FIRToMemRefTypeConverter.h - FIR type conversion to MemRef ------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file defines `FIRToMemRefTypeConverter`, a helper used by the
10// FIR-to-MemRef conversion pass to convert FIR types (scalars, arrays,
11// descriptors) into MemRef types suitable for the MemRef dialect.
12//
13//===----------------------------------------------------------------------===//
14
15#ifndef FORTRAN_OPTIMIZER_TRANSFORMS_FIRTOMEMREFTYPECONVERTER_H
16#define FORTRAN_OPTIMIZER_TRANSFORMS_FIRTOMEMREFTYPECONVERTER_H
17
18#include "flang/Optimizer/Dialect/FIRDialect.h"
19#include "flang/Optimizer/Dialect/FIROps.h"
20#include "flang/Optimizer/Dialect/FIRType.h"
21#include "flang/Optimizer/Dialect/Support/FIRContext.h"
22#include "flang/Optimizer/Dialect/Support/KindMapping.h"
23#include "mlir/IR/BuiltinAttributes.h"
24#include "mlir/IR/BuiltinTypes.h"
25#include "mlir/Transforms/DialectConversion.h"
26
27namespace fir {
28
29class FIRToMemRefTypeConverter : public mlir::TypeConverter {
30private:
31 KindMapping kindMapping;
32 bool convertComplexTypes = false;
33 bool convertScalarTypesOnly = false;
34
35 mlir::MemRefType convertMemrefBaseType(mlir::Type baseTy) const {
36 if (auto charTy = mlir::dyn_cast<fir::CharacterType>(baseTy)) {
37 unsigned kind = charTy.getFKind();
38 unsigned bitWidth = kindMapping.getCharacterBitsize(kind);
39 mlir::Type elTy = mlir::IntegerType::get(charTy.getContext(), bitWidth);
40
41 if (charTy.hasConstantLen() && charTy.getLen() == 1)
42 return mlir::MemRefType::get({}, elTy);
43 if (charTy.hasConstantLen())
44 return mlir::MemRefType::get({charTy.getLen()}, elTy);
45 return mlir::MemRefType::get({mlir::ShapedType::kDynamic}, elTy);
46 }
47
48 if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(baseTy)) {
49 mlir::Type ty = convertType(seqTy.getElementType());
50 llvm::ArrayRef<int64_t> firShape = seqTy.getShape();
52 for (auto it = firShape.rbegin(); it != firShape.rend(); ++it)
53 shape.push_back(*it);
54 assert(mlir::BaseMemRefType::isValidElementType(ty) &&
55 "got invalid memref element type from array fir type");
56 return mlir::MemRefType::get(shape, ty);
57 }
58
59 mlir::Type ty = convertType(baseTy);
60 assert(mlir::BaseMemRefType::isValidElementType(ty) &&
61 "got invalid memref element type from scalar fir type");
62 return mlir::MemRefType::get({}, ty);
63 }
64
65public:
66 explicit FIRToMemRefTypeConverter(mlir::ModuleOp mod)
67 : kindMapping(fir::getKindMapping(mod)) {
68 addConversion([](mlir::Type type) { return type; });
69
70 addConversion([&](fir::LogicalType type) -> mlir::Type {
71 return mlir::IntegerType::get(
72 type.getContext(), kindMapping.getLogicalBitsize(type.getFKind()));
73 });
74
75 addSourceMaterialization([](mlir::OpBuilder &builder, mlir::Type type,
76 mlir::ValueRange inputs,
77 mlir::Location loc) -> mlir::Value {
78 assert(!inputs.empty() && "expected a single input for materialization");
79 builder.setInsertionPointAfter(inputs[0].getDefiningOp());
80 return fir::ConvertOp::create(builder, loc, type, inputs[0]);
81 });
82
83 addTargetMaterialization([](mlir::OpBuilder &builder, mlir::Type type,
84 mlir::ValueRange inputs,
85 mlir::Location loc) -> mlir::Value {
86 return fir::ConvertOp::create(builder, loc, type, inputs[0]);
87 });
88 }
89
91 void setConvertComplexTypes(bool value) { convertComplexTypes = value; }
92
94 void setConvertScalarTypesOnly(bool value) { convertScalarTypesOnly = value; }
95
101 bool convertibleMemrefType(mlir::Type ty) {
102 if (mlir::Type pointee = fir::dyn_cast_ptrEleTy(ty))
103 ty = pointee;
104 else if (auto boxTy = mlir::dyn_cast<fir::BoxType>(ty))
105 return convertibleMemrefType(boxTy.getElementType());
106
107 // convertMemrefType peels only one pointer wrapper. A remaining pointer
108 // or box is not a valid memref element.
110 return false;
111
112 if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(ty))
113 ty = seqTy.getElementType();
115 return false;
116
118 bool result = convertibleType(ty);
120 return result;
121 }
122
125 bool isEmptyArray(mlir::Type ty) const {
126 if (auto refTy = mlir::dyn_cast<fir::ReferenceType>(ty))
127 return isEmptyArray(refTy.getElementType());
128 else if (auto pointerTy = mlir::dyn_cast<fir::PointerType>(ty))
129 return isEmptyArray(pointerTy.getElementType());
130 else if (auto heapTy = mlir::dyn_cast<fir::HeapType>(ty))
131 return isEmptyArray(heapTy.getElementType());
132 else if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(ty)) {
133 llvm::ArrayRef<int64_t> firShape = seqTy.getShape();
134 for (auto shape : firShape)
135 if (shape == 0)
136 return true;
137 return false;
138 }
139 return false;
140 }
141
144 bool convertibleType(mlir::Type type) const {
145 if (!convertScalarTypesOnly) {
146 if (auto refTy = mlir::dyn_cast<fir::ReferenceType>(type)) {
147 auto elTy = refTy.getElementType();
148 if (mlir::isa<fir::SequenceType>(elTy))
149 return false;
150 return convertibleType(elTy);
151 }
152
153 if (auto seqTy = mlir::dyn_cast<fir::SequenceType>(type))
154 return convertibleType(seqTy.getElementType());
155 }
156
157 if (fir::isa_fir_type(type)) {
158 if (mlir::isa<fir::LogicalType>(type))
159 return true;
160 return false;
161 }
162
163 if (type.isUnsignedInteger())
164 return false;
165
166 if (mlir::isa<mlir::ComplexType>(type))
167 return convertComplexTypes;
168
169 if (mlir::isa<mlir::FunctionType>(type))
170 return false;
171
172 if (mlir::isa<mlir::TupleType>(type))
173 return false;
174
175 return true;
176 }
177
179 mlir::MemRefType convertMemrefType(mlir::Type firTy) const {
180 if (mlir::Type pointee = fir::dyn_cast_ptrEleTy(firTy))
181 return convertMemrefBaseType(pointee);
182
183 if (auto boxTy = mlir::dyn_cast<fir::BoxType>(firTy)) {
184 mlir::MemRefType memRefTy = convertMemrefType(boxTy.getElementType());
185 return mlir::MemRefType::Builder(memRefTy).setLayout(
186 mlir::StridedLayoutAttr::get(
187 memRefTy.getContext(), mlir::ShapedType::kDynamic,
188 llvm::SmallVector<int64_t>(memRefTy.getRank(),
189 mlir::ShapedType::kDynamic)));
190 }
191
192 return convertMemrefBaseType(firTy);
193 }
194};
195
196} // namespace fir
197
198#endif // FORTRAN_OPTIMIZER_TRANSFORMS_FIRTOMEMREFTYPECONVERTER_H
mlir::MemRefType convertMemrefType(mlir::Type firTy) const
Convert a FIR element / aggregate type to a MemRef descriptor type.
Definition FIRToMemRefTypeConverter.h:179
bool isEmptyArray(mlir::Type ty) const
Definition FIRToMemRefTypeConverter.h:125
bool convertibleMemrefType(mlir::Type ty)
Definition FIRToMemRefTypeConverter.h:101
void setConvertComplexTypes(bool value)
Control whether complex types are considered convertible.
Definition FIRToMemRefTypeConverter.h:91
void setConvertScalarTypesOnly(bool value)
Control whether only scalar types are considered during convertibleType.
Definition FIRToMemRefTypeConverter.h:94
bool convertibleType(mlir::Type type) const
Definition FIRToMemRefTypeConverter.h:144
Definition KindMapping.h:48
Definition FIRType.h:106
Definition OpenACC.h:20
Definition AbstractConverter.h:37
bool isa_box_type(mlir::Type t)
Is t a boxed type?
Definition FIRType.h:141
KindMapping getKindMapping(mlir::ModuleOp mod)
Definition FIRContext.cpp:43
bool isa_ref_type(mlir::Type t)
Is t a FIR dialect type that implies a memory (de)reference?
Definition FIRType.h:135
mlir::Type dyn_cast_ptrEleTy(mlir::Type t)
Definition FIRType.cpp:257
bool isa_fir_type(mlir::Type t)
Is t any of the FIR dialect types?
Definition FIRType.cpp:218