27#include "llvm/ADT/ArrayRef.h"
28#include "llvm/ADT/SmallVector.h"
29#include "llvm/ADT/TypeSwitch.h"
30#include "llvm/Frontend/OpenMP/OMPConstants.h"
31#include "llvm/Frontend/OpenMP/OMPIRBuilder.h"
32#include "llvm/IR/Constants.h"
33#include "llvm/IR/DebugInfoMetadata.h"
34#include "llvm/IR/DerivedTypes.h"
35#include "llvm/IR/IRBuilder.h"
36#include "llvm/IR/MDBuilder.h"
37#include "llvm/IR/ReplaceConstant.h"
38#include "llvm/Support/AMDGPUAddrSpace.h"
39#include "llvm/Support/FileSystem.h"
40#include "llvm/Support/MathExtras.h"
41#include "llvm/Support/NVPTXAddrSpace.h"
42#include "llvm/Support/VirtualFileSystem.h"
43#include "llvm/TargetParser/Triple.h"
44#include "llvm/Transforms/Utils/ModuleUtils.h"
55static llvm::omp::ScheduleKind
56convertToScheduleKind(std::optional<omp::ClauseScheduleKind> schedKind) {
57 if (!schedKind.has_value())
58 return llvm::omp::OMP_SCHEDULE_Default;
59 switch (schedKind.value()) {
60 case omp::ClauseScheduleKind::Static:
61 return llvm::omp::OMP_SCHEDULE_Static;
62 case omp::ClauseScheduleKind::Dynamic:
63 return llvm::omp::OMP_SCHEDULE_Dynamic;
64 case omp::ClauseScheduleKind::Guided:
65 return llvm::omp::OMP_SCHEDULE_Guided;
66 case omp::ClauseScheduleKind::Auto:
67 return llvm::omp::OMP_SCHEDULE_Auto;
68 case omp::ClauseScheduleKind::Runtime:
69 return llvm::omp::OMP_SCHEDULE_Runtime;
70 case omp::ClauseScheduleKind::Distribute:
71 return llvm::omp::OMP_SCHEDULE_Distribute;
73 llvm_unreachable(
"unhandled schedule clause argument");
78class OpenMPAllocStackFrame
83 explicit OpenMPAllocStackFrame(
84 llvm::OpenMPIRBuilder::InsertPointTy allocaIP,
85 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks)
86 : allocInsertPoint(allocaIP), deallocBlocks(deallocBlocks) {}
87 llvm::OpenMPIRBuilder::InsertPointTy allocInsertPoint;
88 llvm::SmallVector<llvm::BasicBlock *> deallocBlocks;
94class OpenMPLoopInfoStackFrame
98 llvm::CanonicalLoopInfo *loopInfo =
nullptr;
117class PreviouslyReportedError
118 :
public llvm::ErrorInfo<PreviouslyReportedError> {
120 void log(raw_ostream &)
const override {
124 std::error_code convertToErrorCode()
const override {
126 "PreviouslyReportedError doesn't support ECError conversion");
133char PreviouslyReportedError::ID = 0;
144class LinearClauseProcessor {
147 SmallVector<llvm::Value *> linearPreconditionVars;
148 SmallVector<llvm::Value *> linearLoopBodyTemps;
149 SmallVector<llvm::Value *> linearOrigVal;
150 SmallVector<llvm::Value *> linearSteps;
151 SmallVector<llvm::Type *> linearVarTypes;
152 llvm::BasicBlock *linearFinalizationBB;
153 llvm::BasicBlock *linearExitBB;
154 llvm::BasicBlock *linearLastIterExitBB;
159 void registerType(LLVM::ModuleTranslation &moduleTranslation,
160 mlir::Attribute &ty) {
161 linearVarTypes.push_back(moduleTranslation.
convertType(
162 mlir::cast<mlir::TypeAttr>(ty).getValue()));
166 void createLinearVar(llvm::IRBuilderBase &builder,
167 LLVM::ModuleTranslation &moduleTranslation,
168 llvm::Value *linearVar,
int idx) {
169 linearPreconditionVars.push_back(
170 builder.CreateAlloca(linearVarTypes[idx],
nullptr,
".linear_var"));
171 llvm::Value *linearLoopBodyTemp =
172 builder.CreateAlloca(linearVarTypes[idx],
nullptr,
".linear_result");
173 linearOrigVal.push_back(linearVar);
174 linearLoopBodyTemps.push_back(linearLoopBodyTemp);
178 inline void initLinearStep(LLVM::ModuleTranslation &moduleTranslation,
179 mlir::Value &linearStep) {
180 linearSteps.push_back(moduleTranslation.
lookupValue(linearStep));
184 void initLinearVar(llvm::IRBuilderBase &builder,
185 LLVM::ModuleTranslation &moduleTranslation,
186 llvm::BasicBlock *loopPreHeader) {
187 builder.SetInsertPoint(loopPreHeader->getTerminator());
188 for (
size_t index = 0; index < linearOrigVal.size(); index++) {
189 llvm::LoadInst *linearVarLoad =
190 builder.CreateLoad(linearVarTypes[index], linearOrigVal[index]);
191 builder.CreateStore(linearVarLoad, linearPreconditionVars[index]);
196 LogicalResult initLinearIV(omp::SimdOp simdOp) {
197 auto loopOp = cast<omp::LoopNestOp>(simdOp.getWrappedLoop());
199 if (loopOp.getIVs().size() != 1)
207 BlockArgument arg = loopOp.getIVs().front();
208 for (
const Operation *user : arg.
getUsers()) {
209 if (
auto storeOp = dyn_cast<LLVM::StoreOp>(user)) {
210 for (Value linearVar : simdOp.getLinearVars()) {
211 if (linearVar == storeOp.getAddr()) {
212 if (linearLoopIV && linearLoopIV != linearVar)
213 return simdOp.emitError(
214 "Could not determine the linear variable associated with the "
215 "loop nest induction variable");
216 linearLoopIV = linearVar;
225 void updateLinearVar(llvm::IRBuilderBase &builder, llvm::BasicBlock *loopBody,
226 llvm::Value *loopInductionVar) {
227 builder.SetInsertPoint(loopBody->getTerminator());
228 for (
size_t index = 0; index < linearPreconditionVars.size(); index++) {
229 llvm::Type *linearVarType = linearVarTypes[index];
230 llvm::Value *iv = loopInductionVar;
231 llvm::Value *step = linearSteps[index];
233 if (!iv->getType()->isIntegerTy())
234 llvm_unreachable(
"OpenMP loop induction variable must be an integer "
237 if (linearVarType->isIntegerTy()) {
239 iv = builder.CreateSExtOrTrunc(iv, linearVarType);
240 step = builder.CreateSExtOrTrunc(step, linearVarType);
242 llvm::LoadInst *linearVarStart =
243 builder.CreateLoad(linearVarType, linearPreconditionVars[index]);
244 llvm::Value *mulInst = builder.CreateMul(iv, step);
245 llvm::Value *addInst = builder.CreateAdd(linearVarStart, mulInst);
246 builder.CreateStore(addInst, linearLoopBodyTemps[index]);
247 }
else if (linearVarType->isFloatingPointTy()) {
249 step = builder.CreateSExtOrTrunc(step, iv->getType());
250 llvm::Value *mulInst = builder.CreateMul(iv, step);
252 llvm::LoadInst *linearVarStart =
253 builder.CreateLoad(linearVarType, linearPreconditionVars[index]);
254 llvm::Value *mulFp = builder.CreateSIToFP(mulInst, linearVarType);
255 llvm::Value *addInst = builder.CreateFAdd(linearVarStart, mulFp);
256 builder.CreateStore(addInst, linearLoopBodyTemps[index]);
259 "Linear variable must be of integer or floating-point type");
265 void updateLinearIV(llvm::IRBuilderBase &builder,
266 LLVM::ModuleTranslation &moduleTranslation) {
269 llvm::Value *linearIV = moduleTranslation.
lookupValue(linearLoopIV);
273 for (index = 0; index < linearOrigVal.size(); index++)
274 if (linearIV == linearOrigVal[index])
276 if (index == linearOrigVal.size())
280 llvm::Type *varType = linearVarTypes[index];
281 llvm::Value *var = linearLoopBodyTemps[index];
282 llvm::Value *step = linearSteps[index];
283 if (!varType->isIntegerTy())
284 llvm_unreachable(
"Linear iteration variable must be of integer type");
286 step = builder.CreateSExtOrTrunc(step, varType);
287 llvm::Value *val = builder.CreateLoad(varType, var);
288 llvm::Value *addInst = builder.CreateAdd(val, step);
289 builder.CreateStore(addInst, var);
294 void splitLinearFiniBB(llvm::IRBuilderBase &builder,
295 llvm::BasicBlock *loopExit) {
296 linearFinalizationBB = loopExit->splitBasicBlock(
297 loopExit->getTerminator(),
"omp_loop.linear_finalization");
298 linearExitBB = linearFinalizationBB->splitBasicBlock(
299 linearFinalizationBB->getTerminator(),
"omp_loop.linear_exit");
300 linearLastIterExitBB = linearFinalizationBB->splitBasicBlock(
301 linearFinalizationBB->getTerminator(),
"omp_loop.linear_lastiter_exit");
305 llvm::OpenMPIRBuilder::InsertPointOrErrorTy
306 finalizeLinearVar(llvm::IRBuilderBase &builder,
307 LLVM::ModuleTranslation &moduleTranslation,
308 llvm::Value *lastIter) {
310 builder.SetInsertPoint(linearFinalizationBB->getTerminator());
311 llvm::Value *loopLastIterLoad = builder.CreateLoad(
312 llvm::Type::getInt32Ty(builder.getContext()), lastIter);
313 llvm::Value *isLast =
314 builder.CreateCmp(llvm::CmpInst::ICMP_NE, loopLastIterLoad,
315 llvm::ConstantInt::get(
316 llvm::Type::getInt32Ty(builder.getContext()), 0));
318 builder.SetInsertPoint(linearLastIterExitBB->getTerminator());
319 for (
size_t index = 0; index < linearOrigVal.size(); index++) {
320 llvm::LoadInst *linearVarTemp =
321 builder.CreateLoad(linearVarTypes[index], linearLoopBodyTemps[index]);
322 builder.CreateStore(linearVarTemp, linearOrigVal[index]);
328 builder.SetInsertPoint(linearFinalizationBB->getTerminator());
329 builder.CreateCondBr(isLast, linearLastIterExitBB, linearExitBB);
330 linearFinalizationBB->getTerminator()->eraseFromParent();
332 builder.SetInsertPoint(linearExitBB->getTerminator());
334 builder, llvm::omp::OMPD_barrier);
339 void emitStoresForLinearVar(llvm::IRBuilderBase &builder) {
340 for (
size_t index = 0; index < linearOrigVal.size(); index++) {
341 llvm::LoadInst *linearVarTemp =
342 builder.CreateLoad(linearVarTypes[index], linearLoopBodyTemps[index]);
343 builder.CreateStore(linearVarTemp, linearOrigVal[index]);
349 void rewriteInPlace(llvm::IRBuilderBase &builder, llvm::BasicBlock *startBB,
350 llvm::BasicBlock *endBB,
size_t varIndex) {
351 llvm::SmallVector<llvm::BasicBlock *, 32> worklist;
352 llvm::SmallPtrSet<llvm::BasicBlock *, 32> collectedBBs;
354 assert(startBB && endBB &&
"Invalid startBB/endBB");
357 worklist.push_back(startBB);
358 collectedBBs.insert(startBB);
360 while (!worklist.empty()) {
361 llvm::BasicBlock *bb = worklist.pop_back_val();
366 for (llvm::BasicBlock *succ : llvm::successors(bb)) {
367 if (collectedBBs.insert(succ).second)
368 worklist.push_back(succ);
373 llvm::SmallVector<llvm::User *> users(linearOrigVal[varIndex]->users());
374 for (
auto *user : users) {
375 if (
auto *userInst = dyn_cast<llvm::Instruction>(user)) {
376 if (collectedBBs.contains(userInst->getParent()))
377 user->replaceUsesOfWith(linearOrigVal[varIndex],
378 linearLoopBodyTemps[varIndex]);
389 SymbolRefAttr symbolName) {
390 omp::PrivateClauseOp privatizer =
393 assert(privatizer &&
"privatizer not found in the symbol table");
404 auto todo = [&op](StringRef clauseName) {
405 return op.
emitError() <<
"not yet implemented: Unhandled clause "
406 << clauseName <<
" in " << op.
getName()
410 auto checkAllocate = [&todo](
auto op, LogicalResult &
result) {
411 if (!op.getAllocateVars().empty() || !op.getAllocatorVars().empty())
412 result = todo(
"allocate");
414 auto checkBare = [&todo](
auto op, LogicalResult &
result) {
415 if (op.getKernelType() == omp::TargetExecMode::bare)
416 result = todo(
"ompx_bare");
418 auto checkDepend = [&todo](
auto op, LogicalResult &
result) {
419 if (!op.getDependVars().empty() || op.getDependKinds())
422 auto checkHint = [](
auto op, LogicalResult &) {
426 auto checkInReduction = [&todo](
auto op, LogicalResult &
result) {
427 if (isa<omp::TargetOp, omp::TaskOp, omp::TaskloopContextOp>(
428 op.getOperation())) {
429 if (
auto byrefAttr = op.getInReductionByref()) {
430 for (
bool isByRef : *byrefAttr) {
432 result = todo(
"in_reduction with byref modifier");
437 if (isa<omp::TargetOp>(op.getOperation())) {
438 if (
auto inReductionSyms = op.getInReductionSyms()) {
440 (*inReductionSyms).template getAsRange<SymbolRefAttr>()) {
445 "symbol resolution should be guaranteed by the op verifier");
446 if (decl.getInitializerRegion().front().getNumArguments() != 1) {
447 result = todo(
"in_reduction with two-argument initializer");
450 if (!decl.getCleanupRegion().empty()) {
451 result = todo(
"in_reduction with cleanup region");
457 }
else if (!op.getInReductionVars().empty() || op.getInReductionByref() ||
458 op.getInReductionSyms()) {
459 result = todo(
"in_reduction");
462 auto checkNowait = [&todo](
auto op, LogicalResult &
result) {
466 auto checkOrder = [&todo](
auto op, LogicalResult &
result) {
467 if (op.getOrder() || op.getOrderMod())
470 auto checkPrivate = [&todo](
auto op, LogicalResult &
result) {
471 if (!op.getPrivateVars().empty() || op.getPrivateSyms())
472 result = todo(
"privatization");
474 auto checkReduction = [&todo](
auto op, LogicalResult &
result) {
475 if (isa<omp::TeamsOp>(op))
476 if (!op.getReductionVars().empty() || op.getReductionByref() ||
477 op.getReductionSyms())
478 result = todo(
"reduction");
479 if (op.getReductionMod() &&
480 op.getReductionMod().value() != omp::ReductionModifier::defaultmod) {
481 omp::ReductionModifier mod = op.getReductionMod().value();
485 bool taskModifierSupported =
486 mod == omp::ReductionModifier::task &&
487 isa<omp::ParallelOp, omp::WsloopOp, omp::SectionsOp>(op);
488 if (!taskModifierSupported) {
489 result = todo(
"reduction with modifier");
490 }
else if (
auto byref = op.getReductionByref()) {
493 for (
bool isByRef : *byref)
495 result = todo(
"task reduction modifier with by-ref reduction");
501 auto checkTaskReductionByref = [&todo](
auto op, LogicalResult &
result) {
502 if (
auto byrefAttr = op.getTaskReductionByref())
503 for (
bool isByRef : *byrefAttr)
505 result = todo(
"task_reduction with byref modifier");
509 auto checkReductionByref = [&todo](
auto op, LogicalResult &
result) {
510 if (
auto byrefAttr = op.getReductionByref())
511 for (
bool isByRef : *byrefAttr)
513 result = todo(
"reduction with byref modifier");
517 auto checkNumTeams = [&todo](
auto op, LogicalResult &
result) {
518 if (op.hasNumTeamsMultiDim())
519 result = todo(
"num_teams with multi-dimensional values");
521 auto checkNumThreads = [&todo](
auto op, LogicalResult &
result) {
522 if (op.hasNumThreadsMultiDim())
523 result = todo(
"num_threads with multi-dimensional values");
526 auto checkThreadLimit = [&todo](
auto op, LogicalResult &
result) {
527 if (op.hasThreadLimitMultiDim())
528 result = todo(
"thread_limit with multi-dimensional values");
530 auto checkMap = [&todo](
auto op, LogicalResult &
result) {
531 if (!op.getMapIterated().empty())
532 result = todo(
"map/motion clause with iterator modifier");
535 auto checkDynGroupprivate = [&todo](
auto op, LogicalResult &
result) {
536 if (op.getDynGroupprivateSize())
537 result = todo(
"dyn_groupprivate");
542 .Case([&](omp::DistributeOp op) {
543 checkAllocate(op,
result);
546 .Case([&](omp::SectionsOp op) {
547 checkAllocate(op,
result);
549 checkReduction(op,
result);
551 .Case([&](omp::ScopeOp op) { checkReduction(op,
result); })
552 .Case([&](omp::SingleOp op) {
553 checkAllocate(op,
result);
556 .Case([&](omp::TeamsOp op) {
557 checkAllocate(op,
result);
559 checkNumTeams(op,
result);
560 checkThreadLimit(op,
result);
561 checkDynGroupprivate(op,
result);
563 .Case([&](omp::TaskOp op) {
564 checkAllocate(op,
result);
565 checkInReduction(op,
result);
567 .Case([&](omp::TaskgroupOp op) {
568 checkAllocate(op,
result);
569 checkTaskReductionByref(op,
result);
571 .Case([&](omp::DispatchOp op) {
580 .Case([&](omp::TaskloopContextOp op) {
581 checkAllocate(op,
result);
582 checkInReduction(op,
result);
583 checkReduction(op,
result);
584 checkReductionByref(op,
result);
586 .Case([&](omp::WsloopOp op) {
587 checkAllocate(op,
result);
589 checkReduction(op,
result);
591 .Case([&](omp::ParallelOp op) {
592 checkReduction(op,
result);
593 checkNumThreads(op,
result);
595 .Case([&](omp::SimdOp op) { checkReduction(op,
result); })
596 .Case<omp::AtomicReadOp, omp::AtomicWriteOp, omp::AtomicUpdateOp,
597 omp::AtomicCaptureOp>([&](
auto op) { checkHint(op,
result); })
598 .Case([&](omp::AtomicCompareOp op) {
604 auto structTy = dyn_cast<LLVM::LLVMStructType>(argType);
610 result = todo(
"compare for complex types wider than 128 bits");
612 .Case<omp::TargetEnterDataOp, omp::TargetExitDataOp>([&](
auto op) {
616 .Case([&](omp::TargetUpdateOp op) {
620 .Case([&](omp::TargetOp op) {
621 checkAllocate(op,
result);
623 checkInReduction(op,
result);
625 checkThreadLimit(op,
result);
627 .Case([&](omp::TargetDataOp op) { checkMap(op,
result); })
628 .Case([&](omp::DeclareMapperInfoOp op) { checkMap(op,
result); })
639 llvm::handleAllErrors(
641 [&](
const PreviouslyReportedError &) {
result = failure(); },
642 [&](
const llvm::ErrorInfoBase &err) {
665 llvm::OpenMPIRBuilder::InsertPointTy allocInsertPoint;
668 [&](OpenMPAllocStackFrame &frame) {
669 allocInsertPoint = frame.allocInsertPoint;
670 deallocInsertPoints = frame.deallocBlocks;
678 allocInsertPoint.getNodeParent()->getParent() ==
679 builder.GetInsertBlock()->getParent()) {
681 deallocBlocks->insert(deallocBlocks->end(), deallocInsertPoints.begin(),
682 deallocInsertPoints.end());
683 return allocInsertPoint;
693 if (builder.GetInsertBlock() ==
694 &builder.GetInsertBlock()->getParent()->getEntryBlock()) {
695 assert(builder.GetInsertPoint() == builder.GetInsertBlock()->end() &&
696 "Assuming end of basic block");
697 llvm::BasicBlock *entryBB = llvm::BasicBlock::Create(
698 builder.getContext(),
"entry", builder.GetInsertBlock()->getParent(),
699 builder.GetInsertBlock()->getNextNode());
700 builder.CreateBr(entryBB);
701 builder.SetInsertPoint(entryBB);
707 for (llvm::BasicBlock &block : *builder.GetInsertBlock()->getParent()) {
711 llvm::Instruction *terminator = block.getTerminatorOrNull();
712 if (isa_and_present<llvm::ReturnInst>(terminator))
713 deallocBlocks->emplace_back(&block);
717 llvm::BasicBlock &funcEntryBlock =
718 builder.GetInsertBlock()->getParent()->getEntryBlock();
719 return funcEntryBlock.getFirstInsertionPt();
725static llvm::CanonicalLoopInfo *
727 llvm::CanonicalLoopInfo *loopInfo =
nullptr;
728 moduleTranslation.
stackWalk<OpenMPLoopInfoStackFrame>(
729 [&](OpenMPLoopInfoStackFrame &frame) {
730 loopInfo = frame.loopInfo;
742 Region ®ion, StringRef blockName, llvm::IRBuilderBase &builder,
745 bool isLoopWrapper = isa<omp::LoopWrapperInterface>(region.
getParentOp());
747 llvm::BasicBlock *continuationBlock =
748 splitBB(builder,
true,
"omp.region.cont");
749 llvm::BasicBlock *sourceBlock = builder.GetInsertBlock();
751 llvm::LLVMContext &llvmContext = builder.getContext();
752 for (
Block &bb : region) {
753 llvm::BasicBlock *llvmBB = llvm::BasicBlock::Create(
754 llvmContext, blockName, builder.GetInsertBlock()->getParent(),
755 builder.GetInsertBlock()->getNextNode());
756 moduleTranslation.
mapBlock(&bb, llvmBB);
759 llvm::Instruction *sourceTerminator = sourceBlock->getTerminator();
766 unsigned numYields = 0;
768 if (!isLoopWrapper) {
769 bool operandsProcessed =
false;
771 if (omp::YieldOp yield = dyn_cast<omp::YieldOp>(bb.getTerminator())) {
772 if (!operandsProcessed) {
773 for (
unsigned i = 0, e = yield->getNumOperands(); i < e; ++i) {
774 continuationBlockPHITypes.push_back(
775 moduleTranslation.
convertType(yield->getOperand(i).getType()));
777 operandsProcessed =
true;
779 assert(continuationBlockPHITypes.size() == yield->getNumOperands() &&
780 "mismatching number of values yielded from the region");
781 for (
unsigned i = 0, e = yield->getNumOperands(); i < e; ++i) {
782 llvm::Type *operandType =
783 moduleTranslation.
convertType(yield->getOperand(i).getType());
785 assert(continuationBlockPHITypes[i] == operandType &&
786 "values of mismatching types yielded from the region");
796 if (!continuationBlockPHITypes.empty())
798 continuationBlockPHIs &&
799 "expected continuation block PHIs if converted regions yield values");
800 if (continuationBlockPHIs) {
801 llvm::IRBuilderBase::InsertPointGuard guard(builder);
802 continuationBlockPHIs->reserve(continuationBlockPHITypes.size());
803 builder.SetInsertPoint(continuationBlock->begin());
804 for (llvm::Type *ty : continuationBlockPHITypes)
805 continuationBlockPHIs->push_back(builder.CreatePHI(ty, numYields));
811 for (
Block *bb : blocks) {
812 llvm::BasicBlock *llvmBB = moduleTranslation.
lookupBlock(bb);
815 if (bb->isEntryBlock()) {
816 assert(sourceTerminator->getNumSuccessors() == 1 &&
817 "provided entry block has multiple successors");
818 assert(sourceTerminator->getSuccessor(0) == continuationBlock &&
819 "ContinuationBlock is not the successor of the entry block");
820 sourceTerminator->setSuccessor(0, llvmBB);
823 llvm::IRBuilderBase::InsertPointGuard guard(builder);
825 moduleTranslation.
convertBlock(*bb, bb->isEntryBlock(), builder)))
826 return llvm::make_error<PreviouslyReportedError>();
831 builder.CreateBr(continuationBlock);
842 Operation *terminator = bb->getTerminator();
843 if (isa<omp::TerminatorOp, omp::YieldOp>(terminator)) {
844 builder.CreateBr(continuationBlock);
846 for (
unsigned i = 0, e = terminator->
getNumOperands(); i < e; ++i)
847 (*continuationBlockPHIs)[i]->addIncoming(
861 return continuationBlock;
867 case omp::ClauseProcBindKind::Close:
868 return llvm::omp::ProcBindKind::OMP_PROC_BIND_close;
869 case omp::ClauseProcBindKind::Master:
870 return llvm::omp::ProcBindKind::OMP_PROC_BIND_master;
871 case omp::ClauseProcBindKind::Primary:
872 return llvm::omp::ProcBindKind::OMP_PROC_BIND_primary;
873 case omp::ClauseProcBindKind::Spread:
874 return llvm::omp::ProcBindKind::OMP_PROC_BIND_spread;
876 llvm_unreachable(
"Unknown ClauseProcBindKind kind");
883 auto dispatchOp = cast<omp::DispatchOp>(opInst);
888 auto ®ion = dispatchOp.getRegion();
893 builder.SetInsertPoint(*
result);
901 auto maskedOp = cast<omp::MaskedOp>(opInst);
902 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
907 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
910 auto ®ion = maskedOp.getRegion();
911 builder.restoreIP(codeGenIP);
919 auto finiCB = [&](InsertPointTy codeGenIP) {
return llvm::Error::success(); };
921 llvm::Value *filterVal =
nullptr;
922 if (
auto filterVar = maskedOp.getFilteredThreadId()) {
923 filterVal = moduleTranslation.
lookupValue(filterVar);
925 llvm::LLVMContext &llvmContext = builder.getContext();
927 llvm::ConstantInt::get(llvm::Type::getInt32Ty(llvmContext), 0);
929 assert(filterVal !=
nullptr);
930 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
931 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
938 builder.restoreIP(*afterIP);
946 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
947 auto masterOp = cast<omp::MasterOp>(opInst);
952 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
955 auto ®ion = masterOp.getRegion();
956 builder.restoreIP(codeGenIP);
964 auto finiCB = [&](InsertPointTy codeGenIP) {
return llvm::Error::success(); };
966 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
967 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
974 builder.restoreIP(*afterIP);
982 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
983 auto criticalOp = cast<omp::CriticalOp>(opInst);
988 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
991 auto ®ion = cast<omp::CriticalOp>(opInst).getRegion();
992 builder.restoreIP(codeGenIP);
1000 auto finiCB = [&](InsertPointTy codeGenIP) {
return llvm::Error::success(); };
1002 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
1003 llvm::LLVMContext &llvmContext = moduleTranslation.
getLLVMContext();
1004 llvm::Constant *hint =
nullptr;
1007 if (criticalOp.getNameAttr()) {
1010 auto symbolRef = cast<SymbolRefAttr>(criticalOp.getNameAttr());
1011 auto criticalDeclareOp =
1015 llvm::ConstantInt::get(llvm::Type::getInt32Ty(llvmContext),
1016 static_cast<int>(criticalDeclareOp.getHint()));
1018 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
1020 ompLoc, bodyGenCB, finiCB, criticalOp.getName().value_or(
""), hint);
1025 builder.restoreIP(*afterIP);
1037 template <
typename OP>
1040 cast<
omp::BlockArgOpenMPOpInterface>(*op).getPrivateBlockArgs()) {
1043 collectPrivatizationDecls<OP>(op);
1060 void collectPrivatizationDecls(OP op) {
1061 std::optional<ArrayAttr> attr = op.getPrivateSyms();
1066 for (
auto symbolRef : attr->getAsRange<SymbolRefAttr>()) {
1073template <
typename T>
1077 std::optional<ArrayAttr> attr = op.getReductionSyms();
1081 reductions.reserve(reductions.size() + op.getNumReductionVars());
1082 for (
auto symbolRef : attr->getAsRange<SymbolRefAttr>()) {
1083 reductions.push_back(
1098 Operation *contextOp, std::optional<ArrayAttr> syms, StringRef opName,
1102 out.reserve(out.size() + syms->size());
1103 for (
auto sym : syms->getAsRange<SymbolRefAttr>()) {
1108 <<
"failed to resolve " << clauseName
1109 <<
" declare_reduction symbol " << sym.getRootReference() <<
" in "
1111 if (decl.getInitializerRegion().front().getNumArguments() != 1)
1113 <<
"not yet implemented: " << clauseName
1114 <<
" with two-argument initializer in " << opName;
1115 if (!decl.getCleanupRegion().empty())
1116 return contextOp->
emitError() <<
"not yet implemented: " << clauseName
1117 <<
" with cleanup region in " << opName;
1118 if (decl.getReductionRegion().empty())
1120 << clauseName <<
" declare_reduction is missing a combiner region";
1121 out.push_back(decl);
1132 Region ®ion, StringRef blockName, llvm::IRBuilderBase &builder,
1141 llvm::Instruction *potentialTerminator =
1142 builder.GetInsertBlock()->empty() ?
nullptr
1143 : &builder.GetInsertBlock()->back();
1145 if (potentialTerminator && potentialTerminator->isTerminator())
1146 potentialTerminator->removeFromParent();
1147 moduleTranslation.
mapBlock(®ion.
front(), builder.GetInsertBlock());
1150 region.
front(),
true, builder)))
1154 if (continuationBlockArgs)
1156 *continuationBlockArgs,
1163 if (potentialTerminator && potentialTerminator->isTerminator()) {
1164 llvm::BasicBlock *block = builder.GetInsertBlock();
1165 if (block->empty()) {
1171 potentialTerminator->insertInto(block, block->begin());
1173 potentialTerminator->insertAfter(&block->back());
1187 if (continuationBlockArgs)
1188 llvm::append_range(*continuationBlockArgs, phis);
1189 builder.SetInsertPoint((*continuationBlock)->getFirstInsertionPt());
1196using OwningReductionGen =
1197 std::function<llvm::OpenMPIRBuilder::InsertPointOrErrorTy(
1198 llvm::OpenMPIRBuilder::InsertPointTy, llvm::Value *, llvm::Value *,
1200using OwningAtomicReductionGen =
1201 std::function<llvm::OpenMPIRBuilder::InsertPointOrErrorTy(
1202 llvm::OpenMPIRBuilder::InsertPointTy, llvm::Type *, llvm::Value *,
1204using OwningDataPtrPtrReductionGen =
1205 std::function<llvm::OpenMPIRBuilder::InsertPointOrErrorTy(
1206 llvm::OpenMPIRBuilder::InsertPointTy, llvm::Value *, llvm::Value *&)>;
1212static OwningReductionGen
1218 OwningReductionGen gen =
1219 [&, decl](llvm::OpenMPIRBuilder::InsertPointTy insertPoint,
1220 llvm::Value *lhs, llvm::Value *rhs,
1221 llvm::Value *&
result)
mutable
1222 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
1223 moduleTranslation.
mapValue(decl.getReductionLhsArg(), lhs);
1224 moduleTranslation.
mapValue(decl.getReductionRhsArg(), rhs);
1225 builder.restoreIP(insertPoint);
1228 "omp.reduction.nonatomic.body", builder,
1229 moduleTranslation, &phis)))
1230 return llvm::createStringError(
1231 "failed to inline `combiner` region of `omp.declare_reduction`");
1232 result = llvm::getSingleElement(phis);
1233 return builder.saveIP();
1242static OwningAtomicReductionGen
1244 llvm::IRBuilderBase &builder,
1246 if (decl.getAtomicReductionRegion().empty())
1247 return OwningAtomicReductionGen();
1253 [&, decl](llvm::OpenMPIRBuilder::InsertPointTy insertPoint, llvm::Type *,
1254 llvm::Value *lhs, llvm::Value *rhs)
mutable
1255 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
1256 moduleTranslation.
mapValue(decl.getAtomicReductionLhsArg(), lhs);
1257 moduleTranslation.
mapValue(decl.getAtomicReductionRhsArg(), rhs);
1258 builder.restoreIP(insertPoint);
1261 "omp.reduction.atomic.body", builder,
1262 moduleTranslation, &phis)))
1263 return llvm::createStringError(
1264 "failed to inline `atomic` region of `omp.declare_reduction`");
1265 assert(phis.empty());
1266 return builder.saveIP();
1275static OwningDataPtrPtrReductionGen
1278 if (!isByRef || decl.getDataPtrPtrRegion().empty())
1279 return OwningDataPtrPtrReductionGen();
1281 OwningDataPtrPtrReductionGen refDataPtrGen =
1282 [&, decl](llvm::OpenMPIRBuilder::InsertPointTy insertPoint,
1283 llvm::Value *byRefVal, llvm::Value *&
result)
mutable
1284 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
1285 moduleTranslation.
mapValue(decl.getDataPtrPtrRegionArg(), byRefVal);
1286 builder.restoreIP(insertPoint);
1289 "omp.data_ptr_ptr.body", builder,
1290 moduleTranslation, &phis)))
1291 return llvm::createStringError(
1292 "failed to inline `data_ptr_ptr` region of `omp.declare_reduction`");
1293 result = llvm::getSingleElement(phis);
1294 return builder.saveIP();
1297 return refDataPtrGen;
1304 auto orderedOp = cast<omp::OrderedOp>(opInst);
1309 omp::ClauseDepend dependType = *orderedOp.getDoacrossDependType();
1310 bool isDependSource = dependType == omp::ClauseDepend::dependsource;
1311 unsigned numLoops = *orderedOp.getDoacrossNumLoops();
1313 moduleTranslation.
lookupValues(orderedOp.getDoacrossDependVars());
1315 size_t indexVecValues = 0;
1316 while (indexVecValues < vecValues.size()) {
1318 storeValues.reserve(numLoops);
1319 for (
unsigned i = 0; i < numLoops; i++) {
1320 storeValues.push_back(vecValues[indexVecValues]);
1323 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
1325 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
1326 builder.restoreIP(moduleTranslation.
getOpenMPBuilder()->createOrderedDepend(
1327 ompLoc, allocaIP, numLoops, storeValues,
".cnt.addr", isDependSource));
1337 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
1338 auto orderedRegionOp = cast<omp::OrderedRegionOp>(opInst);
1343 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
1346 auto ®ion = cast<omp::OrderedRegionOp>(opInst).getRegion();
1347 builder.restoreIP(codeGenIP);
1355 auto finiCB = [&](InsertPointTy codeGenIP) {
return llvm::Error::success(); };
1357 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
1358 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
1360 ompLoc, bodyGenCB, finiCB, !orderedRegionOp.getParLevelSimd());
1365 builder.restoreIP(*afterIP);
1371struct DeferredStore {
1372 DeferredStore(llvm::Value *value, llvm::Value *address)
1373 : value(value), address(address) {}
1376 llvm::Value *address;
1383template <
typename T>
1386 llvm::IRBuilderBase &builder,
1388 const llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1394 llvm::IRBuilderBase::InsertPointGuard guard(builder);
1395 builder.SetInsertPoint(allocaIP.getNodeParent()->getTerminator());
1401 deferredStores.reserve(op.getNumReductionVars());
1403 for (std::size_t i = 0; i < op.getNumReductionVars(); ++i) {
1404 Region &allocRegion = reductionDecls[i].getAllocRegion();
1406 if (allocRegion.
empty())
1411 builder, moduleTranslation, &phis)))
1412 return op.emitError(
1413 "failed to inline `alloc` region of `omp.declare_reduction`");
1415 assert(phis.size() == 1 &&
"expected one allocation to be yielded");
1416 builder.SetInsertPoint(allocaIP.getNodeParent()->getTerminator());
1420 llvm::Type *ptrTy = builder.getPtrTy();
1424 if (useDeviceSharedMem) {
1425 var = ompBuilder->createOMPAllocShared(builder, varTy);
1427 var = builder.CreateAlloca(varTy);
1428 var = builder.CreatePointerBitCastOrAddrSpaceCast(var, ptrTy);
1431 llvm::Value *castPhi =
1432 builder.CreatePointerBitCastOrAddrSpaceCast(phis[0], ptrTy);
1434 deferredStores.emplace_back(castPhi, var);
1436 privateReductionVariables[i] = var;
1437 moduleTranslation.
mapValue(reductionArgs[i], castPhi);
1438 reductionVariableMap.try_emplace(op.getReductionVars()[i], castPhi);
1440 assert(allocRegion.
empty() &&
1441 "allocaction is implicit for by-val reduction");
1443 llvm::Type *ptrTy = builder.getPtrTy();
1447 if (useDeviceSharedMem) {
1448 var = ompBuilder->createOMPAllocShared(builder, varTy);
1450 var = builder.CreateAlloca(varTy);
1451 var = builder.CreatePointerBitCastOrAddrSpaceCast(var, ptrTy);
1454 moduleTranslation.
mapValue(reductionArgs[i], var);
1455 privateReductionVariables[i] = var;
1456 reductionVariableMap.try_emplace(op.getReductionVars()[i], var);
1464template <
typename T>
1467 llvm::IRBuilderBase &builder,
1472 mlir::omp::DeclareReductionOp &reduction = reductionDecls[i];
1473 Region &initializerRegion = reduction.getInitializerRegion();
1476 mlir::Value mlirSource = loop.getReductionVars()[i];
1477 llvm::Value *llvmSource = moduleTranslation.
lookupValue(mlirSource);
1478 llvm::Value *origVal = llvmSource;
1480 if (!isa<LLVM::LLVMPointerType>(
1481 reduction.getInitializerMoldArg().getType()) &&
1482 isa<LLVM::LLVMPointerType>(mlirSource.
getType())) {
1485 reduction.getInitializerMoldArg().getType()),
1486 llvmSource,
"omp_orig");
1488 moduleTranslation.
mapValue(reduction.getInitializerMoldArg(), origVal);
1491 llvm::Value *allocation =
1492 reductionVariableMap.lookup(loop.getReductionVars()[i]);
1493 moduleTranslation.
mapValue(reduction.getInitializerAllocArg(), allocation);
1499 llvm::BasicBlock *block =
nullptr) {
1500 if (block ==
nullptr)
1501 block = builder.GetInsertBlock();
1503 if (!block->hasTerminator())
1504 builder.SetInsertPoint(block);
1506 builder.SetInsertPoint(block->getTerminator());
1514template <
typename OP>
1517 llvm::IRBuilderBase &builder,
1519 llvm::BasicBlock *latestAllocaBlock,
1525 if (op.getNumReductionVars() == 0)
1531 llvm::BasicBlock *initBlock = splitBB(builder,
true,
"omp.reduction.init");
1532 auto allocaIP = latestAllocaBlock->getTerminator()->getIterator();
1533 builder.restoreIP(allocaIP);
1536 for (
unsigned i = 0; i < op.getNumReductionVars(); ++i) {
1538 if (!reductionDecls[i].getAllocRegion().empty())
1546 if (useDeviceSharedMem)
1547 byRefVars[i] = ompBuilder->createOMPAllocShared(builder, varTy);
1549 byRefVars[i] = builder.CreateAlloca(varTy);
1557 for (
auto [data, addr] : deferredStores)
1558 builder.CreateStore(data, addr);
1563 for (
unsigned i = 0; i < op.getNumReductionVars(); ++i) {
1568 reductionVariableMap, i);
1576 "omp.reduction.neutral", builder,
1577 moduleTranslation, &phis)))
1580 assert(phis.size() == 1 &&
"expected one value to be yielded from the "
1581 "reduction neutral element declaration region");
1586 if (!reductionDecls[i].getAllocRegion().empty())
1595 builder.CreateStore(phis[0], byRefVars[i]);
1597 privateReductionVariables[i] = byRefVars[i];
1598 moduleTranslation.
mapValue(reductionArgs[i], phis[0]);
1599 reductionVariableMap.try_emplace(op.getReductionVars()[i], phis[0]);
1602 builder.CreateStore(phis[0], privateReductionVariables[i]);
1609 moduleTranslation.
forgetMapping(reductionDecls[i].getInitializerRegion());
1616template <
typename T>
1617static void collectReductionInfo(
1618 T loop, llvm::IRBuilderBase &builder,
1627 unsigned numReductions = loop.getNumReductionVars();
1629 for (
unsigned i = 0; i < numReductions; ++i) {
1632 owningAtomicReductionGens.push_back(
1635 reductionDecls[i], builder, moduleTranslation, isByRef[i]));
1639 reductionInfos.reserve(numReductions);
1640 for (
unsigned i = 0; i < numReductions; ++i) {
1641 llvm::OpenMPIRBuilder::ReductionGenAtomicCBTy
atomicGen =
nullptr;
1642 if (owningAtomicReductionGens[i])
1643 atomicGen = owningAtomicReductionGens[i];
1644 llvm::Value *variable =
1645 moduleTranslation.
lookupValue(loop.getReductionVars()[i]);
1648 if (
auto alloca = mlir::dyn_cast<LLVM::AllocaOp>(op)) {
1649 allocatedType = alloca.getElemType();
1656 reductionInfos.push_back(
1658 privateReductionVariables[i],
1659 llvm::OpenMPIRBuilder::EvalKind::Scalar,
1663 allocatedType ? moduleTranslation.
convertType(allocatedType) :
nullptr,
1664 reductionDecls[i].getByrefElementType()
1666 *reductionDecls[i].getByrefElementType())
1676 llvm::IRBuilderBase &builder, StringRef regionName,
1677 bool shouldLoadCleanupRegionArg =
true) {
1678 for (
auto [i, cleanupRegion] : llvm::enumerate(cleanupRegions)) {
1679 if (cleanupRegion->empty())
1685 llvm::Instruction *potentialTerminator =
1686 builder.GetInsertBlock()->empty() ?
nullptr
1687 : &builder.GetInsertBlock()->back();
1688 if (potentialTerminator && potentialTerminator->isTerminator())
1689 builder.SetInsertPoint(potentialTerminator);
1690 llvm::Value *privateVarValue =
1691 shouldLoadCleanupRegionArg
1692 ? builder.CreateLoad(
1694 privateVariables[i])
1695 : privateVariables[i];
1700 moduleTranslation)))
1713 OP op, llvm::IRBuilderBase &builder,
1715 llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1718 bool isNowait =
false,
bool isTeamsReduction =
false) {
1720 if (op.getNumReductionVars() == 0)
1732 collectReductionInfo(op, builder, moduleTranslation, reductionDecls,
1734 owningReductionGenRefDataPtrGens,
1735 privateReductionVariables, reductionInfos, isByRef);
1740 llvm::UnreachableInst *tempTerminator = builder.CreateUnreachable();
1741 builder.SetInsertPoint(tempTerminator);
1742 llvm::DebugLoc reductionLoc = builder.getCurrentDebugLocation();
1743 llvm::OpenMPIRBuilder::InsertPointOrErrorTy contInsertPoint =
1744 ompBuilder->createReductions(builder, allocaIP, reductionInfos, isByRef,
1745 isNowait, isTeamsReduction);
1750 if (!contInsertPoint->isValid())
1751 return op->emitOpError() <<
"failed to convert reductions";
1753 llvm::OpenMPIRBuilder::InsertPointTy afterIP = *contInsertPoint;
1754 if (!isTeamsReduction) {
1755 llvm::OpenMPIRBuilder::InsertPointOrErrorTy barrierIP =
1756 ompBuilder->createBarrier({*contInsertPoint, reductionLoc},
1757 llvm::omp::OMPD_for);
1761 afterIP = *barrierIP;
1764 tempTerminator->eraseFromParent();
1765 builder.restoreIP(afterIP);
1769 llvm::transform(reductionDecls, std::back_inserter(reductionRegions),
1770 [](omp::DeclareReductionOp reductionDecl) {
1771 return &reductionDecl.getCleanupRegion();
1774 reductionRegions, privateReductionVariables, moduleTranslation, builder,
1775 "omp.reduction.cleanup");
1778 if (useDeviceSharedMem) {
1779 for (
auto [var, reductionDecl] :
1780 llvm::zip_equal(privateReductionVariables, reductionDecls))
1781 ompBuilder->createOMPFreeShared(
1782 builder, var, moduleTranslation.
convertType(reductionDecl.getType()));
1795template <
typename OP>
1799 llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1804 if (op.getNumReductionVars() == 0)
1810 allocaIP, reductionDecls,
1811 privateReductionVariables, reductionVariableMap,
1812 deferredStores, isByRef)))
1815 return initReductionVars(op, reductionArgs, builder, moduleTranslation,
1816 allocaIP.getNodeParent(), reductionDecls,
1817 privateReductionVariables, reductionVariableMap,
1818 isByRef, deferredStores);
1832 if (mappedPrivateVars ==
nullptr || !mappedPrivateVars->contains(privateVar))
1835 Value blockArg = (*mappedPrivateVars)[privateVar];
1838 assert(isa<LLVM::LLVMPointerType>(blockArgType) &&
1839 "A block argument corresponding to a mapped var should have "
1842 if (privVarType == blockArgType)
1849 if (!isa<LLVM::LLVMPointerType>(privVarType))
1850 return builder.CreateLoad(moduleTranslation.
convertType(privVarType),
1867 llvm::Type *regionArgType =
1869 if (regionArgType->isPointerTy() || !value->getType()->isPointerTy())
1872 return builder.CreateLoad(regionArgType, value);
1882 omp::PrivateClauseOp &privDecl, llvm::Value *nonPrivateVar,
1884 llvm::BasicBlock *privInitBlock,
1886 Region &initRegion = privDecl.getInitRegion();
1887 if (initRegion.
empty())
1888 return llvmPrivateVar;
1890 assert(nonPrivateVar);
1891 moduleTranslation.
mapValue(privDecl.getInitMoldArg(), nonPrivateVar);
1892 moduleTranslation.
mapValue(privDecl.getInitPrivateArg(), llvmPrivateVar);
1897 moduleTranslation, &phis)))
1898 return llvm::createStringError(
1899 "failed to inline `init` region of `omp.private`");
1901 assert(phis.size() == 1 &&
"expected one allocation to be yielded");
1918 llvm::Value *llvmPrivateVar, llvm::BasicBlock *privInitBlock,
1921 builder, moduleTranslation, privDecl,
1924 blockArg, llvmPrivateVar, privInitBlock, mappedPrivateVars);
1933 return llvm::Error::success();
1935 llvm::BasicBlock *privInitBlock = splitBB(builder,
true,
"omp.private.init");
1938 for (
auto [idx, zip] : llvm::enumerate(llvm::zip_equal(
1941 auto [privDecl, mlirPrivVar, blockArg, llvmPrivateVar] = zip;
1943 builder, moduleTranslation, privDecl, mlirPrivVar, blockArg,
1944 llvmPrivateVar, privInitBlock, mappedPrivateVars);
1947 return privVarOrErr.takeError();
1949 llvmPrivateVar = privVarOrErr.get();
1950 moduleTranslation.
mapValue(blockArg, llvmPrivateVar);
1955 return llvm::Error::success();
1960 llvm::IRBuilderBase &builder,
1963 for (
Value allocatorVar : allocatorVars) {
1967 llvm::Value *allocator = moduleTranslation.
lookupValue(allocatorVar);
1969 return op.
emitError(
"failed to translate OpenMP allocator operand");
1970 if (allocator->getType()->isIntegerTy())
1971 allocator = builder.CreateIntToPtr(allocator, builder.getPtrTy());
1972 else if (allocator->getType()->isPointerTy())
1973 allocator = builder.CreatePointerBitCastOrAddrSpaceCast(
1974 allocator, builder.getPtrTy());
1977 "OpenMP allocator operand must have integer or pointer type");
1987template <
typename T>
1989 T op, llvm::IRBuilderBase &builder,
1992 llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1994 std::optional<llvm::OpenMPIRBuilder::InsertPointTy> allocatorIP =
1999 llvm::BasicBlock *allocaBB = allocaIP.getNodeParent();
2000 llvm::Instruction *allocaTerminator = allocaBB->getTerminator();
2001 splitBB(allocaTerminator->getIterator(),
true,
2002 allocaTerminator->getStableDebugLoc(),
"omp.region.after_alloca");
2004 allocaTerminator = allocaBB->getTerminator();
2006 assert(allocaTerminator->getNumSuccessors() == 1 &&
2007 "This is an unconditional branch created by splitBB");
2008 allocaIP = allocaTerminator->getIterator();
2010 llvm::Instruction *allocatorTerminator =
nullptr;
2011 llvm::BasicBlock *afterAllocatorAllocations =
nullptr;
2013 llvm::BasicBlock *allocatorBB = allocatorIP->getNodeParent();
2014 allocatorTerminator = allocatorBB->getTerminator();
2015 afterAllocatorAllocations = splitBB(
2016 allocatorTerminator->getIterator(),
true,
2017 allocatorTerminator->getStableDebugLoc(),
"omp.region.after_allocate");
2018 allocatorTerminator = allocatorBB->getTerminator();
2019 assert(allocatorTerminator->getNumSuccessors() == 1 &&
2020 "This is an unconditional branch created by splitBB");
2023 std::optional<llvm::IRBuilderBase::InsertPointGuard> guard;
2025 guard.emplace(builder);
2026 builder.SetInsertPoint(allocaTerminator);
2028 llvm::DataLayout dataLayout = builder.GetInsertBlock()->getDataLayout();
2029 llvm::BasicBlock *afterAllocas = allocaTerminator->getSuccessor(0);
2033 unsigned int allocaAS =
2034 moduleTranslation.
getLLVMModule()->getDataLayout().getAllocaAddrSpace();
2037 .getProgramAddressSpace();
2043 if constexpr (std::is_same_v<T, omp::ParallelOp> ||
2044 std::is_same_v<T, omp::ScopeOp>) {
2045 allocatorVars = op.getAllocatorVars();
2046 allocateAlignments = op.getAllocateAlignmentsAttr();
2047 if (
auto privateIndices = op.getAllocatePrivateIndicesAttr())
2048 for (
auto [allocateIndex, privateIndex] :
2049 llvm::enumerate(privateIndices.asArrayRef()))
2050 allocateItemForPrivate[privateIndex] = allocateIndex;
2053 for (
auto [privateIndex, tuple] : llvm::enumerate(llvm::zip_equal(
2056 auto [privDecl, mlirPrivVar, blockArg] = tuple;
2057 llvm::Type *llvmAllocType =
2058 moduleTranslation.
convertType(privDecl.getType());
2059 llvm::Value *llvmPrivateVar =
nullptr;
2060 int64_t allocateIndex = allocateItemForPrivate[privateIndex];
2061 builder.SetInsertPoint(allocateIndex >= 0 && allocatorTerminator
2062 ? allocatorTerminator
2063 : allocaTerminator);
2064 if (allocateIndex >= 0) {
2065 if (ompBuilder->Config.isTargetDevice() ||
2066 op->template getParentOfType<omp::TargetOp>())
2067 return llvm::createStringError(
2068 "allocate clause in an OpenMP device context is not supported");
2069 if (!llvmAllocType->isSized())
2070 return llvm::createStringError(
2071 "allocate clause private type must have a fixed size");
2072 llvm::TypeSize size = dataLayout.getTypeAllocSize(llvmAllocType);
2073 if (size.isScalable())
2074 return llvm::createStringError(
2075 "allocate clause private type must have a fixed size");
2076 llvm::IntegerType *sizeTy =
2077 moduleTranslation.
getLLVMModule()->getDataLayout().getIntPtrType(
2079 if (!llvm::isUIntN(sizeTy->getBitWidth(), size.getFixedValue()))
2080 return llvm::createStringError(
2081 "OpenMP allocation size cannot be represented by the target size "
2083 llvm::Value *sizeValue =
2084 llvm::ConstantInt::get(sizeTy, size.getFixedValue());
2086 Value allocatorVar = allocatorVars[allocateIndex];
2089 return llvm::createStringError(
2090 "failed to find converted OpenMP allocator operand");
2091 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2093 allocateAlignments ? allocateAlignments[allocateIndex] : 0;
2094 if (alignment != 0) {
2098 uint64_t alignmentValue = std::max<uint64_t>(
2099 static_cast<uint64_t
>(alignment),
2100 dataLayout.getABITypeAlign(llvmAllocType).value());
2101 if (!llvm::isUIntN(sizeTy->getBitWidth(), alignmentValue))
2102 return llvm::createStringError(
2103 "OpenMP allocation alignment cannot be represented by the "
2104 "target size type");
2105 llvmPrivateVar = ompBuilder->createOMPAlignedAlloc(
2106 ompLoc, llvm::ConstantInt::get(sizeTy, alignmentValue), sizeValue,
2107 allocator->second,
"omp.private.alloc");
2109 llvmPrivateVar = ompBuilder->createOMPAlloc(
2110 ompLoc, sizeValue, allocator->second,
"omp.private.alloc");
2112 if (!llvmPrivateVar)
2113 return llvm::createStringError(
2114 "failed to create OpenMP private allocation");
2116 {llvmPrivateVar, allocator->second});
2117 }
else if (mightUseDeviceSharedMem &&
2119 llvmPrivateVar = ompBuilder->createOMPAllocShared(builder, llvmAllocType);
2121 llvmPrivateVar = builder.CreateAlloca(
2122 llvmAllocType,
nullptr,
"omp.private.alloc");
2123 if (allocaAS != defaultAS)
2124 llvmPrivateVar = builder.CreateAddrSpaceCast(
2125 llvmPrivateVar, builder.getPtrTy(defaultAS));
2128 privateVarsInfo.
llvmVars.push_back(llvmPrivateVar);
2131 return afterAllocatorAllocations ? afterAllocatorAllocations : afterAllocas;
2139 if (mlir::isa<omp::SingleOp, omp::CriticalOp>(parent))
2148 if (mlir::isa<omp::ParallelOp>(parent))
2162 bool needsFirstprivate =
2163 llvm::any_of(privateDecls, [](omp::PrivateClauseOp &privOp) {
2164 return privOp.getDataSharingType() ==
2165 omp::DataSharingClauseType::FirstPrivate;
2168 if (!needsFirstprivate)
2171 llvm::BasicBlock *copyBlock =
2172 splitBB(builder,
true,
"omp.private.copy");
2175 for (
auto [decl, moldVar, llvmVar] :
2176 llvm::zip_equal(privateDecls, moldVars, llvmPrivateVars)) {
2177 if (decl.getDataSharingType() != omp::DataSharingClauseType::FirstPrivate)
2181 Region ©Region = decl.getCopyRegion();
2184 builder, moduleTranslation, decl.getCopyMoldArg(), moldVar);
2186 builder, moduleTranslation, decl.getCopyPrivateArg(), llvmVar);
2188 moduleTranslation.
mapValue(decl.getCopyMoldArg(), copyMoldVar);
2191 moduleTranslation.
mapValue(decl.getCopyPrivateArg(), copyPrivateVar);
2195 moduleTranslation)))
2196 return decl.emitError(
"failed to inline `copy` region of `omp.private`");
2210 llvm::OpenMPIRBuilder::InsertPointOrErrorTy res =
2211 ompBuilder->createBarrier(builder, llvm::omp::OMPD_barrier);
2227 llvm::transform(mlirPrivateVars, moldVars.begin(), [&](
mlir::Value mlirVar) {
2229 llvm::Value *moldVar = findAssociatedValue(
2230 mlirVar, builder, moduleTranslation, mappedPrivateVars);
2235 llvmPrivateVars, privateDecls, insertBarrier,
2239template <
typename T>
2247 std::back_inserter(privateCleanupRegions),
2248 [](omp::PrivateClauseOp privatizer) {
2249 return &privatizer.getDeallocRegion();
2253 privateVarsInfo.
llvmVars, moduleTranslation,
2254 builder,
"omp.private.dealloc",
2256 return mlir::emitError(loc,
"failed to inline `dealloc` region of an "
2257 "`omp.private` op in");
2262 for (
auto [privDecl, llvmPrivVar, blockArg] :
2266 ompBuilder->createOMPFreeShared(
2267 builder, llvmPrivVar,
2268 moduleTranslation.
convertType(privDecl.getType()));
2272 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2275 ompBuilder->createOMPFree(ompLoc, allocation.allocatedPtr,
2276 allocation.allocator);
2288 if (mlir::isa<omp::CancelOp, omp::CancellationPointOp>(child))
2305 llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder::InsertPointTy allocaIP,
2307 bool isWorksharing =
false);
2315 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2316 using StorableBodyGenCallbackTy =
2317 llvm::OpenMPIRBuilder::StorableBodyGenCallbackTy;
2319 auto sectionsOp = cast<omp::SectionsOp>(opInst);
2325 assert(isByRef.size() == sectionsOp.getNumReductionVars());
2329 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
2333 sectionsOp.getNumReductionVars());
2337 cast<omp::BlockArgOpenMPOpInterface>(opInst).getReductionBlockArgs();
2340 sectionsOp, reductionArgs, builder, moduleTranslation, allocaIP,
2341 reductionDecls, privateReductionVariables, reductionVariableMap,
2345 bool isTaskReductionMod =
2346 sectionsOp.getReductionMod() == omp::ReductionModifier::task &&
2347 sectionsOp.getNumReductionVars() > 0;
2352 auto sectionOp = dyn_cast<omp::SectionOp>(op);
2356 Region ®ion = sectionOp.getRegion();
2357 auto sectionCB = [§ionsOp, ®ion, &builder, &moduleTranslation](
2358 InsertPointTy allocaIP, InsertPointTy codeGenIP,
2360 builder.restoreIP(codeGenIP);
2367 sectionsOp.getRegion().getNumArguments());
2368 for (
auto [sectionsArg, sectionArg] : llvm::zip_equal(
2369 sectionsOp.getRegion().getArguments(), region.
getArguments())) {
2370 llvm::Value *llvmVal = moduleTranslation.
lookupValue(sectionsArg);
2372 moduleTranslation.
mapValue(sectionArg, llvmVal);
2379 sectionCBs.push_back(sectionCB);
2385 if (sectionCBs.empty())
2393 if (isTaskReductionMod &&
2395 "__omp_taskred_mod_", builder, allocaIP,
2396 moduleTranslation,
true,
2398 return sectionsOp.emitError(
2399 "failed to emit task reduction modifier initialization");
2401 assert(isa<omp::SectionOp>(*sectionsOp.getRegion().op_begin()));
2406 auto privCB = [&](InsertPointTy, InsertPointTy codeGenIP, llvm::Value &,
2407 llvm::Value &vPtr, llvm::Value *&replacementValue)
2408 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
2409 replacementValue = &vPtr;
2415 auto finiCB = [&](InsertPointTy codeGenIP) {
return llvm::Error::success(); };
2419 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2420 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2422 ompLoc, allocaIP, sectionCBs, privCB, finiCB, isCancellable,
2423 sectionsOp.getNowait());
2428 builder.restoreIP(*afterIP);
2431 if (isTaskReductionMod)
2437 sectionsOp, builder, moduleTranslation, allocaIP, reductionDecls,
2438 privateReductionVariables, isByRef, sectionsOp.getNowait());
2445 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2452 assert(isByRef.size() == scopeOp.getNumReductionVars());
2456 moduleTranslation, privateVarsInfo)))
2461 InsertPointTy privateAllocaIP =
2465 scopeOp.getNumReductionVars());
2469 cast<omp::BlockArgOpenMPOpInterface>(*scopeOp).getReductionBlockArgs();
2471 if (scopeOp.getAllocateVars().empty()) {
2473 scopeOp, builder, moduleTranslation, privateVarsInfo, privateAllocaIP);
2479 scopeOp, reductionArgs, builder, moduleTranslation, privateAllocaIP,
2480 reductionDecls, privateReductionVariables, reductionVariableMap,
2485 [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
2487 if (!scopeOp.getAllocateVars().empty()) {
2491 scopeOp, builder, moduleTranslation, privateVarsInfo, privateAllocaIP,
2492 nullptr, codeGenIP);
2494 return llvm::make_error<PreviouslyReportedError>();
2495 builder.SetInsertPoint(afterAllocas.get()->getTerminator());
2497 builder.restoreIP(codeGenIP);
2504 return llvm::make_error<PreviouslyReportedError>();
2507 scopeOp, builder, moduleTranslation, privateVarsInfo.
mlirVars,
2509 scopeOp.getPrivateNeedsBarrier())))
2510 return llvm::make_error<PreviouslyReportedError>();
2517 auto finiCB = [&](InsertPointTy codeGenIP) -> llvm::Error {
2518 InsertPointTy oldIP = builder.saveIP();
2519 builder.restoreIP(codeGenIP);
2521 scopeOp.getLoc(), privateVarsInfo)))
2522 return llvm::make_error<PreviouslyReportedError>();
2523 builder.restoreIP(oldIP);
2524 return llvm::Error::success();
2527 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2528 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2529 ompBuilder->createScope(ompLoc, bodyCB, finiCB, scopeOp.getNowait());
2534 builder.restoreIP(*afterIP);
2538 scopeOp, builder, moduleTranslation, privateAllocaIP, reductionDecls,
2539 privateReductionVariables, isByRef, scopeOp.getNowait(),
2547 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2548 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2553 auto bodyCB = [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
2555 builder.restoreIP(codegenIP);
2557 builder, moduleTranslation)
2560 auto finiCB = [&](InsertPointTy codeGenIP) {
return llvm::Error::success(); };
2564 std::optional<ArrayAttr> cpFuncs = singleOp.getCopyprivateSyms();
2567 for (
size_t i = 0, e = cpVars.size(); i < e; ++i) {
2568 llvmCPVars.push_back(moduleTranslation.
lookupValue(cpVars[i]));
2570 singleOp, cast<SymbolRefAttr>((*cpFuncs)[i]));
2571 llvmCPFuncs.push_back(
2575 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2577 ompLoc, bodyCB, finiCB, singleOp.getNowait(), llvmCPVars,
2583 builder.restoreIP(*afterIP);
2587static omp::DistributeOp
2591 omp::DistributeOp distOp;
2592 WalkResult walk = teamsOp.getRegion().walk([&](omp::DistributeOp op) {
2598 if (walk.wasInterrupted() || !distOp)
2602 llvm::cast<mlir::omp::BlockArgOpenMPOpInterface>(teamsOp.getOperation());
2606 for (
auto ra : iface.getReductionBlockArgs())
2607 for (
auto &use : ra.getUses()) {
2608 auto *useOp = use.getOwner();
2610 if (mlir::isa<LLVM::DbgDeclareOp, LLVM::DbgValueOp>(useOp)) {
2611 debugUses.push_back(useOp);
2614 if (!distOp->isProperAncestor(useOp))
2621 for (
auto *use : debugUses)
2630 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2635 unsigned numReductionVars = op.getNumReductionVars();
2639 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
2645 if (doTeamsReduction) {
2646 isByRef =
getIsByRef(op.getReductionByref());
2648 assert(isByRef.size() == op.getNumReductionVars());
2651 llvm::cast<omp::BlockArgOpenMPOpInterface>(*op).getReductionBlockArgs();
2656 op, reductionArgs, builder, moduleTranslation, allocaIP,
2657 reductionDecls, privateReductionVariables, reductionVariableMap,
2662 auto bodyCB = [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
2665 moduleTranslation, allocaIP, deallocBlocks);
2666 builder.restoreIP(codegenIP);
2672 llvm::Value *numTeamsLower =
nullptr;
2673 if (
Value numTeamsLowerVar = op.getNumTeamsLower())
2674 numTeamsLower = moduleTranslation.
lookupValue(numTeamsLowerVar);
2676 llvm::Value *numTeamsUpper =
nullptr;
2677 if (!op.getNumTeamsUpperVars().empty())
2678 numTeamsUpper = moduleTranslation.
lookupValue(op.getNumTeams(0));
2680 llvm::Value *threadLimit =
nullptr;
2681 if (!op.getThreadLimitVars().empty())
2682 threadLimit = moduleTranslation.
lookupValue(op.getThreadLimit(0));
2684 llvm::Value *ifExpr =
nullptr;
2685 if (
Value ifVar = op.getIfExpr())
2688 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2689 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2691 ompLoc, bodyCB, numTeamsLower, numTeamsUpper, threadLimit, ifExpr);
2696 builder.restoreIP(*afterIP);
2697 if (doTeamsReduction) {
2700 op, builder, moduleTranslation, allocaIP, reductionDecls,
2701 privateReductionVariables, isByRef,
2707static llvm::omp::RTLDependenceKindTy
2710 case mlir::omp::ClauseTaskDepend::taskdependin:
2711 return llvm::omp::RTLDependenceKindTy::DepIn;
2715 case mlir::omp::ClauseTaskDepend::taskdependout:
2716 case mlir::omp::ClauseTaskDepend::taskdependinout:
2717 return llvm::omp::RTLDependenceKindTy::DepInOut;
2718 case mlir::omp::ClauseTaskDepend::taskdependmutexinoutset:
2719 return llvm::omp::RTLDependenceKindTy::DepMutexInOutSet;
2720 case mlir::omp::ClauseTaskDepend::taskdependinoutset:
2721 return llvm::omp::RTLDependenceKindTy::DepInOutSet;
2723 llvm_unreachable(
"unhandled depend kind");
2727 std::optional<ArrayAttr> dependKinds,
OperandRange dependVars,
2730 if (dependVars.empty())
2732 for (
auto dep : llvm::zip(dependVars, dependKinds->getValue())) {
2734 cast<mlir::omp::ClauseTaskDependAttr>(std::get<1>(dep)).getValue();
2736 llvm::Value *depVal = moduleTranslation.
lookupValue(std::get<0>(dep));
2737 llvm::OpenMPIRBuilder::DependData dd(type, depVal->getType(), depVal);
2738 dds.emplace_back(dd);
2750 llvm::IRBuilderBase &llvmBuilder, llvm::OpenMPIRBuilder &ompBuilder,
2752 auto finiCB = [&](llvm::OpenMPIRBuilder::InsertPointTy ip) -> llvm::Error {
2753 llvm::IRBuilderBase::InsertPointGuard guard(llvmBuilder);
2757 llvmBuilder.restoreIP(ip);
2763 cancelTerminators.push_back(llvmBuilder.CreateBr(ip.getNodeParent()));
2764 return llvm::Error::success();
2769 ompBuilder.pushFinalizationCB(
2779 llvm::OpenMPIRBuilder &ompBuilder,
2780 llvm::BasicBlock *afterBB) {
2781 ompBuilder.popFinalizationCB();
2782 llvm::BasicBlock *constructFini = afterBB->getSinglePredecessor();
2783 for (llvm::UncondBrInst *cancelBranch : cancelTerminators)
2784 cancelBranch->setSuccessor(constructFini);
2790class TaskContextStructManager {
2792 TaskContextStructManager(llvm::IRBuilderBase &builder,
2793 LLVM::ModuleTranslation &moduleTranslation,
2794 MutableArrayRef<omp::PrivateClauseOp> privateDecls)
2795 : builder{builder}, moduleTranslation{moduleTranslation},
2796 privateDecls{privateDecls} {}
2802 void generateTaskContextStruct();
2808 void createGEPsToPrivateVars();
2814 SmallVector<llvm::Value *>
2815 createGEPsToPrivateVars(llvm::Value *altStructPtr)
const;
2818 void freeStructPtr();
2820 MutableArrayRef<llvm::Value *> getLLVMPrivateVarGEPs() {
2821 return llvmPrivateVarGEPs;
2824 llvm::Value *getStructPtr() {
return structPtr; }
2827 llvm::IRBuilderBase &builder;
2828 LLVM::ModuleTranslation &moduleTranslation;
2829 MutableArrayRef<omp::PrivateClauseOp> privateDecls;
2832 SmallVector<llvm::Type *> privateVarTypes;
2836 SmallVector<llvm::Value *> llvmPrivateVarGEPs;
2839 llvm::Value *structPtr =
nullptr;
2841 llvm::Type *structTy =
nullptr;
2852 llvm::SmallVector<llvm::Value *> lowerBounds;
2853 llvm::SmallVector<llvm::Value *> upperBounds;
2854 llvm::SmallVector<llvm::Value *> steps;
2855 llvm::SmallVector<llvm::Value *> trips;
2857 llvm::Value *totalTrips;
2859 llvm::Value *lookUpAsI64(mlir::Value val,
const LLVM::ModuleTranslation &mt,
2860 llvm::IRBuilderBase &builder) {
2864 if (v->getType()->isIntegerTy(64))
2866 if (v->getType()->isIntegerTy())
2867 return builder.CreateSExtOrTrunc(v, builder.getInt64Ty());
2872 IteratorInfo(mlir::omp::IteratorOp itersOp,
2873 mlir::LLVM::ModuleTranslation &moduleTranslation,
2874 llvm::IRBuilderBase &builder) {
2875 dims = itersOp.getLoopLowerBounds().size();
2876 lowerBounds.resize(dims);
2877 upperBounds.resize(dims);
2881 for (
unsigned d = 0; d < dims; ++d) {
2882 llvm::Value *lb = lookUpAsI64(itersOp.getLoopLowerBounds()[d],
2883 moduleTranslation, builder);
2884 llvm::Value *ub = lookUpAsI64(itersOp.getLoopUpperBounds()[d],
2885 moduleTranslation, builder);
2887 lookUpAsI64(itersOp.getLoopSteps()[d], moduleTranslation, builder);
2888 assert(lb && ub && st &&
2889 "Expect lowerBounds, upperBounds, and steps in IteratorOp");
2890 assert((!llvm::isa<llvm::ConstantInt>(st) ||
2891 !llvm::cast<llvm::ConstantInt>(st)->isZero()) &&
2892 "Expect non-zero step in IteratorOp");
2894 lowerBounds[d] = lb;
2895 upperBounds[d] = ub;
2899 llvm::Value *diff = builder.CreateSub(ub, lb);
2900 llvm::Value *
div = builder.CreateSDiv(diff, st);
2901 trips[d] = builder.CreateAdd(
2902 div, llvm::ConstantInt::get(builder.getInt64Ty(), 1));
2905 totalTrips = llvm::ConstantInt::get(builder.getInt64Ty(), 1);
2906 for (
unsigned d = 0; d < dims; ++d)
2907 totalTrips = builder.CreateMul(totalTrips, trips[d]);
2910 unsigned getDims()
const {
return dims; }
2911 llvm::ArrayRef<llvm::Value *> getLowerBounds()
const {
return lowerBounds; }
2912 llvm::ArrayRef<llvm::Value *> getUpperBounds()
const {
return upperBounds; }
2913 llvm::ArrayRef<llvm::Value *> getSteps()
const {
return steps; }
2914 llvm::ArrayRef<llvm::Value *> getTrips()
const {
return trips; }
2915 llvm::Value *getTotalTrips()
const {
return totalTrips; }
2920void TaskContextStructManager::generateTaskContextStruct() {
2921 if (privateDecls.empty())
2923 privateVarTypes.reserve(privateDecls.size());
2925 for (omp::PrivateClauseOp &privOp : privateDecls) {
2928 if (!privOp.readsFromMold())
2930 Type mlirType = privOp.getType();
2931 privateVarTypes.push_back(moduleTranslation.
convertType(mlirType));
2934 if (privateVarTypes.empty())
2937 structTy = llvm::StructType::get(moduleTranslation.
getLLVMContext(),
2940 llvm::DataLayout dataLayout =
2941 builder.GetInsertBlock()->getModule()->getDataLayout();
2942 llvm::Type *intPtrTy = builder.getIntPtrTy(dataLayout);
2943 llvm::Value *allocSize =
2944 builder.CreateTypeSize(intPtrTy, dataLayout.getTypeAllocSize(structTy));
2947 structPtr = builder.CreateMalloc(intPtrTy, allocSize,
2949 "omp.task.context_ptr");
2952SmallVector<llvm::Value *> TaskContextStructManager::createGEPsToPrivateVars(
2953 llvm::Value *altStructPtr)
const {
2954 SmallVector<llvm::Value *> ret;
2957 ret.reserve(privateDecls.size());
2958 llvm::Value *zero = builder.getInt32(0);
2960 for (
auto privDecl : privateDecls) {
2961 if (!privDecl.readsFromMold()) {
2963 ret.push_back(
nullptr);
2966 llvm::Value *iVal = builder.getInt32(i);
2967 llvm::Value *gep = builder.CreateGEP(structTy, altStructPtr, {zero, iVal});
2974void TaskContextStructManager::createGEPsToPrivateVars() {
2976 assert(privateVarTypes.empty());
2980 llvmPrivateVarGEPs = createGEPsToPrivateVars(structPtr);
2983void TaskContextStructManager::freeStructPtr() {
2987 llvm::IRBuilderBase::InsertPointGuard guard{builder};
2989 builder.SetInsertPoint(builder.GetInsertBlock()->getTerminator());
2990 builder.CreateFree(structPtr);
2994 llvm::OpenMPIRBuilder &ompBuilder,
2995 llvm::Value *affinityList, llvm::Value *
index,
2996 llvm::Value *addr, llvm::Value *len) {
2997 llvm::StructType *kmpTaskAffinityInfoTy =
2998 ompBuilder.getKmpTaskAffinityInfoTy();
2999 llvm::Value *entry = builder.CreateInBoundsGEP(
3000 kmpTaskAffinityInfoTy, affinityList,
index,
"omp.affinity.entry");
3002 addr = builder.CreatePtrToInt(addr, kmpTaskAffinityInfoTy->getElementType(0));
3003 len = builder.CreateIntCast(len, kmpTaskAffinityInfoTy->getElementType(1),
3005 llvm::Value *flags = builder.getInt32(0);
3007 builder.CreateStore(addr,
3008 builder.CreateStructGEP(kmpTaskAffinityInfoTy, entry, 0));
3009 builder.CreateStore(len,
3010 builder.CreateStructGEP(kmpTaskAffinityInfoTy, entry, 1));
3011 builder.CreateStore(flags,
3012 builder.CreateStructGEP(kmpTaskAffinityInfoTy, entry, 2));
3016 llvm::IRBuilderBase &builder,
3018 llvm::Value *affinityList) {
3019 for (
auto [i, affinityVar] : llvm::enumerate(affinityVars)) {
3020 auto entryOp = affinityVar.getDefiningOp<mlir::omp::AffinityEntryOp>();
3021 assert(entryOp &&
"affinity item must be omp.affinity_entry");
3023 llvm::Value *addr = moduleTranslation.
lookupValue(entryOp.getAddr());
3024 llvm::Value *len = moduleTranslation.
lookupValue(entryOp.getLen());
3025 assert(addr && len &&
"expect affinity addr and len to be non-null");
3027 affinityList, builder.getInt64(i), addr, len);
3031static mlir::LogicalResult
3034 llvm::IRBuilderBase &builder,
3036 llvm::Value *tmp = linearIV;
3037 for (
int d = (
int)iterInfo.getDims() - 1; d >= 0; --d) {
3038 llvm::Value *trip = iterInfo.getTrips()[d];
3040 llvm::Value *idx = builder.CreateURem(tmp, trip);
3042 tmp = builder.CreateUDiv(tmp, trip);
3045 llvm::Value *physIV = builder.CreateAdd(
3046 iterInfo.getLowerBounds()[d],
3047 builder.CreateMul(idx, iterInfo.getSteps()[d]),
"omp.it.phys_iv");
3053 moduleTranslation.
mapBlock(&iteratorRegionBlock, builder.GetInsertBlock());
3054 if (mlir::failed(moduleTranslation.
convertBlock(iteratorRegionBlock,
3057 return mlir::failure();
3059 return mlir::success();
3065static mlir::LogicalResult
3068 IteratorInfo &iterInfo, llvm::StringRef loopName,
3073 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
3075 auto bodyGen = [&](llvm::OpenMPIRBuilder::InsertPointTy bodyIP,
3076 llvm::Value *linearIV) -> llvm::Error {
3077 llvm::IRBuilderBase::InsertPointGuard guard(builder);
3078 builder.restoreIP(bodyIP);
3081 builder, moduleTranslation))) {
3082 return llvm::make_error<llvm::StringError>(
3083 "failed to convert iterator region", llvm::inconvertibleErrorCode());
3087 mlir::dyn_cast<mlir::omp::YieldOp>(iteratorRegionBlock.
getTerminator());
3088 assert(yield && yield.getResults().size() == 1 &&
3089 "expect omp.yield in iterator region to have one result");
3091 genStoreEntry(linearIV, yield);
3097 return llvm::Error::success();
3100 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
3102 loc, iterInfo.getTotalTrips(), bodyGen, loopName);
3106 builder.restoreIP(*afterIP);
3108 return mlir::success();
3111static mlir::LogicalResult
3114 llvm::OpenMPIRBuilder::AffinityData &ad) {
3116 if (taskOp.getAffinityVars().empty() && taskOp.getIterated().empty()) {
3119 return mlir::success();
3123 llvm::StructType *kmpTaskAffinityInfoTy =
3126 auto allocateAffinityList = [&](llvm::Value *count) -> llvm::Value * {
3127 llvm::IRBuilderBase::InsertPointGuard guard(builder);
3128 if (llvm::isa<llvm::Constant>(count) || llvm::isa<llvm::Argument>(count))
3130 return builder.CreateAlloca(kmpTaskAffinityInfoTy, count,
3131 "omp.affinity_list");
3134 auto createAffinity =
3135 [&](llvm::Value *count,
3136 llvm::Value *info) -> llvm::OpenMPIRBuilder::AffinityData {
3137 llvm::OpenMPIRBuilder::AffinityData ad{};
3138 ad.Count = builder.CreateTrunc(count, builder.getInt32Ty());
3140 builder.CreatePointerBitCastOrAddrSpaceCast(info, builder.getPtrTy(0));
3144 if (!taskOp.getAffinityVars().empty()) {
3145 llvm::Value *count = llvm::ConstantInt::get(
3146 builder.getInt64Ty(), taskOp.getAffinityVars().size());
3147 llvm::Value *list = allocateAffinityList(count);
3150 ads.emplace_back(createAffinity(count, list));
3153 if (!taskOp.getIterated().empty()) {
3154 for (
auto [i, iter] : llvm::enumerate(taskOp.getIterated())) {
3155 auto itersOp = iter.getDefiningOp<omp::IteratorOp>();
3156 assert(itersOp &&
"iterated value must be defined by omp.iterator");
3157 IteratorInfo iterInfo(itersOp, moduleTranslation, builder);
3158 llvm::Value *affList = allocateAffinityList(iterInfo.getTotalTrips());
3160 itersOp, builder, moduleTranslation, iterInfo,
"iterator",
3161 [&](llvm::Value *linearIV, mlir::omp::YieldOp yield) {
3162 auto entryOp = yield.getResults()[0]
3163 .getDefiningOp<mlir::omp::AffinityEntryOp>();
3164 assert(entryOp &&
"expect yield produce an affinity entry");
3171 affList, linearIV, addr, len);
3173 return llvm::failure();
3174 ads.emplace_back(createAffinity(iterInfo.getTotalTrips(), affList));
3178 llvm::Value *totalAffinityCount = builder.getInt32(0);
3179 for (
const auto &affinity : ads)
3180 totalAffinityCount = builder.CreateAdd(
3182 builder.CreateIntCast(affinity.Count, builder.getInt32Ty(),
3185 llvm::Value *affinityInfo = ads.front().Info;
3186 if (ads.size() > 1) {
3187 llvm::StructType *kmpTaskAffinityInfoTy =
3189 llvm::Value *affinityInfoElemSize = builder.getInt64(
3190 moduleTranslation.
getLLVMModule()->getDataLayout().getTypeAllocSize(
3191 kmpTaskAffinityInfoTy));
3193 llvm::Value *packedAffinityInfo = allocateAffinityList(totalAffinityCount);
3194 llvm::Value *packedAffinityInfoOffset = builder.getInt32(0);
3195 for (
const auto &affinity : ads) {
3196 llvm::Value *affinityCount = builder.CreateIntCast(
3197 affinity.Count, builder.getInt32Ty(),
false);
3198 llvm::Value *affinityCountInt64 = builder.CreateIntCast(
3199 affinityCount, builder.getInt64Ty(),
false);
3200 llvm::Value *affinityInfoSize =
3201 builder.CreateMul(affinityCountInt64, affinityInfoElemSize);
3203 llvm::Value *packedAffinityInfoIndex = builder.CreateIntCast(
3204 packedAffinityInfoOffset, kmpTaskAffinityInfoTy->getElementType(0),
3206 packedAffinityInfoIndex = builder.CreateInBoundsGEP(
3207 kmpTaskAffinityInfoTy, packedAffinityInfo, packedAffinityInfoIndex);
3209 builder.CreateMemCpy(
3210 packedAffinityInfoIndex, llvm::Align(1),
3211 builder.CreatePointerBitCastOrAddrSpaceCast(
3212 affinity.Info, builder.getPtrTy(packedAffinityInfoIndex->getType()
3213 ->getPointerAddressSpace())),
3214 llvm::Align(1), affinityInfoSize);
3216 packedAffinityInfoOffset =
3217 builder.CreateAdd(packedAffinityInfoOffset, affinityCount);
3220 affinityInfo = packedAffinityInfo;
3223 ad.Count = totalAffinityCount;
3224 ad.Info = affinityInfo;
3226 return mlir::success();
3232static mlir::LogicalResult
3235 std::optional<ArrayAttr> dependIteratedKinds,
3236 llvm::IRBuilderBase &builder,
3238 llvm::OpenMPIRBuilder::DependenciesInfo &taskDeps) {
3239 if (dependIterated.empty()) {
3242 return mlir::success();
3246 llvm::Type *dependInfoTy = ompBuilder.DependInfo;
3247 unsigned numLocator = dependVars.size();
3250 llvm::Value *totalCount =
3251 llvm::ConstantInt::get(builder.getInt64Ty(), numLocator);
3254 for (
auto iter : dependIterated) {
3255 auto itersOp = iter.getDefiningOp<mlir::omp::IteratorOp>();
3256 assert(itersOp &&
"depend_iterated value must be defined by omp.iterator");
3257 iterInfos.emplace_back(itersOp, moduleTranslation, builder);
3259 builder.CreateAdd(totalCount, iterInfos.back().getTotalTrips());
3264 llvm::DataLayout dataLayout =
3265 builder.GetInsertBlock()->getModule()->getDataLayout();
3266 llvm::Value *allocSize = builder.CreateTypeSize(
3267 ompBuilder.SizeTy, dataLayout.getTypeAllocSize(dependInfoTy));
3268 llvm::Value *depArray =
3269 builder.CreateMalloc(ompBuilder.SizeTy, allocSize, totalCount,
3270 nullptr,
".dep.arr.addr");
3273 if (numLocator > 0) {
3276 for (
auto [i, dd] : llvm::enumerate(dds)) {
3277 llvm::Value *idx = llvm::ConstantInt::get(builder.getInt64Ty(), i);
3278 llvm::Value *entry =
3279 builder.CreateInBoundsGEP(dependInfoTy, depArray, idx);
3280 ompBuilder.emitTaskDependency(builder, entry, dd);
3285 llvm::Value *offset =
3286 llvm::ConstantInt::get(builder.getInt64Ty(), numLocator);
3287 for (
auto [i, iterInfo] : llvm::enumerate(iterInfos)) {
3288 auto kindAttr = cast<mlir::omp::ClauseTaskDependAttr>(
3289 dependIteratedKinds->getValue()[i]);
3290 llvm::omp::RTLDependenceKindTy rtlKind =
3293 auto itersOp = dependIterated[i].getDefiningOp<mlir::omp::IteratorOp>();
3295 itersOp, builder, moduleTranslation, iterInfo,
"dep_iterator",
3296 [&](llvm::Value *linearIV, mlir::omp::YieldOp yield) {
3298 moduleTranslation.
lookupValue(yield.getResults()[0]);
3299 llvm::Value *idx = builder.CreateAdd(offset, linearIV);
3300 llvm::Value *entry =
3301 builder.CreateInBoundsGEP(dependInfoTy, depArray, idx);
3302 ompBuilder.emitTaskDependency(
3304 llvm::OpenMPIRBuilder::DependData{rtlKind, addr->getType(),
3307 return mlir::failure();
3310 offset = builder.CreateAdd(offset, iterInfo.getTotalTrips());
3313 taskDeps.DepArray = depArray;
3314 taskDeps.NumDeps = builder.CreateTrunc(totalCount, builder.getInt32Ty());
3315 return mlir::success();
3322 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
3327 TaskContextStructManager taskStructMgr{builder, moduleTranslation,
3339 InsertPointTy allocaIP =
3344 assert(builder.GetInsertPoint() == builder.GetInsertBlock()->end());
3345 llvm::BasicBlock *taskStartBlock = llvm::BasicBlock::Create(
3346 builder.getContext(),
"omp.task.start",
3347 builder.GetInsertBlock()->getParent());
3348 llvm::Instruction *branchToTaskStartBlock = builder.CreateBr(taskStartBlock);
3349 builder.SetInsertPoint(branchToTaskStartBlock);
3352 llvm::BasicBlock *copyBlock =
3353 splitBB(builder,
true,
"omp.private.copy");
3354 llvm::BasicBlock *initBlock =
3355 splitBB(builder,
true,
"omp.private.init");
3371 moduleTranslation, allocaIP, deallocBlocks);
3374 builder.SetInsertPoint(initBlock->getTerminator());
3377 taskStructMgr.generateTaskContextStruct();
3384 taskStructMgr.createGEPsToPrivateVars();
3386 for (
auto [privDecl, mlirPrivVar, blockArg, llvmPrivateVarAlloc] :
3389 taskStructMgr.getLLVMPrivateVarGEPs())) {
3391 if (!privDecl.readsFromMold())
3393 assert(llvmPrivateVarAlloc &&
3394 "reads from mold so shouldn't have been skipped");
3397 initPrivateVar(builder, moduleTranslation, privDecl, mlirPrivVar,
3398 blockArg, llvmPrivateVarAlloc, initBlock);
3399 if (!privateVarOrErr)
3400 return handleError(privateVarOrErr, *taskOp.getOperation());
3409 [[maybe_unused]] llvm::Value *llvmPrivateVar = llvmPrivateVarAlloc;
3410 if ((privateVarOrErr.get() != llvmPrivateVarAlloc) &&
3411 !mlir::isa<LLVM::LLVMPointerType>(blockArg.getType())) {
3412 builder.CreateStore(privateVarOrErr.get(), llvmPrivateVarAlloc);
3414 llvmPrivateVar = builder.CreateLoad(privateVarOrErr.get()->getType(),
3415 llvmPrivateVarAlloc);
3417 assert(llvmPrivateVar->getType() ==
3418 moduleTranslation.
convertType(blockArg.getType()));
3428 taskOp, builder, moduleTranslation, privateVarsInfo.
mlirVars,
3429 taskStructMgr.getLLVMPrivateVarGEPs(), privateVarsInfo.
privatizers,
3430 taskOp.getPrivateNeedsBarrier())))
3431 return llvm::failure();
3433 llvm::OpenMPIRBuilder::AffinityData ad;
3435 return llvm::failure();
3445 taskOp.getOperation(), taskOp.getInReductionSyms(),
"omp.task",
3446 "in_reduction", inRedDecls)))
3449 inRedOrigPtrs.reserve(inRedDecls.size());
3450 for (
Value v : taskOp.getInReductionVars())
3451 inRedOrigPtrs.push_back(moduleTranslation.
lookupValue(v));
3454 builder.SetInsertPoint(taskStartBlock);
3457 [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
3462 moduleTranslation, allocaIP, deallocBlocks);
3465 builder.restoreIP(codegenIP);
3467 llvm::BasicBlock *privInitBlock =
nullptr;
3469 for (
auto [i, zip] : llvm::enumerate(llvm::zip_equal(
3472 auto [blockArg, privDecl, mlirPrivVar] = zip;
3474 if (privDecl.readsFromMold())
3477 llvm::IRBuilderBase::InsertPointGuard guard(builder);
3478 llvm::Type *llvmAllocType =
3479 moduleTranslation.
convertType(privDecl.getType());
3480 builder.SetInsertPoint(allocaIP.getNodeParent()->getTerminator());
3481 llvm::Value *llvmPrivateVar = builder.CreateAlloca(
3482 llvmAllocType,
nullptr,
"omp.private.alloc");
3485 initPrivateVar(builder, moduleTranslation, privDecl, mlirPrivVar,
3486 blockArg, llvmPrivateVar, privInitBlock);
3487 if (!privateVarOrError)
3488 return privateVarOrError.takeError();
3489 moduleTranslation.
mapValue(blockArg, privateVarOrError.get());
3490 privateVarsInfo.
llvmVars[i] = privateVarOrError.get();
3493 taskStructMgr.createGEPsToPrivateVars();
3494 for (
auto [i, llvmPrivVar] :
3495 llvm::enumerate(taskStructMgr.getLLVMPrivateVarGEPs())) {
3497 assert(privateVarsInfo.
llvmVars[i] &&
3498 "This is added in the loop above");
3501 privateVarsInfo.
llvmVars[i] = llvmPrivVar;
3506 for (
auto [blockArg, llvmPrivateVar, privateDecl] :
3510 if (!privateDecl.readsFromMold())
3513 if (!mlir::isa<LLVM::LLVMPointerType>(blockArg.getType())) {
3514 llvmPrivateVar = builder.CreateLoad(
3515 moduleTranslation.
convertType(blockArg.getType()), llvmPrivateVar);
3517 assert(llvmPrivateVar->getType() ==
3518 moduleTranslation.
convertType(blockArg.getType()));
3519 moduleTranslation.
mapValue(blockArg, llvmPrivateVar);
3530 if (!inRedDecls.empty()) {
3531 auto iface = cast<omp::BlockArgOpenMPOpInterface>(taskOp.getOperation());
3534 llvm::LLVMContext &llvmCtx = m->getContext();
3535 llvm::OpenMPIRBuilder::LocationDescription bodyLoc(builder);
3536 uint32_t srcLocSize;
3537 llvm::Constant *srcLocStr =
3538 ompB.getOrCreateSrcLocStr(bodyLoc, srcLocSize);
3539 llvm::Value *bodyIdent = ompB.getOrCreateIdent(srcLocStr, srcLocSize);
3542 ompB.updateToLocation(bodyLoc);
3543 llvm::Value *bodyGtid = ompB.getOrCreateThreadID(bodyIdent);
3544 llvm::FunctionCallee getThData = ompB.getOrCreateRuntimeFunction(
3545 *m, llvm::omp::OMPRTL___kmpc_task_reduction_get_th_data);
3546 llvm::Type *ptrTy = llvm::PointerType::getUnqual(llvmCtx);
3547 llvm::Value *nullDesc = llvm::ConstantPointerNull::get(ptrTy);
3549 for (
auto [blockArg, origPtr] :
3550 llvm::zip_equal(inRedBlockArgs, inRedOrigPtrs)) {
3557 llvm::Value *lookupPtr = origPtr;
3558 if (
auto *origPtrTy =
3559 llvm::dyn_cast<llvm::PointerType>(lookupPtr->getType());
3560 origPtrTy && origPtrTy->getAddressSpace() != 0)
3561 lookupPtr = builder.CreateAddrSpaceCast(lookupPtr, ptrTy);
3562 llvm::Value *priv = builder.CreateCall(
3563 getThData, {bodyGtid, nullDesc, lookupPtr},
"omp.inred.priv");
3564 if (
auto *argPtrTy = llvm::dyn_cast<llvm::PointerType>(
3565 moduleTranslation.
convertType(blockArg.getType()));
3566 argPtrTy && argPtrTy->getAddressSpace() != 0)
3567 priv = builder.CreateAddrSpaceCast(priv, argPtrTy);
3568 moduleTranslation.
mapValue(blockArg, priv);
3573 taskOp.getRegion(),
"omp.task.region", builder, moduleTranslation);
3574 if (failed(
handleError(continuationBlockOrError, *taskOp)))
3575 return llvm::make_error<PreviouslyReportedError>();
3577 builder.SetInsertPoint(continuationBlockOrError.get()->getTerminator());
3580 taskOp.getLoc(), privateVarsInfo)))
3581 return llvm::make_error<PreviouslyReportedError>();
3584 taskStructMgr.freeStructPtr();
3586 return llvm::Error::success();
3595 llvm::omp::Directive::OMPD_taskgroup);
3597 llvm::OpenMPIRBuilder::DependenciesInfo dependencies;
3598 if (failed(
buildDependData(taskOp.getDependVars(), taskOp.getDependKinds(),
3599 taskOp.getDependIterated(),
3600 taskOp.getDependIteratedKinds(), builder,
3601 moduleTranslation, dependencies)))
3604 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
3605 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
3607 ompLoc, allocaIP, deallocBlocks, bodyCB, !taskOp.getUntied(),
3609 moduleTranslation.
lookupValue(taskOp.getIfExpr()), dependencies, ad,
3610 taskOp.getMergeable(),
3611 moduleTranslation.
lookupValue(taskOp.getEventHandle()),
3612 moduleTranslation.
lookupValue(taskOp.getPriority()),
3613 taskOp.getThreadset() == omp::ThreadsetPolicy::omp_pool);
3620 afterIP->getNodeParent());
3622 builder.restoreIP(*afterIP);
3624 if (dependencies.DepArray)
3625 builder.CreateFree(dependencies.DepArray);
3634 llvm::IRBuilderBase &builder,
3642 loopWrapperOp.getRegion(),
"omp.taskloop.wrapper.region", builder,
3645 if (failed(
handleError(continuationBlockOrError, opInst)))
3648 builder.SetInsertPoint(continuationBlockOrError.get());
3656static llvm::Expected<llvm::Value *>
3659 llvm::IRBuilderBase &builder) {
3660 if (llvm::Value *mapped = moduleTranslation.
lookupValue(value))
3665 return llvm::make_error<llvm::StringError>(
3666 "value is a block argument and is not mapped",
3667 llvm::inconvertibleErrorCode());
3669 return llvm::make_error<llvm::StringError>(
3670 "unsupported op defining taskloop loop bound",
3671 llvm::inconvertibleErrorCode());
3681 if (!operandOrError)
3682 return operandOrError.takeError();
3683 moduleTranslation.
mapValue(operand, *operandOrError);
3684 mappingsToRemove.push_back(operand);
3688 return llvm::make_error<llvm::StringError>(
3689 "failed to convert op defining taskloop loop bound",
3690 llvm::inconvertibleErrorCode());
3693 assert(
result &&
"expected conversion of loop bound op to produce a value");
3697 mappingsToRemove.push_back(resultValue);
3699 for (
Value mappedValue : mappingsToRemove)
3708 llvm::Value *&lbVal, llvm::Value *&ubVal,
3709 llvm::Value *&stepVal) {
3717 return firstLbOrErr.takeError();
3719 llvm::Type *boundType = (*firstLbOrErr)->getType();
3720 ubVal = builder.getIntN(boundType->getIntegerBitWidth(), 1);
3721 if (loopOp.getCollapseNumLoops() > 1) {
3739 for (uint64_t i = 0; i < loopOp.getCollapseNumLoops(); i++) {
3741 i == 0 ? std::move(firstLbOrErr)
3745 return lbOrErr.takeError();
3747 upperBounds[i], moduleTranslation, builder);
3749 return ubOrErr.takeError();
3753 return stepOrErr.takeError();
3755 llvm::Value *loopLb = *lbOrErr;
3756 llvm::Value *loopUb = *ubOrErr;
3757 llvm::Value *loopStep = *stepOrErr;
3763 llvm::Value *loopLbMinusOne = builder.CreateSub(
3764 loopLb, builder.getIntN(boundType->getIntegerBitWidth(), 1));
3765 llvm::Value *loopUbMinusOne = builder.CreateSub(
3766 loopUb, builder.getIntN(boundType->getIntegerBitWidth(), 1));
3767 llvm::Value *boundsCmp = builder.CreateICmpSLT(loopLb, loopUb);
3768 llvm::Value *ubMinusLb = builder.CreateSub(loopUb, loopLbMinusOne);
3769 llvm::Value *lbMinusUb = builder.CreateSub(loopLb, loopUbMinusOne);
3770 llvm::Value *loopTripCount =
3771 builder.CreateSelect(boundsCmp, ubMinusLb, lbMinusUb);
3772 loopTripCount = builder.CreateBinaryIntrinsic(
3773 llvm::Intrinsic::abs, loopTripCount, builder.getFalse());
3777 llvm::Value *loopTripCountDivStep =
3778 builder.CreateSDiv(loopTripCount, loopStep);
3779 loopTripCountDivStep = builder.CreateBinaryIntrinsic(
3780 llvm::Intrinsic::abs, loopTripCountDivStep, builder.getFalse());
3781 llvm::Value *loopTripCountRem =
3782 builder.CreateSRem(loopTripCount, loopStep);
3783 loopTripCountRem = builder.CreateBinaryIntrinsic(
3784 llvm::Intrinsic::abs, loopTripCountRem, builder.getFalse());
3785 llvm::Value *needsRoundUp = builder.CreateICmpNE(
3787 builder.getIntN(loopTripCountRem->getType()->getIntegerBitWidth(),
3790 builder.CreateAdd(loopTripCountDivStep,
3791 builder.CreateZExtOrTrunc(
3792 needsRoundUp, loopTripCountDivStep->getType()));
3793 ubVal = builder.CreateMul(ubVal, loopTripCount);
3795 lbVal = builder.getIntN(boundType->getIntegerBitWidth(), 1);
3796 stepVal = builder.getIntN(boundType->getIntegerBitWidth(), 1);
3801 return ubOrErr.takeError();
3805 return stepOrErr.takeError();
3806 lbVal = *firstLbOrErr;
3808 stepVal = *stepOrErr;
3811 assert(lbVal !=
nullptr &&
"Expected value for lbVal");
3812 assert(ubVal !=
nullptr &&
"Expected value for ubVal");
3813 assert(stepVal !=
nullptr &&
"Expected value for stepVal");
3814 return llvm::Error::success();
3820 llvm::IRBuilderBase &builder,
3822 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
3824 omp::TaskloopWrapperOp loopWrapperOp = contextOp.getLoopOp();
3832 TaskContextStructManager taskStructMgr{builder, moduleTranslation,
3836 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
3839 assert(builder.GetInsertPoint() == builder.GetInsertBlock()->end());
3840 llvm::BasicBlock *taskloopStartBlock = llvm::BasicBlock::Create(
3841 builder.getContext(),
"omp.taskloop.wrapper.start",
3842 builder.GetInsertBlock()->getParent());
3843 llvm::Instruction *branchToTaskloopStartBlock =
3844 builder.CreateBr(taskloopStartBlock);
3845 builder.SetInsertPoint(branchToTaskloopStartBlock);
3847 llvm::BasicBlock *copyBlock =
3848 splitBB(builder,
true,
"omp.private.copy");
3849 llvm::BasicBlock *initBlock =
3850 splitBB(builder,
true,
"omp.private.init");
3853 moduleTranslation, allocaIP, deallocBlocks);
3856 builder.SetInsertPoint(initBlock->getTerminator());
3859 taskStructMgr.generateTaskContextStruct();
3860 taskStructMgr.createGEPsToPrivateVars();
3862 llvmFirstPrivateVars.resize(privateVarsInfo.
blockArgs.size());
3864 for (
auto [i, zip] : llvm::enumerate(llvm::zip_equal(
3866 privateVarsInfo.
blockArgs, taskStructMgr.getLLVMPrivateVarGEPs()))) {
3867 auto [privDecl, mlirPrivVar, blockArg, llvmPrivateVarAlloc] = zip;
3869 if (!privDecl.readsFromMold())
3871 assert(llvmPrivateVarAlloc &&
3872 "reads from mold so shouldn't have been skipped");
3875 initPrivateVar(builder, moduleTranslation, privDecl, mlirPrivVar,
3876 blockArg, llvmPrivateVarAlloc, initBlock);
3877 if (!privateVarOrErr)
3878 return handleError(privateVarOrErr, *contextOp.getOperation());
3880 llvmFirstPrivateVars[i] = privateVarOrErr.get();
3882 llvm::IRBuilderBase::InsertPointGuard guard(builder);
3883 builder.SetInsertPoint(builder.GetInsertBlock()->getTerminator());
3885 [[maybe_unused]] llvm::Value *llvmPrivateVar = llvmPrivateVarAlloc;
3886 if ((privateVarOrErr.get() != llvmPrivateVarAlloc) &&
3887 !mlir::isa<LLVM::LLVMPointerType>(blockArg.getType())) {
3888 builder.CreateStore(privateVarOrErr.get(), llvmPrivateVarAlloc);
3890 llvmPrivateVar = builder.CreateLoad(privateVarOrErr.get()->getType(),
3891 llvmPrivateVarAlloc);
3893 assert(llvmPrivateVar->getType() ==
3894 moduleTranslation.
convertType(blockArg.getType()));
3900 contextOp, builder, moduleTranslation, privateVarsInfo.
mlirVars,
3901 taskStructMgr.getLLVMPrivateVarGEPs(), privateVarsInfo.
privatizers,
3902 contextOp.getPrivateNeedsBarrier())))
3903 return llvm::failure();
3913 contextOp.getOperation(), contextOp.getReductionSyms(),
3914 "omp.taskloop.context",
"reduction", redDecls)))
3918 contextOp.getOperation(), contextOp.getInReductionSyms(),
3919 "omp.taskloop.context",
"in_reduction", inRedDecls)))
3925 redOrigPtrs.reserve(redDecls.size());
3926 for (
Value v : contextOp.getReductionVars())
3927 redOrigPtrs.push_back(moduleTranslation.
lookupValue(v));
3929 inRedOrigPtrs.reserve(inRedDecls.size());
3930 for (
Value v : contextOp.getInReductionVars())
3931 inRedOrigPtrs.push_back(moduleTranslation.
lookupValue(v));
3935 builder.SetInsertPoint(taskloopStartBlock);
3937 llvm::OpenMPIRBuilder &ompBuilderRef = *moduleTranslation.
getOpenMPBuilder();
3944 bool implicitTaskgroup = !redDecls.empty();
3945 llvm::Value *redDesc =
nullptr;
3946 if (implicitTaskgroup) {
3947 llvm::OpenMPIRBuilder::LocationDescription redLoc(builder);
3948 uint32_t srcLocSize;
3949 llvm::Constant *srcLocStr =
3950 ompBuilderRef.getOrCreateSrcLocStr(redLoc, srcLocSize);
3951 llvm::Value *ident = ompBuilderRef.getOrCreateIdent(srcLocStr, srcLocSize);
3954 ompBuilderRef.updateToLocation(redLoc);
3955 llvm::Value *outerGtid = ompBuilderRef.getOrCreateThreadID(ident);
3956 llvm::FunctionCallee taskgroupFn = ompBuilderRef.getOrCreateRuntimeFunction(
3957 *module, llvm::omp::OMPRTL___kmpc_taskgroup);
3958 builder.CreateCall(taskgroupFn, {ident, outerGtid});
3961 "__omp_taskloop_taskred_", builder,
3962 allocaIP, moduleTranslation);
3967 auto loopOp = cast<omp::LoopNestOp>(loopWrapperOp.getWrappedLoop());
3968 llvm::Value *lbVal =
nullptr;
3969 llvm::Value *ubVal =
nullptr;
3970 llvm::Value *stepVal =
nullptr;
3972 loopOp, builder, moduleTranslation, lbVal, ubVal, stepVal))
3976 [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
3981 moduleTranslation, allocaIP, deallocBlocks);
3984 builder.restoreIP(codegenIP);
3986 llvm::BasicBlock *privInitBlock =
nullptr;
3988 for (
auto [i, zip] : llvm::enumerate(llvm::zip_equal(
3991 auto [blockArg, privDecl, mlirPrivVar] = zip;
3993 if (privDecl.readsFromMold())
3996 llvm::IRBuilderBase::InsertPointGuard guard(builder);
3997 llvm::Type *llvmAllocType =
3998 moduleTranslation.
convertType(privDecl.getType());
3999 builder.SetInsertPoint(allocaIP.getNodeParent()->getTerminator());
4000 llvm::Value *llvmPrivateVar = builder.CreateAlloca(
4001 llvmAllocType,
nullptr,
"omp.private.alloc");
4004 initPrivateVar(builder, moduleTranslation, privDecl, mlirPrivVar,
4005 blockArg, llvmPrivateVar, privInitBlock);
4006 if (!privateVarOrError)
4007 return privateVarOrError.takeError();
4008 moduleTranslation.
mapValue(blockArg, privateVarOrError.get());
4009 privateVarsInfo.
llvmVars[i] = privateVarOrError.get();
4012 taskStructMgr.createGEPsToPrivateVars();
4013 for (
auto [i, llvmPrivVar] :
4014 llvm::enumerate(taskStructMgr.getLLVMPrivateVarGEPs())) {
4016 assert(privateVarsInfo.
llvmVars[i] &&
4017 "This is added in the loop above");
4020 privateVarsInfo.
llvmVars[i] = llvmPrivVar;
4025 for (
auto [blockArg, llvmPrivateVar, privateDecl] :
4029 if (!privateDecl.readsFromMold())
4032 if (!mlir::isa<LLVM::LLVMPointerType>(blockArg.getType())) {
4033 llvmPrivateVar = builder.CreateLoad(
4034 moduleTranslation.
convertType(blockArg.getType()), llvmPrivateVar);
4036 assert(llvmPrivateVar->getType() ==
4037 moduleTranslation.
convertType(blockArg.getType()));
4038 moduleTranslation.
mapValue(blockArg, llvmPrivateVar);
4050 if (!redDecls.empty() || !inRedDecls.empty()) {
4052 cast<omp::BlockArgOpenMPOpInterface>(contextOp.getOperation());
4055 llvm::LLVMContext &llvmCtx = m->getContext();
4056 llvm::OpenMPIRBuilder::LocationDescription bodyLoc(builder);
4057 uint32_t srcLocSize;
4058 llvm::Constant *srcLocStr =
4059 ompB.getOrCreateSrcLocStr(bodyLoc, srcLocSize);
4060 llvm::Value *bodyIdent = ompB.getOrCreateIdent(srcLocStr, srcLocSize);
4063 ompB.updateToLocation(bodyLoc);
4064 llvm::Value *bodyGtid = ompB.getOrCreateThreadID(bodyIdent);
4065 llvm::FunctionCallee getThData = ompB.getOrCreateRuntimeFunction(
4066 *m, llvm::omp::OMPRTL___kmpc_task_reduction_get_th_data);
4067 llvm::Type *ptrTy = llvm::PointerType::getUnqual(llvmCtx);
4077 auto remapReductionArg = [&](
BlockArgument blockArg, llvm::Value *desc,
4078 llvm::Value *origPtr,
4079 const llvm::Twine &name) {
4080 if (
auto *origPtrTy =
4081 llvm::dyn_cast<llvm::PointerType>(origPtr->getType());
4082 origPtrTy && origPtrTy->getAddressSpace() != 0)
4083 origPtr = builder.CreateAddrSpaceCast(origPtr, ptrTy);
4085 builder.CreateCall(getThData, {bodyGtid, desc, origPtr}, name);
4086 if (
auto *argPtrTy = llvm::dyn_cast<llvm::PointerType>(
4088 argPtrTy && argPtrTy->getAddressSpace() != 0)
4089 priv = builder.CreateAddrSpaceCast(priv, argPtrTy);
4090 moduleTranslation.
mapValue(blockArg, priv);
4094 for (
auto [blockArg, origPtr] :
4095 llvm::zip_equal(redBlockArgs, redOrigPtrs))
4096 remapReductionArg(blockArg, redDesc, origPtr,
"omp.taskred.priv");
4098 llvm::Value *nullDesc = llvm::ConstantPointerNull::get(ptrTy);
4099 for (
auto [blockArg, origPtr] :
4100 llvm::zip_equal(inRedBlockArgs, inRedOrigPtrs))
4101 remapReductionArg(blockArg, nullDesc, origPtr,
"omp.inred.priv");
4107 contextOp.getRegion(),
"omp.taskloop.context.region", builder,
4110 if (failed(
handleError(continuationBlockOrError, opInst)))
4111 return llvm::make_error<PreviouslyReportedError>();
4113 builder.SetInsertPoint(continuationBlockOrError.get()->getTerminator());
4121 contextOp.getLoc(), privateVarsInfo)))
4122 return llvm::make_error<PreviouslyReportedError>();
4125 taskStructMgr.freeStructPtr();
4127 return llvm::Error::success();
4133 auto taskDupCB = [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
4134 llvm::Value *destPtr, llvm::Value *srcPtr)
4136 llvm::IRBuilderBase::InsertPointGuard guard(builder);
4137 builder.restoreIP(codegenIP);
4140 builder.getPtrTy(srcPtr->getType()->getPointerAddressSpace());
4142 builder.CreateLoad(ptrTy, srcPtr,
"omp.taskloop.context.src");
4144 TaskContextStructManager &srcStructMgr = taskStructMgr;
4145 TaskContextStructManager destStructMgr(builder, moduleTranslation,
4147 destStructMgr.generateTaskContextStruct();
4148 llvm::Value *dest = destStructMgr.getStructPtr();
4149 dest->setName(
"omp.taskloop.context.dest");
4150 builder.CreateStore(dest, destPtr);
4153 srcStructMgr.createGEPsToPrivateVars(src);
4155 destStructMgr.createGEPsToPrivateVars(dest);
4158 for (
auto [privDecl, mold, blockArg, llvmPrivateVarAlloc] :
4159 llvm::zip_equal(privateVarsInfo.
privatizers, srcGEPs,
4162 if (!privDecl.readsFromMold())
4164 assert(llvmPrivateVarAlloc &&
4165 "reads from mold so shouldn't have been skipped");
4168 builder, moduleTranslation, privDecl.getInitMoldArg(), mold);
4170 builder, moduleTranslation, privDecl, moldArg, blockArg,
4171 llvmPrivateVarAlloc, builder.GetInsertBlock());
4172 if (!privateVarOrErr)
4173 return privateVarOrErr.takeError();
4182 [[maybe_unused]] llvm::Value *llvmPrivateVar = llvmPrivateVarAlloc;
4183 if ((privateVarOrErr.get() != llvmPrivateVarAlloc) &&
4184 !mlir::isa<LLVM::LLVMPointerType>(blockArg.getType())) {
4185 builder.CreateStore(privateVarOrErr.get(), llvmPrivateVarAlloc);
4187 llvmPrivateVar = builder.CreateLoad(privateVarOrErr.get()->getType(),
4188 llvmPrivateVarAlloc);
4190 assert(llvmPrivateVar->getType() ==
4191 moduleTranslation.
convertType(blockArg.getType()));
4199 moduleTranslation, srcGEPs, destGEPs,
4201 contextOp.getPrivateNeedsBarrier())))
4202 return llvm::make_error<PreviouslyReportedError>();
4204 return builder.saveIP();
4212 llvm::Value *ifCond =
nullptr;
4213 llvm::Value *grainsize =
nullptr;
4215 mlir::Value grainsizeVal = contextOp.getGrainsize();
4216 mlir::Value numTasksVal = contextOp.getNumTasks();
4217 if (
Value ifVar = contextOp.getIfExpr())
4220 grainsize = moduleTranslation.
lookupValue(grainsizeVal);
4222 }
else if (numTasksVal) {
4223 grainsize = moduleTranslation.
lookupValue(numTasksVal);
4227 llvm::OpenMPIRBuilder::TaskDupCallbackTy taskDupOrNull =
nullptr;
4228 if (taskStructMgr.getStructPtr())
4229 taskDupOrNull = taskDupCB;
4239 llvm::omp::Directive::OMPD_taskgroup);
4241 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4242 bool effectiveNoGroup = contextOp.getNogroup() || implicitTaskgroup;
4243 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
4245 ompLoc, allocaIP, deallocBlocks, bodyCB, loopInfo, lbVal, ubVal,
4246 stepVal, contextOp.getUntied(), ifCond, grainsize, effectiveNoGroup,
4247 sched, moduleTranslation.
lookupValue(contextOp.getFinal()),
4248 contextOp.getMergeable(),
4249 moduleTranslation.
lookupValue(contextOp.getPriority()),
4250 loopOp.getCollapseNumLoops(), taskDupOrNull,
4251 taskStructMgr.getStructPtr(),
4252 contextOp.getThreadset() == omp::ThreadsetPolicy::omp_pool);
4258 afterIP->getNodeParent());
4260 builder.restoreIP(*afterIP);
4264 if (implicitTaskgroup) {
4265 llvm::OpenMPIRBuilder::LocationDescription endLoc(builder);
4266 uint32_t srcLocSize;
4267 llvm::Constant *srcLocStr =
4268 ompBuilder.getOrCreateSrcLocStr(endLoc, srcLocSize);
4269 llvm::Value *ident = ompBuilder.getOrCreateIdent(srcLocStr, srcLocSize);
4272 ompBuilder.updateToLocation(endLoc);
4273 llvm::Value *outerGtid = ompBuilder.getOrCreateThreadID(ident);
4274 llvm::FunctionCallee endTgFn = ompBuilder.getOrCreateRuntimeFunction(
4276 llvm::omp::OMPRTL___kmpc_end_taskgroup);
4277 builder.CreateCall(endTgFn, {ident, outerGtid});
4288static llvm::Function *
4291 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
4292 llvm::LLVMContext &ctx = llvmModule->getContext();
4293 llvm::Type *voidTy = llvm::Type::getVoidTy(ctx);
4294 llvm::Type *ptrTy = llvm::PointerType::getUnqual(ctx);
4295 llvm::FunctionType *fty =
4296 llvm::FunctionType::get(voidTy, {ptrTy, ptrTy},
false);
4297 llvm::Function *fn =
4298 llvm::Function::Create(fty, llvm::GlobalValue::InternalLinkage,
4299 baseName +
".red.init", llvmModule);
4300 fn->setDoesNotRecurse();
4301 fn->getArg(0)->setName(
"priv");
4302 fn->getArg(1)->setName(
"orig");
4304 llvm::BasicBlock *entry = llvm::BasicBlock::Create(ctx,
"entry", fn);
4305 llvm::IRBuilder<>
b(entry);
4312 Value moldArg = decl.getInitializerMoldArg();
4313 llvm::Value *origVal = fn->getArg(1);
4314 if (!isa<LLVM::LLVMPointerType>(moldArg.
getType()))
4316 fn->getArg(1),
"omp.orig");
4317 moduleTranslation.
mapValue(moldArg, origVal);
4320 "omp.taskred.init",
b, moduleTranslation,
4322 fn->eraseFromParent();
4325 assert(phis.size() == 1 &&
4326 "expected one value yielded from reduction initializer");
4327 b.CreateStore(phis[0], fn->getArg(0));
4330 moduleTranslation.
forgetMapping(decl.getInitializerRegion());
4338static llvm::Function *
4341 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
4342 llvm::LLVMContext &ctx = llvmModule->getContext();
4343 llvm::Type *voidTy = llvm::Type::getVoidTy(ctx);
4344 llvm::Type *ptrTy = llvm::PointerType::getUnqual(ctx);
4345 llvm::FunctionType *fty =
4346 llvm::FunctionType::get(voidTy, {ptrTy, ptrTy},
false);
4347 llvm::Function *fn =
4348 llvm::Function::Create(fty, llvm::GlobalValue::InternalLinkage,
4349 baseName +
".red.comb", llvmModule);
4350 fn->setDoesNotRecurse();
4351 fn->getArg(0)->setName(
"lhs");
4352 fn->getArg(1)->setName(
"rhs");
4354 llvm::BasicBlock *entry = llvm::BasicBlock::Create(ctx,
"entry", fn);
4355 llvm::IRBuilder<>
b(entry);
4357 llvm::Type *elemTy = moduleTranslation.
convertType(decl.getType());
4358 Block &combBlock = decl.getReductionRegion().
front();
4360 "expected two arguments in declare_reduction combiner");
4361 llvm::Value *lhsVal =
b.CreateLoad(elemTy, fn->getArg(0),
"omp.lhs");
4362 llvm::Value *rhsVal =
b.CreateLoad(elemTy, fn->getArg(1),
"omp.rhs");
4368 "omp.taskred.comb",
b, moduleTranslation,
4370 fn->eraseFromParent();
4373 assert(phis.size() == 1 &&
4374 "expected one value yielded from reduction combiner");
4375 b.CreateStore(phis[0], fn->getArg(0));
4401 llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder::InsertPointTy allocaIP,
4403 bool isWorksharing) {
4404 assert(redDecls.size() == origPtrs.size() &&
4405 "expected one orig pointer per reduction decl");
4407 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
4408 llvm::LLVMContext &ctx = llvmModule->getContext();
4409 const llvm::DataLayout &dl = llvmModule->getDataLayout();
4411 llvm::Type *ptrTy = llvm::PointerType::getUnqual(ctx);
4412 llvm::Type *i32Ty = llvm::Type::getInt32Ty(ctx);
4413 llvm::Type *sizeTy =
4414 llvm::Type::getIntNTy(ctx, dl.getPointerSizeInBits(0));
4418 llvm::StructType *redInputTy =
4419 llvm::StructType::getTypeByName(ctx,
"kmp_taskred_input_t");
4421 redInputTy = llvm::StructType::create(
4422 ctx, {ptrTy, ptrTy, sizeTy, ptrTy, ptrTy, ptrTy, i32Ty},
4423 "kmp_taskred_input_t");
4425 unsigned n = redDecls.size();
4426 llvm::ArrayType *arrTy = llvm::ArrayType::get(redInputTy, n);
4429 llvm::AllocaInst *arrAlloca;
4431 llvm::IRBuilderBase::InsertPointGuard guard(builder);
4432 builder.restoreIP(allocaIP);
4434 builder.CreateAlloca(arrTy,
nullptr,
".taskred.input");
4438 llvm::Value *zero = builder.getInt32(0);
4439 for (
unsigned i = 0; i < n; ++i) {
4440 omp::DeclareReductionOp decl = redDecls[i];
4441 llvm::Value *orig = origPtrs[i];
4442 if (
auto *origPtrTy = llvm::dyn_cast<llvm::PointerType>(orig->getType());
4443 origPtrTy && origPtrTy->getAddressSpace() != 0)
4444 orig = builder.CreateAddrSpaceCast(orig, ptrTy);
4445 llvm::Type *elemTy = moduleTranslation.
convertType(decl.getType());
4446 uint64_t size = dl.getTypeAllocSize(elemTy).getFixedValue();
4448 std::string baseName =
4449 (llvm::Twine(helperNamePrefix) + decl.getSymName()).str();
4450 llvm::Function *initFn =
4452 llvm::Function *combFn =
4454 if (!initFn || !combFn)
4456 llvm::Value *elemPtr = builder.CreateInBoundsGEP(
4457 arrTy, arrAlloca, {zero, builder.getInt32(i)},
".taskred.elem");
4458 auto storeField = [&](
unsigned fieldIdx, llvm::Value *val) {
4459 llvm::Value *fieldPtr =
4460 builder.CreateStructGEP(redInputTy, elemPtr, fieldIdx);
4461 builder.CreateStore(val, fieldPtr);
4463 storeField(0, orig);
4464 storeField(1, orig);
4465 storeField(2, llvm::ConstantInt::get(sizeTy, size));
4466 storeField(3, initFn);
4467 storeField(4, llvm::ConstantPointerNull::get(ptrTy));
4468 storeField(5, combFn);
4469 storeField(6, llvm::ConstantInt::get(i32Ty, 0));
4473 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4474 uint32_t srcLocSize;
4475 llvm::Constant *srcLocStr =
4476 ompBuilder->getOrCreateSrcLocStr(ompLoc, srcLocSize);
4477 llvm::Value *ident = ompBuilder->getOrCreateIdent(srcLocStr, srcLocSize);
4478 ompBuilder->updateToLocation(ompLoc);
4479 llvm::Value *gtid = ompBuilder->getOrCreateThreadID(ident);
4483 llvm::FunctionCallee modInit = ompBuilder->getOrCreateRuntimeFunction(
4484 *llvmModule, llvm::omp::OMPRTL___kmpc_taskred_modifier_init);
4485 return builder.CreateCall(modInit,
4487 builder.getInt32(isWorksharing ? 1 : 0),
4488 builder.getInt32(n), arrAlloca},
4492 llvm::FunctionCallee taskredInit = ompBuilder->getOrCreateRuntimeFunction(
4493 *llvmModule, llvm::omp::OMPRTL___kmpc_taskred_init);
4494 return builder.CreateCall(taskredInit, {gtid, builder.getInt32(n), arrAlloca},
4505 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
4506 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4507 uint32_t srcLocSize;
4508 llvm::Constant *srcLocStr =
4509 ompBuilder->getOrCreateSrcLocStr(ompLoc, srcLocSize);
4510 llvm::Value *ident = ompBuilder->getOrCreateIdent(srcLocStr, srcLocSize);
4511 ompBuilder->updateToLocation(ompLoc);
4512 llvm::Value *gtid = ompBuilder->getOrCreateThreadID(ident);
4513 llvm::FunctionCallee fini = ompBuilder->getOrCreateRuntimeFunction(
4514 *llvmModule, llvm::omp::OMPRTL___kmpc_task_reduction_modifier_fini);
4515 builder.CreateCall(fini,
4516 {ident, gtid, builder.getInt32(isWorksharing ? 1 : 0)});
4523 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
4532 if (
auto syms = tgOp.getTaskReductionSyms()) {
4533 redDecls.reserve(syms->size());
4534 for (
auto sym : syms->getAsRange<SymbolRefAttr>()) {
4538 return tgOp.emitError()
4539 <<
"failed to resolve task_reduction declare_reduction symbol "
4540 << sym.getRootReference() <<
" in omp.taskgroup";
4541 if (decl.getInitializerRegion().front().getNumArguments() != 1)
4542 return tgOp.emitError(
"not yet implemented: task_reduction with "
4543 "two-argument initializer in omp.taskgroup");
4544 if (!decl.getCleanupRegion().empty())
4545 return tgOp.emitError(
"not yet implemented: task_reduction with "
4546 "cleanup region in omp.taskgroup");
4547 if (decl.getReductionRegion().empty())
4548 return tgOp.emitError(
"task_reduction declare_reduction is missing a "
4550 redDecls.push_back(decl);
4555 [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
4557 builder.restoreIP(codegenIP);
4559 if (!redDecls.empty()) {
4561 origPtrs.reserve(redDecls.size());
4562 for (
Value v : tgOp.getTaskReductionVars())
4563 origPtrs.push_back(moduleTranslation.
lookupValue(v));
4565 builder, allocaIP, moduleTranslation))
4566 return llvm::createStringError(
4567 llvm::inconvertibleErrorCode(),
4568 "failed to emit task_reduction initialization for omp.taskgroup");
4576 for (
auto [i, blockArg] :
4577 llvm::enumerate(tgOp.getRegion().getArguments())) {
4579 moduleTranslation.
lookupValue(tgOp.getTaskReductionVars()[i]);
4580 moduleTranslation.
mapValue(blockArg, orig);
4584 builder, moduleTranslation)
4589 InsertPointTy allocaIP =
4591 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4592 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
4594 ompLoc, allocaIP, deallocBlocks, bodyCB);
4599 builder.restoreIP(*afterIP);
4606 if (!initOp.getDependVars().empty() || initOp.getDependKinds() ||
4607 !initOp.getDependIterated().empty() || initOp.getDependIteratedKinds())
4608 return initOp.emitError()
4609 <<
"not yet implemented: Unhandled clause depend in "
4610 << omp::InteropInitOp::getOperationName() <<
" operation";
4613 llvm::Value *interopVar =
4614 moduleTranslation.
lookupValue(initOp.getInteropVar());
4615 llvm::Value *device = initOp.getDevice()
4616 ? moduleTranslation.
lookupValue(initOp.getDevice())
4620 llvm::Value *numDeps = llvm::ConstantInt::get(builder.getInt32Ty(), 0);
4621 llvm::Value *depArray = llvm::ConstantPointerNull::get(builder.getPtrTy());
4622 bool hasNowait = initOp.getNowait();
4629 bool hasTarget =
false, hasTargetSync =
false;
4631 switch (cast<omp::InteropTypeAttr>(typeAttr).getValue()) {
4632 case omp::InteropType::target:
4635 case omp::InteropType::targetsync:
4636 hasTargetSync =
true;
4640 llvm::omp::OMPInteropType interopType =
4641 (!hasTarget && hasTargetSync) ? llvm::omp::OMPInteropType::TargetSync
4642 : llvm::omp::OMPInteropType::Target;
4643 ompBuilder->createOMPInteropInit(builder, interopVar, interopType, device,
4644 numDeps, depArray, hasNowait);
4650 llvm::IRBuilderBase &builder,
4652 if (!destroyOp.getDependVars().empty() || destroyOp.getDependKinds() ||
4653 !destroyOp.getDependIterated().empty() ||
4654 destroyOp.getDependIteratedKinds())
4655 return destroyOp.emitError()
4656 <<
"not yet implemented: Unhandled clause depend in "
4657 << omp::InteropDestroyOp::getOperationName() <<
" operation";
4660 llvm::Value *interopVar =
4661 moduleTranslation.
lookupValue(destroyOp.getInteropVar());
4662 llvm::Value *device =
4663 destroyOp.getDevice()
4664 ? moduleTranslation.
lookupValue(destroyOp.getDevice())
4667 llvm::Value *numDeps = llvm::ConstantInt::get(builder.getInt32Ty(), 0);
4668 llvm::Value *depArray = llvm::ConstantPointerNull::get(builder.getPtrTy());
4669 bool hasNowait = destroyOp.getNowait();
4671 ompBuilder->createOMPInteropDestroy(builder, interopVar, device, numDeps,
4672 depArray, hasNowait);
4679 if (!useOp.getDependVars().empty() || useOp.getDependKinds() ||
4680 !useOp.getDependIterated().empty() || useOp.getDependIteratedKinds())
4681 return useOp.emitError()
4682 <<
"not yet implemented: Unhandled clause depend in "
4683 << omp::InteropUseOp::getOperationName() <<
" operation";
4686 llvm::Value *interopVar =
4687 moduleTranslation.
lookupValue(useOp.getInteropVar());
4688 llvm::Value *device = useOp.getDevice()
4689 ? moduleTranslation.
lookupValue(useOp.getDevice())
4692 llvm::Value *numDeps = llvm::ConstantInt::get(builder.getInt32Ty(), 0);
4693 llvm::Value *depArray = llvm::ConstantPointerNull::get(builder.getPtrTy());
4694 bool hasNowait = useOp.getNowait();
4696 ompBuilder->createOMPInteropUse(builder, interopVar, device, numDeps,
4697 depArray, hasNowait);
4707 llvm::OpenMPIRBuilder::DependenciesInfo dds;
4709 twOp.getDependVars(), twOp.getDependKinds(), twOp.getDependIterated(),
4710 twOp.getDependIteratedKinds(), builder, moduleTranslation, dds))) {
4717 builder.CreateFree(dds.DepArray);
4728 auto wsloopOp = cast<omp::WsloopOp>(opInst);
4732 auto loopOp = cast<omp::LoopNestOp>(wsloopOp.getWrappedLoop());
4734 assert(isByRef.size() == wsloopOp.getNumReductionVars());
4738 wsloopOp.getScheduleKind().value_or(omp::ClauseScheduleKind::Static);
4741 llvm::Value *step = moduleTranslation.
lookupValue(loopOp.getLoopSteps()[0]);
4742 llvm::Type *ivType = step->getType();
4743 llvm::Value *chunk =
nullptr;
4744 if (wsloopOp.getScheduleChunk()) {
4745 llvm::Value *chunkVar =
4746 moduleTranslation.
lookupValue(wsloopOp.getScheduleChunk());
4747 chunk = builder.CreateSExtOrTrunc(chunkVar, ivType);
4750 omp::DistributeOp distributeOp =
nullptr;
4751 llvm::Value *distScheduleChunk =
nullptr;
4752 bool hasDistSchedule =
false;
4753 if (llvm::isa_and_present<omp::DistributeOp>(opInst.
getParentOp())) {
4754 distributeOp = cast<omp::DistributeOp>(opInst.
getParentOp());
4755 hasDistSchedule = distributeOp.getDistScheduleStatic();
4756 if (distributeOp.getDistScheduleChunkSize()) {
4757 llvm::Value *chunkVar = moduleTranslation.
lookupValue(
4758 distributeOp.getDistScheduleChunkSize());
4759 distScheduleChunk = builder.CreateSExtOrTrunc(chunkVar, ivType);
4768 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
4772 wsloopOp.getNumReductionVars());
4775 wsloopOp, builder, moduleTranslation, privateVarsInfo, allocaIP);
4782 cast<omp::BlockArgOpenMPOpInterface>(opInst).getReductionBlockArgs();
4787 moduleTranslation, allocaIP, reductionDecls,
4788 privateReductionVariables, reductionVariableMap,
4789 deferredStores, isByRef)))
4798 wsloopOp, builder, moduleTranslation, privateVarsInfo.
mlirVars,
4800 wsloopOp.getPrivateNeedsBarrier())))
4803 assert(afterAllocas.get()->getSinglePredecessor());
4804 if (failed(initReductionVars(wsloopOp, reductionArgs, builder,
4806 afterAllocas.get()->getSinglePredecessor(),
4807 reductionDecls, privateReductionVariables,
4808 reductionVariableMap, isByRef, deferredStores)))
4814 bool isTaskReductionMod =
4815 wsloopOp.getReductionMod() == omp::ReductionModifier::task &&
4816 wsloopOp.getNumReductionVars() > 0;
4817 if (isTaskReductionMod &&
4819 "__omp_taskred_mod_", builder, allocaIP,
4820 moduleTranslation,
true,
4822 return wsloopOp.emitError(
4823 "failed to emit task reduction modifier initialization");
4826 bool isOrdered = wsloopOp.getOrdered().has_value();
4827 std::optional<omp::ScheduleModifier> scheduleMod = wsloopOp.getScheduleMod();
4828 bool isSimd = wsloopOp.getScheduleSimd();
4829 bool loopNeedsBarrier = !wsloopOp.getNowait();
4834 llvm::omp::WorksharingLoopType workshareLoopType =
4835 llvm::isa_and_present<omp::DistributeOp>(opInst.
getParentOp())
4836 ? llvm::omp::WorksharingLoopType::DistributeForStaticLoop
4837 : llvm::omp::WorksharingLoopType::ForStaticLoop;
4841 llvm::omp::Directive::OMPD_for);
4843 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4846 LinearClauseProcessor linearClauseProcessor;
4848 if (!wsloopOp.getLinearVars().empty()) {
4849 auto linearVarTypes = wsloopOp.getLinearVarTypes().value();
4851 linearClauseProcessor.registerType(moduleTranslation, linearVarType);
4853 for (
auto [idx, linearVar] : llvm::enumerate(wsloopOp.getLinearVars()))
4854 linearClauseProcessor.createLinearVar(
4855 builder, moduleTranslation, moduleTranslation.
lookupValue(linearVar),
4857 for (
mlir::Value linearStep : wsloopOp.getLinearStepVars())
4858 linearClauseProcessor.initLinearStep(moduleTranslation, linearStep);
4862 wsloopOp.getRegion(),
"omp.wsloop.region", builder, moduleTranslation);
4870 if (!wsloopOp.getLinearVars().empty()) {
4871 linearClauseProcessor.initLinearVar(builder, moduleTranslation,
4872 loopInfo->getPreheader());
4873 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterBarrierIP =
4875 builder, llvm::omp::OMPD_barrier);
4878 builder.restoreIP(*afterBarrierIP);
4879 linearClauseProcessor.updateLinearVar(builder, loopInfo->getBody(),
4880 loopInfo->getIndVar());
4881 linearClauseProcessor.splitLinearFiniBB(builder, loopInfo->getExit());
4884 builder.SetInsertPoint((*regionBlock)->begin());
4887 bool noLoopMode =
false;
4888 omp::TargetOp targetOp = wsloopOp->getParentOfType<mlir::omp::TargetOp>();
4890 targetOp.getKernelType() == omp::TargetExecMode::spmd_no_loop) {
4892 cast<omp::ComposableOpInterface>(*targetOp).findCapturedOp();
4896 if (loopOp == targetCapturedOp)
4900 for (
size_t index = 0;
index < wsloopOp.getLinearVars().size();
index++)
4901 linearClauseProcessor.rewriteInPlace(builder, loopInfo->getBody(),
4902 loopInfo->getLatch(),
index);
4904 llvm::OpenMPIRBuilder::InsertPointOrErrorTy wsloopIP =
4905 ompBuilder->applyWorkshareLoop(
4906 ompLoc.DL, loopInfo, allocaIP, loopNeedsBarrier,
4907 convertToScheduleKind(schedule), chunk, isSimd,
4908 scheduleMod == omp::ScheduleModifier::monotonic,
4909 scheduleMod == omp::ScheduleModifier::nonmonotonic, isOrdered,
4910 workshareLoopType, noLoopMode, hasDistSchedule, distScheduleChunk);
4917 llvm::BasicBlock *wsloopContinuationBB = wsloopIP->getNodeParent();
4920 if (!wsloopOp.getLinearVars().empty()) {
4921 llvm::OpenMPIRBuilder::InsertPointTy oldIP = builder.saveIP();
4922 assert(loopInfo->getLastIter() &&
4923 "`lastiter` in CanonicalLoopInfo is nullptr");
4924 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterBarrierIP =
4925 linearClauseProcessor.finalizeLinearVar(builder, moduleTranslation,
4926 loopInfo->getLastIter());
4930 builder.restoreIP(oldIP);
4937 if (isTaskReductionMod)
4943 wsloopOp, builder, moduleTranslation, allocaIP, reductionDecls,
4944 privateReductionVariables, isByRef, wsloopOp.getNowait(),
4949 wsloopOp.getLoc(), privateVarsInfo);
4956 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
4958 assert(isByRef.size() == opInst.getNumReductionVars());
4967 moduleTranslation, privateVarsInfo)))
4974 opInst.getNumReductionVars());
4980 bool isTaskReductionMod =
4981 opInst.getReductionMod() == omp::ReductionModifier::task &&
4982 opInst.getNumReductionVars() > 0;
4985 [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
4988 opInst, builder, moduleTranslation, privateVarsInfo, allocaIP);
4990 return llvm::make_error<PreviouslyReportedError>();
4996 cast<omp::BlockArgOpenMPOpInterface>(*opInst).getReductionBlockArgs();
4998 allocaIP = allocaIP.getNodeParent()->getTerminator()->getIterator();
5001 opInst, reductionArgs, builder, moduleTranslation, allocaIP,
5002 reductionDecls, privateReductionVariables, reductionVariableMap,
5003 deferredStores, isByRef)))
5004 return llvm::make_error<PreviouslyReportedError>();
5006 assert(afterAllocas.get()->getSinglePredecessor());
5007 builder.restoreIP(codeGenIP);
5013 return llvm::make_error<PreviouslyReportedError>();
5016 opInst, builder, moduleTranslation, privateVarsInfo.
mlirVars,
5018 opInst.getPrivateNeedsBarrier())))
5019 return llvm::make_error<PreviouslyReportedError>();
5022 initReductionVars(opInst, reductionArgs, builder, moduleTranslation,
5023 afterAllocas.get()->getSinglePredecessor(),
5024 reductionDecls, privateReductionVariables,
5025 reductionVariableMap, isByRef, deferredStores)))
5026 return llvm::make_error<PreviouslyReportedError>();
5031 if (isTaskReductionMod &&
5033 "__omp_taskred_mod_", builder, allocaIP,
5034 moduleTranslation,
true,
5036 return llvm::createStringError(
5037 "failed to emit task reduction modifier initialization");
5042 moduleTranslation, allocaIP, deallocBlocks);
5046 opInst.getRegion(),
"omp.par.region", builder, moduleTranslation);
5048 return regionBlock.takeError();
5051 if (opInst.getNumReductionVars() > 0) {
5056 owningReductionGenRefDataPtrGens;
5058 collectReductionInfo(opInst, builder, moduleTranslation, reductionDecls,
5060 owningReductionGenRefDataPtrGens,
5061 privateReductionVariables, reductionInfos, isByRef);
5064 builder.SetInsertPoint((*regionBlock)->getTerminator());
5068 if (isTaskReductionMod)
5073 llvm::UnreachableInst *tempTerminator = builder.CreateUnreachable();
5074 builder.SetInsertPoint(tempTerminator);
5076 llvm::OpenMPIRBuilder::InsertPointOrErrorTy contInsertPoint =
5077 ompBuilder->createReductions(builder, allocaIP, reductionInfos,
5081 if (!contInsertPoint)
5082 return contInsertPoint.takeError();
5084 if (!contInsertPoint->isValid())
5085 return llvm::make_error<PreviouslyReportedError>();
5087 tempTerminator->eraseFromParent();
5088 builder.restoreIP(*contInsertPoint);
5091 return llvm::Error::success();
5094 auto privCB = [](InsertPointTy allocaIP, InsertPointTy codeGenIP,
5095 llvm::Value &, llvm::Value &val, llvm::Value *&replVal) {
5104 auto finiCB = [&](InsertPointTy codeGenIP) -> llvm::Error {
5105 InsertPointTy oldIP = builder.saveIP();
5106 builder.restoreIP(codeGenIP);
5111 llvm::transform(reductionDecls, std::back_inserter(reductionCleanupRegions),
5112 [](omp::DeclareReductionOp reductionDecl) {
5113 return &reductionDecl.getCleanupRegion();
5116 reductionCleanupRegions, privateReductionVariables,
5117 moduleTranslation, builder,
"omp.reduction.cleanup")))
5118 return llvm::createStringError(
5119 "failed to inline `cleanup` region of `omp.declare_reduction`");
5122 opInst.getLoc(), privateVarsInfo)))
5123 return llvm::make_error<PreviouslyReportedError>();
5127 if (isCancellable) {
5128 auto IPOrErr = ompBuilder->createBarrier(
5129 llvm::OpenMPIRBuilder::LocationDescription(builder),
5130 llvm::omp::Directive::OMPD_unknown,
5134 return IPOrErr.takeError();
5137 builder.restoreIP(oldIP);
5138 return llvm::Error::success();
5141 llvm::Value *ifCond =
nullptr;
5142 if (
auto ifVar = opInst.getIfExpr())
5144 llvm::Value *numThreads =
nullptr;
5145 if (!opInst.getNumThreadsVars().empty())
5146 numThreads = moduleTranslation.
lookupValue(opInst.getNumThreads(0));
5147 auto pbKind = llvm::omp::OMP_PROC_BIND_default;
5148 if (
auto bind = opInst.getProcBindKind())
5152 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5154 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5156 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
5157 ompBuilder->createParallel(ompLoc, allocaIP, deallocBlocks, bodyGenCB,
5158 privCB, finiCB, ifCond, numThreads, pbKind,
5164 builder.restoreIP(*afterIP);
5169static llvm::omp::OrderKind
5172 return llvm::omp::OrderKind::OMP_ORDER_unknown;
5174 case omp::ClauseOrderKind::Concurrent:
5175 return llvm::omp::OrderKind::OMP_ORDER_concurrent;
5177 llvm_unreachable(
"Unknown ClauseOrderKind kind");
5185 auto simdOp = cast<omp::SimdOp>(opInst);
5193 cast<omp::BlockArgOpenMPOpInterface>(opInst).getReductionBlockArgs();
5196 simdOp.getNumReductionVars());
5201 assert(isByRef.size() == simdOp.getNumReductionVars());
5203 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5207 simdOp, builder, moduleTranslation, privateVarsInfo, allocaIP);
5212 LinearClauseProcessor linearClauseProcessor;
5213 if (linearClauseProcessor.initLinearIV(simdOp).failed())
5216 if (!simdOp.getLinearVars().empty()) {
5217 auto linearVarTypes = simdOp.getLinearVarTypes().value();
5219 linearClauseProcessor.registerType(moduleTranslation, linearVarType);
5220 for (
auto [idx, linearVar] : llvm::enumerate(simdOp.getLinearVars())) {
5221 bool isImplicit =
false;
5222 for (
auto [mlirPrivVar, llvmPrivateVar] : llvm::zip_equal(
5226 if (linearVar == mlirPrivVar) {
5228 linearClauseProcessor.createLinearVar(builder, moduleTranslation,
5229 llvmPrivateVar, idx);
5235 linearClauseProcessor.createLinearVar(
5236 builder, moduleTranslation,
5239 for (
mlir::Value linearStep : simdOp.getLinearStepVars())
5240 linearClauseProcessor.initLinearStep(moduleTranslation, linearStep);
5244 moduleTranslation, allocaIP, reductionDecls,
5245 privateReductionVariables, reductionVariableMap,
5246 deferredStores, isByRef)))
5257 assert(afterAllocas.get()->getSinglePredecessor());
5258 if (failed(initReductionVars(simdOp, reductionArgs, builder,
5260 afterAllocas.get()->getSinglePredecessor(),
5261 reductionDecls, privateReductionVariables,
5262 reductionVariableMap, isByRef, deferredStores)))
5265 llvm::ConstantInt *simdlen =
nullptr;
5266 if (std::optional<uint64_t> simdlenVar = simdOp.getSimdlen())
5267 simdlen = builder.getInt64(simdlenVar.value());
5269 llvm::ConstantInt *safelen =
nullptr;
5270 if (std::optional<uint64_t> safelenVar = simdOp.getSafelen())
5271 safelen = builder.getInt64(safelenVar.value());
5273 llvm::MapVector<llvm::Value *, llvm::Value *> alignedVars;
5276 llvm::BasicBlock *sourceBlock = builder.GetInsertBlock();
5277 std::optional<ArrayAttr> alignmentValues = simdOp.getAlignments();
5279 for (
size_t i = 0; i < operands.size(); ++i) {
5280 llvm::Value *alignment =
nullptr;
5281 llvm::Value *llvmVal = moduleTranslation.
lookupValue(operands[i]);
5282 llvm::Type *ty = llvmVal->getType();
5284 auto intAttr = cast<IntegerAttr>((*alignmentValues)[i]);
5285 alignment = builder.getInt64(intAttr.getInt());
5286 assert(ty->isPointerTy() &&
"Invalid type for aligned variable");
5287 assert(alignment &&
"Invalid alignment value");
5291 if (!intAttr.getValue().isPowerOf2())
5294 auto curInsert = builder.saveIP();
5295 builder.SetInsertPoint(sourceBlock);
5296 llvmVal = builder.CreateLoad(ty, llvmVal);
5297 builder.restoreIP(curInsert);
5298 alignedVars[llvmVal] = alignment;
5302 simdOp.getRegion(),
"omp.simd.region", builder, moduleTranslation);
5309 if (simdOp.getLinearVars().size()) {
5310 linearClauseProcessor.initLinearVar(builder, moduleTranslation,
5311 loopInfo->getPreheader());
5313 linearClauseProcessor.updateLinearVar(builder, loopInfo->getBody(),
5314 loopInfo->getIndVar());
5316 builder.SetInsertPoint((*regionBlock)->begin());
5318 for (
size_t index = 0;
index < simdOp.getLinearVars().size();
index++)
5319 linearClauseProcessor.rewriteInPlace(builder, loopInfo->getBody(),
5320 loopInfo->getLatch(),
index);
5322 ompBuilder->applySimd(loopInfo, alignedVars,
5324 ? moduleTranslation.
lookupValue(simdOp.getIfExpr())
5326 order, simdlen, safelen);
5328 linearClauseProcessor.updateLinearIV(builder, moduleTranslation);
5329 linearClauseProcessor.emitStoresForLinearVar(builder);
5335 for (
auto [i, tuple] : llvm::enumerate(
5336 llvm::zip(reductionDecls, isByRef, simdOp.getReductionVars(),
5337 privateReductionVariables))) {
5338 auto [decl, byRef, reductionVar, privateReductionVar] = tuple;
5340 OwningReductionGen gen =
makeReductionGen(decl, builder, moduleTranslation);
5341 llvm::Value *originalVariable = moduleTranslation.
lookupValue(reductionVar);
5342 llvm::Type *reductionType = moduleTranslation.
convertType(decl.getType());
5346 llvm::Value *redValue = originalVariable;
5349 builder.CreateLoad(reductionType, redValue,
"red.value." + Twine(i));
5350 llvm::Value *privateRedValue = builder.CreateLoad(
5351 reductionType, privateReductionVar,
"red.private.value." + Twine(i));
5352 llvm::Value *reduced;
5354 auto res = gen(builder.saveIP(), redValue, privateRedValue, reduced);
5357 builder.restoreIP(res.get());
5361 builder.CreateStore(reduced, originalVariable);
5366 llvm::transform(reductionDecls, std::back_inserter(reductionRegions),
5367 [](omp::DeclareReductionOp reductionDecl) {
5368 return &reductionDecl.getCleanupRegion();
5371 moduleTranslation, builder,
5372 "omp.reduction.cleanup")))
5384 auto loopOp = cast<omp::LoopNestOp>(opInst);
5390 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5395 auto bodyGen = [&](llvm::OpenMPIRBuilder::InsertPointTy ip,
5396 llvm::Value *iv) -> llvm::Error {
5399 loopOp.getRegion().front().getArgument(loopInfos.size()), iv);
5404 bodyInsertPoints.push_back(ip);
5406 if (loopInfos.size() != loopOp.getNumLoops() - 1)
5407 return llvm::Error::success();
5410 builder.restoreIP(ip);
5412 loopOp.getRegion(),
"omp.loop_nest.region", builder, moduleTranslation);
5414 return regionBlock.takeError();
5416 builder.SetInsertPoint((*regionBlock)->begin());
5417 return llvm::Error::success();
5425 for (
unsigned i = 0, e = loopOp.getNumLoops(); i < e; ++i) {
5426 llvm::Value *lowerBound =
5427 moduleTranslation.
lookupValue(loopOp.getLoopLowerBounds()[i]);
5428 llvm::Value *upperBound =
5429 moduleTranslation.
lookupValue(loopOp.getLoopUpperBounds()[i]);
5430 llvm::Value *step = moduleTranslation.
lookupValue(loopOp.getLoopSteps()[i]);
5435 llvm::OpenMPIRBuilder::LocationDescription loc = ompLoc;
5436 llvm::OpenMPIRBuilder::InsertPointTy computeIP = ompLoc.IP;
5438 loc = llvm::OpenMPIRBuilder::LocationDescription(bodyInsertPoints.back(),
5440 computeIP = loopInfos.front()->getPreheaderIP();
5444 ompBuilder->createCanonicalLoop(
5445 loc, bodyGen, lowerBound, upperBound, step,
5446 true, loopOp.getLoopInclusive(), computeIP);
5451 loopInfos.push_back(*loopResult);
5454 llvm::OpenMPIRBuilder::InsertPointTy afterIP =
5455 loopInfos.front()->getAfterIP();
5458 if (
const auto &tiles = loopOp.getTileSizes()) {
5459 llvm::Type *ivType = loopInfos.front()->getIndVarType();
5462 for (
auto tile : tiles.value()) {
5463 llvm::Value *tileVal = llvm::ConstantInt::get(ivType,
tile);
5464 tileSizes.push_back(tileVal);
5467 std::vector<llvm::CanonicalLoopInfo *> newLoops =
5468 ompBuilder->tileLoops(ompLoc.DL, loopInfos, tileSizes);
5472 llvm::BasicBlock *afterBB = newLoops.front()->getAfter();
5473 llvm::BasicBlock *afterAfterBB = afterBB->getSingleSuccessor();
5474 afterIP = afterAfterBB->begin();
5478 for (
const auto &newLoop : newLoops)
5479 loopInfos.push_back(newLoop);
5483 const auto &numCollapse = loopOp.getCollapseNumLoops();
5485 loopInfos.begin(), loopInfos.begin() + (numCollapse));
5487 auto newTopLoopInfo =
5488 ompBuilder->collapseLoops(ompLoc.DL, collapseLoopInfos, {});
5490 assert(newTopLoopInfo &&
"New top loop information is missing");
5491 moduleTranslation.
stackWalk<OpenMPLoopInfoStackFrame>(
5492 [&](OpenMPLoopInfoStackFrame &frame) {
5493 frame.loopInfo = newTopLoopInfo;
5501 builder.restoreIP(afterIP);
5511 llvm::OpenMPIRBuilder::LocationDescription loopLoc(builder);
5512 Value loopIV = op.getInductionVar();
5513 Value loopTC = op.getTripCount();
5515 llvm::Value *llvmTC = moduleTranslation.
lookupValue(loopTC);
5518 ompBuilder->createCanonicalLoop(
5520 [&](llvm::OpenMPIRBuilder::InsertPointTy ip, llvm::Value *llvmIV) {
5523 moduleTranslation.
mapValue(loopIV, llvmIV);
5525 builder.restoreIP(ip);
5530 return bodyGenStatus.takeError();
5532 llvmTC,
"omp.loop");
5534 return op.emitError(llvm::toString(llvmOrError.takeError()));
5536 llvm::CanonicalLoopInfo *llvmCLI = *llvmOrError;
5537 llvm::IRBuilderBase::InsertPoint afterIP = llvmCLI->getAfterIP();
5538 builder.restoreIP(afterIP);
5541 if (
Value cli = op.getCli())
5554 Value applyee = op.getApplyee();
5555 assert(applyee &&
"Loop to apply unrolling on required");
5557 llvm::CanonicalLoopInfo *consBuilderCLI =
5559 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5560 ompBuilder->unrollLoopHeuristic(loc.DL, consBuilderCLI);
5573 Value applyee = op.getApplyee();
5574 assert(applyee &&
"Loop to apply unrolling on required");
5576 llvm::CanonicalLoopInfo *consBuilderCLI =
5578 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5579 ompBuilder->unrollLoopFull(loc.DL, consBuilderCLI);
5592 Value applyee = op.getApplyee();
5593 assert(applyee &&
"Loop to apply unrolling on required");
5595 llvm::CanonicalLoopInfo *consBuilderCLI =
5597 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5601 int32_t factor =
static_cast<int32_t
>(op.getUnrollFactor());
5602 ompBuilder->unrollLoopPartial(loc.DL, consBuilderCLI, factor,
5611static LogicalResult
applyTile(omp::TileOp op, llvm::IRBuilderBase &builder,
5614 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5619 for (
Value size : op.getSizes()) {
5620 llvm::Value *translatedSize = moduleTranslation.
lookupValue(size);
5621 assert(translatedSize &&
5622 "sizes clause arguments must already be translated");
5623 translatedSizes.push_back(translatedSize);
5626 for (
Value applyee : op.getApplyees()) {
5627 llvm::CanonicalLoopInfo *consBuilderCLI =
5629 assert(applyee &&
"Canonical loop must already been translated");
5630 translatedLoops.push_back(consBuilderCLI);
5633 auto generatedLoops =
5634 ompBuilder->tileLoops(loc.DL, translatedLoops, translatedSizes);
5635 if (!op.getGeneratees().empty()) {
5636 for (
auto [mlirLoop,
genLoop] :
5637 zip_equal(op.getGeneratees(), generatedLoops))
5642 for (
Value applyee : op.getApplyees())
5650static LogicalResult
applyFuse(omp::FuseOp op, llvm::IRBuilderBase &builder,
5653 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5657 for (
size_t i = 0; i < op.getApplyees().size(); i++) {
5658 Value applyee = op.getApplyees()[i];
5659 llvm::CanonicalLoopInfo *consBuilderCLI =
5661 assert(applyee &&
"Canonical loop must already been translated");
5662 if (op.getFirst().has_value() && i < op.getFirst().value() - 1)
5663 beforeFuse.push_back(consBuilderCLI);
5664 else if (op.getCount().has_value() &&
5665 i >= op.getFirst().value() + op.getCount().value() - 1)
5666 afterFuse.push_back(consBuilderCLI);
5668 toFuse.push_back(consBuilderCLI);
5671 (op.getGeneratees().empty() ||
5672 beforeFuse.size() + afterFuse.size() + 1 == op.getGeneratees().size()) &&
5673 "Wrong number of generatees");
5676 auto generatedLoop = ompBuilder->fuseLoops(loc.DL, toFuse);
5677 if (!op.getGeneratees().empty()) {
5679 for (; i < beforeFuse.size(); i++)
5680 moduleTranslation.
mapOmpLoop(op.getGeneratees()[i], beforeFuse[i]);
5681 moduleTranslation.
mapOmpLoop(op.getGeneratees()[i++], generatedLoop);
5682 for (; i < afterFuse.size(); i++)
5683 moduleTranslation.
mapOmpLoop(op.getGeneratees()[i], afterFuse[i]);
5687 for (
Value applyee : op.getApplyees())
5694static llvm::AtomicOrdering
5697 return llvm::AtomicOrdering::Monotonic;
5700 case omp::ClauseMemoryOrderKind::Seq_cst:
5701 return llvm::AtomicOrdering::SequentiallyConsistent;
5702 case omp::ClauseMemoryOrderKind::Acq_rel:
5703 return llvm::AtomicOrdering::AcquireRelease;
5704 case omp::ClauseMemoryOrderKind::Acquire:
5705 return llvm::AtomicOrdering::Acquire;
5706 case omp::ClauseMemoryOrderKind::Release:
5707 return llvm::AtomicOrdering::Release;
5708 case omp::ClauseMemoryOrderKind::Relaxed:
5709 return llvm::AtomicOrdering::Monotonic;
5711 llvm_unreachable(
"Unknown ClauseMemoryOrderKind kind");
5718static llvm::AtomicOrdering
5720 llvm::AtomicOrdering atomicOrdering) {
5721 if (atomicCompareOp.getFailMemoryOrder())
5723 return llvm::AtomicCmpXchgInst::getStrongestFailureOrdering(atomicOrdering);
5730 auto readOp = cast<omp::AtomicReadOp>(opInst);
5735 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5738 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5741 llvm::Value *x = moduleTranslation.
lookupValue(readOp.getX());
5742 llvm::Value *v = moduleTranslation.
lookupValue(readOp.getV());
5744 llvm::Type *elementType =
5745 moduleTranslation.
convertType(readOp.getElementType());
5747 llvm::OpenMPIRBuilder::AtomicOpValue V = {v, elementType,
false,
false};
5748 llvm::OpenMPIRBuilder::AtomicOpValue X = {x, elementType,
false,
false};
5749 builder.restoreIP(ompBuilder->createAtomicRead(ompLoc, X, V, AO, allocaIP));
5757 auto writeOp = cast<omp::AtomicWriteOp>(opInst);
5762 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5765 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5767 llvm::Value *expr = moduleTranslation.
lookupValue(writeOp.getExpr());
5768 llvm::Value *dest = moduleTranslation.
lookupValue(writeOp.getX());
5769 llvm::Type *ty = moduleTranslation.
convertType(writeOp.getExpr().getType());
5770 llvm::OpenMPIRBuilder::AtomicOpValue x = {dest, ty,
false,
5773 ompBuilder->createAtomicWrite(ompLoc, x, expr, ao, allocaIP));
5781 .Case([&](LLVM::AddOp) {
return llvm::AtomicRMWInst::BinOp::Add; })
5782 .Case([&](LLVM::SubOp) {
return llvm::AtomicRMWInst::BinOp::Sub; })
5783 .Case([&](LLVM::AndOp) {
return llvm::AtomicRMWInst::BinOp::And; })
5784 .Case([&](LLVM::OrOp) {
return llvm::AtomicRMWInst::BinOp::Or; })
5785 .Case([&](LLVM::XOrOp) {
return llvm::AtomicRMWInst::BinOp::Xor; })
5786 .Case([&](LLVM::UMaxOp) {
return llvm::AtomicRMWInst::BinOp::UMax; })
5787 .Case([&](LLVM::UMinOp) {
return llvm::AtomicRMWInst::BinOp::UMin; })
5788 .Case([&](LLVM::FAddOp) {
return llvm::AtomicRMWInst::BinOp::FAdd; })
5789 .Case([&](LLVM::FSubOp) {
return llvm::AtomicRMWInst::BinOp::FSub; })
5790 .Default(llvm::AtomicRMWInst::BinOp::BAD_BINOP);
5794 bool &isIgnoreDenormalMode,
5795 bool &isFineGrainedMemory,
5796 bool &isRemoteMemory) {
5797 isIgnoreDenormalMode =
false;
5798 isFineGrainedMemory =
false;
5799 isRemoteMemory =
false;
5800 if (atomicUpdateOp && atomicUpdateOp.getAtomicControlAttr()) {
5801 mlir::omp::AtomicControlAttr atomicControlAttr =
5802 atomicUpdateOp.getAtomicControlAttr();
5803 isIgnoreDenormalMode = atomicControlAttr.getIgnoreDenormalMode();
5804 isFineGrainedMemory = atomicControlAttr.getFineGrainedMemory();
5805 isRemoteMemory = atomicControlAttr.getRemoteMemory();
5812 llvm::IRBuilderBase &builder,
5819 auto &innerOpList = opInst.getRegion().front().getOperations();
5820 bool isXBinopExpr{
false};
5821 llvm::AtomicRMWInst::BinOp binop;
5823 llvm::Value *llvmExpr =
nullptr;
5824 llvm::Value *llvmX =
nullptr;
5825 llvm::Type *llvmXElementType =
nullptr;
5826 if (innerOpList.size() == 2) {
5832 opInst.getRegion().getArgument(0))) {
5833 return opInst.emitError(
"no atomic update operation with region argument"
5834 " as operand found inside atomic.update region");
5837 isXBinopExpr = innerOp.
getOperand(0) == opInst.getRegion().getArgument(0);
5839 llvmExpr = moduleTranslation.
lookupValue(mlirExpr);
5843 binop = llvm::AtomicRMWInst::BinOp::BAD_BINOP;
5845 llvmX = moduleTranslation.
lookupValue(opInst.getX());
5847 opInst.getRegion().getArgument(0).getType());
5848 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicX = {llvmX, llvmXElementType,
5852 llvm::AtomicOrdering atomicOrdering =
5857 [&opInst, &moduleTranslation](
5858 llvm::Value *atomicx,
5861 moduleTranslation.
mapValue(*opInst.getRegion().args_begin(), atomicx);
5862 moduleTranslation.
mapBlock(&bb, builder.GetInsertBlock());
5863 if (failed(moduleTranslation.
convertBlock(bb,
true, builder)))
5864 return llvm::make_error<PreviouslyReportedError>();
5866 omp::YieldOp yieldop = dyn_cast<omp::YieldOp>(bb.
getTerminator());
5867 assert(yieldop && yieldop.getResults().size() == 1 &&
5868 "terminator must be omp.yield op and it must have exactly one "
5870 return moduleTranslation.
lookupValue(yieldop.getResults()[0]);
5873 bool isIgnoreDenormalMode;
5874 bool isFineGrainedMemory;
5875 bool isRemoteMemory;
5880 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5881 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
5882 ompBuilder->createAtomicUpdate(ompLoc, allocaIP, llvmAtomicX, llvmExpr,
5883 atomicOrdering, binop, updateFn,
5884 isXBinopExpr, isIgnoreDenormalMode,
5885 isFineGrainedMemory, isRemoteMemory);
5890 builder.restoreIP(*afterIP);
5896static std::optional<llvm::omp::OMPAtomicCompareOp>
5898 switch (predicate) {
5899 case LLVM::ICmpPredicate::eq:
5900 return llvm::omp::OMPAtomicCompareOp::EQ;
5901 case LLVM::ICmpPredicate::slt:
5902 case LLVM::ICmpPredicate::ult:
5903 return llvm::omp::OMPAtomicCompareOp::MIN;
5904 case LLVM::ICmpPredicate::sgt:
5905 case LLVM::ICmpPredicate::ugt:
5906 return llvm::omp::OMPAtomicCompareOp::MAX;
5908 return std::nullopt;
5914static std::optional<llvm::omp::OMPAtomicCompareOp>
5916 switch (predicate) {
5917 case LLVM::FCmpPredicate::oeq:
5918 case LLVM::FCmpPredicate::ueq:
5919 return llvm::omp::OMPAtomicCompareOp::EQ;
5920 case LLVM::FCmpPredicate::olt:
5921 case LLVM::FCmpPredicate::ult:
5922 return llvm::omp::OMPAtomicCompareOp::MIN;
5923 case LLVM::FCmpPredicate::ogt:
5924 case LLVM::FCmpPredicate::ugt:
5925 return llvm::omp::OMPAtomicCompareOp::MAX;
5927 return std::nullopt;
5954 if (
auto extractOp = v.getDefiningOp<LLVM::ExtractValueOp>())
5955 return extractOp.getContainer();
5959 if (!isa<LLVM::AndOp, LLVM::OrOp>(op))
5963 if (!lhsFcmp || !rhsFcmp)
5965 mlir::Value lhsAgg0 = traceToAggregate(lhsFcmp.getOperand(0));
5966 mlir::Value lhsAgg1 = traceToAggregate(lhsFcmp.getOperand(1));
5967 bool lhsXIsOp0 = (lhsAgg0 == block.
getArgument(0));
5968 bool lhsXIsOp1 = (lhsAgg1 == block.
getArgument(0));
5969 if (!lhsXIsOp0 && !lhsXIsOp1)
5971 mlir::Value eAggregate = lhsXIsOp0 ? lhsAgg1 : lhsAgg0;
5975 result.isNE = isa<LLVM::OrOp>(op);
5976 result.eAggregate = eAggregate;
5977 result.isXBinopExpr = lhsXIsOp0;
5989 llvm::Value *llvmX, llvm::Type *complexTy,
5990 llvm::Value *eVal, llvm::Value *dVal,
5991 llvm::AtomicOrdering atomicOrdering,
5992 llvm::AtomicOrdering failOrdering,
5993 bool isWeak, llvm::Value *&oldComplex,
5994 llvm::Value *&cmpOk) {
5995 const llvm::DataLayout &DL =
5996 builder.GetInsertBlock()->getModule()->getDataLayout();
5997 unsigned totalBits = DL.getTypeStoreSizeInBits(complexTy).getFixedValue();
5998 llvm::IntegerType *intTy =
5999 llvm::IntegerType::get(builder.getContext(), totalBits);
6000 llvm::Align complexAlign = DL.getABITypeAlign(complexTy);
6001 llvm::Align intAlign = DL.getABITypeAlign(intTy);
6002 llvm::Align maxAlign = std::max(complexAlign, intAlign);
6005 llvm::AllocaInst *dAlloca =
6006 builder.CreateAlloca(complexTy,
nullptr,
"cmplx.d");
6007 dAlloca->setAlignment(maxAlign);
6008 builder.CreateAlignedStore(dVal, dAlloca, maxAlign);
6010 builder.CreateAlignedLoad(intTy, dAlloca, maxAlign,
"cmplx.d.int");
6015 llvm::LoadInst *xCurr =
6016 builder.CreateAlignedLoad(intTy, llvmX, maxAlign,
"cmplx.x.load");
6017 xCurr->setAtomic(failOrdering);
6018 llvm::AllocaInst *xAlloca =
6019 builder.CreateAlloca(complexTy,
nullptr,
"cmplx.x");
6020 xAlloca->setAlignment(maxAlign);
6021 builder.CreateAlignedStore(xCurr, xAlloca, maxAlign);
6022 llvm::Value *xStruct =
6023 builder.CreateAlignedLoad(complexTy, xAlloca, maxAlign,
"cmplx.x.val");
6028 llvm::Value *reX = builder.CreateExtractValue(xStruct, 0);
6029 llvm::Value *imX = builder.CreateExtractValue(xStruct, 1);
6030 llvm::Value *reE = builder.CreateExtractValue(eVal, 0);
6031 llvm::Value *imE = builder.CreateExtractValue(eVal, 1);
6032 llvm::Value *reEq = builder.CreateFCmpOEQ(reX, reE,
"cmplx.re.eq");
6033 llvm::Value *imEq = builder.CreateFCmpOEQ(imX, imE,
"cmplx.im.eq");
6034 llvm::Value *fpEqual = builder.CreateAnd(reEq, imEq,
"cmplx.eq");
6039 llvm::BasicBlock *curBB = builder.GetInsertBlock();
6040 llvm::Function *fn = curBB->getParent();
6041 llvm::BasicBlock *swapBB =
6042 llvm::BasicBlock::Create(builder.getContext(),
"cmplx.atomic.swap", fn);
6043 llvm::BasicBlock *exitBB =
6044 llvm::BasicBlock::Create(builder.getContext(),
"cmplx.atomic.exit", fn);
6045 builder.CreateCondBr(fpEqual, swapBB, exitBB);
6047 builder.SetInsertPoint(swapBB);
6048 llvm::AtomicCmpXchgInst *cmpXchg = builder.CreateAtomicCmpXchg(
6049 llvmX, xCurr, dInt, maxAlign, atomicOrdering, failOrdering);
6050 cmpXchg->setWeak(isWeak);
6051 llvm::Value *oldSwap = builder.CreateExtractValue(cmpXchg, 0);
6052 llvm::Value *okSwap = builder.CreateExtractValue(cmpXchg, 1);
6053 builder.CreateBr(exitBB);
6056 builder.SetInsertPoint(exitBB);
6057 llvm::PHINode *oldIntPHI = builder.CreatePHI(intTy, 2,
"cmplx.old.int");
6058 oldIntPHI->addIncoming(oldSwap, swapBB);
6059 oldIntPHI->addIncoming(xCurr, curBB);
6060 llvm::PHINode *okPHI = builder.CreatePHI(builder.getInt1Ty(), 2,
"cmplx.ok");
6061 okPHI->addIncoming(okSwap, swapBB);
6062 okPHI->addIncoming(builder.getFalse(), curBB);
6065 llvm::AllocaInst *oldAlloca =
6066 builder.CreateAlloca(complexTy,
nullptr,
"cmplx.old");
6067 oldAlloca->setAlignment(maxAlign);
6068 builder.CreateAlignedStore(oldIntPHI, oldAlloca, maxAlign);
6069 oldComplex = builder.CreateAlignedLoad(complexTy, oldAlloca, maxAlign,
6077 llvm::omp::OMPAtomicCompareOp
compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6097 return atomicCompareOp.emitError(
6098 "unsupported comparison predicate (NE) for complex atomic compare");
6099 info.
compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6101 info.
eVal = materializeValue(cplx.eAggregate);
6103 if (
auto selectOp = dyn_cast<LLVM::SelectOp>(op)) {
6104 info.
dVal = materializeValue(selectOp.getTrueValue());
6114 if (
auto icmpOp = dyn_cast<LLVM::ICmpOp>(op);
6115 icmpOp && icmpOp.getOperand(0) != block.
getArgument(0) &&
6121 .Case<LLVM::ICmpOp>([&](LLVM::ICmpOp icmpOp) -> LogicalResult {
6125 return atomicCompareOp.emitError(
6126 "unsupported comparison predicate in atomic compare");
6128 LLVM::ICmpPredicate pred = icmpOp.getPredicate();
6129 info.
isSigned = (pred == LLVM::ICmpPredicate::slt ||
6130 pred == LLVM::ICmpPredicate::sgt ||
6131 pred == LLVM::ICmpPredicate::sle ||
6132 pred == LLVM::ICmpPredicate::sge);
6136 : icmpOp.getOperand(0);
6137 info.
eVal = materializeValue(eOperand);
6140 .Case<LLVM::FCmpOp>([&](LLVM::FCmpOp fcmpOp) -> LogicalResult {
6144 return atomicCompareOp.emitError(
6145 "unsupported comparison predicate in atomic compare");
6150 : fcmpOp.getOperand(0);
6151 info.
eVal = materializeValue(eOperand);
6154 .Case<LLVM::SelectOp>([&](LLVM::SelectOp selectOp) {
6156 info.
dVal = materializeValue(selectOp.getTrueValue());
6159 .Case<mlir::arith::MaxSIOp, mlir::arith::MinSIOp,
6160 mlir::arith::MaxUIOp, mlir::arith::MinUIOp,
6161 mlir::arith::MaximumFOp, mlir::arith::MinimumFOp,
6162 LLVM::SMaxOp, LLVM::SMinOp, LLVM::UMaxOp, LLVM::UMinOp,
6163 LLVM::MaxNumOp, LLVM::MinNumOp>([&](
Operation *) {
6168 bool isMax = isa<mlir::arith::MaxSIOp, mlir::arith::MaxUIOp,
6169 mlir::arith::MaximumFOp, LLVM::SMaxOp,
6170 LLVM::UMaxOp, LLVM::MaxNumOp>(op);
6171 info.
compareOp = isMax ? llvm::omp::OMPAtomicCompareOp::MIN
6172 : llvm::omp::OMPAtomicCompareOp::MAX;
6173 info.
isSigned = isa<mlir::arith::MaxSIOp, mlir::arith::MinSIOp,
6174 LLVM::SMaxOp, LLVM::SMinOp>(op);
6178 info.
eVal = materializeValue(eOperand);
6192 llvm::IRBuilderBase &builder,
6198 omp::AtomicUpdateOp atomicUpdateOp = atomicCaptureOp.getAtomicUpdateOp();
6199 omp::AtomicWriteOp atomicWriteOp = atomicCaptureOp.getAtomicWriteOp();
6200 omp::AtomicCompareOp atomicCompareOp = atomicCaptureOp.getAtomicCompareOp();
6204 if (atomicCompareOp) {
6205 omp::AtomicReadOp atomicReadOp = atomicCaptureOp.getAtomicReadOp();
6206 assert(atomicReadOp &&
"expected atomic.read in capture+compare");
6208 Region ®ion = atomicCompareOp.getRegion();
6211 llvm::Type *llvmXElementType =
6213 llvm::Value *llvmX = moduleTranslation.
lookupValue(atomicCompareOp.getX());
6214 llvm::Value *llvmV = moduleTranslation.
lookupValue(atomicReadOp.getV());
6216 bool isSigned =
false;
6217 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicX = {
6218 llvmX, llvmXElementType, isSigned,
false};
6219 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicV = {
6220 llvmV, llvmXElementType,
false,
false};
6221 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicR = {
nullptr,
nullptr,
false,
6224 llvm::AtomicOrdering atomicOrdering =
6228 auto isAtomicComparePatternOp = [](
Operation &op) {
6229 return llvm::isa<LLVM::ICmpOp, LLVM::FCmpOp, LLVM::SelectOp, LLVM::AndOp,
6230 LLVM::OrOp, LLVM::SMaxOp, LLVM::SMinOp, LLVM::UMaxOp,
6231 LLVM::UMinOp, LLVM::MaxNumOp, LLVM::MinNumOp,
6232 mlir::arith::MaxSIOp, mlir::arith::MinSIOp,
6233 mlir::arith::MaxUIOp, mlir::arith::MinUIOp,
6234 mlir::arith::MaximumFOp, mlir::arith::MinimumFOp>(op);
6237 if (isAtomicComparePatternOp(op))
6239 bool allOperandsMapped =
6241 return moduleTranslation.lookupValue(v) != nullptr;
6243 if (!allOperandsMapped)
6246 return atomicCompareOp.emitError(
6247 "failed to translate operation inside atomic compare region");
6250 auto materializeValue = [&](
mlir::Value val) -> llvm::Value * {
6251 if (llvm::Value *existing = moduleTranslation.
lookupValue(val))
6254 if (loadOp->getParentRegion() == ®ion) {
6255 llvm::Value *loadAddr =
6259 llvm::Type *loadType =
6260 moduleTranslation.
convertType(loadOp.getResult().getType());
6261 return builder.CreateLoad(loadType, loadAddr);
6270 atomicCompareOp, patternInfo)))
6273 llvm::omp::OMPAtomicCompareOp compareOp = patternInfo.
compareOp;
6274 llvm::Value *eVal = patternInfo.
eVal;
6275 llvm::Value *dVal = patternInfo.
dVal;
6280 return atomicCompareOp.emitError(
6281 "failed to extract expected value (e) from atomic compare region");
6284 if (yieldOp.getResults().empty())
6285 return atomicCompareOp.emitError(
6286 "failed to extract desired value (d) from atomic compare region");
6287 dVal = materializeValue(yieldOp.getResults()[0]);
6290 llvmAtomicX.IsSigned = isSigned;
6292 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6293 bool isReadFirst = isa<omp::AtomicReadOp>(atomicCaptureOp.getFirstOp());
6294 bool isPostfixCapture = !isReadFirst;
6295 bool isFailOnly = atomicCaptureOp.getFailOnly();
6303 if (llvmXElementType->isStructTy()) {
6304 llvm::Value *oldComplex =
nullptr;
6305 llvm::Value *cmpOk =
nullptr;
6306 llvm::AtomicOrdering failOrdering =
6309 atomicOrdering, failOrdering,
6310 atomicCompareOp.getWeak(), oldComplex, cmpOk);
6314 llvm::Value *cmpFailed = builder.CreateNot(cmpOk);
6315 llvm::BasicBlock *curBB = builder.GetInsertBlock();
6316 llvm::Function *fn = curBB->getParent();
6317 llvm::BasicBlock *contBB = llvm::BasicBlock::Create(
6318 builder.getContext(),
"omp.atomic.cont", fn);
6319 llvm::BasicBlock *exitBB = llvm::BasicBlock::Create(
6320 builder.getContext(),
"omp.atomic.exit", fn);
6321 builder.CreateCondBr(cmpFailed, contBB, exitBB);
6322 builder.SetInsertPoint(contBB);
6323 builder.CreateStore(oldComplex, llvmAtomicV.Var,
6324 llvmAtomicV.IsVolatile);
6325 builder.CreateBr(exitBB);
6326 builder.SetInsertPoint(exitBB);
6327 }
else if (isPostfixCapture) {
6329 llvm::Value *newComplex = builder.CreateSelect(cmpOk, dVal, oldComplex);
6330 builder.CreateStore(newComplex, llvmAtomicV.Var,
6331 llvmAtomicV.IsVolatile);
6334 builder.CreateStore(oldComplex, llvmAtomicV.Var,
6335 llvmAtomicV.IsVolatile);
6339 if (atomicOrdering == llvm::AtomicOrdering::Release ||
6340 atomicOrdering == llvm::AtomicOrdering::AcquireRelease ||
6341 atomicOrdering == llvm::AtomicOrdering::SequentiallyConsistent) {
6342 llvm::OpenMPIRBuilder::LocationDescription flushLoc(builder);
6343 ompBuilder->createFlush(flushLoc);
6354 bool isMinMax = compareOp != llvm::omp::OMPAtomicCompareOp::EQ;
6356 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicVForCall = llvmAtomicV;
6363 bool minMaxManualCapture = isMinMax && (isPostfixCapture || isFailOnly);
6364 bool eqPostfixManualCapture = !isMinMax && isPostfixCapture && !isFailOnly;
6365 if (minMaxManualCapture || eqPostfixManualCapture)
6366 llvmAtomicVForCall = {
nullptr,
nullptr,
false,
false};
6370 bool builderFailOnly = isFailOnly && !isMinMax;
6377 bool isPostfixUpdate = !builderFailOnly;
6379 bool isWeak = atomicCompareOp.getWeak();
6380 bool savedHandleFPNegZero = ompBuilder->setHandleFPNegZero(
true);
6381 llvm::AtomicOrdering failureOrdering =
6383 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6384 ompBuilder->createAtomicCompare(
6385 ompLoc, llvmAtomicX, llvmAtomicVForCall, llvmAtomicR, eVal, dVal,
6386 atomicOrdering, compareOp, isXBinopExpr, isPostfixUpdate,
6387 builderFailOnly, failureOrdering, isWeak);
6388 ompBuilder->setHandleFPNegZero(savedHandleFPNegZero);
6390 if (failed(
handleError(afterIP, *atomicCaptureOp)))
6393 builder.restoreIP(*afterIP);
6402 if (isMinMax && (isPostfixCapture || isFailOnly)) {
6403 llvm::BasicBlock *curBB = builder.GetInsertBlock();
6404 llvm::AtomicRMWInst *rmw =
nullptr;
6405 for (
auto &inst : llvm::reverse(*curBB)) {
6406 if (
auto *r = dyn_cast<llvm::AtomicRMWInst>(&inst)) {
6411 assert(rmw &&
"expected atomicrmw for min/max compare capture");
6412 llvm::Value *oldVal = rmw;
6413 llvm::Value *rhs = rmw->getValOperand();
6419 llvm::CmpInst::Predicate updatePred;
6420 switch (rmw->getOperation()) {
6421 case llvm::AtomicRMWInst::Min:
6422 updatePred = llvm::CmpInst::ICMP_SGT;
6424 case llvm::AtomicRMWInst::Max:
6425 updatePred = llvm::CmpInst::ICMP_SLT;
6427 case llvm::AtomicRMWInst::UMin:
6428 updatePred = llvm::CmpInst::ICMP_UGT;
6430 case llvm::AtomicRMWInst::UMax:
6431 updatePred = llvm::CmpInst::ICMP_ULT;
6433 case llvm::AtomicRMWInst::FMin:
6434 updatePred = llvm::CmpInst::FCMP_OGT;
6436 case llvm::AtomicRMWInst::FMax:
6437 updatePred = llvm::CmpInst::FCMP_OLT;
6441 "unexpected atomicrmw op for min/max compare capture");
6443 llvm::Value *updated = builder.CreateCmp(updatePred, oldVal, rhs);
6444 llvm::Value *failed = builder.CreateNot(updated);
6445 llvm::Function *fn = curBB->getParent();
6446 llvm::BasicBlock *contBB = llvm::BasicBlock::Create(
6447 builder.getContext(),
"omp.atomic.cont", fn);
6448 llvm::BasicBlock *exitBB = llvm::BasicBlock::Create(
6449 builder.getContext(),
"omp.atomic.exit", fn);
6450 builder.CreateCondBr(failed, contBB, exitBB);
6451 builder.SetInsertPoint(contBB);
6452 builder.CreateStore(oldVal, llvmAtomicV.Var, llvmAtomicV.IsVolatile);
6453 builder.CreateBr(exitBB);
6454 builder.SetInsertPoint(exitBB);
6456 llvm::Intrinsic::ID id;
6457 switch (rmw->getOperation()) {
6458 case llvm::AtomicRMWInst::Min:
6459 id = llvm::Intrinsic::smin;
6461 case llvm::AtomicRMWInst::Max:
6462 id = llvm::Intrinsic::smax;
6464 case llvm::AtomicRMWInst::UMin:
6465 id = llvm::Intrinsic::umin;
6467 case llvm::AtomicRMWInst::UMax:
6468 id = llvm::Intrinsic::umax;
6470 case llvm::AtomicRMWInst::FMin:
6471 id = llvm::Intrinsic::minnum;
6473 case llvm::AtomicRMWInst::FMax:
6474 id = llvm::Intrinsic::maxnum;
6478 "unexpected atomicrmw op for min/max compare capture");
6480 llvm::Value *newVal = builder.CreateBinaryIntrinsic(
id, oldVal, rhs);
6481 builder.CreateStore(newVal, llvmAtomicV.Var, llvmAtomicV.IsVolatile);
6487 if (!isMinMax && isPostfixCapture && !isFailOnly) {
6488 llvm::BasicBlock *curBB = builder.GetInsertBlock();
6489 llvm::Value *oldVal =
nullptr;
6490 llvm::Value *successVal =
nullptr;
6494 for (
auto &inst : llvm::reverse(*curBB)) {
6495 if (isa<llvm::AtomicCmpXchgInst>(&inst)) {
6496 oldVal = builder.CreateExtractValue(&inst, 0);
6497 successVal = builder.CreateExtractValue(&inst, 1);
6508 for (
auto &inst : *curBB) {
6509 auto *phi = dyn_cast<llvm::PHINode>(&inst);
6512 if (phi->getType()->isIntegerTy(1))
6515 for (
auto &inst : *curBB) {
6516 if (
auto *bc = dyn_cast<llvm::BitCastInst>(&inst)) {
6523 assert(oldVal &&
"expected cmpxchg or PHI+bitcast for compare capture");
6524 assert(successVal &&
"expected success flag for compare capture");
6525 llvm::Value *newVal = builder.CreateSelect(successVal, dVal, oldVal);
6526 builder.CreateStore(newVal, llvmAtomicV.Var, llvmAtomicV.IsVolatile);
6533 bool isXBinopExpr =
false, isPostfixUpdate =
false;
6534 llvm::AtomicRMWInst::BinOp binop = llvm::AtomicRMWInst::BinOp::BAD_BINOP;
6536 assert((atomicUpdateOp || atomicWriteOp) &&
6537 "internal op must be an atomic.update or atomic.write op");
6539 if (atomicWriteOp) {
6540 isPostfixUpdate =
true;
6541 mlirExpr = atomicWriteOp.getExpr();
6543 isPostfixUpdate = atomicCaptureOp.getSecondOp() ==
6544 atomicCaptureOp.getAtomicUpdateOp().getOperation();
6545 auto &innerOpList = atomicUpdateOp.getRegion().front().getOperations();
6548 if (innerOpList.size() == 2) {
6551 atomicUpdateOp.getRegion().getArgument(0))) {
6552 return atomicUpdateOp.emitError(
6553 "no atomic update operation with region argument"
6554 " as operand found inside atomic.update region");
6558 innerOp.
getOperand(0) == atomicUpdateOp.getRegion().getArgument(0);
6561 binop = llvm::AtomicRMWInst::BinOp::BAD_BINOP;
6565 llvm::Value *llvmExpr = moduleTranslation.
lookupValue(mlirExpr);
6566 llvm::Value *llvmX =
6567 moduleTranslation.
lookupValue(atomicCaptureOp.getAtomicReadOp().getX());
6568 llvm::Value *llvmV =
6569 moduleTranslation.
lookupValue(atomicCaptureOp.getAtomicReadOp().getV());
6570 llvm::Type *llvmXElementType = moduleTranslation.
convertType(
6571 atomicCaptureOp.getAtomicReadOp().getElementType());
6572 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicX = {llvmX, llvmXElementType,
6575 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicV = {llvmV, llvmXElementType,
6579 llvm::AtomicOrdering atomicOrdering =
6583 [&](llvm::Value *atomicx,
6586 return moduleTranslation.
lookupValue(atomicWriteOp.getExpr());
6587 Block &bb = *atomicUpdateOp.getRegion().
begin();
6588 moduleTranslation.
mapValue(*atomicUpdateOp.getRegion().args_begin(),
6590 moduleTranslation.
mapBlock(&bb, builder.GetInsertBlock());
6591 if (failed(moduleTranslation.
convertBlock(bb,
true, builder)))
6592 return llvm::make_error<PreviouslyReportedError>();
6594 omp::YieldOp yieldop = dyn_cast<omp::YieldOp>(bb.
getTerminator());
6595 assert(yieldop && yieldop.getResults().size() == 1 &&
6596 "terminator must be omp.yield op and it must have exactly one "
6598 return moduleTranslation.
lookupValue(yieldop.getResults()[0]);
6601 bool isIgnoreDenormalMode;
6602 bool isFineGrainedMemory;
6603 bool isRemoteMemory;
6605 isFineGrainedMemory, isRemoteMemory);
6608 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6609 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6610 ompBuilder->createAtomicCapture(
6611 ompLoc, allocaIP, llvmAtomicX, llvmAtomicV, llvmExpr, atomicOrdering,
6612 binop, updateFn, atomicUpdateOp, isPostfixUpdate, isXBinopExpr,
6613 isIgnoreDenormalMode, isFineGrainedMemory, isRemoteMemory);
6615 if (failed(
handleError(afterIP, *atomicCaptureOp)))
6618 builder.restoreIP(*afterIP);
6640 llvm::IRBuilderBase &builder,
6646 Region ®ion = atomicCompareOp.getRegion();
6650 llvm::Type *llvmXElementType =
6652 if (!llvmXElementType)
6653 return atomicCompareOp.emitError(
6654 "unable to determine element type for atomic compare");
6656 llvm::Value *llvmX = moduleTranslation.
lookupValue(atomicCompareOp.getX());
6661 bool isSigned =
false;
6662 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicX = {llvmX, llvmXElementType,
6666 llvm::AtomicOrdering atomicOrdering =
6669 auto isAtomicComparePatternOp = [](
Operation &op) {
6670 return llvm::isa<LLVM::ICmpOp, LLVM::FCmpOp, LLVM::SelectOp, LLVM::AndOp,
6671 LLVM::OrOp, LLVM::SMaxOp, LLVM::SMinOp, LLVM::UMaxOp,
6672 LLVM::UMinOp, LLVM::MaxNumOp, LLVM::MinNumOp,
6673 mlir::arith::MaxSIOp, mlir::arith::MinSIOp,
6674 mlir::arith::MaxUIOp, mlir::arith::MinUIOp,
6675 mlir::arith::MaximumFOp, mlir::arith::MinimumFOp>(op);
6695 if (isAtomicComparePatternOp(op))
6700 return moduleTranslation.lookupValue(v) != nullptr;
6702 if (!allOperandsMapped)
6706 return atomicCompareOp.emitError(
6707 "failed to translate operation inside atomic compare region");
6712 auto materializeValue = [&](
mlir::Value val) -> llvm::Value * {
6714 if (llvm::Value *existing = moduleTranslation.
lookupValue(val))
6719 if (loadOp->getParentRegion() == ®ion) {
6720 llvm::Value *loadAddr = moduleTranslation.
lookupValue(loadOp.getAddr());
6723 llvm::Type *loadType =
6724 moduleTranslation.
convertType(loadOp.getResult().getType());
6725 return builder.CreateLoad(loadType, loadAddr);
6733 llvm::omp::OMPAtomicCompareOp compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6734 llvm::Value *eVal =
nullptr;
6735 llvm::Value *dVal =
nullptr;
6736 bool isXBinopExpr =
false;
6742 if (isComplexPattern) {
6745 return atomicCompareOp.emitError(
6746 "unsupported comparison predicate (NE) for complex atomic compare");
6747 compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6752 if (isComplexPattern) {
6755 if (
auto selectOp = dyn_cast<LLVM::SelectOp>(op)) {
6756 dVal = materializeValue(selectOp.getTrueValue());
6762 if (yieldOp.getResults().empty())
6763 return atomicCompareOp.emitError(
6764 "failed to extract desired value (d) from atomic compare region");
6765 dVal = materializeValue(yieldOp.getResults()[0]);
6768 llvm::Value *oldComplex =
nullptr;
6769 llvm::Value *cmpOk =
nullptr;
6770 llvm::AtomicOrdering failOrdering =
6773 atomicOrdering, failOrdering,
6774 atomicCompareOp.getWeak(), oldComplex, cmpOk);
6780 if (atomicOrdering == llvm::AtomicOrdering::Release ||
6781 atomicOrdering == llvm::AtomicOrdering::AcquireRelease ||
6782 atomicOrdering == llvm::AtomicOrdering::SequentiallyConsistent) {
6783 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6784 ompBuilder->createFlush(ompLoc);
6790 atomicCompareOp, patternInfo)))
6793 eVal = patternInfo.
eVal;
6794 dVal = patternInfo.
dVal;
6800 return atomicCompareOp.emitError(
6801 "failed to extract expected value (e) from atomic compare region");
6805 if (yieldOp.getResults().empty())
6806 return atomicCompareOp.emitError(
6807 "failed to extract desired value (d) from atomic compare region");
6808 dVal = materializeValue(yieldOp.getResults()[0]);
6811 llvmAtomicX.IsSigned = isSigned;
6813 llvm::OpenMPIRBuilder::AtomicOpValue vOpVal = {
nullptr,
nullptr,
false,
6815 llvm::OpenMPIRBuilder::AtomicOpValue rOpVal = {
nullptr,
nullptr,
false,
6817 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6819 bool isWeak = atomicCompareOp.getWeak();
6821 bool savedHandleFPNegZero = ompBuilder->setHandleFPNegZero(
true);
6822 llvm::AtomicOrdering failureOrdering =
6824 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6825 ompBuilder->createAtomicCompare(
6826 ompLoc, llvmAtomicX, vOpVal, rOpVal, eVal, dVal, atomicOrdering,
6827 compareOp, isXBinopExpr,
false,
6828 false, failureOrdering, isWeak);
6829 ompBuilder->setHandleFPNegZero(savedHandleFPNegZero);
6831 if (failed(
handleError(afterIP, *atomicCompareOp)))
6834 builder.restoreIP(*afterIP);
6839 omp::ClauseCancellationConstructType directive) {
6840 switch (directive) {
6841 case omp::ClauseCancellationConstructType::Loop:
6842 return llvm::omp::Directive::OMPD_for;
6843 case omp::ClauseCancellationConstructType::Parallel:
6844 return llvm::omp::Directive::OMPD_parallel;
6845 case omp::ClauseCancellationConstructType::Sections:
6846 return llvm::omp::Directive::OMPD_sections;
6847 case omp::ClauseCancellationConstructType::Taskgroup:
6848 return llvm::omp::Directive::OMPD_taskgroup;
6850 llvm_unreachable(
"Unhandled cancellation construct type");
6859 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6862 llvm::Value *ifCond =
nullptr;
6863 if (
Value ifVar = op.getIfExpr())
6866 llvm::omp::Directive cancelledDirective =
6869 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6870 ompBuilder->createCancel(ompLoc, ifCond, cancelledDirective);
6872 if (failed(
handleError(afterIP, *op.getOperation())))
6875 builder.restoreIP(afterIP.get());
6882 llvm::IRBuilderBase &builder,
6887 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6890 llvm::omp::Directive cancelledDirective =
6893 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6894 ompBuilder->createCancellationPoint(ompLoc, cancelledDirective);
6896 if (failed(
handleError(afterIP, *op.getOperation())))
6899 builder.restoreIP(afterIP.get());
6909 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6911 auto threadprivateOp = cast<omp::ThreadprivateOp>(opInst);
6916 Value symAddr = threadprivateOp.getSymAddr();
6919 if (
auto asCast = dyn_cast<LLVM::AddrSpaceCastOp>(symOp))
6922 if (!isa<LLVM::AddressOfOp>(symOp))
6923 return opInst.
emitError(
"Addressing symbol not found");
6924 LLVM::AddressOfOp addressOfOp = dyn_cast<LLVM::AddressOfOp>(symOp);
6926 LLVM::GlobalOp global =
6927 addressOfOp.getGlobal(moduleTranslation.
symbolTable());
6928 llvm::GlobalValue *globalValue = moduleTranslation.
lookupGlobal(global);
6929 llvm::Type *type = globalValue->getValueType();
6930 llvm::TypeSize typeSize =
6931 builder.GetInsertBlock()->getModule()->getDataLayout().getTypeStoreSize(
6933 llvm::ConstantInt *size = builder.getInt64(typeSize.getFixedValue());
6934 llvm::Value *callInst = ompBuilder->createCachedThreadPrivate(
6935 ompLoc, globalValue, size, global.getSymName() +
".cache");
6941static llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseKind
6943 switch (deviceClause) {
6944 case mlir::omp::DeclareTargetDeviceType::host:
6945 return llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseHost;
6947 case mlir::omp::DeclareTargetDeviceType::nohost:
6948 return llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseNoHost;
6950 case mlir::omp::DeclareTargetDeviceType::any:
6951 return llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseAny;
6954 llvm_unreachable(
"unhandled device clause");
6957static llvm::OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind
6959 mlir::omp::DeclareTargetCaptureClause captureClause) {
6960 switch (captureClause) {
6961 case mlir::omp::DeclareTargetCaptureClause::to:
6962 return llvm::OffloadEntriesInfoManager::OMPTargetGlobalVarEntryTo;
6963 case mlir::omp::DeclareTargetCaptureClause::link:
6964 return llvm::OffloadEntriesInfoManager::OMPTargetGlobalVarEntryLink;
6965 case mlir::omp::DeclareTargetCaptureClause::enter:
6966 return llvm::OffloadEntriesInfoManager::OMPTargetGlobalVarEntryEnter;
6967 case mlir::omp::DeclareTargetCaptureClause::none:
6968 return llvm::OffloadEntriesInfoManager::OMPTargetGlobalVarEntryNone;
6970 llvm_unreachable(
"unhandled capture clause");
6975 if (
auto addrCast = dyn_cast_if_present<LLVM::AddrSpaceCastOp>(op))
6977 if (
auto addressOfOp = dyn_cast_if_present<LLVM::AddressOfOp>(op)) {
6978 auto modOp = addressOfOp->getParentOfType<mlir::ModuleOp>();
6979 return modOp.lookupSymbol(addressOfOp.getGlobalName());
6986 if (
auto addrCast = dyn_cast_if_present<LLVM::AddrSpaceCastOp>(op))
6987 value = addrCast.getOperand();
7011 if (!llvmVarTy->isPointerTy())
7015 if (
auto gop = dyn_cast<LLVM::GlobalOp>(globalOp))
7016 return moduleTranslation.
convertType(gop.getGlobalType());
7019 dyn_cast_if_present<LLVM::AllocaOp>(baseVar.
getDefiningOp()))
7020 return moduleTranslation.
convertType(allocaOp.getElemType());
7022 if (llvm::Value *baseLlvm = moduleTranslation.
lookupValue(baseVar))
7023 if (
auto *allocaInst = dyn_cast<llvm::AllocaInst>(baseLlvm))
7024 return allocaInst->getAllocatedType();
7033 llvm::IRBuilderBase &builder,
const llvm::DataLayout &dataLayout) {
7035 dyn_cast_if_present<LLVM::AllocaOp>(baseVar.
getDefiningOp())) {
7036 if (
Value arraySize = allocaOp.getArraySize()) {
7037 llvm::Type *elemTy =
7038 moduleTranslation.
convertType(allocaOp.getElemType());
7039 llvm::Value *numElems = moduleTranslation.
lookupValue(arraySize);
7040 if (!numElems->getType()->isIntegerTy(64))
7041 numElems = builder.CreateZExt(numElems, builder.getInt64Ty());
7042 uint64_t elemSize = dataLayout.getTypeAllocSize(elemTy).getFixedValue();
7043 return builder.CreateMul(numElems, builder.getInt64(elemSize));
7046 if (llvm::Value *baseLlvm = moduleTranslation.
lookupValue(baseVar)) {
7047 if (
auto *allocaInst = dyn_cast<llvm::AllocaInst>(baseLlvm)) {
7048 if (allocaInst->isArrayAllocation() &&
7049 !llvm::isa<llvm::ArrayType>(allocaInst->getAllocatedType())) {
7051 dataLayout.getTypeAllocSize(allocaInst->getAllocatedType())
7053 return builder.CreateMul(allocaInst->getArraySize(),
7054 builder.getInt64(elemSize));
7058 return std::nullopt;
7061static llvm::SmallString<64>
7063 llvm::OpenMPIRBuilder &ompBuilder,
7064 llvm::vfs::FileSystem &vfs) {
7066 llvm::raw_svector_ostream os(suffix);
7069 auto fileInfoCallBack = [&loc]() {
7070 return std::pair<std::string, uint64_t>(
7071 llvm::StringRef(loc.getFilename()), loc.getLine());
7076 ompBuilder.getTargetEntryUniqueInfo(fileInfoCallBack, vfs).FileID);
7078 os <<
"_decl_tgt_ref_ptr";
7084 if (
auto declareTargetGlobal =
7085 dyn_cast_if_present<omp::DeclareTargetInterface>(
7087 omp::DeclareTargetAttr declareTargetAttr =
7088 declareTargetGlobal.getDeclareTarget();
7089 if (declareTargetAttr && declareTargetAttr.getCaptureClause() ==
7090 omp::DeclareTargetCaptureClause::link)
7097 if (
auto declareTargetGlobal =
7098 dyn_cast_if_present<omp::DeclareTargetInterface>(
7100 omp::DeclareTargetAttr declareTargetAttr =
7101 declareTargetGlobal.getDeclareTarget();
7102 if (declareTargetAttr && (declareTargetAttr.getCaptureClause() ==
7103 omp::DeclareTargetCaptureClause::to ||
7104 declareTargetAttr.getCaptureClause() ==
7105 omp::DeclareTargetCaptureClause::enter))
7124 ompBuilder->Config.hasRequiresUnifiedSharedMemory())) {
7128 if (gOp.getSymName().contains(suffix))
7133 (gOp.getSymName().str() + suffix.str()).str());
7141struct MapInfosTy : llvm::OpenMPIRBuilder::MapInfosTy {
7142 SmallVector<Operation *, 4> Mappers;
7145 void append(MapInfosTy &curInfo) {
7146 Mappers.append(curInfo.Mappers.begin(), curInfo.Mappers.end());
7147 llvm::OpenMPIRBuilder::MapInfosTy::append(curInfo);
7156struct MapInfoData : MapInfosTy {
7157 llvm::SmallVector<bool, 4> IsDeclareTarget;
7158 llvm::SmallVector<bool, 4> IsAMember;
7160 llvm::SmallVector<bool, 4> IsAMapping;
7161 llvm::SmallVector<mlir::Operation *, 4> MapClause;
7162 llvm::SmallVector<llvm::Value *, 4> OriginalValue;
7165 llvm::SmallVector<llvm::Type *, 4> BaseType;
7168 void append(MapInfoData &CurInfo) {
7169 IsDeclareTarget.append(CurInfo.IsDeclareTarget.begin(),
7170 CurInfo.IsDeclareTarget.end());
7171 MapClause.append(CurInfo.MapClause.begin(), CurInfo.MapClause.end());
7172 OriginalValue.append(CurInfo.OriginalValue.begin(),
7173 CurInfo.OriginalValue.end());
7174 BaseType.append(CurInfo.BaseType.begin(), CurInfo.BaseType.end());
7175 MapInfosTy::append(CurInfo);
7179enum class TargetDirectiveEnumTy : uint32_t {
7183 TargetEnterData = 3,
7188static TargetDirectiveEnumTy getTargetDirectiveEnumTyFromOp(Operation *op) {
7189 return llvm::TypeSwitch<Operation *, TargetDirectiveEnumTy>(op)
7190 .Case([](omp::TargetDataOp) {
return TargetDirectiveEnumTy::TargetData; })
7191 .Case([](omp::TargetEnterDataOp) {
7192 return TargetDirectiveEnumTy::TargetEnterData;
7194 .Case([&](omp::TargetExitDataOp) {
7195 return TargetDirectiveEnumTy::TargetExitData;
7197 .Case([&](omp::TargetUpdateOp) {
7198 return TargetDirectiveEnumTy::TargetUpdate;
7200 .Case([&](omp::TargetOp) {
return TargetDirectiveEnumTy::Target; })
7201 .Default([&](Operation *op) {
return TargetDirectiveEnumTy::None; });
7208 if (
auto nestedArrTy = llvm::dyn_cast_if_present<LLVM::LLVMArrayType>(
7209 arrTy.getElementType()))
7223 if (mapOp.getVarPtrPtr())
7240 return bitEnumContainsAll(mapType, omp::ClauseMapFlags::priv |
7241 omp::ClauseMapFlags::target_param |
7242 omp::ClauseMapFlags::attach);
7257 llvm::Value *basePointer,
7258 llvm::Type *baseType,
7259 llvm::IRBuilderBase &builder,
7261 if (
auto memberClause =
7262 mlir::dyn_cast_if_present<mlir::omp::MapInfoOp>(clauseOp)) {
7267 if (!memberClause.getBounds().empty()) {
7268 llvm::Value *elementCount = builder.getInt64(1);
7269 for (
auto bounds : memberClause.getBounds()) {
7270 if (
auto boundOp = mlir::dyn_cast_if_present<mlir::omp::MapBoundsOp>(
7271 bounds.getDefiningOp())) {
7276 elementCount = builder.CreateMul(
7280 moduleTranslation.
lookupValue(boundOp.getUpperBound()),
7281 moduleTranslation.
lookupValue(boundOp.getLowerBound())),
7282 builder.getInt64(1)));
7289 if (
auto arrTy = llvm::dyn_cast_if_present<LLVM::LLVMArrayType>(type))
7297 llvm::Value *sizeCalc = builder.CreateMul(
7298 elementCount, builder.getInt64(underlyingTypeSzInBits / 8),
7336 return builder.CreateSelect(
7337 builder.CreateICmpEQ(sizeCalc, builder.getInt64(0)),
7338 builder.getInt64(1), sizeCalc);
7352static llvm::omp::OpenMPOffloadMappingFlags
7354 const bool hasExplicitMap =
7355 (mlirFlags &
~omp::ClauseMapFlags::is_device_ptr) !=
7356 omp::ClauseMapFlags::none;
7358 llvm::omp::OpenMPOffloadMappingFlags mapType =
7359 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE;
7361 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::to))
7362 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TO;
7364 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::from))
7365 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_FROM;
7367 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::always))
7368 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ALWAYS;
7370 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::del))
7371 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_DELETE;
7373 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::return_param))
7374 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
7376 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::priv))
7377 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_PRIVATE;
7379 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::literal))
7380 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL;
7382 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::implicit))
7383 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_IMPLICIT;
7385 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::close))
7386 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_CLOSE;
7388 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::present))
7389 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_PRESENT;
7391 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::ompx_hold))
7392 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_OMPX_HOLD;
7394 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::attach))
7395 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH;
7397 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::target_param))
7398 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TARGET_PARAM;
7400 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::is_device_ptr)) {
7401 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TARGET_PARAM;
7402 if (!hasExplicitMap)
7403 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL;
7413 ArrayRef<Value> useDevAddrOperands = {},
7414 ArrayRef<Value> hasDevAddrOperands = {}) {
7416 auto checkRefPtrOrPteeMapWithAttach = [](omp::ClauseMapFlags mapType) {
7418 bitEnumContainsAll(mapType, omp::ClauseMapFlags::ref_ptr) ||
7419 bitEnumContainsAll(mapType, omp::ClauseMapFlags::ref_ptee);
7420 return hasRefType &&
7421 bitEnumContainsAll(mapType, omp::ClauseMapFlags::attach);
7424 auto checkIsAMember = [](
const auto &mapVars,
auto mapOp) {
7432 for (Value mapValue : mapVars) {
7433 auto map = cast<omp::MapInfoOp>(mapValue.getDefiningOp());
7434 for (
auto member : map.getMembers())
7435 if (member == mapOp)
7442 for (Value mapValue : mapVars) {
7443 auto mapOp = cast<omp::MapInfoOp>(mapValue.getDefiningOp());
7444 bool isAttachStyleMap =
7445 checkRefPtrOrPteeMapWithAttach(mapOp.getMapType()) ||
7447 Value offloadPtr = (mapOp.getVarPtrPtr() && !isAttachStyleMap)
7448 ? mapOp.getVarPtrPtr()
7449 : mapOp.getVarPtr();
7450 mapData.OriginalValue.push_back(moduleTranslation.
lookupValue(offloadPtr));
7451 mapData.Pointers.push_back(
7452 isAttachStyleMap ? moduleTranslation.
lookupValue(mapOp.getVarPtrPtr())
7453 : mapData.OriginalValue.back());
7455 if (llvm::Value *refPtr =
7457 mapData.IsDeclareTarget.push_back(
true);
7458 mapData.BasePointers.push_back(refPtr);
7460 mapData.IsDeclareTarget.push_back(
true);
7461 mapData.BasePointers.push_back(mapData.OriginalValue.back());
7463 mapData.IsDeclareTarget.push_back(
false);
7464 mapData.BasePointers.push_back(mapData.OriginalValue.back());
7470 mapData.BaseType.push_back(moduleTranslation.
convertType(
7471 mapOp.getVarPtrPtr() ? mapOp.getVarPtrPtrType().value()
7472 : mapOp.getVarPtrType()));
7479 mlir::Type sizeType = (isAttachStyleMap || !mapOp.getVarPtrPtr())
7480 ? mapOp.getVarPtrType()
7481 : mapOp.getVarPtrPtrType().value();
7483 dl, sizeType, isAttachStyleMap ?
nullptr : mapOp,
7484 mapData.Pointers.back(), moduleTranslation.
convertType(sizeType),
7485 builder, moduleTranslation));
7486 mapData.MapClause.push_back(mapOp.getOperation());
7489 mapData.HasAttachPtr.push_back(
false);
7490 mapData.Names.push_back(LLVM::createMappingInformation(
7492 mapData.DevicePointers.push_back(llvm::OpenMPIRBuilder::DeviceInfoTy::None);
7493 if (mapOp.getMapperId())
7494 mapData.Mappers.push_back(
7496 mapOp, mapOp.getMapperIdAttr()));
7498 mapData.Mappers.push_back(
nullptr);
7499 mapData.IsAMapping.push_back(
true);
7500 mapData.IsAMember.push_back(checkIsAMember(mapVars, mapOp));
7503 auto findMapInfo = [&mapData](llvm::Value *val,
7504 llvm::OpenMPIRBuilder::DeviceInfoTy devInfoTy,
7505 size_t memberCount) {
7508 for (llvm::Value *basePtr : mapData.OriginalValue) {
7509 auto mapOp = cast<omp::MapInfoOp>(mapData.MapClause[index]);
7520 (mapData.Types[index] &
7521 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH) ==
7522 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH;
7523 if (!isAttachMap && basePtr == val && mapData.IsAMapping[index] &&
7524 memberCount == mapOp.getMembers().size()) {
7526 mapData.Types[index] |=
7527 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
7528 mapData.DevicePointers[index] = devInfoTy;
7536 auto addDevInfos = [&](
const llvm::ArrayRef<Value> &useDevOperands,
7537 llvm::OpenMPIRBuilder::DeviceInfoTy devInfoTy) {
7538 for (Value mapValue : useDevOperands) {
7539 auto mapOp = cast<omp::MapInfoOp>(mapValue.getDefiningOp());
7541 mapOp.getVarPtrPtr() ? mapOp.getVarPtrPtr() : mapOp.getVarPtr();
7542 llvm::Value *origValue = moduleTranslation.
lookupValue(offloadPtr);
7545 if (!findMapInfo(origValue, devInfoTy, mapOp.getMembers().size())) {
7546 mapData.OriginalValue.push_back(origValue);
7547 mapData.Pointers.push_back(mapData.OriginalValue.back());
7548 mapData.IsDeclareTarget.push_back(
false);
7549 mapData.BasePointers.push_back(mapData.OriginalValue.back());
7550 mlir::Type baseTy = mapOp.getVarPtrPtr()
7551 ? mapOp.getVarPtrPtrType().value()
7552 : mapOp.getVarPtrType();
7553 mapData.BaseType.push_back(moduleTranslation.
convertType(baseTy));
7554 mapData.Sizes.push_back(builder.getInt64(0));
7555 mapData.MapClause.push_back(mapOp.getOperation());
7556 mapData.Types.push_back(
7557 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM);
7559 mapData.HasAttachPtr.push_back(
false);
7560 mapData.Names.push_back(LLVM::createMappingInformation(
7562 mapData.DevicePointers.push_back(devInfoTy);
7563 mapData.Mappers.push_back(
nullptr);
7564 mapData.IsAMapping.push_back(
false);
7565 mapData.IsAMember.push_back(checkIsAMember(useDevOperands, mapOp));
7570 addDevInfos(useDevAddrOperands, llvm::OpenMPIRBuilder::DeviceInfoTy::Address);
7571 addDevInfos(useDevPtrOperands, llvm::OpenMPIRBuilder::DeviceInfoTy::Pointer);
7573 for (Value mapValue : hasDevAddrOperands) {
7574 auto mapOp = cast<omp::MapInfoOp>(mapValue.getDefiningOp());
7576 mapOp.getVarPtrPtr() ? mapOp.getVarPtrPtr() : mapOp.getVarPtr();
7577 llvm::Value *origValue = moduleTranslation.
lookupValue(offloadPtr);
7579 auto mapTypeAlways = llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ALWAYS;
7581 (mapOp.getMapType() & omp::ClauseMapFlags::is_device_ptr) !=
7582 omp::ClauseMapFlags::none;
7584 mapData.OriginalValue.push_back(origValue);
7585 mapData.BasePointers.push_back(origValue);
7586 mapData.Pointers.push_back(origValue);
7587 mapData.IsDeclareTarget.push_back(
false);
7589 mlir::Type baseTy = mapOp.getVarPtrPtr() ? mapOp.getVarPtrPtrType().value()
7590 : mapOp.getVarPtrType();
7591 mapData.BaseType.push_back(moduleTranslation.
convertType(baseTy));
7592 mapData.Sizes.push_back(builder.getInt64(dl.
getTypeSize(baseTy)));
7594 mapData.MapClause.push_back(mapOp.getOperation());
7595 if (llvm::to_underlying(mapType & mapTypeAlways)) {
7599 mapData.Types.push_back(mapType);
7601 mapData.HasAttachPtr.push_back(
false);
7605 if (mapOp.getMapperId()) {
7606 mapData.Mappers.push_back(
7608 mapOp, mapOp.getMapperIdAttr()));
7610 mapData.Mappers.push_back(
nullptr);
7615 mapData.Types.push_back(
7616 isDevicePtr ? mapType
7617 : llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL);
7619 mapData.HasAttachPtr.push_back(
false);
7620 mapData.Mappers.push_back(
nullptr);
7622 mapData.Names.push_back(LLVM::createMappingInformation(
7624 mapData.DevicePointers.push_back(
7625 isDevicePtr ? llvm::OpenMPIRBuilder::DeviceInfoTy::Pointer
7626 : llvm::OpenMPIRBuilder::DeviceInfoTy::Address);
7627 mapData.IsAMapping.push_back(
false);
7628 mapData.IsAMember.push_back(checkIsAMember(hasDevAddrOperands, mapOp));
7633 auto *res = llvm::find(mapData.MapClause, memberOp);
7634 assert(res != mapData.MapClause.end() &&
7635 "MapInfoOp for member not found in MapData, cannot return index");
7636 return std::distance(mapData.MapClause.begin(), res);
7640 omp::MapInfoOp mapInfo,
bool first =
true) {
7641 ArrayAttr indexAttr = mapInfo.getMembersIndexAttr();
7651 auto memberIndicesA = cast<ArrayAttr>(indexAttr[a]);
7652 auto memberIndicesB = cast<ArrayAttr>(indexAttr[b]);
7654 for (auto it : llvm::zip(memberIndicesA, memberIndicesB)) {
7655 int64_t aIndex = mlir::cast<IntegerAttr>(std::get<0>(it)).getInt();
7656 int64_t bIndex = mlir::cast<IntegerAttr>(std::get<1>(it)).getInt();
7658 if (aIndex == bIndex)
7661 if (aIndex < bIndex)
7664 if (aIndex > bIndex)
7671 bool memberAParent = memberIndicesA.size() < memberIndicesB.size();
7673 occludedChildren.push_back(
b);
7675 occludedChildren.push_back(a);
7676 return memberAParent;
7679 for (
auto v : occludedChildren)
7686 ArrayAttr indexAttr = mapInfo.getMembersIndexAttr();
7688 if (indexAttr.size() == 1)
7689 return cast<omp::MapInfoOp>(mapInfo.getMembers()[0].getDefiningOp());
7693 return llvm::cast<omp::MapInfoOp>(
7694 mapInfo.getMembers()[
indices.front()].getDefiningOp());
7717static std::vector<llvm::Value *>
7719 llvm::IRBuilderBase &builder,
bool isArrayTy,
7721 std::vector<llvm::Value *> idx;
7732 idx.push_back(builder.getInt64(0));
7733 for (
int i = bounds.size() - 1; i >= 0; --i) {
7734 if (
auto boundOp = dyn_cast_if_present<omp::MapBoundsOp>(
7735 bounds[i].getDefiningOp())) {
7736 idx.push_back(moduleTranslation.
lookupValue(boundOp.getLowerBound()));
7754 for (
int i = bounds.size() - 1; i >= 0; --i) {
7755 if (
auto boundOp = dyn_cast_if_present<omp::MapBoundsOp>(
7756 bounds[i].getDefiningOp())) {
7757 if (i == ((
int)bounds.size() - 1))
7759 moduleTranslation.
lookupValue(boundOp.getLowerBound()));
7761 idx.back() = builder.CreateAdd(
7762 builder.CreateMul(idx.back(), moduleTranslation.
lookupValue(
7763 boundOp.getExtent())),
7764 moduleTranslation.
lookupValue(boundOp.getLowerBound()));
7773 llvm::transform(values, std::back_inserter(ints), [](
Attribute value) {
7774 return cast<IntegerAttr>(value).getInt();
7782 omp::MapInfoOp parentOp) {
7784 if (parentOp.getMembers().empty())
7788 if (parentOp.getMembers().size() == 1) {
7789 overlapMapDataIdxs.push_back(0);
7793 ArrayAttr indexAttr = parentOp.getMembersIndexAttr();
7794 size_t numMembers = indexAttr.size();
7798 for (
auto [i, indicesAttr] : llvm::enumerate(indexAttr))
7799 getAsIntegers(cast<ArrayAttr>(indicesAttr), memberIndices[i]);
7805 llvm::SmallDenseSet<size_t> skipIndices;
7806 for (
size_t i = 0; i < numMembers; ++i) {
7807 const auto &iIndices = memberIndices[i];
7808 for (
size_t j = 0;
j < numMembers; ++
j) {
7811 const auto &jIndices = memberIndices[
j];
7813 if (jIndices.size() < iIndices.size() &&
7814 std::equal(jIndices.begin(), jIndices.end(), iIndices.begin())) {
7815 skipIndices.insert(i);
7822 for (
size_t i = 0; i < numMembers; ++i)
7823 if (!skipIndices.contains(i))
7824 overlapMapDataIdxs.push_back(i);
7838 llvm::OpenMPIRBuilder &ompBuilder, MapInfoData &mapData,
7839 size_t mapDataIdx, MapInfosTy &combinedInfo,
7840 TargetDirectiveEnumTy targetDirective,
7841 llvm::omp::OpenMPOffloadMappingFlags memberOfFlag =
7842 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE,
7843 bool isTargetParam =
true,
int mapDataParentIdx = -1) {
7844 auto mapFlag = mapData.Types[mapDataIdx];
7845 auto mapInfoOp = llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataIdx]);
7849 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH) ==
7850 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH);
7856 if (isTargetParam &&
7857 (targetDirective == TargetDirectiveEnumTy::Target &&
7858 !mapData.IsDeclareTarget[mapDataIdx]) &&
7860 mapFlag |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TARGET_PARAM;
7862 if (mapInfoOp.getMapCaptureType() == omp::VariableCaptureKind::ByCopy &&
7864 mapFlag |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL;
7873 if (memberOfFlag != llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE) {
7874 if (!isPtrTy && !isAttachMap)
7875 ompBuilder.setCorrectMemberOfFlag(mapFlag, memberOfFlag);
7882 mapFlag &=
~llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
7892 if (isPtrTy && !isAttachMap && mapData.IsDeclareTarget[mapDataIdx])
7893 mapFlag |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_PTR_AND_OBJ;
7902 !bitEnumContainsAll(mapInfoOp.getMapType(),
7903 omp::ClauseMapFlags::ref_ptr) &&
7904 bitEnumContainsAll(mapInfoOp.getMapType(), omp::ClauseMapFlags::ref_ptee);
7905 bool isRefPtrPtee = bitEnumContainsAll(mapInfoOp.getMapType(),
7906 omp::ClauseMapFlags::ref_ptr |
7907 omp::ClauseMapFlags::ref_ptee);
7909 if (!mapInfoOp->getParentOfType<omp::DeclareMapperOp>() &&
7910 mapDataParentIdx >= 0 && !(isRefPtee || (isRefPtrPtee && isPtrTy))) {
7911 combinedInfo.BasePointers.emplace_back(
7912 mapData.BasePointers[mapDataParentIdx]);
7914 combinedInfo.BasePointers.emplace_back(mapData.BasePointers[mapDataIdx]);
7917 combinedInfo.Pointers.emplace_back(mapData.Pointers[mapDataIdx]);
7918 combinedInfo.DevicePointers.emplace_back(
7919 memberOfFlag != llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE
7920 ? llvm::OpenMPIRBuilder::DeviceInfoTy::None
7921 : mapData.DevicePointers[mapDataIdx]);
7922 combinedInfo.Mappers.emplace_back(mapData.Mappers[mapDataIdx]);
7923 combinedInfo.Names.emplace_back(mapData.Names[mapDataIdx]);
7924 combinedInfo.Types.emplace_back(mapFlag);
7926 combinedInfo.HasAttachPtr.emplace_back(
false);
7930 combinedInfo.Sizes.emplace_back(
7932 ? builder.CreateSelect(
7933 builder.CreateIsNull(mapData.Pointers[mapDataIdx]),
7934 builder.getInt64(0), mapData.Sizes[mapDataIdx])
7935 : mapData.Sizes[mapDataIdx]);
7955 llvm::OpenMPIRBuilder &ompBuilder,
DataLayout &dl, MapInfosTy &combinedInfo,
7956 MapInfoData &mapData, uint64_t mapDataIndex,
7957 llvm::omp::OpenMPOffloadMappingFlags memberOfFlag,
7958 TargetDirectiveEnumTy targetDirective) {
7959 using MapFlags = llvm::omp::OpenMPOffloadMappingFlags;
7960 assert(!ompBuilder.Config.isTargetDevice() &&
7961 "function only supported for host device codegen");
7963 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
7964 auto *parentMapper = mapData.Mappers[mapDataIndex];
7970 MapFlags baseFlag = (targetDirective == TargetDirectiveEnumTy::Target &&
7971 !mapData.IsDeclareTarget[mapDataIndex])
7972 ? MapFlags::OMP_MAP_TARGET_PARAM
7973 : MapFlags::OMP_MAP_NONE;
7979 MapFlags parentFlags = mapData.Types[mapDataIndex];
7980 MapFlags preserve = MapFlags::OMP_MAP_TO | MapFlags::OMP_MAP_FROM |
7981 MapFlags::OMP_MAP_ALWAYS | MapFlags::OMP_MAP_CLOSE |
7982 MapFlags::OMP_MAP_PRESENT |
7983 MapFlags::OMP_MAP_OMPX_HOLD |
7984 MapFlags::OMP_MAP_IMPLICIT;
7985 baseFlag |= (parentFlags & preserve);
7987 MapFlags parentFlags = mapData.Types[mapDataIndex];
7988 MapFlags preserve = MapFlags::OMP_MAP_TO | MapFlags::OMP_MAP_FROM |
7989 MapFlags::OMP_MAP_PRESENT |
7990 MapFlags::OMP_MAP_RETURN_PARAM |
7991 MapFlags::OMP_MAP_IMPLICIT;
7992 baseFlag |= (parentFlags & preserve);
7995 combinedInfo.Types.emplace_back(baseFlag);
7997 combinedInfo.HasAttachPtr.emplace_back(
false);
7998 combinedInfo.DevicePointers.emplace_back(
7999 mapData.DevicePointers[mapDataIndex]);
8003 combinedInfo.Mappers.emplace_back(
8004 parentMapper && !parentClause.getPartialMap() ? parentMapper :
nullptr);
8006 mapData.MapClause[mapDataIndex]->getLoc(), ompBuilder));
8007 combinedInfo.BasePointers.emplace_back(mapData.BasePointers[mapDataIndex]);
8016 llvm::Value *lowAddr, *highAddr;
8017 if (!parentClause.getPartialMap()) {
8018 lowAddr = builder.CreatePointerCast(mapData.Pointers[mapDataIndex],
8019 builder.getPtrTy());
8020 highAddr = builder.CreatePointerCast(
8021 builder.CreateConstGEP1_32(mapData.BaseType[mapDataIndex],
8022 mapData.Pointers[mapDataIndex], 1),
8023 builder.getPtrTy());
8024 combinedInfo.Pointers.emplace_back(mapData.Pointers[mapDataIndex]);
8026 auto mapOp = dyn_cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
8029 lowAddr = builder.CreatePointerCast(mapData.BasePointers[firstMemberIdx],
8030 builder.getPtrTy());
8034 auto lastMemberMapInfo =
8035 cast<omp::MapInfoOp>(mapData.MapClause[lastMemberIdx]);
8044 bool isRefPteeMap = bitEnumContainsAll(lastMemberMapInfo.getMapType(),
8045 omp::ClauseMapFlags::ref_ptee) &&
8046 !bitEnumContainsAll(lastMemberMapInfo.getMapType(),
8047 omp::ClauseMapFlags::ref_ptr);
8048 llvm::Type *castType = mapData.BaseType[lastMemberIdx];
8051 moduleTranslation.
convertType(lastMemberMapInfo.getVarPtrType());
8052 highAddr = builder.CreatePointerCast(
8053 builder.CreateGEP(castType, mapData.BasePointers[lastMemberIdx],
8054 builder.getInt64(1)),
8055 builder.getPtrTy());
8056 combinedInfo.Pointers.emplace_back(mapData.BasePointers[firstMemberIdx]);
8059 llvm::Value *size = builder.CreateIntCast(
8060 builder.CreatePtrDiff(builder.getInt8Ty(), highAddr, lowAddr),
8061 builder.getInt64Ty(),
8063 combinedInfo.Sizes.push_back(size);
8071 if (!parentClause.getPartialMap()) {
8076 MapFlags mapFlag = mapData.Types[mapDataIndex];
8077 bool hasMapClose = (MapFlags(mapFlag) & MapFlags::OMP_MAP_CLOSE) ==
8078 MapFlags::OMP_MAP_CLOSE;
8079 ompBuilder.setCorrectMemberOfFlag(mapFlag, memberOfFlag);
8095 if (targetDirective == TargetDirectiveEnumTy::TargetUpdate || hasMapClose ||
8096 overlapIdxs.size() == 1) {
8097 combinedInfo.Types.emplace_back(mapFlag);
8099 combinedInfo.HasAttachPtr.emplace_back(
false);
8100 combinedInfo.DevicePointers.emplace_back(
8101 mapData.DevicePointers[mapDataIndex]);
8103 mapData.MapClause[mapDataIndex]->getLoc(), ompBuilder));
8104 combinedInfo.BasePointers.emplace_back(
8105 mapData.BasePointers[mapDataIndex]);
8106 combinedInfo.Pointers.emplace_back(mapData.Pointers[mapDataIndex]);
8107 combinedInfo.Sizes.emplace_back(mapData.Sizes[mapDataIndex]);
8108 combinedInfo.Mappers.emplace_back(
nullptr);
8114 lowAddr = builder.CreatePointerCast(mapData.Pointers[mapDataIndex],
8115 builder.getPtrTy());
8116 highAddr = builder.CreatePointerCast(
8117 builder.CreateConstGEP1_32(mapData.BaseType[mapDataIndex],
8118 mapData.Pointers[mapDataIndex], 1),
8119 builder.getPtrTy());
8126 mapFlag &=
~llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
8133 for (
auto v : overlapIdxs) {
8136 cast<omp::MapInfoOp>(parentClause.getMembers()[v].getDefiningOp()));
8138 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataOverlapIdx]));
8139 combinedInfo.Types.emplace_back(mapFlag);
8141 combinedInfo.HasAttachPtr.emplace_back(
false);
8142 combinedInfo.DevicePointers.emplace_back(
8143 llvm::OpenMPIRBuilder::DeviceInfoTy::None);
8145 mapData.MapClause[mapDataIndex]->getLoc(), ompBuilder));
8146 combinedInfo.BasePointers.emplace_back(
8147 mapData.BasePointers[mapDataIndex]);
8148 combinedInfo.Mappers.emplace_back(
nullptr);
8149 combinedInfo.Pointers.emplace_back(lowAddr);
8150 auto sizeCalc = builder.CreateIntCast(
8151 builder.CreatePtrDiff(builder.getInt8Ty(),
8152 mapData.OriginalValue[mapDataOverlapIdx],
8154 builder.getInt64Ty(),
true);
8159 llvm::DataLayout dataLayout = builder.GetInsertBlock()->getDataLayout();
8160 auto sizeSel = builder.CreateSelect(
8161 builder.CreateICmpNE(builder.getInt64(0), sizeCalc), sizeCalc,
8162 isPtrMap ? builder.getInt64(dataLayout.getPointerSize())
8163 : mapData.Sizes[mapDataOverlapIdx]);
8164 combinedInfo.Sizes.emplace_back(sizeSel);
8165 lowAddr = builder.CreateConstGEP1_32(
8166 isPtrMap ? builder.getPtrTy() : mapData.BaseType[mapDataOverlapIdx],
8167 mapData.BasePointers[mapDataOverlapIdx], 1);
8170 combinedInfo.Types.emplace_back(mapFlag);
8172 combinedInfo.HasAttachPtr.emplace_back(
false);
8173 combinedInfo.DevicePointers.emplace_back(
8174 llvm::OpenMPIRBuilder::DeviceInfoTy::None);
8176 mapData.MapClause[mapDataIndex]->getLoc(), ompBuilder));
8177 combinedInfo.BasePointers.emplace_back(
8178 mapData.BasePointers[mapDataIndex]);
8179 combinedInfo.Mappers.emplace_back(
nullptr);
8180 combinedInfo.Pointers.emplace_back(lowAddr);
8181 combinedInfo.Sizes.emplace_back(builder.CreateIntCast(
8182 builder.CreatePtrDiff(builder.getInt8Ty(), highAddr, lowAddr),
8183 builder.getInt64Ty(),
true));
8189 llvm::IRBuilderBase &builder,
8190 llvm::OpenMPIRBuilder &ompBuilder,
8192 MapInfoData &mapData, uint64_t mapDataIndex,
8193 TargetDirectiveEnumTy targetDirective) {
8194 assert(!ompBuilder.Config.isTargetDevice() &&
8195 "function only supported for host device codegen");
8198 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
8203 if (parentClause.getMembers().size() == 1 && parentClause.getPartialMap()) {
8204 auto memberClause = llvm::cast<omp::MapInfoOp>(
8205 parentClause.getMembers()[0].getDefiningOp());
8218 builder, ompBuilder, mapData, memberDataIdx, combinedInfo,
8220 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE,
8221 true, mapDataIndex);
8225 auto collectMapInfoIdxs =
8228 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
8230 for (
auto member : parentClause.getMembers())
8232 mapData, llvm::cast<omp::MapInfoOp>(member.getDefiningOp())));
8236 collectMapInfoIdxs(mapInfoIdx);
8238 llvm::omp::OpenMPOffloadMappingFlags memberOfFlag =
8239 ompBuilder.getMemberOfFlag(combinedInfo.Types.size());
8249 bool parentIsPrivatizeableAttach =
8251 for (
auto [i, idx] : llvm::enumerate(mapInfoIdx)) {
8252 bool emitParentMap = i == 0 && !parentIsPrivatizeableAttach;
8253 if (emitParentMap) {
8255 combinedInfo, mapData, idx, memberOfFlag,
8259 builder, ompBuilder, mapData, idx, combinedInfo, targetDirective,
8260 parentIsPrivatizeableAttach
8261 ? llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE
8263 false, mapDataIndex);
8275 llvm::IRBuilderBase &builder) {
8277 "function only supported for host device codegen");
8278 for (
size_t i = 0; i < mapData.MapClause.size(); ++i) {
8279 auto mapOp = cast<omp::MapInfoOp>(mapData.MapClause[i]);
8282 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH) ==
8283 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH);
8288 if (!mapData.IsDeclareTarget[i] ||
8289 (mapData.IsDeclareTarget[i] && isAttachMap)) {
8290 omp::VariableCaptureKind captureKind = mapOp.getMapCaptureType();
8300 switch (captureKind) {
8301 case omp::VariableCaptureKind::ByRef: {
8302 llvm::Value *newV = mapData.Pointers[i];
8304 moduleTranslation, builder, mapData.BaseType[i]->isArrayTy(),
8307 newV = builder.CreateLoad(builder.getPtrTy(), newV);
8309 if (!offsetIdx.empty())
8310 newV = builder.CreateInBoundsGEP(mapData.BaseType[i], newV, offsetIdx,
8312 mapData.Pointers[i] = newV;
8314 case omp::VariableCaptureKind::ByCopy: {
8315 llvm::Type *type = mapData.BaseType[i];
8317 if (mapData.Pointers[i]->getType()->isPointerTy())
8318 newV = builder.CreateLoad(type, mapData.Pointers[i]);
8320 newV = mapData.Pointers[i];
8323 auto curInsert = builder.saveIP();
8324 llvm::DebugLoc DbgLoc = builder.getCurrentDebugLocation();
8326 auto *memTempAlloc =
8327 builder.CreateAlloca(builder.getPtrTy(),
nullptr,
".casted");
8328 builder.SetCurrentDebugLocation(DbgLoc);
8329 builder.restoreIP(curInsert);
8331 builder.CreateStore(newV, memTempAlloc);
8332 newV = builder.CreateLoad(builder.getPtrTy(), memTempAlloc);
8335 mapData.Pointers[i] = newV;
8336 mapData.BasePointers[i] = newV;
8338 case omp::VariableCaptureKind::This:
8339 case omp::VariableCaptureKind::VLAType:
8340 mapData.MapClause[i]->emitOpError(
"Unhandled capture kind");
8351 MapInfoData &mapData,
8352 TargetDirectiveEnumTy targetDirective) {
8354 "function only supported for host device codegen");
8375 for (
size_t i = 0; i < mapData.MapClause.size(); ++i) {
8376 if (mapData.IsAMember[i])
8379 auto mapInfoOp = dyn_cast<omp::MapInfoOp>(mapData.MapClause[i]);
8380 if (!mapInfoOp.getMembers().empty()) {
8382 combinedInfo, mapData, i, targetDirective);
8391static llvm::Expected<llvm::Function *>
8393 LLVM::ModuleTranslation &moduleTranslation,
8394 llvm::StringRef mapperFuncName,
8395 TargetDirectiveEnumTy targetDirective);
8397static llvm::Expected<llvm::Function *>
8400 TargetDirectiveEnumTy targetDirective) {
8402 "function only supported for host device codegen");
8403 auto declMapperOp = cast<omp::DeclareMapperOp>(op);
8404 std::string mapperFuncName =
8406 {
"omp_mapper", declMapperOp.getSymName()});
8408 if (
auto *lookupFunc = moduleTranslation.
lookupFunction(mapperFuncName))
8416 if (llvm::Function *existingFunc =
8417 moduleTranslation.
getLLVMModule()->getFunction(mapperFuncName)) {
8418 moduleTranslation.
mapFunction(mapperFuncName, existingFunc);
8419 return existingFunc;
8423 mapperFuncName, targetDirective);
8426static llvm::Expected<llvm::Function *>
8429 llvm::StringRef mapperFuncName,
8430 TargetDirectiveEnumTy targetDirective) {
8432 "function only supported for host device codegen");
8433 auto declMapperOp = cast<omp::DeclareMapperOp>(op);
8434 auto declMapperInfoOp = declMapperOp.getDeclareMapperInfo();
8436 return llvm::make_error<PreviouslyReportedError>();
8440 llvm::Type *varType = moduleTranslation.
convertType(declMapperOp.getType());
8443 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
8446 MapInfosTy combinedInfo;
8448 [&](InsertPointTy codeGenIP, llvm::Value *ptrPHI,
8449 llvm::Value *unused2) -> llvm::OpenMPIRBuilder::MapInfosOrErrorTy {
8450 builder.restoreIP(codeGenIP);
8451 moduleTranslation.
mapValue(declMapperOp.getSymVal(), ptrPHI);
8452 moduleTranslation.
mapBlock(&declMapperOp.getRegion().front(),
8453 builder.GetInsertBlock());
8454 if (failed(moduleTranslation.
convertBlock(declMapperOp.getRegion().front(),
8457 return llvm::make_error<PreviouslyReportedError>();
8458 MapInfoData mapData;
8461 genMapInfos(builder, moduleTranslation, dl, combinedInfo, mapData,
8467 return combinedInfo;
8471 if (!combinedInfo.Mappers[i])
8474 moduleTranslation, targetDirective);
8478 genMapInfoCB, varType, mapperFuncName, customMapperCB,
8481 return newFn.takeError();
8482 if ([[maybe_unused]] llvm::Function *mappedFunc =
8484 assert(mappedFunc == *newFn &&
8485 "mapper function mapping disagrees with emitted function");
8487 moduleTranslation.
mapFunction(mapperFuncName, *newFn);
8493 llvm::OpenMPIRBuilder &ompBuilder,
8499 llvm::Function *parentFn = builder.GetInsertBlock()->getParent();
8500 llvm::StringRef fnName = parentFn ? parentFn->getName() :
"";
8502 fileLoc, ompBuilder, fnName, strSize);
8503 return ompBuilder.getOrCreateIdent(srcStr, strSize);
8509 llvm::Value *ifCond =
nullptr;
8510 llvm::Value *deviceID = builder.getInt64(llvm::omp::OMP_DEVICEID_UNDEF);
8514 llvm::omp::RuntimeFunction RTLFn;
8516 TargetDirectiveEnumTy targetDirective = getTargetDirectiveEnumTyFromOp(op);
8519 llvm::OpenMPIRBuilder::TargetDataInfo info(
8523 if (ompBuilder->Config.isTargetDevice())
8524 return op->
emitOpError() <<
"not allowed in a target device";
8526 bool isOffloadEntry = !ompBuilder->Config.TargetTriples.empty();
8528 auto getDeviceID = [&](
mlir::Value dev) -> llvm::Value * {
8529 llvm::Value *v = moduleTranslation.
lookupValue(dev);
8530 return builder.CreateIntCast(v, builder.getInt64Ty(),
true);
8535 .Case([&](omp::TargetDataOp dataOp) {
8539 if (
auto ifVar = dataOp.getIfExpr())
8543 deviceID = getDeviceID(devId);
8545 mapVars = dataOp.getMapVars();
8546 useDevicePtrVars = dataOp.getUseDevicePtrVars();
8547 useDeviceAddrVars = dataOp.getUseDeviceAddrVars();
8550 .Case([&](omp::TargetEnterDataOp enterDataOp) -> LogicalResult {
8554 if (
auto ifVar = enterDataOp.getIfExpr())
8558 deviceID = getDeviceID(devId);
8561 enterDataOp.getNowait()
8562 ? llvm::omp::OMPRTL___tgt_target_data_begin_nowait_mapper
8563 : llvm::omp::OMPRTL___tgt_target_data_begin_mapper;
8564 mapVars = enterDataOp.getMapVars();
8565 info.HasNoWait = enterDataOp.getNowait();
8568 .Case([&](omp::TargetExitDataOp exitDataOp) -> LogicalResult {
8572 if (
auto ifVar = exitDataOp.getIfExpr())
8576 deviceID = getDeviceID(devId);
8578 RTLFn = exitDataOp.getNowait()
8579 ? llvm::omp::OMPRTL___tgt_target_data_end_nowait_mapper
8580 : llvm::omp::OMPRTL___tgt_target_data_end_mapper;
8581 mapVars = exitDataOp.getMapVars();
8582 info.HasNoWait = exitDataOp.getNowait();
8585 .Case([&](omp::TargetUpdateOp updateDataOp) -> LogicalResult {
8589 if (
auto ifVar = updateDataOp.getIfExpr())
8593 deviceID = getDeviceID(devId);
8596 updateDataOp.getNowait()
8597 ? llvm::omp::OMPRTL___tgt_target_data_update_nowait_mapper
8598 : llvm::omp::OMPRTL___tgt_target_data_update_mapper;
8599 mapVars = updateDataOp.getMapVars();
8600 info.HasNoWait = updateDataOp.getNowait();
8603 .DefaultUnreachable(
"unexpected operation");
8608 if (!isOffloadEntry)
8609 ifCond = builder.getFalse();
8611 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
8612 MapInfoData mapData;
8614 builder, useDevicePtrVars, useDeviceAddrVars);
8617 MapInfosTy combinedInfo;
8618 auto genMapInfoCB = [&](InsertPointTy codeGenIP) -> MapInfosTy & {
8619 builder.restoreIP(codeGenIP);
8620 genMapInfos(builder, moduleTranslation, DL, combinedInfo, mapData,
8622 return combinedInfo;
8628 [&moduleTranslation](
8629 llvm::OpenMPIRBuilder::DeviceInfoTy type,
8633 for (
auto [arg, useDevVar] :
8634 llvm::zip_equal(blockArgs, useDeviceVars)) {
8636 auto getMapBasePtr = [](omp::MapInfoOp mapInfoOp) {
8637 return mapInfoOp.getVarPtrPtr() ? mapInfoOp.getVarPtrPtr()
8638 : mapInfoOp.getVarPtr();
8641 auto useDevMap = cast<omp::MapInfoOp>(useDevVar.getDefiningOp());
8642 for (
auto [mapClause, devicePointer, basePointer] : llvm::zip_equal(
8643 mapInfoData.MapClause, mapInfoData.DevicePointers,
8644 mapInfoData.BasePointers)) {
8645 auto mapOp = cast<omp::MapInfoOp>(mapClause);
8646 if (getMapBasePtr(mapOp) != getMapBasePtr(useDevMap) ||
8647 devicePointer != type)
8650 if (llvm::Value *devPtrInfoMap =
8651 mapper ? mapper(basePointer) : basePointer) {
8652 moduleTranslation.
mapValue(arg, devPtrInfoMap);
8659 using BodyGenTy = llvm::OpenMPIRBuilder::BodyGenTy;
8660 auto bodyGenCB = [&](InsertPointTy codeGenIP, BodyGenTy bodyGenType)
8661 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
8664 builder.restoreIP(codeGenIP);
8665 assert(isa<omp::TargetDataOp>(op) &&
8666 "BodyGen requested for non TargetDataOp");
8667 auto blockArgIface = cast<omp::BlockArgOpenMPOpInterface>(op);
8668 Region ®ion = cast<omp::TargetDataOp>(op).getRegion();
8669 switch (bodyGenType) {
8670 case BodyGenTy::Priv:
8672 if (!info.DevicePtrInfoMap.empty()) {
8673 mapUseDevice(llvm::OpenMPIRBuilder::DeviceInfoTy::Address,
8674 blockArgIface.getUseDeviceAddrBlockArgs(),
8675 useDeviceAddrVars, mapData,
8676 [&](llvm::Value *basePointer) -> llvm::Value * {
8677 if (!info.DevicePtrInfoMap[basePointer].second)
8679 return builder.CreateLoad(
8681 info.DevicePtrInfoMap[basePointer].second);
8683 mapUseDevice(llvm::OpenMPIRBuilder::DeviceInfoTy::Pointer,
8684 blockArgIface.getUseDevicePtrBlockArgs(), useDevicePtrVars,
8685 mapData, [&](llvm::Value *basePointer) {
8686 return info.DevicePtrInfoMap[basePointer].second;
8690 moduleTranslation)))
8691 return llvm::make_error<PreviouslyReportedError>();
8694 case BodyGenTy::DupNoPriv:
8695 if (info.DevicePtrInfoMap.empty()) {
8698 mapUseDevice(llvm::OpenMPIRBuilder::DeviceInfoTy::Address,
8699 blockArgIface.getUseDeviceAddrBlockArgs(),
8700 useDeviceAddrVars, mapData);
8701 mapUseDevice(llvm::OpenMPIRBuilder::DeviceInfoTy::Pointer,
8702 blockArgIface.getUseDevicePtrBlockArgs(), useDevicePtrVars,
8706 case BodyGenTy::NoPriv:
8708 if (info.DevicePtrInfoMap.empty()) {
8710 moduleTranslation)))
8711 return llvm::make_error<PreviouslyReportedError>();
8715 return builder.saveIP();
8718 auto customMapperCB =
8720 if (!combinedInfo.Mappers[i])
8722 info.HasMapper =
true;
8724 moduleTranslation, targetDirective);
8727 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
8729 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
8736 llvm::Value *srcLocOverride =
8740 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP = [&]() {
8741 if (isa<omp::TargetDataOp>(op))
8742 return ompBuilder->createTargetData(
8743 ompLoc, allocaIP, builder.saveIP(), deallocBlocks, deviceID, ifCond,
8744 info, genMapInfoCB, customMapperCB,
8746 nullptr, srcLocOverride);
8747 return ompBuilder->createTargetData(
8748 ompLoc, allocaIP, builder.saveIP(), deallocBlocks, deviceID, ifCond,
8749 info, genMapInfoCB, customMapperCB, &RTLFn,
8751 nullptr, srcLocOverride);
8757 builder.restoreIP(*afterIP);
8765 auto distributeOp = cast<omp::DistributeOp>(opInst);
8772 bool doDistributeReduction =
8776 unsigned numReductionVars = teamsOp ? teamsOp.getNumReductionVars() : 0;
8781 if (doDistributeReduction) {
8782 isByRef =
getIsByRef(teamsOp.getReductionByref());
8783 assert(isByRef.size() == teamsOp.getNumReductionVars());
8786 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
8790 llvm::cast<omp::BlockArgOpenMPOpInterface>(*teamsOp)
8791 .getReductionBlockArgs();
8794 teamsOp, reductionArgs, builder, moduleTranslation, allocaIP,
8795 reductionDecls, privateReductionVariables, reductionVariableMap,
8800 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
8802 [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
8805 builder.restoreIP(codeGenIP);
8809 distributeOp, builder, moduleTranslation, privVarsInfo, allocaIP);
8811 return llvm::make_error<PreviouslyReportedError>();
8817 moduleTranslation, allocaIP, deallocBlocks);
8822 return llvm::make_error<PreviouslyReportedError>();
8825 distributeOp, builder, moduleTranslation, privVarsInfo.
mlirVars,
8827 distributeOp.getPrivateNeedsBarrier())))
8828 return llvm::make_error<PreviouslyReportedError>();
8831 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
8834 builder, moduleTranslation);
8836 return regionBlock.takeError();
8837 builder.SetInsertPoint((*regionBlock)->begin());
8842 if (!isa_and_present<omp::WsloopOp>(distributeOp.getNestedWrapper())) {
8845 bool hasDistSchedule = distributeOp.getDistScheduleStatic();
8846 auto schedule = hasDistSchedule ? omp::ClauseScheduleKind::Distribute
8847 : omp::ClauseScheduleKind::Static;
8849 bool isOrdered = hasDistSchedule;
8850 std::optional<omp::ScheduleModifier> scheduleMod;
8851 bool isSimd =
false;
8852 llvm::omp::WorksharingLoopType workshareLoopType =
8853 llvm::omp::WorksharingLoopType::DistributeStaticLoop;
8854 bool loopNeedsBarrier =
false;
8855 llvm::Value *chunk = moduleTranslation.
lookupValue(
8856 distributeOp.getDistScheduleChunkSize());
8857 llvm::CanonicalLoopInfo *loopInfo =
8859 llvm::OpenMPIRBuilder::InsertPointOrErrorTy wsloopIP =
8860 ompBuilder->applyWorkshareLoop(
8861 ompLoc.DL, loopInfo, allocaIP, loopNeedsBarrier,
8862 convertToScheduleKind(schedule), chunk, isSimd,
8863 scheduleMod == omp::ScheduleModifier::monotonic,
8864 scheduleMod == omp::ScheduleModifier::nonmonotonic, isOrdered,
8865 workshareLoopType,
false, hasDistSchedule, chunk);
8868 return wsloopIP.takeError();
8871 distributeOp.getLoc(), privVarsInfo)))
8872 return llvm::make_error<PreviouslyReportedError>();
8874 return llvm::Error::success();
8878 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
8880 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
8881 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
8882 ompBuilder->createDistribute(ompLoc, allocaIP, deallocBlocks, bodyGenCB);
8887 builder.restoreIP(*afterIP);
8889 if (doDistributeReduction) {
8892 teamsOp, builder, moduleTranslation, allocaIP, reductionDecls,
8893 privateReductionVariables, isByRef,
8905 auto offloadMod = dyn_cast<omp::OffloadModuleInterface>(op);
8907 return op->
emitOpError() <<
"omp flags attached to non offload module op";
8911 if (offloadMod.getIsTargetDevice())
8912 ompBuilder->M.addModuleFlag(llvm::Module::Max,
"openmp-device",
8913 attribute.getOpenmpDeviceVersion());
8916 if (!offloadMod.getIsGPU())
8919 if (attribute.getNoGpuLib())
8922 ompBuilder->createGlobalFlag(attribute.getDebugKind(),
8923 "__omp_rtl_debug_kind");
8924 ompBuilder->createGlobalFlag(attribute.getAssumeTeamsOversubscription(),
8925 "__omp_rtl_assume_teams_oversubscription");
8926 ompBuilder->createGlobalFlag(attribute.getAssumeThreadsOversubscription(),
8927 "__omp_rtl_assume_threads_oversubscription");
8928 ompBuilder->createGlobalFlag(attribute.getAssumeNoThreadState(),
8929 "__omp_rtl_assume_no_thread_state");
8930 ompBuilder->createGlobalFlag(attribute.getAssumeNoNestedParallelism(),
8931 "__omp_rtl_assume_no_nested_parallelism");
8936 omp::TargetOp targetOp,
8937 llvm::OpenMPIRBuilder &ompBuilder,
8938 llvm::vfs::FileSystem &vfs,
8939 llvm::StringRef parentName =
"") {
8940 auto fileLoc = targetOp.getLoc()->findInstanceOf<
FileLineColLoc>();
8941 assert(fileLoc &&
"No file found from location");
8943 auto fileInfoCallBack = [&fileLoc]() {
8944 return std::pair<std::string, uint64_t>(
8945 llvm::StringRef(fileLoc.getFilename()), fileLoc.getLine());
8949 ompBuilder.getTargetEntryUniqueInfo(fileInfoCallBack, vfs, parentName);
8992 omp::TargetOp targetOp, MapInfoData &mapData, llvm::Argument &arg,
8993 llvm::Value *input, llvm::Value *&retVal, llvm::IRBuilderBase &builder,
8994 llvm::OpenMPIRBuilder &ompBuilder,
8996 llvm::IRBuilderBase::InsertPoint allocaIP,
8997 llvm::IRBuilderBase::InsertPoint codeGenIP,
8999 assert(ompBuilder.Config.isTargetDevice() &&
9000 "function only supported for target device codegen");
9001 builder.restoreIP(allocaIP);
9003 omp::VariableCaptureKind capture = omp::VariableCaptureKind::ByRef;
9005 ompBuilder.M.getContext());
9006 unsigned alignmentValue = 0;
9009 cast<omp::BlockArgOpenMPOpInterface>(*targetOp).getBlockArgsPairs(
9012 for (
size_t i = 0; i < mapData.MapClause.size(); ++i) {
9013 if (mapData.OriginalValue[i] == input) {
9014 auto mapOp = cast<omp::MapInfoOp>(mapData.MapClause[i]);
9015 capture = mapOp.getMapCaptureType();
9018 mapOp.getVarPtrType(), ompBuilder.M.getDataLayout());
9022 for (
auto &[val, arg] : blockArgsPairs) {
9023 if (mapOp.getResult() == val) {
9028 assert(mlirArg &&
"expected to find entry block argument for map clause");
9033 unsigned int allocaAS = ompBuilder.M.getDataLayout().getAllocaAddrSpace();
9034 unsigned int defaultAS =
9035 ompBuilder.M.getDataLayout().getProgramAddressSpace();
9038 llvm::Value *v =
nullptr;
9046 builder.SetInsertPoint(codeGenIP.getNodeParent()->getFirstInsertionPt());
9047 v = ompBuilder.createOMPAllocShared(builder, arg.getType());
9051 llvm::IRBuilderBase::InsertPointGuard guard(builder);
9052 for (
auto deallocIP : deallocIPs) {
9053 builder.SetInsertPoint(deallocIP);
9054 ompBuilder.createOMPFreeShared(builder, v, arg.getType());
9058 v = builder.CreateAlloca(arg.getType(), allocaAS);
9060 if (allocaAS != defaultAS && arg.getType()->isPointerTy())
9061 v = builder.CreateAddrSpaceCast(v, builder.getPtrTy(defaultAS));
9064 builder.CreateStore(&arg, v);
9066 builder.restoreIP(codeGenIP);
9069 case omp::VariableCaptureKind::ByCopy: {
9073 case omp::VariableCaptureKind::ByRef: {
9074 llvm::LoadInst *loadInst = builder.CreateAlignedLoad(
9076 ompBuilder.M.getDataLayout().getPrefTypeAlign(v->getType()));
9091 if (v->getType()->isPointerTy() && alignmentValue) {
9092 llvm::MDBuilder MDB(builder.getContext());
9093 loadInst->setMetadata(
9094 llvm::LLVMContext::MD_align,
9095 llvm::MDNode::get(builder.getContext(),
9096 MDB.createConstant(llvm::ConstantInt::get(
9097 llvm::Type::getInt64Ty(builder.getContext()),
9104 case omp::VariableCaptureKind::This:
9105 case omp::VariableCaptureKind::VLAType:
9108 assert(
false &&
"Currently unsupported capture kind");
9112 return builder.saveIP();
9129 auto blockArgIface = llvm::cast<omp::BlockArgOpenMPOpInterface>(*targetOp);
9130 for (
auto item : llvm::zip_equal(targetOp.getHostEvalVars(),
9131 blockArgIface.getHostEvalBlockArgs())) {
9132 Value hostEvalVar = std::get<0>(item), blockArg = std::get<1>(item);
9136 .Case([&](omp::TeamsOp teamsOp) {
9137 if (teamsOp.getNumTeamsLower() == blockArg)
9138 numTeamsLower = hostEvalVar;
9139 else if (llvm::is_contained(teamsOp.getNumTeamsUpperVars(),
9141 numTeamsUpper = hostEvalVar;
9142 else if (!teamsOp.getThreadLimitVars().empty() &&
9143 teamsOp.getThreadLimit(0) == blockArg)
9144 threadLimit = hostEvalVar;
9146 llvm_unreachable(
"unsupported host_eval use");
9148 .Case([&](omp::ParallelOp parallelOp) {
9149 if (!parallelOp.getNumThreadsVars().empty() &&
9150 parallelOp.getNumThreads(0) == blockArg)
9151 numThreads = hostEvalVar;
9153 llvm_unreachable(
"unsupported host_eval use");
9155 .Case([&](omp::LoopNestOp loopOp) {
9156 auto processBounds =
9160 for (
auto [i, lb] : llvm::enumerate(opBounds)) {
9161 if (lb == blockArg) {
9164 (*outBounds)[i] = hostEvalVar;
9170 processBounds(loopOp.getLoopLowerBounds(), lowerBounds);
9171 found = processBounds(loopOp.getLoopUpperBounds(), upperBounds) ||
9173 found = processBounds(loopOp.getLoopSteps(), steps) || found;
9175 assert(found &&
"unsupported host_eval use");
9177 .DefaultUnreachable(
"unsupported host_eval use");
9189template <
typename OpTy>
9194 if (OpTy casted = dyn_cast<OpTy>(op))
9197 if (immediateParent)
9198 return dyn_cast_if_present<OpTy>(op->
getParentOp());
9207 return std::nullopt;
9210 if (
auto constAttr = dyn_cast<IntegerAttr>(constOp.getValue()))
9211 return constAttr.getInt();
9213 return std::nullopt;
9218 uint64_t sizeInBytes = sizeInBits / 8;
9222template <
typename OpTy>
9224 if (op.getNumReductionVars() > 0) {
9229 members.reserve(reductions.size());
9230 for (omp::DeclareReductionOp &red : reductions) {
9234 if (red.getByrefElementType())
9235 members.push_back(*red.getByrefElementType());
9237 members.push_back(red.getType());
9240 auto structType = mlir::LLVM::LLVMStructType::getLiteral(
9256 llvm::OpenMPIRBuilder::TargetKernelDefaultAttrs &attrs,
9257 bool isTargetDevice,
bool isGPU) {
9260 Value numThreads, numTeamsLower, numTeamsUpper, threadLimit;
9261 if (!isTargetDevice) {
9269 numTeamsLower = teamsOp.getNumTeamsLower();
9271 if (!teamsOp.getNumTeamsUpperVars().empty())
9272 numTeamsUpper = teamsOp.getNumTeams(0);
9273 if (!teamsOp.getThreadLimitVars().empty())
9274 threadLimit = teamsOp.getThreadLimit(0);
9278 if (!parallelOp.getNumThreadsVars().empty())
9279 numThreads = parallelOp.getNumThreads(0);
9285 int32_t minTeamsVal = 1, maxTeamsVal = -1;
9289 if (numTeamsUpper) {
9291 minTeamsVal = maxTeamsVal = *val;
9293 minTeamsVal = maxTeamsVal = 0;
9299 minTeamsVal = maxTeamsVal = 1;
9301 minTeamsVal = maxTeamsVal = -1;
9306 auto setMaxValueFromClause = [](
Value clauseValue, int32_t &
result) {
9320 int32_t targetThreadLimitVal = -1, teamsThreadLimitVal = -1;
9321 if (!targetOp.getThreadLimitVars().empty())
9322 setMaxValueFromClause(targetOp.getThreadLimit(0), targetThreadLimitVal);
9323 setMaxValueFromClause(threadLimit, teamsThreadLimitVal);
9326 int32_t maxThreadsVal = -1;
9328 setMaxValueFromClause(numThreads, maxThreadsVal);
9336 int32_t combinedMaxThreadsVal = targetThreadLimitVal;
9337 if (combinedMaxThreadsVal < 0 ||
9338 (teamsThreadLimitVal >= 0 && teamsThreadLimitVal < combinedMaxThreadsVal))
9339 combinedMaxThreadsVal = teamsThreadLimitVal;
9341 if (combinedMaxThreadsVal < 0 ||
9342 (maxThreadsVal >= 0 && maxThreadsVal < combinedMaxThreadsVal))
9343 combinedMaxThreadsVal = maxThreadsVal;
9345 int32_t reductionDataSize = 0;
9346 if (isGPU && capturedOp) {
9353 omp::TargetExecMode execMode = targetOp.getKernelType();
9355 case omp::TargetExecMode::bare:
9356 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_BARE;
9358 case omp::TargetExecMode::generic:
9359 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_GENERIC;
9361 case omp::TargetExecMode::spmd:
9362 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_SPMD;
9364 case omp::TargetExecMode::spmd_no_loop:
9365 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_SPMD_NO_LOOP;
9368 attrs.MinTeams.front() = minTeamsVal;
9369 attrs.MaxTeams.front() = maxTeamsVal;
9370 attrs.MinThreads.front() = 1;
9371 attrs.MaxThreads.front() = combinedMaxThreadsVal;
9372 attrs.ReductionDataSize = reductionDataSize;
9384 omp::TargetOp targetOp,
Operation *capturedOp,
9385 llvm::OpenMPIRBuilder::TargetKernelRuntimeAttrs &attrs) {
9387 unsigned numLoops = loopOp ? loopOp.getNumLoops() : 0;
9389 Value numThreads, numTeamsLower, numTeamsUpper, teamsThreadLimit;
9393 teamsThreadLimit, &lowerBounds, &upperBounds, &steps);
9396 if (!targetOp.getThreadLimitVars().empty()) {
9397 Value targetThreadLimit = targetOp.getThreadLimit(0);
9398 attrs.TargetThreadLimit.front() =
9406 attrs.MinTeams.front() = builder.CreateSExtOrTrunc(
9407 moduleTranslation.
lookupValue(numTeamsLower), builder.getInt32Ty());
9410 attrs.MaxTeams.front() = builder.CreateSExtOrTrunc(
9411 moduleTranslation.
lookupValue(numTeamsUpper), builder.getInt32Ty());
9413 if (teamsThreadLimit)
9414 attrs.TeamsThreadLimit.front() = builder.CreateSExtOrTrunc(
9415 moduleTranslation.
lookupValue(teamsThreadLimit), builder.getInt32Ty());
9418 attrs.MaxThreads.front() = moduleTranslation.
lookupValue(numThreads);
9420 if (targetOp.hasHostEvalTripCount()) {
9422 attrs.LoopTripCount =
nullptr;
9427 for (
auto [loopLower, loopUpper, loopStep] :
9428 llvm::zip_equal(lowerBounds, upperBounds, steps)) {
9429 llvm::Value *lowerBound = moduleTranslation.
lookupValue(loopLower);
9430 llvm::Value *upperBound = moduleTranslation.
lookupValue(loopUpper);
9431 llvm::Value *step = moduleTranslation.
lookupValue(loopStep);
9433 if (!lowerBound || !upperBound || !step) {
9434 attrs.LoopTripCount =
nullptr;
9438 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
9439 llvm::Value *tripCount = ompBuilder->calculateCanonicalLoopTripCount(
9440 loc, lowerBound, upperBound, step,
true,
9441 loopOp.getLoopInclusive());
9443 if (!attrs.LoopTripCount) {
9444 attrs.LoopTripCount = tripCount;
9449 attrs.LoopTripCount = builder.CreateMul(attrs.LoopTripCount, tripCount,
9454 attrs.DeviceID = builder.getInt64(llvm::omp::OMP_DEVICEID_UNDEF);
9456 attrs.DeviceID = moduleTranslation.
lookupValue(devId);
9458 builder.CreateSExtOrTrunc(attrs.DeviceID, builder.getInt64Ty());
9462static llvm::omp::OMPDynGroupprivateFallbackType
9464 omp::FallbackModifier fb = fallbackAttr ? fallbackAttr.getValue()
9465 : omp::FallbackModifier::default_mem;
9467 case omp::FallbackModifier::abort:
9468 return llvm::omp::OMPDynGroupprivateFallbackType::Abort;
9469 case omp::FallbackModifier::null:
9470 return llvm::omp::OMPDynGroupprivateFallbackType::Null;
9471 case omp::FallbackModifier::default_mem:
9472 return llvm::omp::OMPDynGroupprivateFallbackType::DefaultMem;
9475 llvm_unreachable(
"unexpected dyn_groupprivate fallback type");
9481 auto targetOp = cast<omp::TargetOp>(opInst);
9486 llvm::DebugLoc outlinedFnLoc = builder.getCurrentDebugLocation();
9495 llvm::BasicBlock *parentBB = builder.GetInsertBlock();
9496 assert(parentBB &&
"No insert block is set for the builder");
9497 llvm::Function *parentLLVMFn = parentBB->getParent();
9498 assert(parentLLVMFn &&
"Parent Function must be valid");
9499 if (llvm::DISubprogram *SP = parentLLVMFn->getSubprogram())
9500 builder.SetCurrentDebugLocation(llvm::DILocation::get(
9501 parentLLVMFn->getContext(), outlinedFnLoc.getLine(),
9502 outlinedFnLoc.getCol(), SP, outlinedFnLoc.getInlinedAt()));
9510 llvm::DebugLoc outlinedFnDbgLoc;
9511 if (outlinedFnLoc && parentLLVMFn->getSubprogram())
9512 outlinedFnDbgLoc = outlinedFnLoc;
9515 bool isTargetDevice = ompBuilder->Config.isTargetDevice();
9516 bool isGPU = ompBuilder->Config.isGPU();
9519 auto argIface = cast<omp::BlockArgOpenMPOpInterface>(opInst);
9520 auto &targetRegion = targetOp.getRegion();
9537 llvm::Function *llvmOutlinedFn =
nullptr;
9538 TargetDirectiveEnumTy targetDirective =
9539 getTargetDirectiveEnumTyFromOp(&opInst);
9543 bool isOffloadEntry =
9544 isTargetDevice || !ompBuilder->Config.TargetTriples.empty();
9564 if (!targetOp.getInReductionVars().empty() && !isTargetDevice) {
9565 inRedOrigPtrs.reserve(targetOp.getInReductionVars().size());
9566 inRedMapArgIdx.reserve(targetOp.getInReductionVars().size());
9567 for (
Value v : targetOp.getInReductionVars()) {
9572 std::optional<unsigned> matchIdx;
9573 for (
auto [idx, mapV] : llvm::enumerate(targetOp.getMapVars())) {
9574 auto mapInfo = mapV.getDefiningOp<omp::MapInfoOp>();
9575 if (v != mapInfo.getVarPtr())
9578 return targetOp.emitError()
9579 <<
"in_reduction variable on omp.target has multiple matching "
9580 "map_entries entries; the redirect target is ambiguous";
9586 "TargetOp verifier guarantees a matching map_entries entry for "
9587 "each in_reduction variable");
9588 inRedMapArgIdx.push_back(*matchIdx);
9591 inRedOrigPtrs.push_back(moduleTranslation.
lookupValue(v));
9600 if (!targetOp.getPrivateVars().empty() && !targetOp.getMapVars().empty()) {
9602 std::optional<ArrayAttr> privateSyms = targetOp.getPrivateSyms();
9603 std::optional<DenseI64ArrayAttr> privateMapIndices =
9604 targetOp.getPrivateMapsAttr();
9606 for (
auto [privVarIdx, privVarSymPair] :
9607 llvm::enumerate(llvm::zip_equal(privateVars, *privateSyms))) {
9608 auto privVar = std::get<0>(privVarSymPair);
9609 auto privSym = std::get<1>(privVarSymPair);
9611 SymbolRefAttr privatizerName = llvm::cast<SymbolRefAttr>(privSym);
9612 omp::PrivateClauseOp privatizer =
9615 if (!privatizer.needsMap())
9619 targetOp.getMappedValueForPrivateVar(privVarIdx);
9620 assert(mappedValue &&
"Expected to find mapped value for a privatized "
9621 "variable that needs mapping");
9626 auto mapInfoOp = mappedValue.
getDefiningOp<omp::MapInfoOp>();
9627 [[maybe_unused]]
Type varType = mapInfoOp.getVarPtrType();
9631 if (!isa<LLVM::LLVMPointerType>(privVar.getType()))
9633 varType == privVar.getType() &&
9634 "Type of private var doesn't match the type of the mapped value");
9638 mappedPrivateVars.insert(
9640 targetRegion.getArgument(argIface.getMapBlockArgsStart() +
9641 (*privateMapIndices)[privVarIdx])});
9645 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
9646 auto bodyCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
9648 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
9649 llvm::IRBuilderBase::InsertPointGuard guard(builder);
9650 builder.SetCurrentDebugLocation(llvm::DebugLoc());
9653 llvm::Function *llvmParentFn =
9655 llvmOutlinedFn = codeGenIP.getNodeParent()->getParent();
9656 assert(llvmParentFn && llvmOutlinedFn &&
9657 "Both parent and outlined functions must exist at this point");
9659 if (outlinedFnLoc && llvmParentFn->getSubprogram())
9660 llvmOutlinedFn->setSubprogram(outlinedFnLoc->getScope()->getSubprogram());
9662 if (
auto attr = llvmParentFn->getFnAttribute(
"target-cpu");
9663 attr.isStringAttribute())
9664 llvmOutlinedFn->addFnAttr(attr);
9666 if (
auto attr = llvmParentFn->getFnAttribute(
"target-features");
9667 attr.isStringAttribute())
9668 llvmOutlinedFn->addFnAttr(attr);
9670 for (
auto [idx, arg] : llvm::enumerate(mapBlockArgs)) {
9676 if (llvm::is_contained(inRedMapArgIdx, idx))
9678 auto mapInfoOp = cast<omp::MapInfoOp>(mapVars[idx].getDefiningOp());
9679 llvm::Value *mapOpValue =
9680 moduleTranslation.
lookupValue(mapInfoOp.getVarPtr());
9681 moduleTranslation.
mapValue(arg, mapOpValue);
9683 for (
auto [arg, mapOp] : llvm::zip_equal(hdaBlockArgs, hdaVars)) {
9684 auto mapInfoOp = cast<omp::MapInfoOp>(mapOp.getDefiningOp());
9685 llvm::Value *mapOpValue =
9686 moduleTranslation.
lookupValue(mapInfoOp.getVarPtr());
9687 moduleTranslation.
mapValue(arg, mapOpValue);
9696 privateVarsInfo, allocaIP, &mappedPrivateVars);
9699 return llvm::make_error<PreviouslyReportedError>();
9701 builder.restoreIP(codeGenIP);
9703 &mappedPrivateVars),
9706 return llvm::make_error<PreviouslyReportedError>();
9709 targetOp, builder, moduleTranslation, privateVarsInfo.
mlirVars,
9711 targetOp.getPrivateNeedsBarrier(), &mappedPrivateVars)))
9712 return llvm::make_error<PreviouslyReportedError>();
9723 if (!inRedOrigPtrs.empty()) {
9729 inRedResultPtrTys.reserve(inRedMapArgIdx.size());
9730 for (
unsigned mapArgIdx : inRedMapArgIdx)
9731 inRedResultPtrTys.push_back(
9732 moduleTranslation.
convertType(mapBlockArgs[mapArgIdx].getType()));
9734 llvm::OpenMPIRBuilder::LocationDescription bodyLoc(builder);
9735 llvm::OpenMPIRBuilder::InsertPointTy redIP =
9736 ompBuilder->createTargetInReduction(
9737 bodyLoc, inRedOrigPtrs, inRedResultPtrTys,
9738 [&](
unsigned idx, llvm::Value *priv) {
9739 moduleTranslation.
mapValue(mapBlockArgs[inRedMapArgIdx[idx]],
9742 builder.restoreIP(redIP);
9746 moduleTranslation, allocaIP, deallocBlocks);
9748 targetRegion,
"omp.target", builder, moduleTranslation);
9751 return llvm::make_error<PreviouslyReportedError>();
9753 builder.SetInsertPoint(exitBlock.get()->getTerminator());
9756 targetOp.getLoc(), privateVarsInfo)))
9757 return llvm::make_error<PreviouslyReportedError>();
9759 return builder.saveIP();
9762 StringRef parentName = parentFn.getName();
9764 llvm::TargetRegionEntryInfo entryInfo;
9770 MapInfoData mapData;
9775 MapInfosTy combinedInfos;
9777 [&](llvm::OpenMPIRBuilder::InsertPointTy codeGenIP) -> MapInfosTy & {
9778 builder.restoreIP(codeGenIP);
9779 genMapInfos(builder, moduleTranslation, dl, combinedInfos, mapData,
9784 auto *nullPtr = llvm::Constant::getNullValue(builder.getPtrTy());
9785 combinedInfos.BasePointers.push_back(nullPtr);
9786 combinedInfos.Pointers.push_back(nullPtr);
9787 combinedInfos.DevicePointers.push_back(
9788 llvm::OpenMPIRBuilder::DeviceInfoTy::None);
9789 combinedInfos.Sizes.push_back(builder.getInt64(0));
9790 combinedInfos.Types.push_back(
9791 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TARGET_PARAM |
9792 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL);
9794 combinedInfos.HasAttachPtr.push_back(
false);
9795 if (!combinedInfos.Names.empty())
9796 combinedInfos.Names.push_back(nullPtr);
9797 combinedInfos.Mappers.push_back(
nullptr);
9799 return combinedInfos;
9802 auto argAccessorCB = [&](llvm::Argument &arg, llvm::Value *input,
9803 llvm::Value *&retVal, InsertPointTy allocaIP,
9804 InsertPointTy codeGenIP,
9806 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
9807 llvm::IRBuilderBase::InsertPointGuard guard(builder);
9808 builder.SetCurrentDebugLocation(llvm::DebugLoc());
9814 if (!isTargetDevice) {
9815 retVal = cast<llvm::Value>(&arg);
9820 builder, *ompBuilder, moduleTranslation,
9821 allocaIP, codeGenIP, deallocIPs);
9824 llvm::OpenMPIRBuilder::TargetKernelRuntimeAttrs runtimeAttrs;
9825 llvm::OpenMPIRBuilder::TargetKernelDefaultAttrs defaultAttrs;
9827 cast<omp::ComposableOpInterface>(*targetOp).findCapturedOp();
9829 isTargetDevice, isGPU);
9833 if (!isTargetDevice)
9835 targetCapturedOp, runtimeAttrs);
9843 for (
auto [arg, var] : llvm::zip_equal(hostEvalBlockArgs, hostEvalVars)) {
9844 llvm::Value *value = moduleTranslation.
lookupValue(var);
9845 moduleTranslation.
mapValue(arg, value);
9847 if (!llvm::isa<llvm::Constant>(value))
9848 kernelInput.push_back(value);
9851 for (
size_t i = 0, e = mapData.OriginalValue.size(); i != e; ++i) {
9861 using MapFlags = llvm::omp::OpenMPOffloadMappingFlags;
9862 bool isAttachMap = (mapData.Types[i] & MapFlags::OMP_MAP_ATTACH) ==
9863 MapFlags::OMP_MAP_ATTACH;
9864 bool isPrivateTargetParam =
9866 (MapFlags::OMP_MAP_PRIVATE | MapFlags::OMP_MAP_TARGET_PARAM)) ==
9867 (MapFlags::OMP_MAP_PRIVATE | MapFlags::OMP_MAP_TARGET_PARAM);
9869 if (!mapData.IsDeclareTarget[i] && !mapData.IsAMember[i] &&
9870 (!isAttachMap || (isAttachMap && isPrivateTargetParam)))
9871 kernelInput.push_back(mapData.OriginalValue[i]);
9875 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
9878 llvm::OpenMPIRBuilder::DependenciesInfo dds;
9880 targetOp.getDependVars(), targetOp.getDependKinds(),
9881 targetOp.getDependIterated(), targetOp.getDependIteratedKinds(),
9882 builder, moduleTranslation, dds)))
9885 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
9887 llvm::OpenMPIRBuilder::TargetDataInfo info(
9891 auto customMapperCB =
9893 if (!combinedInfos.Mappers[i])
9895 info.HasMapper =
true;
9897 moduleTranslation, targetDirective);
9900 llvm::Value *ifCond =
nullptr;
9901 if (
Value targetIfCond = targetOp.getIfExpr())
9902 ifCond = moduleTranslation.
lookupValue(targetIfCond);
9904 Value dynGroupPrivateSize = targetOp.getDynGroupprivateSize();
9905 llvm::Value *dynSizeVal =
nullptr;
9906 if (dynGroupPrivateSize) {
9907 dynSizeVal = moduleTranslation.
lookupValue(dynGroupPrivateSize);
9908 dynSizeVal = builder.CreateIntCast(dynSizeVal, builder.getInt32Ty(),
9912 llvm::omp::OMPDynGroupprivateFallbackType fallbackType =
9918 llvm::Value *rtLocOverride =
9919 (!isTargetDevice && isOffloadEntry)
9923 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
9925 ompLoc, isOffloadEntry, allocaIP, builder.saveIP(), deallocBlocks,
9926 info, entryInfo, defaultAttrs, runtimeAttrs, ifCond, kernelInput,
9927 genMapInfoCB, bodyCB, argAccessorCB, customMapperCB, dds,
9928 targetOp.getNowait(), dynSizeVal, fallbackType, outlinedFnDbgLoc,
9934 builder.restoreIP(*afterIP);
9937 builder.CreateFree(dds.DepArray);
9944 llvm::OpenMPIRBuilder *ompBuilder,
9953 if (FunctionOpInterface funcOp = dyn_cast<FunctionOpInterface>(op)) {
9954 if (
auto offloadMod = dyn_cast<omp::OffloadModuleInterface>(
9956 if (!offloadMod.getIsTargetDevice())
9959 omp::DeclareTargetDeviceType declareType = attribute.getDeviceType();
9961 if (declareType == omp::DeclareTargetDeviceType::host) {
9962 llvm::Function *llvmFunc =
9964 llvmFunc->dropAllReferences();
9965 llvmFunc->eraseFromParent();
9969 ompBuilder->Builder.ClearInsertionPoint();
9970 ompBuilder->Builder.SetCurrentDebugLocation(llvm::DebugLoc());
9971 }
else if (llvm::Function *llvmFunc =
9983 if (!llvmFunc->isDeclaration() && llvmFunc->hasExternalLinkage() &&
9984 llvmFunc->getVisibility() == llvm::GlobalValue::DefaultVisibility)
9985 llvmFunc->setVisibility(llvm::GlobalValue::HiddenVisibility);
9991 if (LLVM::GlobalOp gOp = dyn_cast<LLVM::GlobalOp>(op)) {
9992 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
9993 if (
auto *gVal = llvmModule->getNamedValue(gOp.getSymName())) {
9994 auto *gVar = cast<llvm::GlobalVariable>(gVal);
9996 bool isDeclaration = gOp.isDeclaration();
9997 bool isExternallyVisible =
10000 llvm::StringRef mangledName = gOp.getSymName();
10001 mlir::omp::DeclareTargetCaptureClause captureClause =
10002 attribute.getCaptureClause();
10005 llvm::StringRef entryMangledName = mangledName;
10006 llvm::Constant *entryAddr = llvm::cast<llvm::Constant>(gVal);
10007 std::function<llvm::GlobalValue::LinkageTypes()> variableLinkage;
10009 bool requiresUSM = ompBuilder->Config.hasRequiresUnifiedSharedMemory();
10011 captureClause == omp::DeclareTargetCaptureClause::to ||
10012 captureClause == omp::DeclareTargetCaptureClause::enter;
10014 attribute.getDeviceType() == omp::DeclareTargetDeviceType::host;
10019 if (isToOrEnter && !isHostOnly && !requiresUSM &&
10020 gVar->hasLocalLinkage()) {
10021 gVar->setLinkage(llvm::GlobalValue::ExternalLinkage);
10022 isExternallyVisible =
true;
10026 if (ompBuilder->Config.isTargetDevice())
10027 gVar->setDSOLocal(
false);
10032 llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseAny &&
10033 !requiresUSM && !isDeclaration &&
10034 (gVal->hasLocalLinkage() || gVal->hasHiddenVisibility())) {
10038 entryNameStorage = (mangledName + llvm::Twine(
"_decl_tgt_entry")).str();
10039 entryMangledName = entryNameStorage;
10040 if (llvm::GlobalValue *existing =
10041 llvmModule->getNamedValue(entryMangledName)) {
10042 entryAddr = llvm::cast<llvm::Constant>(existing);
10044 entryAddr = llvm::GlobalAlias::create(
10045 gVal->getValueType(), gVal->getAddressSpace(),
10046 llvm::GlobalValue::WeakAnyLinkage, entryMangledName, entryAddr,
10048 llvm::cast<llvm::GlobalAlias>(entryAddr)->setVisibility(
10049 llvm::GlobalValue::DefaultVisibility);
10051 variableLinkage = [] {
return llvm::GlobalValue::WeakAnyLinkage; };
10055 std::vector<llvm::GlobalVariable *> generatedRefs;
10057 std::vector<llvm::Triple> targetTriple;
10058 auto targetTripleAttr = dyn_cast_or_null<mlir::StringAttr>(
10060 LLVM::LLVMDialect::getTargetTripleAttrName()));
10061 if (targetTripleAttr)
10062 targetTriple.emplace_back(targetTripleAttr.data());
10064 auto fileInfoCallBack = [&loc]() {
10065 std::string filename =
"";
10066 std::uint64_t lineNo = 0;
10069 filename = loc.getFilename().str();
10070 lineNo = loc.getLine();
10073 return std::pair<std::string, std::uint64_t>(llvm::StringRef(filename),
10077 llvm::vfs::FileSystem &vfs = moduleTranslation.
getFileSystem();
10078 ompBuilder->registerTargetGlobalVariable(
10079 captureClauseKind, deviceClause, isDeclaration, isExternallyVisible,
10080 ompBuilder->getTargetEntryUniqueInfo(fileInfoCallBack, vfs),
10081 entryMangledName, generatedRefs,
false, targetTriple,
10082 nullptr, variableLinkage, gVal->getType(),
10085 if (ompBuilder->Config.isTargetDevice() &&
10086 (captureClause == omp::DeclareTargetCaptureClause::link ||
10091 llvm::Type *ptrTy = llvm::PointerType::get(llvmModule->getContext(), 0);
10092 llvm::Constant *refPtr = ompBuilder->getAddrOfDeclareTargetVar(
10093 captureClauseKind, deviceClause, isDeclaration, isExternallyVisible,
10094 ompBuilder->getTargetEntryUniqueInfo(fileInfoCallBack, vfs),
10095 mangledName, generatedRefs,
false, targetTriple,
10106 if (!gVar->isDeclaration())
10107 gVar->setLinkage(llvm::GlobalValue::InternalLinkage);
10113 dyn_cast<llvm::GlobalValue>(refPtr->stripPointerCasts()))
10114 ompBuilder->registerDeclareTargetGlobalReplacement(gVal, newGV);
10121 if (ompBuilder->Config.isTargetDevice() && isHostOnly && isToOrEnter) {
10122 gVar->setLinkage(llvm::GlobalValue::ExternalLinkage);
10123 gVar->setInitializer(
nullptr);
10135class OpenMPDialectLLVMIRTranslationInterface
10136 :
public LLVMTranslationDialectInterface {
10138 using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;
10143 convertOperation(Operation *op, llvm::IRBuilderBase &builder,
10144 LLVM::ModuleTranslation &moduleTranslation)
const final;
10149 amendOperation(Operation *op, ArrayRef<llvm::Instruction *> instructions,
10150 NamedAttribute attribute,
10151 LLVM::ModuleTranslation &moduleTranslation)
const final;
10156 void registerAllocatedPtr(Value var, llvm::Value *ptr)
const {
10157 ompAllocatedPtrs[var] = ptr;
10162 llvm::Value *lookupAllocatedPtr(Value var)
const {
10163 auto it = ompAllocatedPtrs.find(var);
10164 return it != ompAllocatedPtrs.end() ? it->second :
nullptr;
10176LogicalResult OpenMPDialectLLVMIRTranslationInterface::amendOperation(
10177 Operation *op, ArrayRef<llvm::Instruction *> instructions,
10178 NamedAttribute attribute,
10179 LLVM::ModuleTranslation &moduleTranslation)
const {
10180 return llvm::StringSwitch<llvm::function_ref<LogicalResult(Attribute)>>(
10182 .Case(
"omp.is_target_device",
10183 [&](Attribute attr) {
10184 if (
auto deviceAttr = dyn_cast<BoolAttr>(attr)) {
10185 llvm::OpenMPIRBuilderConfig &config =
10187 config.setIsTargetDevice(deviceAttr.getValue());
10192 .Case(
"omp.is_gpu",
10193 [&](Attribute attr) {
10194 if (
auto gpuAttr = dyn_cast<BoolAttr>(attr)) {
10195 llvm::OpenMPIRBuilderConfig &config =
10197 config.setIsGPU(gpuAttr.getValue());
10202 .Case(
"omp.host_ir_filepath",
10203 [&](Attribute attr) {
10204 if (
auto filepathAttr = dyn_cast<StringAttr>(attr)) {
10205 llvm::OpenMPIRBuilder *ompBuilder =
10207 ompBuilder->loadOffloadInfoMetadata(
10208 moduleTranslation.
getFileSystem(), filepathAttr.getValue());
10214 [&](Attribute attr) {
10215 if (
auto rtlAttr = dyn_cast<omp::FlagsAttr>(attr))
10219 .Case(
"omp.version",
10220 [&](Attribute attr) {
10221 if (
auto versionAttr = dyn_cast<omp::VersionAttr>(attr)) {
10222 llvm::OpenMPIRBuilder *ompBuilder =
10224 ompBuilder->M.addModuleFlag(llvm::Module::Max,
"openmp",
10225 versionAttr.getVersion());
10230 .Case(
"omp.declare_target",
10231 [&](Attribute attr) {
10232 if (
auto declareTargetAttr =
10233 dyn_cast<omp::DeclareTargetAttr>(attr)) {
10234 llvm::OpenMPIRBuilder *ompBuilder =
10237 ompBuilder, moduleTranslation);
10241 .Case(
"omp.requires",
10242 [&](Attribute attr) {
10243 if (
auto requiresAttr = dyn_cast<omp::ClauseRequiresAttr>(attr)) {
10244 using Requires = omp::ClauseRequires;
10245 Requires flags = requiresAttr.getValue();
10246 llvm::OpenMPIRBuilderConfig &config =
10248 config.setHasRequiresReverseOffload(
10249 bitEnumContainsAll(flags, Requires::reverse_offload));
10250 config.setHasRequiresUnifiedAddress(
10251 bitEnumContainsAll(flags, Requires::unified_address));
10252 config.setHasRequiresUnifiedSharedMemory(
10253 bitEnumContainsAll(flags, Requires::unified_shared_memory));
10254 config.setHasRequiresDynamicAllocators(
10255 bitEnumContainsAll(flags, Requires::dynamic_allocators));
10260 .Case(
"omp.target_triples",
10261 [&](Attribute attr) {
10262 if (
auto triplesAttr = dyn_cast<ArrayAttr>(attr)) {
10263 llvm::OpenMPIRBuilderConfig &config =
10265 config.TargetTriples.clear();
10266 config.TargetTriples.reserve(triplesAttr.size());
10267 for (Attribute tripleAttr : triplesAttr) {
10268 if (
auto tripleStrAttr = dyn_cast<StringAttr>(tripleAttr))
10269 config.TargetTriples.emplace_back(tripleStrAttr.getValue());
10277 .Case(
"omp.integer_wrap_around",
10278 [&](Attribute attr) {
10279 if (
auto wrapAttr = dyn_cast<omp::IntegerWrapAroundAttr>(attr)) {
10280 llvm::OpenMPIRBuilderConfig &config =
10282 config.setNoSignedWrap(!wrapAttr.getIntegerWrapAround());
10287 .Default([](Attribute) {
10303 if (
auto declareTargetIface =
10304 llvm::dyn_cast<mlir::omp::DeclareTargetInterface>(
10305 parentFn.getOperation())) {
10306 omp::DeclareTargetAttr declareTargetAttr =
10307 declareTargetIface.getDeclareTarget();
10308 if (declareTargetAttr && declareTargetAttr.getDeviceType() !=
10309 mlir::omp::DeclareTargetDeviceType::host)
10320 llvm::Module *llvmModule) {
10321 llvm::Type *i64Ty = builder.getInt64Ty();
10322 llvm::Type *i32Ty = builder.getInt32Ty();
10323 llvm::Type *returnType = builder.getPtrTy(0);
10324 llvm::FunctionType *fnType =
10325 llvm::FunctionType::get(returnType, {i64Ty, i32Ty},
false);
10326 llvm::Function *
func = cast<llvm::Function>(
10327 llvmModule->getOrInsertFunction(
"omp_target_alloc", fnType).getCallee());
10331template <
typename T>
10332static llvm::Value *
10335 llvm::DataLayout dataLayout =
10337 llvm::Type *llvmHeapTy =
10338 moduleTranslation.
convertType(op.getMemElemTypeAttr().getValue());
10340 auto alignment = op.getMemAlignment();
10341 llvm::TypeSize typeSize = llvm::alignTo(
10342 dataLayout.getTypeStoreSize(llvmHeapTy),
10343 alignment ? *alignment : dataLayout.getABITypeAlign(llvmHeapTy).value());
10345 llvm::Value *allocSize = builder.getInt64(typeSize.getFixedValue());
10346 return builder.CreateMul(
10348 builder.CreateIntCast(moduleTranslation.
lookupValue(op.getMemArraySize()),
10349 builder.getInt64Ty(),
10356 omp::TargetAllocMemOp op) {
10357 llvm::DataLayout dataLayout =
10359 llvm::Type *llvmHeapTy = moduleTranslation.
convertType(op.getAllocatedType());
10360 llvm::TypeSize typeSize = dataLayout.getTypeAllocSize(llvmHeapTy);
10361 llvm::Value *allocSize = builder.getInt64(typeSize.getFixedValue());
10362 for (
auto typeParam : op.getTypeparams()) {
10363 allocSize = builder.CreateMul(
10365 builder.CreateIntCast(moduleTranslation.
lookupValue(typeParam),
10366 builder.getInt64Ty(),
10372static LogicalResult
10375 auto allocMemOp = cast<omp::TargetAllocMemOp>(opInst);
10380 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
10384 llvm::Value *llvmDeviceNum = moduleTranslation.
lookupValue(deviceNum);
10386 llvm::Value *allocSize =
10389 llvm::CallInst *call =
10390 builder.CreateCall(ompTargetAllocFunc, {allocSize, llvmDeviceNum});
10391 llvm::Value *resultI64 = builder.CreatePtrToInt(call, builder.getInt64Ty());
10394 moduleTranslation.
mapValue(allocMemOp.getResult(), resultI64);
10398static LogicalResult
10400 llvm::IRBuilderBase &builder,
10404 moduleTranslation.
mapValue(allocMemOp.getResult(),
10405 ompBuilder->createOMPAllocShared(builder, size));
10409static LogicalResult
10412 const OpenMPDialectLLVMIRTranslationInterface &ompIface) {
10413 auto allocateDirOp = cast<omp::AllocateDirOp>(opInst);
10416 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
10417 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
10418 llvm::DataLayout dataLayout = llvmModule->getDataLayout();
10420 std::optional<int64_t> alignAttr = allocateDirOp.getAlign();
10422 llvm::Value *allocator;
10423 if (
auto allocatorVar = allocateDirOp.getAllocator()) {
10424 allocator = moduleTranslation.
lookupValue(allocatorVar);
10425 if (allocator->getType()->isIntegerTy())
10426 allocator = builder.CreateIntToPtr(allocator, builder.getPtrTy());
10427 else if (allocator->getType()->isPointerTy())
10428 allocator = builder.CreatePointerBitCastOrAddrSpaceCast(
10429 allocator, builder.getPtrTy());
10431 allocator = llvm::ConstantPointerNull::get(builder.getPtrTy());
10434 for (
Value var : vars) {
10436 llvm::Type *typeToInspect =
10441 var, baseVar, moduleTranslation, builder, dataLayout)) {
10442 size = *dynamicSize;
10443 }
else if (typeToInspect->isArrayTy()) {
10444 size = builder.getInt64(
10445 dataLayout.getTypeAllocSize(typeToInspect).getFixedValue());
10447 size = builder.getInt64(
10448 dataLayout.getTypeAllocSize(typeToInspect).getFixedValue());
10451 uint64_t alignValue =
10452 alignAttr ? alignAttr.value()
10453 : dataLayout.getABITypeAlign(typeToInspect).value();
10454 llvm::Value *alignConst = builder.getInt64(alignValue);
10456 size = builder.CreateAdd(size, builder.getInt64(alignValue - 1),
"",
true);
10457 size = builder.CreateUDiv(size, alignConst);
10458 size = builder.CreateMul(size, alignConst,
"",
true);
10460 std::string allocName =
10461 ompBuilder->createPlatformSpecificName({
".void.addr"});
10462 llvm::CallInst *allocCall;
10463 if (alignAttr.has_value()) {
10464 allocCall = ompBuilder->createOMPAlignedAlloc(
10465 ompLoc, builder.getInt64(alignAttr.value()), size, allocator,
10469 ompBuilder->createOMPAlloc(ompLoc, size, allocator, allocName);
10472 ompIface.registerAllocatedPtr(var, allocCall);
10474 if (llvm::Value *baseLlvm = moduleTranslation.
lookupValue(baseVar)) {
10475 llvm::Value *boundPtr = builder.CreatePointerBitCastOrAddrSpaceCast(
10476 allocCall, baseLlvm->getType());
10478 }
else if (llvm::Value *varLlvm = moduleTranslation.
lookupValue(var)) {
10479 llvm::Value *boundPtr = builder.CreatePointerBitCastOrAddrSpaceCast(
10480 allocCall, varLlvm->getType());
10488static LogicalResult
10491 const OpenMPDialectLLVMIRTranslationInterface &ompIface) {
10492 auto freeOp = cast<omp::AllocateFreeOp>(opInst);
10494 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
10496 llvm::Value *allocator;
10497 if (
auto allocatorVar = freeOp.getAllocator()) {
10498 allocator = moduleTranslation.
lookupValue(allocatorVar);
10499 if (allocator->getType()->isIntegerTy())
10500 allocator = builder.CreateIntToPtr(allocator, builder.getPtrTy());
10501 else if (allocator->getType()->isPointerTy())
10502 allocator = builder.CreatePointerBitCastOrAddrSpaceCast(
10503 allocator, builder.getPtrTy());
10505 allocator = llvm::ConstantPointerNull::get(builder.getPtrTy());
10510 for (
Value var : llvm::reverse(vars)) {
10511 llvm::Value *allocPtr = ompIface.lookupAllocatedPtr(var);
10513 return opInst.
emitError(
"omp.allocate_free: no allocation recorded");
10514 ompBuilder->createOMPFree(ompLoc, allocPtr, allocator,
"");
10521 llvm::Module *llvmModule) {
10522 llvm::Type *ptrTy = builder.getPtrTy(0);
10523 llvm::Type *i32Ty = builder.getInt32Ty();
10524 llvm::Type *voidTy = builder.getVoidTy();
10525 llvm::FunctionType *fnType =
10526 llvm::FunctionType::get(voidTy, {ptrTy, i32Ty},
false);
10527 llvm::Function *
func = dyn_cast<llvm::Function>(
10528 llvmModule->getOrInsertFunction(
"omp_target_free", fnType).getCallee());
10532static LogicalResult
10535 auto freeMemOp = cast<omp::TargetFreeMemOp>(opInst);
10540 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
10541 llvm::Function *ompTragetFreeFunc =
getOmpTargetFree(builder, llvmModule);
10544 llvm::Value *llvmDeviceNum = moduleTranslation.
lookupValue(deviceNum);
10547 llvm::Value *llvmHeapref = moduleTranslation.
lookupValue(heapref);
10549 llvm::Value *intToPtr =
10550 builder.CreateIntToPtr(llvmHeapref, builder.getPtrTy(0));
10551 builder.CreateCall(ompTragetFreeFunc, {intToPtr, llvmDeviceNum});
10555static LogicalResult
10557 llvm::IRBuilderBase &builder,
10561 ompBuilder->createOMPFreeShared(
10562 builder, moduleTranslation.
lookupValue(freeMemOp.getHeapref()), size);
10567static LogicalResult
10571 auto groupprivateOp = cast<omp::GroupprivateOp>(opInst);
10576 bool isTargetDevice = ompBuilder->Config.isTargetDevice();
10580 bool shouldAllocate =
true;
10581 switch (groupprivateOp.getDeviceType().value_or(
10582 mlir::omp::DeclareTargetDeviceType::any)) {
10583 case mlir::omp::DeclareTargetDeviceType::host:
10584 shouldAllocate = !isTargetDevice;
10586 case mlir::omp::DeclareTargetDeviceType::nohost:
10587 shouldAllocate = isTargetDevice;
10589 case mlir::omp::DeclareTargetDeviceType::any:
10590 shouldAllocate =
true;
10596 &opInst, groupprivateOp.getSymNameAttr());
10599 <<
"expected symbol '" << groupprivateOp.getSymName()
10600 <<
"' to reference an LLVM global variable";
10602 llvm::GlobalValue *globalValue = moduleTranslation.
lookupGlobal(global);
10603 llvm::Type *varType = moduleTranslation.
convertType(global.getType());
10604 std::string varName = globalValue->getName().str();
10606 llvm::Value *resultPtr;
10607 if (shouldAllocate && isTargetDevice) {
10608 llvm::Module *llvmModule = moduleTranslation.
getLLVMModule();
10609 llvm::Triple targetTriple(llvmModule->getTargetTriple());
10610 unsigned sharedAddressSpace;
10611 if (targetTriple.isAMDGCN())
10612 sharedAddressSpace = llvm::AMDGPUAS::LOCAL_ADDRESS;
10613 else if (targetTriple.isNVPTX())
10614 sharedAddressSpace = llvm::NVPTXAS::ADDRESS_SPACE_SHARED;
10616 return opInst.
emitError() <<
"groupprivate is not supported for target: "
10617 << targetTriple.str();
10618 llvm::GlobalVariable *sharedVar =
new llvm::GlobalVariable(
10619 *llvmModule, varType,
false,
10620 llvm::GlobalValue::InternalLinkage, llvm::PoisonValue::get(varType),
10621 varName,
nullptr, llvm::GlobalValue::NotThreadLocal,
10622 sharedAddressSpace,
10624 resultPtr = sharedVar;
10626 if (shouldAllocate && !isTargetDevice)
10627 opInst.
emitWarning(
"groupprivate directive is currently ignored on the "
10628 "host, using original global");
10629 resultPtr = globalValue;
10638LogicalResult OpenMPDialectLLVMIRTranslationInterface::convertOperation(
10639 Operation *op, llvm::IRBuilderBase &builder,
10640 LLVM::ModuleTranslation &moduleTranslation)
const {
10643 if (ompBuilder->Config.isTargetDevice() &&
10644 !isa<omp::TargetOp, omp::MapInfoOp, omp::TerminatorOp, omp::YieldOp>(
10647 return op->
emitOpError() <<
"unsupported host op found in device";
10655 bool isOutermostLoopWrapper =
10656 isa_and_present<omp::LoopWrapperInterface>(op) &&
10657 !dyn_cast_if_present<omp::LoopWrapperInterface>(op->
getParentOp());
10666 if (isa<omp::TaskloopContextOp>(op))
10667 isOutermostLoopWrapper =
true;
10668 else if (isa<omp::TaskloopWrapperOp>(op))
10669 isOutermostLoopWrapper =
false;
10671 if (isOutermostLoopWrapper)
10672 moduleTranslation.
stackPush<OpenMPLoopInfoStackFrame>();
10675 llvm::TypeSwitch<Operation *, LogicalResult>(op)
10676 .Case([&](omp::BarrierOp op) -> LogicalResult {
10680 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
10681 ompBuilder->createBarrier(builder, llvm::omp::OMPD_barrier);
10683 if (res.succeeded()) {
10686 builder.restoreIP(*afterIP);
10690 .Case([&](omp::TaskyieldOp op) {
10694 ompBuilder->createTaskyield(builder);
10697 .Case([&](omp::FlushOp op) {
10709 ompBuilder->createFlush(builder);
10712 .Case([&](omp::ErrorOp op) {
10716 llvm::Value *message =
nullptr;
10717 if (mlir::Value messageExpr = op.getMessageExpr())
10718 message = moduleTranslation.
lookupValue(messageExpr);
10719 else if (std::optional<StringRef> msg = op.getMessage();
10720 msg && !msg->empty())
10721 message = builder.CreateGlobalString(*msg);
10722 ompBuilder->createError(
10723 llvm::OpenMPIRBuilder::LocationDescription(builder),
10724 op.getSeverity() == omp::ClauseSeverity::fatal, message);
10727 .Case([&](omp::ParallelOp op) {
10730 .Case([&](omp::DispatchOp) {
10733 .Case([&](omp::MaskedOp) {
10736 .Case([&](omp::MasterOp) {
10739 .Case([&](omp::CriticalOp) {
10742 .Case([&](omp::OrderedRegionOp) {
10745 .Case([&](omp::OrderedOp) {
10748 .Case([&](omp::WsloopOp) {
10751 .Case([&](omp::SimdOp) {
10754 .Case([&](omp::AtomicReadOp) {
10757 .Case([&](omp::AtomicWriteOp) {
10760 .Case([&](omp::AtomicUpdateOp op) {
10763 .Case([&](omp::AtomicCaptureOp op) {
10766 .Case([&](omp::AtomicCompareOp op) {
10769 .Case([&](omp::CancelOp op) {
10772 .Case([&](omp::CancellationPointOp op) {
10775 .Case([&](omp::SectionsOp) {
10778 .Case([&](omp::ScopeOp op) {
10781 .Case([&](omp::SingleOp op) {
10784 .Case([&](omp::TeamsOp op) {
10787 .Case([&](omp::TaskOp op) {
10790 .Case([&](omp::TaskloopWrapperOp op) {
10793 .Case([&](omp::TaskloopContextOp op) {
10796 .Case([&](omp::TaskgroupOp op) {
10799 .Case([&](omp::TaskwaitOp op) {
10802 .Case([&](omp::InteropInitOp op) {
10805 .Case([&](omp::InteropDestroyOp op) {
10808 .Case([&](omp::InteropUseOp op) {
10811 .Case<omp::YieldOp, omp::TerminatorOp, omp::DeclareMapperOp,
10812 omp::DeclareMapperInfoOp, omp::DeclareReductionOp,
10813 omp::CriticalDeclareOp>([](
auto op) {
10826 .Case([&](omp::ThreadprivateOp) {
10829 .Case<omp::TargetDataOp, omp::TargetEnterDataOp,
10830 omp::TargetExitDataOp, omp::TargetUpdateOp>([&](
auto op) {
10833 .Case([&](omp::TargetOp) {
10836 .Case([&](omp::DistributeOp) {
10839 .Case([&](omp::LoopNestOp) {
10842 .Case<omp::MapInfoOp, omp::MapBoundsOp, omp::PrivateClauseOp,
10843 omp::AffinityEntryOp, omp::IteratorOp>([&](
auto op) {
10849 .Case([&](omp::NewCliOp op) {
10854 .Case([&](omp::CanonicalLoopOp op) {
10857 .Case([&](omp::UnrollHeuristicOp op) {
10866 .Case([&](omp::UnrollFullOp op) {
10869 .Case([&](omp::UnrollPartialOp op) {
10872 .Case([&](omp::TileOp op) {
10873 return applyTile(op, builder, moduleTranslation);
10875 .Case([&](omp::FuseOp op) {
10876 return applyFuse(op, builder, moduleTranslation);
10878 .Case([&](omp::TargetAllocMemOp) {
10881 .Case([&](omp::TargetFreeMemOp) {
10884 .Case([&](omp::AllocateDirOp) {
10887 .Case([&](omp::AllocateFreeOp) {
10891 .Case([&](omp::AllocSharedMemOp op) {
10894 .Case([&](omp::FreeSharedMemOp op) {
10897 .Case([&](omp::GroupprivateOp) {
10900 .Default([&](Operation *inst) {
10902 <<
"not yet implemented: " << inst->
getName();
10905 if (isOutermostLoopWrapper)
10912 registry.
insert<omp::OpenMPDialect>();
10914 dialect->addInterfaces<OpenMPDialectLLVMIRTranslationInterface>();
if(failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) return failure()
static mlir::LogicalResult buildDependData(OperandRange dependVars, std::optional< ArrayAttr > dependKinds, OperandRange dependIterated, std::optional< ArrayAttr > dependIteratedKinds, llvm::IRBuilderBase &builder, mlir::LLVM::ModuleTranslation &moduleTranslation, llvm::OpenMPIRBuilder::DependenciesInfo &taskDeps)
static LogicalResult convertOmpAtomicUpdate(omp::AtomicUpdateOp &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP atomic update operation using OpenMPIRBuilder.
static llvm::omp::OrderKind convertOrderKind(std::optional< omp::ClauseOrderKind > o)
Convert Order attribute to llvm::omp::OrderKind.
static void mapParentWithMembers(LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder &ompBuilder, DataLayout &dl, MapInfosTy &combinedInfo, MapInfoData &mapData, uint64_t mapDataIndex, llvm::omp::OpenMPOffloadMappingFlags memberOfFlag, TargetDirectiveEnumTy targetDirective)
static void processIndividualMap(llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder &ompBuilder, MapInfoData &mapData, size_t mapDataIdx, MapInfosTy &combinedInfo, TargetDirectiveEnumTy targetDirective, llvm::omp::OpenMPOffloadMappingFlags memberOfFlag=llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE, bool isTargetParam=true, int mapDataParentIdx=-1)
This function handles the insertion of a single item of map data from MapInfoData into the OMPIRBuild...
static llvm::OpenMPIRBuilder::InsertPointTy findAllocInsertPoints(llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, llvm::SmallVectorImpl< llvm::BasicBlock * > *deallocBlocks=nullptr)
Find the insertion point for allocas given the current insertion point for normal operations in the b...
static void sortMapIndices(llvm::SmallVectorImpl< size_t > &indices, omp::MapInfoOp mapInfo, bool first=true)
static LogicalResult convertOmpAtomicCapture(omp::AtomicCaptureOp atomicCaptureOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
owningDataPtrPtrReductionGens[i]
static LogicalResult convertOmpTaskloopContextOp(omp::TaskloopContextOp contextOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static Operation * getGlobalOpFromValue(Value value)
static llvm::OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind convertToCaptureClauseKind(mlir::omp::DeclareTargetCaptureClause captureClause)
static mlir::LogicalResult convertIteratorRegion(llvm::Value *linearIV, IteratorInfo &iterInfo, mlir::Block &iteratorRegionBlock, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static omp::MapInfoOp getFirstOrLastMappedMemberPtr(omp::MapInfoOp mapInfo, bool first)
static LogicalResult convertOmpDispatch(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Convert 'dispatch' operation into LLVM IR.
static OpTy castOrGetParentOfType(Operation *op, bool immediateParent=false)
If op is of the given type parameter, return it casted to that type. Otherwise, if its immediate pare...
static LogicalResult convertOmpOrderedRegion(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP 'ordered_region' operation into LLVM IR using OpenMPIRBuilder.
static LogicalResult convertFreeSharedMemOp(omp::FreeSharedMemOp freeMemOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static LogicalResult convertTargetFreeMemOp(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static LogicalResult convertOmpAtomicWrite(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an omp.atomic.write operation to LLVM IR.
static OwningAtomicReductionGen makeAtomicReductionGen(omp::DeclareReductionOp decl, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Create an OpenMPIRBuilder-compatible atomic reduction generator for the given reduction declaration.
static OwningDataPtrPtrReductionGen makeRefDataPtrGen(omp::DeclareReductionOp decl, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, bool isByRef)
Create an OpenMPIRBuilder-compatible data_ptr_ptr reduction generator for the given reduction declara...
static llvm::Value * getRefPtrIfDeclareTarget(Value value, LLVM::ModuleTranslation &moduleTranslation)
static llvm::Function * emitTaskReductionCombFn(omp::DeclareReductionOp decl, StringRef baseName, LLVM::ModuleTranslation &moduleTranslation)
Build an outlined combiner helper for a task_reduction declare_reduction op. Signature: void(ptr lhs,...
static LogicalResult convertOmpWsloop(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP workshare loop into LLVM IR using OpenMPIRBuilder.
static LogicalResult applyUnrollHeuristic(omp::UnrollHeuristicOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Apply a #pragma omp unroll / "!$omp unroll" transformation using the OpenMPIRBuilder.
static LogicalResult convertOmpMaster(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP 'master' operation into LLVM IR using OpenMPIRBuilder.
static void getAsIntegers(ArrayAttr values, llvm::SmallVector< int64_t > &ints)
static void emitComplexAtomicCmpXchg(llvm::IRBuilderBase &builder, llvm::Value *llvmX, llvm::Type *complexTy, llvm::Value *eVal, llvm::Value *dVal, llvm::AtomicOrdering atomicOrdering, llvm::AtomicOrdering failOrdering, bool isWeak, llvm::Value *&oldComplex, llvm::Value *&cmpOk)
Emit an IEEE-754-correct cmpxchg for a complex (struct-typed) atomic compare with fcmp oeq....
static llvm::Value * findAssociatedValue(Value privateVar, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, llvm::DenseMap< Value, Value > *mappedPrivateVars=nullptr)
Return the llvm::Value * corresponding to the privateVar that is being privatized....
static ArrayRef< bool > getIsByRef(std::optional< ArrayRef< bool > > attr)
static llvm::Expected< llvm::Value * > lookupOrTranslatePureValue(Value value, LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder)
Look up the given value in the mapping, and if it's not there, translate its defining operation at th...
static LogicalResult allocReductionVars(T op, ArrayRef< BlockArgument > reductionArgs, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, const llvm::OpenMPIRBuilder::InsertPointTy &allocaIP, SmallVectorImpl< omp::DeclareReductionOp > &reductionDecls, SmallVectorImpl< llvm::Value * > &privateReductionVariables, DenseMap< Value, llvm::Value * > &reductionVariableMap, SmallVectorImpl< DeferredStore > &deferredStores, llvm::ArrayRef< bool > isByRefs)
Allocate space for privatized reduction variables.
static void emitTaskReductionModifierFini(bool isWorksharing, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Emits __kmpc_task_reduction_modifier_fini(loc, gtid, is_ws) at the current builder insertion point,...
static LogicalResult convertOmpInteropUseOp(omp::InteropUseOp useOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static LogicalResult convertOmpTaskwaitOp(omp::TaskwaitOp twOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static LogicalResult collectAndValidateTaskloopRedDecls(Operation *contextOp, std::optional< ArrayAttr > syms, StringRef opName, StringRef clauseName, SmallVectorImpl< omp::DeclareReductionOp > &out)
Look up and validate the declare_reduction ops referenced by a reduction-like clause on the omp....
static LogicalResult convertOmpLoopNest(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP loop nest into LLVM IR using OpenMPIRBuilder.
static mlir::LogicalResult fillIteratorLoop(mlir::omp::IteratorOp itersOp, llvm::IRBuilderBase &builder, mlir::LLVM::ModuleTranslation &moduleTranslation, IteratorInfo &iterInfo, llvm::StringRef loopName, IteratorStoreEntryTy genStoreEntry)
static void createAlteredByCaptureMap(MapInfoData &mapData, LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder)
static LogicalResult convertOmpTaskOp(omp::TaskOp taskOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP task construct into LLVM IR using OpenMPIRBuilder.
static void genMapInfos(llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, DataLayout &dl, MapInfosTy &combinedInfo, MapInfoData &mapData, TargetDirectiveEnumTy targetDirective)
static llvm::AtomicOrdering convertAtomicOrdering(std::optional< omp::ClauseMemoryOrderKind > ao)
Convert an Atomic Ordering attribute to llvm::AtomicOrdering.
static LogicalResult convertOmpInteropInitOp(omp::InteropInitOp initOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static void setInsertPointForPossiblyEmptyBlock(llvm::IRBuilderBase &builder, llvm::BasicBlock *block=nullptr)
llvm::function_ref< void(llvm::Value *linearIV, mlir::omp::YieldOp yield)> IteratorStoreEntryTy
static llvm::Function * emitTaskReductionInitFn(omp::DeclareReductionOp decl, StringRef baseName, LLVM::ModuleTranslation &moduleTranslation)
Build an outlined init helper for a task_reduction declare_reduction op. Signature: void(ptr priv,...
static LogicalResult convertOmpSections(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static llvm::Expected< llvm::BasicBlock * > allocatePrivateVars(T op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, PrivateVarsInfo &privateVarsInfo, llvm::OpenMPIRBuilder::InsertPointTy &allocaIP, llvm::DenseMap< Value, Value > *mappedPrivateVars=nullptr, std::optional< llvm::OpenMPIRBuilder::InsertPointTy > allocatorIP=std::nullopt)
Allocate and initialize delayed private variables. Returns the basic block which comes after all of t...
static LogicalResult applyUnrollPartial(omp::UnrollPartialOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Apply a #pragma omp unroll partial / !$omp unroll partial transformation using the OpenMPIRBuilder.
static LogicalResult convertOmpCritical(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP 'critical' operation into LLVM IR using OpenMPIRBuilder.
static LogicalResult convertTargetAllocMemOp(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static omp::DistributeOp getDistributeCapturingTeamsReduction(omp::TeamsOp teamsOp)
static LogicalResult convertOmpCanonicalLoopOp(omp::CanonicalLoopOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Convert an omp.canonical_loop to LLVM-IR.
static LogicalResult convertOmpTargetData(Operation *op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static std::optional< int64_t > extractConstInteger(Value value)
If the given value is defined by an llvm.mlir.constant operation and it is of an integer type,...
static llvm::Expected< llvm::Value * > initPrivateVar(llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, omp::PrivateClauseOp &privDecl, llvm::Value *nonPrivateVar, BlockArgument &blockArg, llvm::Value *llvmPrivateVar, llvm::BasicBlock *privInitBlock, llvm::DenseMap< Value, Value > *mappedPrivateVars=nullptr)
Initialize a single (first)private variable. You probably want to use allocateAndInitPrivateVars inst...
static mlir::LogicalResult buildAffinityData(mlir::omp::TaskOp &taskOp, llvm::IRBuilderBase &builder, mlir::LLVM::ModuleTranslation &moduleTranslation, llvm::OpenMPIRBuilder::AffinityData &ad)
static LogicalResult allocAndInitializeReductionVars(OP op, ArrayRef< BlockArgument > reductionArgs, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, llvm::OpenMPIRBuilder::InsertPointTy &allocaIP, SmallVectorImpl< omp::DeclareReductionOp > &reductionDecls, SmallVectorImpl< llvm::Value * > &privateReductionVariables, DenseMap< Value, llvm::Value * > &reductionVariableMap, llvm::ArrayRef< bool > isByRef)
static LogicalResult convertOmpSimd(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP simd loop into LLVM IR using OpenMPIRBuilder.
static LogicalResult convertOmpInteropDestroyOp(omp::InteropDestroyOp destroyOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static LogicalResult convertOmpDistribute(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static llvm::Value * getAllocationSize(llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, T op)
static llvm::Function * getOmpTargetAlloc(llvm::IRBuilderBase &builder, llvm::Module *llvmModule)
static llvm::omp::OMPDynGroupprivateFallbackType getDynGroupprivateFallbackType(omp::FallbackModifierAttr fallbackAttr)
static llvm::Expected< llvm::Function * > emitUserDefinedMapper(Operation *declMapperOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, llvm::StringRef mapperFuncName, TargetDirectiveEnumTy targetDirective)
static LogicalResult convertOmpOrdered(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP 'ordered' operation into LLVM IR using OpenMPIRBuilder.
static LogicalResult cleanupPrivateVars(T op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, Location loc, PrivateVarsInfo &privateVarsInfo)
static void processMapWithMembersOf(LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder &ompBuilder, DataLayout &dl, MapInfosTy &combinedInfo, MapInfoData &mapData, uint64_t mapDataIndex, TargetDirectiveEnumTy targetDirective)
static LogicalResult convertOmpMasked(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP 'masked' operation into LLVM IR using OpenMPIRBuilder.
static llvm::AtomicRMWInst::BinOp convertBinOpToAtomic(Operation &op)
Converts an LLVM dialect binary operation to the corresponding enum value for atomicrmw supported bin...
static LogicalResult convertOmpCancel(omp::CancelOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static int getMapDataMemberIdx(MapInfoData &mapData, omp::MapInfoOp memberOp)
allocatedType moduleTranslation static convertType(allocatedType) LogicalResult inlineOmpRegionCleanup(llvm::SmallVectorImpl< Region * > &cleanupRegions, llvm::ArrayRef< llvm::Value * > privateVariables, LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder, StringRef regionName, bool shouldLoadCleanupRegionArg=true)
handling of DeclareReductionOp's cleanup region
static LogicalResult applyFuse(omp::FuseOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Apply a #pragma omp fuse / !$omp fuse transformation using the OpenMPIRBuilder.
static llvm::Value * materializeRegionArgValue(llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, BlockArgument regionArg, llvm::Value *value)
static LogicalResult convertOmpScope(omp::ScopeOp &scopeOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP scope construct into LLVM IR.
static bool isPrivatizeableAttachMap(omp::ClauseMapFlags mapType)
static llvm::Value * getSizeInBytes(DataLayout &dl, const mlir::Type &type, Operation *clauseOp, llvm::Value *basePointer, llvm::Type *baseType, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static llvm::Error initPrivateVars(llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, PrivateVarsInfo &privateVarsInfo, llvm::DenseMap< Value, Value > *mappedPrivateVars=nullptr)
static LogicalResult convertAllocSharedMemOp(omp::AllocSharedMemOp allocMemOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static llvm::CanonicalLoopInfo * findCurrentLoopInfo(LLVM::ModuleTranslation &moduleTranslation)
Find the loop information structure for the loop nest being translated.
static OwningReductionGen makeReductionGen(omp::DeclareReductionOp decl, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Create an OpenMPIRBuilder-compatible reduction generator for the given reduction declaration.
static std::vector< llvm::Value * > calculateBoundsOffset(LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder, bool isArrayTy, OperandRange bounds)
This function calculates the array/pointer offset for map data provided with bounds operations,...
static void storeAffinityEntry(llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder &ompBuilder, llvm::Value *affinityList, llvm::Value *index, llvm::Value *addr, llvm::Value *len)
static LogicalResult convertOmpParallel(omp::ParallelOp opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts the OpenMP parallel operation to LLVM IR.
static void pushCancelFinalizationCB(SmallVectorImpl< llvm::UncondBrInst * > &cancelTerminators, llvm::IRBuilderBase &llvmBuilder, llvm::OpenMPIRBuilder &ompBuilder, mlir::Operation *op, llvm::omp::Directive cancelDirective)
Shared implementation of a callback which adds a termiator for the new block created for the branch t...
static LogicalResult inlineConvertOmpRegions(Region ®ion, StringRef blockName, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, SmallVectorImpl< llvm::Value * > *continuationBlockArgs=nullptr)
Translates the blocks contained in the given region and appends them to at the current insertion poin...
static void getTargetEntryUniqueInfo(llvm::TargetRegionEntryInfo &targetInfo, omp::TargetOp targetOp, llvm::OpenMPIRBuilder &ompBuilder, llvm::vfs::FileSystem &vfs, llvm::StringRef parentName="")
static LogicalResult convertOmpThreadprivate(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP Threadprivate operation into LLVM IR using OpenMPIRBuilder.
static omp::PrivateClauseOp findPrivatizer(Operation *from, SymbolRefAttr symbolName)
Looks up from the operation from and returns the PrivateClauseOp with name symbolName.
static LogicalResult applyUnrollFull(omp::UnrollFullOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Apply a #pragma omp unroll full / !$omp unroll full transformation using the OpenMPIRBuilder.
static LogicalResult convertOmpGroupprivate(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP groupprivate operation into LLVM IR.
static llvm::Expected< llvm::Function * > getOrCreateUserDefinedMapperFunc(Operation *op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, TargetDirectiveEnumTy targetDirective)
static llvm::Value * getSourceLocIdentFromOp(llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder &ompBuilder, Operation *op)
static uint64_t getTypeByteSize(mlir::Type type, const DataLayout &dl)
static llvm::SmallString< 64 > getDeclareTargetRefPtrSuffix(LLVM::GlobalOp globalOp, llvm::OpenMPIRBuilder &ompBuilder, llvm::vfs::FileSystem &vfs)
static void extractHostEvalClauses(omp::TargetOp targetOp, Value &numThreads, Value &numTeamsLower, Value &numTeamsUpper, Value &threadLimit, llvm::SmallVectorImpl< Value > *lowerBounds=nullptr, llvm::SmallVectorImpl< Value > *upperBounds=nullptr, llvm::SmallVectorImpl< Value > *steps=nullptr)
Follow uses of host_eval-defined block arguments of the given omp.target operation and populate outpu...
static llvm::Expected< llvm::BasicBlock * > convertOmpOpRegions(Region ®ion, StringRef blockName, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, SmallVectorImpl< llvm::PHINode * > *continuationBlockPHIs=nullptr)
Converts the given region that appears within an OpenMP dialect operation to LLVM IR,...
static LogicalResult extractAtomicComparePattern(Block &block, llvm::function_ref< llvm::Value *(mlir::Value)> materializeValue, omp::AtomicCompareOp atomicCompareOp, AtomicComparePatternInfo &info)
Extract comparison predicate, expected value (e), desired value (d), and related flags from an atomic...
static LogicalResult convertOmpAtomicCompare(omp::AtomicCompareOp atomicCompareOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an omp.atomic.compare operation to LLVM IR.
static LogicalResult copyFirstPrivateVars(mlir::Operation *op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, SmallVectorImpl< llvm::Value * > &moldVars, ArrayRef< llvm::Value * > llvmPrivateVars, SmallVectorImpl< omp::PrivateClauseOp > &privateDecls, bool insertBarrier, llvm::DenseMap< Value, Value > *mappedPrivateVars=nullptr)
static LogicalResult convertAllocateDirOp(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, const OpenMPDialectLLVMIRTranslationInterface &ompIface)
static bool constructIsCancellable(Operation *op)
Returns true if the construct contains omp.cancel or omp.cancellation_point.
static llvm::omp::OpenMPOffloadMappingFlags convertClauseMapFlags(omp::ClauseMapFlags mlirFlags)
static LogicalResult convertAllocatorVars(Operation &op, ValueRange allocatorVars, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, PrivateVarsInfo &privateVarsInfo)
static void buildDependDataLocator(std::optional< ArrayAttr > dependKinds, OperandRange dependVars, LLVM::ModuleTranslation &moduleTranslation, SmallVectorImpl< llvm::OpenMPIRBuilder::DependData > &dds)
static std::optional< llvm::omp::OMPAtomicCompareOp > convertFCmpPredicateToAtomicCompareOp(LLVM::FCmpPredicate predicate)
Helper to extract the OMPAtomicCompareOp from a floating-point comparison predicate....
static llvm::Value * emitTaskReductionInitCall(ArrayRef< omp::DeclareReductionOp > redDecls, ArrayRef< llvm::Value * > origPtrs, StringRef helperNamePrefix, llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder::InsertPointTy allocaIP, LLVM::ModuleTranslation &moduleTranslation, bool isModifier=false, bool isWorksharing=false)
Emit the per-taskgroup task_reduction descriptor array and the __kmpc_taskred_init runtime call....
static void mapInitializationArgs(T loop, LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder, SmallVectorImpl< omp::DeclareReductionOp > &reductionDecls, DenseMap< Value, llvm::Value * > &reductionVariableMap, unsigned i)
Map input arguments to reduction initialization region.
static llvm::omp::ProcBindKind getProcBindKind(omp::ClauseProcBindKind kind)
Convert ProcBindKind from MLIR-generated enum to LLVM enum.
static void fillAffinityLocators(Operation::operand_range affinityVars, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, llvm::Value *affinityList)
static LogicalResult convertOmpTaskloopWrapperOp(omp::TaskloopWrapperOp loopWrapperOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
The correct entry point is convertOmpTaskloopContextOp. This gets called whilst lowering the body of ...
static void getOverlappedMembers(llvm::SmallVectorImpl< size_t > &overlapMapDataIdxs, omp::MapInfoOp parentOp)
static LogicalResult convertOmpSingle(omp::SingleOp &singleOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP single construct into LLVM IR using OpenMPIRBuilder.
static bool isDeclareTargetTo(Value value)
static uint64_t getArrayElementSizeInBits(LLVM::LLVMArrayType arrTy, DataLayout &dl)
static void popCancelFinalizationCB(const ArrayRef< llvm::UncondBrInst * > cancelTerminators, llvm::OpenMPIRBuilder &ompBuilder, llvm::BasicBlock *afterBB)
If we cancelled the construct, we should branch to the finalization block of that construct....
static void collectReductionDecls(T op, SmallVectorImpl< omp::DeclareReductionOp > &reductions)
Populates reductions with reduction declarations used in the given op.
static LogicalResult handleError(llvm::Error error, Operation &op)
static LogicalResult convertOmpTarget(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseKind convertToDeviceClauseKind(mlir::omp::DeclareTargetDeviceType deviceClause)
static std::optional< llvm::omp::OMPAtomicCompareOp > convertICmpPredicateToAtomicCompareOp(LLVM::ICmpPredicate predicate)
Helper to extract the OMPAtomicCompareOp from an integer comparison predicate. Returns std::nullopt f...
static llvm::Error computeTaskloopBounds(omp::LoopNestOp loopOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, llvm::Value *&lbVal, llvm::Value *&ubVal, llvm::Value *&stepVal)
static LogicalResult checkImplementationStatus(Operation &op)
Check whether translation to LLVM IR for the given operation is currently supported.
static llvm::IRBuilderBase::InsertPoint createDeviceArgumentAccessor(omp::TargetOp targetOp, MapInfoData &mapData, llvm::Argument &arg, llvm::Value *input, llvm::Value *&retVal, llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder &ompBuilder, LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase::InsertPoint allocaIP, llvm::IRBuilderBase::InsertPoint codeGenIP, llvm::ArrayRef< llvm::IRBuilderBase::InsertPoint > deallocIPs)
static LogicalResult createReductionsAndCleanup(OP op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, llvm::OpenMPIRBuilder::InsertPointTy &allocaIP, SmallVectorImpl< omp::DeclareReductionOp > &reductionDecls, ArrayRef< llvm::Value * > privateReductionVariables, ArrayRef< bool > isByRef, bool isNowait=false, bool isTeamsReduction=false)
static LogicalResult convertOmpCancellationPoint(omp::CancellationPointOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static bool opIsInSingleThread(mlir::Operation *op)
This can't always be determined statically, but when we can, it is good to avoid generating compiler-...
static uint64_t getReductionDataSize(OpTy &op)
static LogicalResult convertOmpAtomicRead(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Convert omp.atomic.read operation to LLVM IR.
static llvm::omp::Directive convertCancellationConstructType(omp::ClauseCancellationConstructType directive)
static void initTargetDefaultAttrs(omp::TargetOp targetOp, Operation *capturedOp, llvm::OpenMPIRBuilder::TargetKernelDefaultAttrs &attrs, bool isTargetDevice, bool isGPU)
Populate default MinTeams, MaxTeams and MaxThreads to their default values as stated by the correspon...
static llvm::AtomicOrdering getAtomicCompareFailureOrdering(omp::AtomicCompareOp atomicCompareOp, llvm::AtomicOrdering atomicOrdering)
Compute the cmpxchg failure ordering for an atomic compare op: use the fail clause ordering when pres...
static llvm::omp::RTLDependenceKindTy convertDependKind(mlir::omp::ClauseTaskDepend kind)
static void initTargetRuntimeAttrs(llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, omp::TargetOp targetOp, Operation *capturedOp, llvm::OpenMPIRBuilder::TargetKernelRuntimeAttrs &attrs)
Gather LLVM runtime values for all clauses evaluated in the host that are passed to the kernel invoca...
static ComplexComparePattern detectComplexCompareEq(Block &block)
Detect a decomposed complex equality comparison in an atomic compare region: re_x = llvm....
static LogicalResult convertOmpTeams(omp::TeamsOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static Value getBaseValueForTypeLookup(Value value)
static bool isHostDeviceOp(Operation *op)
static LogicalResult convertDeclareTargetAttr(Operation *op, mlir::omp::DeclareTargetAttr attribute, llvm::OpenMPIRBuilder *ompBuilder, LLVM::ModuleTranslation &moduleTranslation)
static bool isDeclareTargetLink(Value value)
static LogicalResult convertFlagsAttr(Operation *op, mlir::omp::FlagsAttr attribute, LLVM::ModuleTranslation &moduleTranslation)
Lowers the FlagsAttr which is applied to the module when offloading. This attribute contains OpenMP R...
static bool checkIfPointerMap(omp::MapInfoOp mapOp)
static llvm::Type * getAllocatedLlvmTypeForVariable(Value var, Value baseVar, LLVM::ModuleTranslation &moduleTranslation)
static LogicalResult applyTile(omp::TileOp op, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Apply a #pragma omp tile / !$omp tile transformation using the OpenMPIRBuilder.
static LogicalResult convertAllocateFreeOp(Operation &opInst, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation, const OpenMPDialectLLVMIRTranslationInterface &ompIface)
static llvm::Function * getOmpTargetFree(llvm::IRBuilderBase &builder, llvm::Module *llvmModule)
static LogicalResult convertOmpTaskgroupOp(omp::TaskgroupOp tgOp, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
Converts an OpenMP taskgroup construct into LLVM IR using OpenMPIRBuilder.
static std::optional< llvm::Value * > getDynamicAllocatedSize(Value var, Value baseVar, LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder, const llvm::DataLayout &dataLayout)
static void collectMapDataFromMapOperands(MapInfoData &mapData, SmallVectorImpl< Value > &mapVars, LLVM::ModuleTranslation &moduleTranslation, DataLayout &dl, llvm::IRBuilderBase &builder, ArrayRef< Value > useDevPtrOperands={}, ArrayRef< Value > useDevAddrOperands={}, ArrayRef< Value > hasDevAddrOperands={})
static void extractAtomicControlFlags(omp::AtomicUpdateOp atomicUpdateOp, bool &isIgnoreDenormalMode, bool &isFineGrainedMemory, bool &isRemoteMemory)
static Operation * genLoop(CodegenEnv &env, OpBuilder &builder, LoopId curr, unsigned numCases, bool needsUniv, ArrayRef< TensorLevel > tidLvls)
Generates a for-loop or a while-loop, depending on whether it implements singleton iteration or co-it...
#define MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(CLASS_NAME)
Attributes are known-constant values of operations.
This class represents an argument of a Block.
Block represents an ordered list of Operations.
BlockArgument getArgument(unsigned i)
unsigned getNumArguments()
OpListType & getOperations()
Operation * getTerminator()
Get the terminator operation of this block.
iterator_range< iterator > without_terminator()
Return an iterator range over the operation within this block excluding the terminator operation at t...
The main mechanism for performing data layout queries.
llvm::TypeSize getTypeSize(Type t) const
Returns the size of the given type in the current scope.
llvm::TypeSize getTypeSizeInBits(Type t) const
Returns the size in bits of the given type in the current scope.
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
An instance of this location represents a tuple of file, line number, and column number.
Implementation class for module translation.
llvm::BasicBlock * lookupBlock(Block *block) const
Finds an LLVM IR basic block that corresponds to the given MLIR block.
WalkResult stackWalk(llvm::function_ref< WalkResult(T &)> callback)
Calls callback for every ModuleTranslation stack frame of type T starting from the top of the stack.
void stackPush(Args &&...args)
Creates a stack frame of type T on ModuleTranslation stack.
LogicalResult convertBlock(Block &bb, bool ignoreArguments, llvm::IRBuilderBase &builder)
Translates the contents of the given block to LLVM IR using this translator.
SmallVector< llvm::Value * > lookupValues(ValueRange values)
Looks up remapped a list of remapped values.
void mapFunction(StringRef name, llvm::Function *func)
Stores the mapping between a function name and its LLVM IR representation.
llvm::Value * lookupValue(Value value) const
Finds an LLVM IR value corresponding to the given MLIR value.
void invalidateOmpLoop(omp::NewCliOp mlir)
Mark an OpenMP loop as having been consumed.
SymbolTableCollection & symbolTable()
llvm::Type * convertType(Type type)
Converts the type from MLIR LLVM dialect to LLVM.
llvm::OpenMPIRBuilder * getOpenMPBuilder()
Returns the OpenMP IR builder associated with the LLVM IR module being constructed.
llvm::vfs::FileSystem & getFileSystem()
Returns the virtual filesystem to use for file operations.
void mapOmpLoop(omp::NewCliOp mlir, llvm::CanonicalLoopInfo *llvm)
Map an MLIR OpenMP dialect CanonicalLoopInfo to its lowered LLVM-IR OpenMPIRBuilder CanonicalLoopInfo...
llvm::GlobalValue * lookupGlobal(Operation *op)
Finds an LLVM IR global value that corresponds to the given MLIR operation defining a global value.
SaveStateStack< T, ModuleTranslation > SaveStack
RAII object calling stackPush/stackPop on construction/destruction.
void remapAllValuesWith(llvm::Value *oldValue, llvm::Value *newValue)
Remap old value with new value in the MLIR-to-LLVM value map so later translations use the replacemen...
LogicalResult convertOperation(Operation &op, llvm::IRBuilderBase &builder)
Converts the given MLIR operation into LLVM IR using this translator.
llvm::Function * lookupFunction(StringRef name) const
Finds an LLVM IR function by its name.
void mapBlock(Block *mlir, llvm::BasicBlock *llvm)
Stores the mapping between an MLIR block and LLVM IR basic block.
llvm::Module * getLLVMModule()
Returns the LLVM module in which the IR is being constructed.
void stackPop()
Pops the last element from the ModuleTranslation stack.
void forgetMapping(Region ®ion)
Removes the mapping for blocks contained in the region and values defined in these blocks.
void mapValue(Value mlir, llvm::Value *llvm)
Stores the mapping between an MLIR value and its LLVM IR counterpart.
llvm::CanonicalLoopInfo * lookupOMPLoop(omp::NewCliOp mlir) const
Find the LLVM-IR loop that represents an MLIR loop.
llvm::LLVMContext & getLLVMContext() const
Returns the LLVM context in which the IR is being constructed.
Utility class to translate MLIR LLVM dialect types to LLVM IR.
unsigned getPreferredAlignment(Type type, const llvm::DataLayout &layout)
Returns the preferred alignment for the type given the data layout.
T findInstanceOf()
Return an instance of the given location type if one is nested under the current location.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
void appendDialectRegistry(const DialectRegistry ®istry)
Append the contents of the given dialect registry to the registry associated with this context.
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
This class implements the operand iterators for the Operation class.
StringAttr getIdentifier() const
Return the name of this operation as a StringAttr.
Operation is the basic unit of execution within MLIR.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Value getOperand(unsigned idx)
InFlightDiagnostic emitWarning(const Twine &message={})
Emit a warning about this operation, reporting up to any diagnostic handlers that may be listening.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
unsigned getNumRegions()
Returns the number of regions held by this operation.
Location getLoc()
The source location the operation was defined or derived from.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
unsigned getNumOperands()
OperandRange operand_range
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
OperationName getName()
The name of an operation is the key identifier for it.
operand_range getOperands()
Returns an iterator on the underlying Value's.
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
user_range getUsers()
Returns a range of all users.
result_range getResults()
MLIRContext * getContext()
Return the context this operation is associated with.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
unsigned getNumResults()
Return the number of results held by this operation.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
BlockArgListType getArguments()
unsigned getNumArguments()
Operation * getParentOp()
Return the parent operation this region is attached to.
BlockListType & getBlocks()
bool hasOneBlock()
Return true if this region has exactly one block.
Concrete CRTP base class for StateStack frames.
@ Private
The symbol is private and may only be referenced by SymbolRefAttrs local to the operations within the...
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
user_range getUsers() const
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
A utility result that is used to signal how to proceed with an ongoing walk:
static WalkResult advance()
bool wasInterrupted() const
Returns true if the walk was interrupted.
static WalkResult interrupt()
The OpAsmOpInterface, see OpAsmInterface.td for more details.
void connectPHINodes(Region ®ion, const ModuleTranslation &state)
For all blocks in the region that were converted to LLVM IR using the given ModuleTranslation,...
llvm::Constant * createMappingInformation(Location loc, llvm::OpenMPIRBuilder &builder)
Create a constant string representing the mapping information extracted from the MLIR location inform...
llvm::Constant * createSourceLocStrFromLocation(Location loc, llvm::OpenMPIRBuilder &builder, StringRef name, uint32_t &strLen, bool ForOffloadMap=false)
Create a constant string location from the MLIR Location information.
int64_t getOpenMPVersionAttribute(ModuleOp module, int64_t fallback=-1)
Returns the value of the omp.version attribute, if present, or the fallback.
bool opInSharedDeviceContext(Operation &op)
Check whether the given operation is located in a context where an allocation to be used by multiple ...
bool allocaUsesRequireSharedMem(Value alloc)
Check whether the value representing an allocation, assumed to have been defined in a shared device c...
auto getDims(VectorType vType)
Returns a range over the dims (size and scalability) of a VectorType.
Include the generated interface declarations.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
SetVector< Block * > getBlocksSortedByDominance(Region ®ion)
Gets a list of blocks that is sorted according to dominance.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
bool isPure(Operation *op)
Returns true if the given operation is pure, i.e., is speculatable that does not touch memory.
void registerOpenMPDialectTranslation(DialectRegistry ®istry)
Register the OpenMP dialect and the translation from it to the LLVM IR in the given registry;.
llvm::SetVector< T, Vector, Set, N > SetVector
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Holds the extracted comparison pattern information from an atomic compare region.
llvm::omp::OMPAtomicCompareOp compareOp
Result of matching the decomposed complex equality pattern inside an atomic compare region.
llvm::Value * allocatedPtr
A util to collect info needed to convert delayed privatizers from MLIR to LLVM.
SmallVector< mlir::Value > mlirVars
SmallVector< omp::PrivateClauseOp > privatizers
MutableArrayRef< BlockArgument > blockArgs
llvm::DenseMap< Value, llvm::Value * > convertedAllocators
SmallVector< llvm::Value * > llvmVars
SmallVector< AllocatorPrivateInfo > allocatorPrivates
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.