23#define GEN_PASS_DEF_ACCDECLARECTORDTORCONVERSION
24#include "mlir/Dialect/OpenACC/Transforms/Passes.h.inc"
32static void collectExistingGlobalCtors(
36 for (
auto globalCtors : mod.getOps<LLVM::GlobalCtorsOp>()) {
37 ctors.append(globalCtors.getCtors().begin(), globalCtors.getCtors().end());
38 for (
Attribute attr : globalCtors.getPriorities())
39 priorities.push_back(cast<IntegerAttr>(attr).getInt());
40 data.append(globalCtors.getData().begin(), globalCtors.getData().end());
41 globalCtorsOps.push_back(globalCtors);
45static void collectExistingGlobalDtors(
49 for (
auto globalDtors : mod.getOps<LLVM::GlobalDtorsOp>()) {
50 dtors.append(globalDtors.getDtors().begin(), globalDtors.getDtors().end());
51 for (
Attribute attr : globalDtors.getPriorities())
52 priorities.push_back(cast<IntegerAttr>(attr).getInt());
53 data.append(globalDtors.getData().begin(), globalDtors.getData().end());
54 globalDtorsOps.push_back(globalDtors);
58static void replaceGlobalCtors(ModuleOp mod,
OpBuilder &builder,
63 for (
auto globalCtors : oldOps)
69 LLVM::GlobalCtorsOp::create(
74static void replaceGlobalDtors(ModuleOp mod,
OpBuilder &builder,
79 for (
auto globalDtors : oldOps)
85 LLVM::GlobalDtorsOp::create(
92static LLVM::LLVMFuncOp createLLVMFunctionFromRegion(StringRef symName,
96 auto llvmVoidTy = LLVM::LLVMVoidType::get(mod.getContext());
97 auto funcTy = LLVM::LLVMFunctionType::get(llvmVoidTy, {},
false);
99 auto newFunc = LLVM::LLVMFuncOp::create(builder, mod.getLoc(), symName,
100 funcTy, LLVM::Linkage::Internal);
102 Block *entry = newFunc.addEntryBlock(builder);
114 LLVM::ReturnOp::create(builder, mod.getLoc(),
ValueRange{});
123static LLVM::LLVMFuncOp
124createExtraConstructorCaller(ModuleOp mod,
OpBuilder &builder,
126 StringRef extraCtorName) {
127 auto llvmVoidTy = LLVM::LLVMVoidType::get(mod.getContext());
128 auto funcTy = LLVM::LLVMFunctionType::get(llvmVoidTy, {},
false);
131 for (
const std::string &name : extraNames) {
132 if (mod.lookupSymbol<LLVM::LLVMFuncOp>(name))
134 LLVM::LLVMFuncOp::create(builder, mod.getLoc(), name, funcTy,
135 LLVM::Linkage::External);
138 auto wrapper = LLVM::LLVMFuncOp::create(builder, mod.getLoc(), extraCtorName,
139 funcTy, LLVM::Linkage::Internal);
140 Block *entry = wrapper.addEntryBlock(builder);
142 for (
const std::string &name : extraNames)
143 LLVM::CallOp::create(builder, mod.getLoc(), funcTy,
145 LLVM::ReturnOp::create(builder, mod.getLoc(),
ValueRange{});
149struct ACCDeclareCtorDtorConversion
150 :
public acc::impl::ACCDeclareCtorDtorConversionBase<
151 ACCDeclareCtorDtorConversion> {
154 void runOnOperation()
override {
155 ModuleOp mod = getOperation();
156 OpBuilder builder{mod.getBodyRegion()};
157 SmallVector<Operation *> worklist;
159 SmallVector<Attribute, 8> allCtors;
160 SmallVector<int32_t, 8> ctorPriorities;
161 SmallVector<Attribute, 8> ctorData;
162 SmallVector<LLVM::GlobalCtorsOp, 4> globalCtorsOps;
163 collectExistingGlobalCtors(mod, allCtors, ctorPriorities, ctorData,
165 size_t existingCtorCount = allCtors.size();
167 SmallVector<Attribute, 8> allDtors;
168 SmallVector<int32_t, 8> dtorPriorities;
169 SmallVector<Attribute, 8> dtorData;
170 SmallVector<LLVM::GlobalDtorsOp, 4> globalDtorsOps;
171 collectExistingGlobalDtors(mod, allDtors, dtorPriorities, dtorData,
173 size_t existingDtorCount = allDtors.size();
175 mod.walk([&](acc::GlobalConstructorOp op) {
176 LLVM::LLVMFuncOp newCtor = createLLVMFunctionFromRegion(
177 op.getSymName(), op.getRegion(), mod, builder);
180 ctorPriorities.push_back(priority);
182 ctorData.push_back(LLVM::ZeroAttr::get(builder.
getContext()));
183 worklist.push_back(op.getOperation());
186 mod.walk([&](acc::GlobalDestructorOp op) {
188 LLVM::LLVMFuncOp newDtor = createLLVMFunctionFromRegion(
189 op.getSymName(), op.getRegion(), mod, builder);
192 dtorPriorities.push_back(priority);
194 dtorData.push_back(LLVM::ZeroAttr::get(builder.
getContext()));
196 worklist.push_back(op.getOperation());
199 bool hasEntryPoint = !entryPointName.empty() &&
200 static_cast<bool>(mod.lookupSymbol(entryPointName));
201 SmallVector<std::string, 4> extraNames;
202 for (
const auto &funcName : extraConstructors)
203 extraNames.push_back(funcName);
204 for (
const auto &funcName : entryOnlyConstructors) {
206 extraNames.push_back(funcName);
208 if (!extraNames.empty()) {
209 LLVM::LLVMFuncOp extraCtor =
210 createExtraConstructorCaller(mod, builder, extraNames, extraCtorName);
213 ctorPriorities.push_back(priority);
214 ctorData.push_back(LLVM::ZeroAttr::get(builder.
getContext()));
217 if (allCtors.size() > existingCtorCount)
218 replaceGlobalCtors(mod, builder, allCtors, ctorPriorities, ctorData,
220 if (allDtors.size() > existingDtorCount)
221 replaceGlobalDtors(mod, builder, allDtors, dtorPriorities, dtorData,
224 for (Operation *op : worklist)
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgListType getArguments()
ArrayAttr getI32ArrayAttr(ArrayRef< int32_t > values)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
static FlatSymbolRefAttr get(StringAttr value)
Construct a symbol reference for the given value name.
This is a utility class for mapping one set of IR entities to another.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
This class helps build Operations.
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
Operation is the basic unit of execution within MLIR.
result_range getResults()
void erase()
Remove this operation from its parent block and delete it.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
This class provides an abstraction over the different types of ranges over Values.
Include the generated interface declarations.