MLIR 24.0.0git
OpenMPToLLVMIRTranslation.cpp
Go to the documentation of this file.
1//===- OpenMPToLLVMIRTranslation.cpp - Translate OpenMP dialect to LLVM IR-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements a translation between the MLIR OpenMP dialect and LLVM
10// IR.
11//
12//===----------------------------------------------------------------------===//
21#include "mlir/IR/Operation.h"
23#include "mlir/Support/LLVM.h"
26
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"
45
46#include <cstdint>
47#include <iterator>
48#include <numeric>
49#include <optional>
50#include <utility>
51
52using namespace mlir;
53
54namespace {
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;
72 }
73 llvm_unreachable("unhandled schedule clause argument");
74}
75
76/// ModuleTranslation stack frame for OpenMP operations. This keeps track of the
77/// insertion points for allocas.
78class OpenMPAllocStackFrame
79 : public StateStackFrameBase<OpenMPAllocStackFrame> {
80public:
82
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;
89};
90
91/// Stack frame to hold a \see llvm::CanonicalLoopInfo representing the
92/// collapsed canonical loop information corresponding to an \c omp.loop_nest
93/// operation.
94class OpenMPLoopInfoStackFrame
95 : public StateStackFrameBase<OpenMPLoopInfoStackFrame> {
96public:
97 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(OpenMPLoopInfoStackFrame)
98 llvm::CanonicalLoopInfo *loopInfo = nullptr;
99};
100
101/// Custom error class to signal translation errors that don't need reporting,
102/// since encountering them will have already triggered relevant error messages.
103///
104/// Its purpose is to serve as the glue between MLIR failures represented as
105/// \see LogicalResult instances and \see llvm::Error instances used to
106/// propagate errors through the \see llvm::OpenMPIRBuilder. Generally, when an
107/// error of the first type is raised, a message is emitted directly (the \see
108/// LogicalResult itself does not hold any information). If we need to forward
109/// this error condition as an \see llvm::Error while avoiding triggering some
110/// redundant error reporting later on, we need a custom \see llvm::ErrorInfo
111/// class to just signal this situation has happened.
112///
113/// For example, this class should be used to trigger errors from within
114/// callbacks passed to the \see OpenMPIRBuilder when they were triggered by the
115/// translation of their own regions. This unclutters the error log from
116/// redundant messages.
117class PreviouslyReportedError
118 : public llvm::ErrorInfo<PreviouslyReportedError> {
119public:
120 void log(raw_ostream &) const override {
121 // Do not log anything.
122 }
123
124 std::error_code convertToErrorCode() const override {
125 llvm_unreachable(
126 "PreviouslyReportedError doesn't support ECError conversion");
127 }
128
129 // Used by ErrorInfo::classID.
130 static char ID;
131};
132
133char PreviouslyReportedError::ID = 0;
134
135/*
136 * Custom class for processing linear clause for omp.wsloop
137 * and omp.simd. Linear clause translation requires setup,
138 * initialization, update, and finalization at varying
139 * basic blocks in the IR. This class helps maintain
140 * internal state to allow consistent translation in
141 * each of these stages.
142 */
143
144class LinearClauseProcessor {
145
146private:
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;
155 Value linearLoopIV;
156
157public:
158 // Register type for the linear variables
159 void registerType(LLVM::ModuleTranslation &moduleTranslation,
160 mlir::Attribute &ty) {
161 linearVarTypes.push_back(moduleTranslation.convertType(
162 mlir::cast<mlir::TypeAttr>(ty).getValue()));
163 }
164
165 // Allocate space for linear variabes
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);
175 }
176
177 // Initialize linear step
178 inline void initLinearStep(LLVM::ModuleTranslation &moduleTranslation,
179 mlir::Value &linearStep) {
180 linearSteps.push_back(moduleTranslation.lookupValue(linearStep));
181 }
182
183 // Emit IR for initialization of linear variables
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]);
192 }
193 }
194
195 // Find linear iteration variable and save it for later updates
196 LogicalResult initLinearIV(omp::SimdOp simdOp) {
197 auto loopOp = cast<omp::LoopNestOp>(simdOp.getWrappedLoop());
198 // NOTE iteration variables can only be linear in non-nested loops.
199 if (loopOp.getIVs().size() != 1)
200 return success();
201 // Currently, frontends using `omp.simd` always generate a store from the
202 // `omp.loop_nest`'s IV to the corresponding iteration variable.
203 // We leverage this to find the linear iteration variable.
204 //
205 // TODO Add an attribute to `omp.loop_nest` that explicitly lists the
206 // variables that correspond to the loop induction variables.
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;
217 }
218 }
219 }
220 }
221 return success();
222 }
223
224 // Emit IR for updating Linear variables
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];
232
233 if (!iv->getType()->isIntegerTy())
234 llvm_unreachable("OpenMP loop induction variable must be an integer "
235 "type");
236
237 if (linearVarType->isIntegerTy()) {
238 // Integer path: normalize all arithmetic to linearVarType
239 iv = builder.CreateSExtOrTrunc(iv, linearVarType);
240 step = builder.CreateSExtOrTrunc(step, linearVarType);
241
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()) {
248 // Float path: perform multiply in integer, then convert to float
249 step = builder.CreateSExtOrTrunc(step, iv->getType());
250 llvm::Value *mulInst = builder.CreateMul(iv, step);
251
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]);
257 } else {
258 llvm_unreachable(
259 "Linear variable must be of integer or floating-point type");
260 }
261 }
262 }
263
264 // Emit IR for updating linear iteration variables on loop exit
265 void updateLinearIV(llvm::IRBuilderBase &builder,
266 LLVM::ModuleTranslation &moduleTranslation) {
267 if (!linearLoopIV)
268 return;
269 llvm::Value *linearIV = moduleTranslation.lookupValue(linearLoopIV);
270
271 // Find linearIV's index
272 size_t index;
273 for (index = 0; index < linearOrigVal.size(); index++)
274 if (linearIV == linearOrigVal[index])
275 break;
276 if (index == linearOrigVal.size())
277 return;
278
279 // Add one more step to the linear iteration variable
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");
285
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);
290 }
291
292 // Linear variable finalization is conditional on the last logical iteration.
293 // Create BB splits to manage the same.
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");
302 }
303
304 // Finalize the linear vars
305 llvm::OpenMPIRBuilder::InsertPointOrErrorTy
306 finalizeLinearVar(llvm::IRBuilderBase &builder,
307 LLVM::ModuleTranslation &moduleTranslation,
308 llvm::Value *lastIter) {
309 // Emit condition to check whether last logical iteration is being executed
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));
317 // Store the linear variable values to original variables.
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]);
323 }
324
325 // Create conditional branch such that the linear variable
326 // values are stored to original variables only at the
327 // last logical iteration
328 builder.SetInsertPoint(linearFinalizationBB->getTerminator());
329 builder.CreateCondBr(isLast, linearLastIterExitBB, linearExitBB);
330 linearFinalizationBB->getTerminator()->eraseFromParent();
331 // Emit barrier
332 builder.SetInsertPoint(linearExitBB->getTerminator());
333 return moduleTranslation.getOpenMPBuilder()->createBarrier(
334 builder, llvm::omp::OMPD_barrier);
335 }
336
337 // Emit stores for linear variables. Useful in case of SIMD
338 // construct.
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]);
344 }
345 }
346
347 // Rewrite all uses of the original variable, in the basic blocks in the
348 // [startBB, endBB] interval, with the linear variable in-place.
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;
353
354 assert(startBB && endBB && "Invalid startBB/endBB");
355
356 // Collect basic blocks from startBB to endBB.
357 worklist.push_back(startBB);
358 collectedBBs.insert(startBB);
359
360 while (!worklist.empty()) {
361 llvm::BasicBlock *bb = worklist.pop_back_val();
362
363 if (bb == endBB)
364 continue;
365
366 for (llvm::BasicBlock *succ : llvm::successors(bb)) {
367 if (collectedBBs.insert(succ).second)
368 worklist.push_back(succ);
369 }
370 }
371
372 // Rewrite all uses in the collected BBs.
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]);
379 }
380 }
381 }
382};
383
384} // namespace
385
386/// Looks up from the operation from and returns the PrivateClauseOp with
387/// name symbolName
388static omp::PrivateClauseOp findPrivatizer(Operation *from,
389 SymbolRefAttr symbolName) {
390 omp::PrivateClauseOp privatizer =
392 symbolName);
393 assert(privatizer && "privatizer not found in the symbol table");
394 return privatizer;
395}
396
397/// Check whether translation to LLVM IR for the given operation is currently
398/// supported. If not, descriptive diagnostics will be emitted to let users know
399/// this is a not-yet-implemented feature.
400///
401/// \returns success if no unimplemented features are needed to translate the
402/// given operation.
403static LogicalResult checkImplementationStatus(Operation &op) {
404 auto todo = [&op](StringRef clauseName) {
405 return op.emitError() << "not yet implemented: Unhandled clause "
406 << clauseName << " in " << op.getName()
407 << " operation";
408 };
409
410 auto checkAllocate = [&todo](auto op, LogicalResult &result) {
411 if (!op.getAllocateVars().empty() || !op.getAllocatorVars().empty())
412 result = todo("allocate");
413 };
414 auto checkBare = [&todo](auto op, LogicalResult &result) {
415 if (op.getKernelType() == omp::TargetExecMode::bare)
416 result = todo("ompx_bare");
417 };
418 auto checkDepend = [&todo](auto op, LogicalResult &result) {
419 if (!op.getDependVars().empty() || op.getDependKinds())
420 result = todo("depend");
421 };
422 auto checkHint = [](auto op, LogicalResult &) {
423 if (op.getHint())
424 op.emitWarning("hint clause discarded");
425 };
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) {
431 if (isByRef) {
432 result = todo("in_reduction with byref modifier");
433 return;
434 }
435 }
436 }
437 if (isa<omp::TargetOp>(op.getOperation())) {
438 if (auto inReductionSyms = op.getInReductionSyms()) {
439 for (auto sym :
440 (*inReductionSyms).template getAsRange<SymbolRefAttr>()) {
441 auto decl =
443 op, sym);
444 assert(decl &&
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");
448 return;
449 }
450 if (!decl.getCleanupRegion().empty()) {
451 result = todo("in_reduction with cleanup region");
452 return;
453 }
454 }
455 }
456 }
457 } else if (!op.getInReductionVars().empty() || op.getInReductionByref() ||
458 op.getInReductionSyms()) {
459 result = todo("in_reduction");
460 }
461 };
462 auto checkNowait = [&todo](auto op, LogicalResult &result) {
463 if (op.getNowait())
464 result = todo("nowait");
465 };
466 auto checkOrder = [&todo](auto op, LogicalResult &result) {
467 if (op.getOrder() || op.getOrderMod())
468 result = todo("order");
469 };
470 auto checkPrivate = [&todo](auto op, LogicalResult &result) {
471 if (!op.getPrivateVars().empty() || op.getPrivateSyms())
472 result = todo("privatization");
473 };
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();
482 // The `task` reduction modifier is supported on the parallel and
483 // worksharing (do/for and sections) constructs. Other modifiers, and the
484 // `task` modifier on other constructs, are not yet implemented.
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()) {
491 // The task reduction modifier lowering only handles non-byref
492 // reductions for now.
493 for (bool isByRef : *byref)
494 if (isByRef) {
495 result = todo("task reduction modifier with by-ref reduction");
496 break;
497 }
498 }
499 }
500 };
501 auto checkTaskReductionByref = [&todo](auto op, LogicalResult &result) {
502 if (auto byrefAttr = op.getTaskReductionByref())
503 for (bool isByRef : *byrefAttr)
504 if (isByRef) {
505 result = todo("task_reduction with byref modifier");
506 return;
507 }
508 };
509 auto checkReductionByref = [&todo](auto op, LogicalResult &result) {
510 if (auto byrefAttr = op.getReductionByref())
511 for (bool isByRef : *byrefAttr)
512 if (isByRef) {
513 result = todo("reduction with byref modifier");
514 return;
515 }
516 };
517 auto checkNumTeams = [&todo](auto op, LogicalResult &result) {
518 if (op.hasNumTeamsMultiDim())
519 result = todo("num_teams with multi-dimensional values");
520 };
521 auto checkNumThreads = [&todo](auto op, LogicalResult &result) {
522 if (op.hasNumThreadsMultiDim())
523 result = todo("num_threads with multi-dimensional values");
524 };
525
526 auto checkThreadLimit = [&todo](auto op, LogicalResult &result) {
527 if (op.hasThreadLimitMultiDim())
528 result = todo("thread_limit with multi-dimensional values");
529 };
530 auto checkMap = [&todo](auto op, LogicalResult &result) {
531 if (!op.getMapIterated().empty())
532 result = todo("map/motion clause with iterator modifier");
533 };
534
535 auto checkDynGroupprivate = [&todo](auto op, LogicalResult &result) {
536 if (op.getDynGroupprivateSize())
537 result = todo("dyn_groupprivate");
538 };
539
540 LogicalResult result = success();
542 .Case([&](omp::DistributeOp op) {
543 checkAllocate(op, result);
544 checkOrder(op, result);
545 })
546 .Case([&](omp::SectionsOp op) {
547 checkAllocate(op, result);
548 checkPrivate(op, result);
549 checkReduction(op, result);
550 })
551 .Case([&](omp::ScopeOp op) { checkReduction(op, result); })
552 .Case([&](omp::SingleOp op) {
553 checkAllocate(op, result);
554 checkPrivate(op, result);
555 })
556 .Case([&](omp::TeamsOp op) {
557 checkAllocate(op, result);
558 checkPrivate(op, result);
559 checkNumTeams(op, result);
560 checkThreadLimit(op, result);
561 checkDynGroupprivate(op, result);
562 })
563 .Case([&](omp::TaskOp op) {
564 checkAllocate(op, result);
565 checkInReduction(op, result);
566 })
567 .Case([&](omp::TaskgroupOp op) {
568 checkAllocate(op, result);
569 checkTaskReductionByref(op, result);
570 })
571 .Case([&](omp::DispatchOp op) {
572 // OpenMP 5.1 dispatch creates an explicit task; nowait controls whether
573 // it is included. Diagnose unsupported asynchronous tasking before 5.2,
574 // where the nowait property has no effect on dispatch.
576 op->getParentOfType<ModuleOp>(), /*fallback=*/51);
577 if (version < 52)
578 checkNowait(op, result);
579 })
580 .Case([&](omp::TaskloopContextOp op) {
581 checkAllocate(op, result);
582 checkInReduction(op, result);
583 checkReduction(op, result);
584 checkReductionByref(op, result);
585 })
586 .Case([&](omp::WsloopOp op) {
587 checkAllocate(op, result);
588 checkOrder(op, result);
589 checkReduction(op, result);
590 })
591 .Case([&](omp::ParallelOp op) {
592 checkReduction(op, result);
593 checkNumThreads(op, result);
594 })
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) {
599 checkHint(op, result);
600 Region &region = op.getRegion();
601 if (region.empty())
602 return;
603 mlir::Type argType = region.front().getArgument(0).getType();
604 auto structTy = dyn_cast<LLVM::LLVMStructType>(argType);
605 if (!structTy)
606 return;
607 DataLayout dl = DataLayout(op->getParentOfType<ModuleOp>());
608 unsigned totalBits = dl.getTypeSizeInBits(structTy);
609 if (totalBits > 128)
610 result = todo("compare for complex types wider than 128 bits");
611 })
612 .Case<omp::TargetEnterDataOp, omp::TargetExitDataOp>([&](auto op) {
613 checkDepend(op, result);
614 checkMap(op, result);
615 })
616 .Case([&](omp::TargetUpdateOp op) {
617 checkDepend(op, result);
618 checkMap(op, result);
619 })
620 .Case([&](omp::TargetOp op) {
621 checkAllocate(op, result);
622 checkBare(op, result);
623 checkInReduction(op, result);
624 checkMap(op, result);
625 checkThreadLimit(op, result);
626 })
627 .Case([&](omp::TargetDataOp op) { checkMap(op, result); })
628 .Case([&](omp::DeclareMapperInfoOp op) { checkMap(op, result); })
629 .Default([](Operation &) {
630 // Assume all clauses for an operation can be translated unless they are
631 // checked above.
632 });
633 return result;
634}
635
636static LogicalResult handleError(llvm::Error error, Operation &op) {
637 LogicalResult result = success();
638 if (error) {
639 llvm::handleAllErrors(
640 std::move(error),
641 [&](const PreviouslyReportedError &) { result = failure(); },
642 [&](const llvm::ErrorInfoBase &err) {
643 result = op.emitError(err.message());
644 });
645 }
646 return result;
647}
648
649template <typename T>
650static LogicalResult handleError(llvm::Expected<T> &result, Operation &op) {
651 if (!result)
652 return handleError(result.takeError(), op);
653
654 return success();
655}
656
657/// Find the insertion point for allocas given the current insertion point for
658/// normal operations in the builder.
659static llvm::OpenMPIRBuilder::InsertPointTy findAllocInsertPoints(
660 llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation,
661 llvm::SmallVectorImpl<llvm::BasicBlock *> *deallocBlocks = nullptr) {
662 // If there is an allocation insertion point on stack, i.e. we are in a nested
663 // operation and a specific point was provided by some surrounding operation,
664 // use it.
665 llvm::OpenMPIRBuilder::InsertPointTy allocInsertPoint;
666 llvm::ArrayRef<llvm::BasicBlock *> deallocInsertPoints;
667 WalkResult walkResult = moduleTranslation.stackWalk<OpenMPAllocStackFrame>(
668 [&](OpenMPAllocStackFrame &frame) {
669 allocInsertPoint = frame.allocInsertPoint;
670 deallocInsertPoints = frame.deallocBlocks;
671 return WalkResult::interrupt();
672 });
673 // In cases with multiple levels of outlining, the tree walk might find an
674 // insertion point that is inside the original function while the builder
675 // insertion point is inside the outlined function. We need to make sure that
676 // we do not use it in those cases.
677 if (walkResult.wasInterrupted() &&
678 allocInsertPoint.getNodeParent()->getParent() ==
679 builder.GetInsertBlock()->getParent()) {
680 if (deallocBlocks)
681 deallocBlocks->insert(deallocBlocks->end(), deallocInsertPoints.begin(),
682 deallocInsertPoints.end());
683 return allocInsertPoint;
684 }
685
686 // Otherwise, insert to the entry block of the surrounding function.
687 // If the current IRBuilder InsertPoint is the function's entry, it cannot
688 // also be used for alloca insertion which would result in insertion order
689 // confusion. Create a new BasicBlock for the Builder and use the entry block
690 // for the allocs.
691 // TODO: Create a dedicated alloca BasicBlock at function creation such that
692 // we do not need to move the current InsertPoint here.
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);
702 }
703
704 // Collect exit blocks, which is where explicit deallocations should happen in
705 // this case.
706 if (deallocBlocks) {
707 for (llvm::BasicBlock &block : *builder.GetInsertBlock()->getParent()) {
708 // TODO: This currently results in no blocks being added to the list when
709 // all exit blocks of the enclosing function have not been lowered before
710 // this is reached.
711 llvm::Instruction *terminator = block.getTerminatorOrNull();
712 if (isa_and_present<llvm::ReturnInst>(terminator))
713 deallocBlocks->emplace_back(&block);
714 }
715 }
716
717 llvm::BasicBlock &funcEntryBlock =
718 builder.GetInsertBlock()->getParent()->getEntryBlock();
719 return funcEntryBlock.getFirstInsertionPt();
720}
721
722/// Find the loop information structure for the loop nest being translated. It
723/// will return a `null` value unless called from the translation function for
724/// a loop wrapper operation after successfully translating its body.
725static llvm::CanonicalLoopInfo *
727 llvm::CanonicalLoopInfo *loopInfo = nullptr;
728 moduleTranslation.stackWalk<OpenMPLoopInfoStackFrame>(
729 [&](OpenMPLoopInfoStackFrame &frame) {
730 loopInfo = frame.loopInfo;
731 return WalkResult::interrupt();
732 });
733 return loopInfo;
734}
735
736/// Converts the given region that appears within an OpenMP dialect operation to
737/// LLVM IR, creating a branch from the `sourceBlock` to the entry block of the
738/// region, and a branch from any block with an successor-less OpenMP terminator
739/// to `continuationBlock`. Populates `continuationBlockPHIs` with the PHI nodes
740/// of the continuation block if provided.
742 Region &region, StringRef blockName, llvm::IRBuilderBase &builder,
743 LLVM::ModuleTranslation &moduleTranslation,
744 SmallVectorImpl<llvm::PHINode *> *continuationBlockPHIs = nullptr) {
745 bool isLoopWrapper = isa<omp::LoopWrapperInterface>(region.getParentOp());
746
747 llvm::BasicBlock *continuationBlock =
748 splitBB(builder, true, "omp.region.cont");
749 llvm::BasicBlock *sourceBlock = builder.GetInsertBlock();
750
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);
757 }
758
759 llvm::Instruction *sourceTerminator = sourceBlock->getTerminator();
760
761 // Terminators (namely YieldOp) may be forwarding values to the region that
762 // need to be available in the continuation block. Collect the types of these
763 // operands in preparation of creating PHI nodes. This is skipped for loop
764 // wrapper operations, for which we know in advance they have no terminators.
765 SmallVector<llvm::Type *> continuationBlockPHITypes;
766 unsigned numYields = 0;
767
768 if (!isLoopWrapper) {
769 bool operandsProcessed = false;
770 for (Block &bb : region.getBlocks()) {
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()));
776 }
777 operandsProcessed = true;
778 } else {
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());
784 (void)operandType;
785 assert(continuationBlockPHITypes[i] == operandType &&
786 "values of mismatching types yielded from the region");
787 }
788 }
789 numYields++;
790 }
791 }
792 }
793
794 // Insert PHI nodes in the continuation block for any values forwarded by the
795 // terminators in this region.
796 if (!continuationBlockPHITypes.empty())
797 assert(
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));
806 }
807
808 // Convert blocks one by one in topological order to ensure
809 // defs are converted before uses.
811 for (Block *bb : blocks) {
812 llvm::BasicBlock *llvmBB = moduleTranslation.lookupBlock(bb);
813 // Retarget the branch of the entry block to the entry block of the
814 // converted region (regions are single-entry).
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);
821 }
822
823 llvm::IRBuilderBase::InsertPointGuard guard(builder);
824 if (failed(
825 moduleTranslation.convertBlock(*bb, bb->isEntryBlock(), builder)))
826 return llvm::make_error<PreviouslyReportedError>();
827
828 // Create a direct branch here for loop wrappers to prevent their lack of a
829 // terminator from causing a crash below.
830 if (isLoopWrapper) {
831 builder.CreateBr(continuationBlock);
832 continue;
833 }
834
835 // Special handling for `omp.yield` and `omp.terminator` (we may have more
836 // than one): they return the control to the parent OpenMP dialect operation
837 // so replace them with the branch to the continuation block. We handle this
838 // here to avoid relying inter-function communication through the
839 // ModuleTranslation class to set up the correct insertion point. This is
840 // also consistent with MLIR's idiom of handling special region terminators
841 // in the same code that handles the region-owning operation.
842 Operation *terminator = bb->getTerminator();
843 if (isa<omp::TerminatorOp, omp::YieldOp>(terminator)) {
844 builder.CreateBr(continuationBlock);
845
846 for (unsigned i = 0, e = terminator->getNumOperands(); i < e; ++i)
847 (*continuationBlockPHIs)[i]->addIncoming(
848 moduleTranslation.lookupValue(terminator->getOperand(i)), llvmBB);
849 }
850 }
851 // After all blocks have been traversed and values mapped, connect the PHI
852 // nodes to the results of preceding blocks.
853 LLVM::detail::connectPHINodes(region, moduleTranslation);
854
855 // Remove the blocks and values defined in this region from the mapping since
856 // they are not visible outside of this region. This allows the same region to
857 // be converted several times, that is cloned, without clashes, and slightly
858 // speeds up the lookups.
859 moduleTranslation.forgetMapping(region);
860
861 return continuationBlock;
862}
863
864/// Convert ProcBindKind from MLIR-generated enum to LLVM enum.
865static llvm::omp::ProcBindKind getProcBindKind(omp::ClauseProcBindKind kind) {
866 switch (kind) {
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;
875 }
876 llvm_unreachable("Unknown ClauseProcBindKind kind");
877}
878
879/// Convert 'dispatch' operation into LLVM IR.
880static LogicalResult
881convertOmpDispatch(Operation &opInst, llvm::IRBuilderBase &builder,
882 LLVM::ModuleTranslation &moduleTranslation) {
883 auto dispatchOp = cast<omp::DispatchOp>(opInst);
884
885 if (failed(checkImplementationStatus(opInst)))
886 return failure();
887
888 auto &region = dispatchOp.getRegion();
889 auto result = convertOmpOpRegions(region, "omp.dispatch.region", builder,
890 moduleTranslation);
891 if (!result)
892 return handleError(result.takeError(), opInst);
893 builder.SetInsertPoint(*result);
894 return success();
895}
896
897/// Converts an OpenMP 'masked' operation into LLVM IR using OpenMPIRBuilder.
898static LogicalResult
899convertOmpMasked(Operation &opInst, llvm::IRBuilderBase &builder,
900 LLVM::ModuleTranslation &moduleTranslation) {
901 auto maskedOp = cast<omp::MaskedOp>(opInst);
902 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
903
904 if (failed(checkImplementationStatus(opInst)))
905 return failure();
906
907 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
909 // MaskedOp has only one region associated with it.
910 auto &region = maskedOp.getRegion();
911 builder.restoreIP(codeGenIP);
912 return convertOmpOpRegions(region, "omp.masked.region", builder,
913 moduleTranslation)
914 .takeError();
915 };
916
917 // TODO: Perform finalization actions for variables. This has to be
918 // called for variables which have destructors/finalizers.
919 auto finiCB = [&](InsertPointTy codeGenIP) { return llvm::Error::success(); };
920
921 llvm::Value *filterVal = nullptr;
922 if (auto filterVar = maskedOp.getFilteredThreadId()) {
923 filterVal = moduleTranslation.lookupValue(filterVar);
924 } else {
925 llvm::LLVMContext &llvmContext = builder.getContext();
926 filterVal =
927 llvm::ConstantInt::get(llvm::Type::getInt32Ty(llvmContext), /*V=*/0);
928 }
929 assert(filterVal != nullptr);
930 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
931 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
932 moduleTranslation.getOpenMPBuilder()->createMasked(ompLoc, bodyGenCB,
933 finiCB, filterVal);
934
935 if (failed(handleError(afterIP, opInst)))
936 return failure();
937
938 builder.restoreIP(*afterIP);
939 return success();
940}
941
942/// Converts an OpenMP 'master' operation into LLVM IR using OpenMPIRBuilder.
943static LogicalResult
944convertOmpMaster(Operation &opInst, llvm::IRBuilderBase &builder,
945 LLVM::ModuleTranslation &moduleTranslation) {
946 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
947 auto masterOp = cast<omp::MasterOp>(opInst);
948
949 if (failed(checkImplementationStatus(opInst)))
950 return failure();
951
952 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
954 // MasterOp has only one region associated with it.
955 auto &region = masterOp.getRegion();
956 builder.restoreIP(codeGenIP);
957 return convertOmpOpRegions(region, "omp.master.region", builder,
958 moduleTranslation)
959 .takeError();
960 };
961
962 // TODO: Perform finalization actions for variables. This has to be
963 // called for variables which have destructors/finalizers.
964 auto finiCB = [&](InsertPointTy codeGenIP) { return llvm::Error::success(); };
965
966 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
967 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
968 moduleTranslation.getOpenMPBuilder()->createMaster(ompLoc, bodyGenCB,
969 finiCB);
970
971 if (failed(handleError(afterIP, opInst)))
972 return failure();
973
974 builder.restoreIP(*afterIP);
975 return success();
976}
977
978/// Converts an OpenMP 'critical' operation into LLVM IR using OpenMPIRBuilder.
979static LogicalResult
980convertOmpCritical(Operation &opInst, llvm::IRBuilderBase &builder,
981 LLVM::ModuleTranslation &moduleTranslation) {
982 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
983 auto criticalOp = cast<omp::CriticalOp>(opInst);
984
985 if (failed(checkImplementationStatus(opInst)))
986 return failure();
987
988 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
990 // CriticalOp has only one region associated with it.
991 auto &region = cast<omp::CriticalOp>(opInst).getRegion();
992 builder.restoreIP(codeGenIP);
993 return convertOmpOpRegions(region, "omp.critical.region", builder,
994 moduleTranslation)
995 .takeError();
996 };
997
998 // TODO: Perform finalization actions for variables. This has to be
999 // called for variables which have destructors/finalizers.
1000 auto finiCB = [&](InsertPointTy codeGenIP) { return llvm::Error::success(); };
1001
1002 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
1003 llvm::LLVMContext &llvmContext = moduleTranslation.getLLVMContext();
1004 llvm::Constant *hint = nullptr;
1005
1006 // If it has a name, it probably has a hint too.
1007 if (criticalOp.getNameAttr()) {
1008 // The verifiers in OpenMP Dialect guarentee that all the pointers are
1009 // non-null
1010 auto symbolRef = cast<SymbolRefAttr>(criticalOp.getNameAttr());
1011 auto criticalDeclareOp =
1013 symbolRef);
1014 hint =
1015 llvm::ConstantInt::get(llvm::Type::getInt32Ty(llvmContext),
1016 static_cast<int>(criticalDeclareOp.getHint()));
1017 }
1018 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
1019 moduleTranslation.getOpenMPBuilder()->createCritical(
1020 ompLoc, bodyGenCB, finiCB, criticalOp.getName().value_or(""), hint);
1021
1022 if (failed(handleError(afterIP, opInst)))
1023 return failure();
1024
1025 builder.restoreIP(*afterIP);
1026 return success();
1027}
1028
1029/// A util to collect info needed to convert delayed privatizers from MLIR to
1030/// LLVM.
1033 llvm::Value *allocatedPtr;
1034 llvm::Value *allocator;
1035 };
1036
1037 template <typename OP>
1039 : blockArgs(
1040 cast<omp::BlockArgOpenMPOpInterface>(*op).getPrivateBlockArgs()) {
1041 mlirVars.reserve(blockArgs.size());
1042 llvmVars.reserve(blockArgs.size());
1043 collectPrivatizationDecls<OP>(op);
1044
1045 for (mlir::Value privateVar : op.getPrivateVars())
1046 mlirVars.push_back(privateVar);
1047 }
1048
1055
1056private:
1057 /// Populates `privatizations` with privatization declarations used for the
1058 /// given op.
1059 template <class OP>
1060 void collectPrivatizationDecls(OP op) {
1061 std::optional<ArrayAttr> attr = op.getPrivateSyms();
1062 if (!attr)
1063 return;
1064
1065 privatizers.reserve(privatizers.size() + attr->size());
1066 for (auto symbolRef : attr->getAsRange<SymbolRefAttr>()) {
1067 privatizers.push_back(findPrivatizer(op, symbolRef));
1068 }
1069 }
1070};
1071
1072/// Populates `reductions` with reduction declarations used in the given op.
1073template <typename T>
1074static void
1077 std::optional<ArrayAttr> attr = op.getReductionSyms();
1078 if (!attr)
1079 return;
1080
1081 reductions.reserve(reductions.size() + op.getNumReductionVars());
1082 for (auto symbolRef : attr->getAsRange<SymbolRefAttr>()) {
1083 reductions.push_back(
1085 op, symbolRef));
1086 }
1087}
1088
1089/// Look up and validate the declare_reduction ops referenced by a
1090/// reduction-like clause on the omp.taskloop.context translation path. Only
1091/// the non-byref, single-init-arg, no-cleanup form is supported in this
1092/// initial cut; richer shapes are rejected here with a diagnostic. \p syms
1093/// is the clause's symbol list (e.g. `getReductionSyms()` or
1094/// `getInReductionSyms()`), \p opName is the textual op name used in
1095/// diagnostics, and \p clauseName distinguishes "reduction" from
1096/// "in_reduction" in those diagnostics.
1098 Operation *contextOp, std::optional<ArrayAttr> syms, StringRef opName,
1099 StringRef clauseName, SmallVectorImpl<omp::DeclareReductionOp> &out) {
1100 if (!syms)
1101 return success();
1102 out.reserve(out.size() + syms->size());
1103 for (auto sym : syms->getAsRange<SymbolRefAttr>()) {
1105 contextOp, sym);
1106 if (!decl)
1107 return contextOp->emitError()
1108 << "failed to resolve " << clauseName
1109 << " declare_reduction symbol " << sym.getRootReference() << " in "
1110 << opName;
1111 if (decl.getInitializerRegion().front().getNumArguments() != 1)
1112 return contextOp->emitError()
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())
1119 return contextOp->emitError()
1120 << clauseName << " declare_reduction is missing a combiner region";
1121 out.push_back(decl);
1122 }
1123 return success();
1124}
1125
1126/// Translates the blocks contained in the given region and appends them to at
1127/// the current insertion point of `builder`. The operations of the entry block
1128/// are appended to the current insertion block. If set, `continuationBlockArgs`
1129/// is populated with translated values that correspond to the values
1130/// omp.yield'ed from the region.
1131static LogicalResult inlineConvertOmpRegions(
1132 Region &region, StringRef blockName, llvm::IRBuilderBase &builder,
1133 LLVM::ModuleTranslation &moduleTranslation,
1134 SmallVectorImpl<llvm::Value *> *continuationBlockArgs = nullptr) {
1135 if (region.empty())
1136 return success();
1137
1138 // Special case for single-block regions that don't create additional blocks:
1139 // insert operations without creating additional blocks.
1140 if (region.hasOneBlock()) {
1141 llvm::Instruction *potentialTerminator =
1142 builder.GetInsertBlock()->empty() ? nullptr
1143 : &builder.GetInsertBlock()->back();
1144
1145 if (potentialTerminator && potentialTerminator->isTerminator())
1146 potentialTerminator->removeFromParent();
1147 moduleTranslation.mapBlock(&region.front(), builder.GetInsertBlock());
1148
1149 if (failed(moduleTranslation.convertBlock(
1150 region.front(), /*ignoreArguments=*/true, builder)))
1151 return failure();
1152
1153 // The continuation arguments are simply the translated terminator operands.
1154 if (continuationBlockArgs)
1155 llvm::append_range(
1156 *continuationBlockArgs,
1157 moduleTranslation.lookupValues(region.front().back().getOperands()));
1158
1159 // Drop the mapping that is no longer necessary so that the same region can
1160 // be processed multiple times.
1161 moduleTranslation.forgetMapping(region);
1162
1163 if (potentialTerminator && potentialTerminator->isTerminator()) {
1164 llvm::BasicBlock *block = builder.GetInsertBlock();
1165 if (block->empty()) {
1166 // this can happen for really simple reduction init regions e.g.
1167 // %0 = llvm.mlir.constant(0 : i32) : i32
1168 // omp.yield(%0 : i32)
1169 // because the llvm.mlir.constant (MLIR op) isn't converted into any
1170 // llvm op
1171 potentialTerminator->insertInto(block, block->begin());
1172 } else {
1173 potentialTerminator->insertAfter(&block->back());
1174 }
1175 }
1176
1177 return success();
1178 }
1179
1181 llvm::Expected<llvm::BasicBlock *> continuationBlock =
1182 convertOmpOpRegions(region, blockName, builder, moduleTranslation, &phis);
1183
1184 if (failed(handleError(continuationBlock, *region.getParentOp())))
1185 return failure();
1186
1187 if (continuationBlockArgs)
1188 llvm::append_range(*continuationBlockArgs, phis);
1189 builder.SetInsertPoint((*continuationBlock)->getFirstInsertionPt());
1190 return success();
1191}
1192
1193namespace {
1194/// Owning equivalents of OpenMPIRBuilder::(Atomic)ReductionGen that are used to
1195/// store lambdas with capture.
1196using OwningReductionGen =
1197 std::function<llvm::OpenMPIRBuilder::InsertPointOrErrorTy(
1198 llvm::OpenMPIRBuilder::InsertPointTy, llvm::Value *, llvm::Value *,
1199 llvm::Value *&)>;
1200using OwningAtomicReductionGen =
1201 std::function<llvm::OpenMPIRBuilder::InsertPointOrErrorTy(
1202 llvm::OpenMPIRBuilder::InsertPointTy, llvm::Type *, llvm::Value *,
1203 llvm::Value *)>;
1204using OwningDataPtrPtrReductionGen =
1205 std::function<llvm::OpenMPIRBuilder::InsertPointOrErrorTy(
1206 llvm::OpenMPIRBuilder::InsertPointTy, llvm::Value *, llvm::Value *&)>;
1207} // namespace
1208
1209/// Create an OpenMPIRBuilder-compatible reduction generator for the given
1210/// reduction declaration. The generator uses `builder` but ignores its
1211/// insertion point.
1212static OwningReductionGen
1213makeReductionGen(omp::DeclareReductionOp decl, llvm::IRBuilderBase &builder,
1214 LLVM::ModuleTranslation &moduleTranslation) {
1215 // The lambda is mutable because we need access to non-const methods of decl
1216 // (which aren't actually mutating it), and we must capture decl by-value to
1217 // avoid the dangling reference after the parent function returns.
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);
1227 if (failed(inlineConvertOmpRegions(decl.getReductionRegion(),
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();
1234 };
1235 return gen;
1236}
1237
1238/// Create an OpenMPIRBuilder-compatible atomic reduction generator for the
1239/// given reduction declaration. The generator uses `builder` but ignores its
1240/// insertion point. Returns null if there is no atomic region available in the
1241/// reduction declaration.
1242static OwningAtomicReductionGen
1243makeAtomicReductionGen(omp::DeclareReductionOp decl,
1244 llvm::IRBuilderBase &builder,
1245 LLVM::ModuleTranslation &moduleTranslation) {
1246 if (decl.getAtomicReductionRegion().empty())
1247 return OwningAtomicReductionGen();
1248
1249 // The lambda is mutable because we need access to non-const methods of decl
1250 // (which aren't actually mutating it), and we must capture decl by-value to
1251 // avoid the dangling reference after the parent function returns.
1252 OwningAtomicReductionGen atomicGen =
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);
1260 if (failed(inlineConvertOmpRegions(decl.getAtomicReductionRegion(),
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();
1267 };
1268 return atomicGen;
1269}
1270
1271/// Create an OpenMPIRBuilder-compatible `data_ptr_ptr` reduction generator for
1272/// the given reduction declaration. The generator uses `builder` but ignores
1273/// its insertion point. Returns null if there is no `data_ptr_ptr` region
1274/// available in the reduction declaration.
1275static OwningDataPtrPtrReductionGen
1276makeRefDataPtrGen(omp::DeclareReductionOp decl, llvm::IRBuilderBase &builder,
1277 LLVM::ModuleTranslation &moduleTranslation, bool isByRef) {
1278 if (!isByRef || decl.getDataPtrPtrRegion().empty())
1279 return OwningDataPtrPtrReductionGen();
1280
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);
1288 if (failed(inlineConvertOmpRegions(decl.getDataPtrPtrRegion(),
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();
1295 };
1296
1297 return refDataPtrGen;
1298}
1299
1300/// Converts an OpenMP 'ordered' operation into LLVM IR using OpenMPIRBuilder.
1301static LogicalResult
1302convertOmpOrdered(Operation &opInst, llvm::IRBuilderBase &builder,
1303 LLVM::ModuleTranslation &moduleTranslation) {
1304 auto orderedOp = cast<omp::OrderedOp>(opInst);
1305
1306 if (failed(checkImplementationStatus(opInst)))
1307 return failure();
1308
1309 omp::ClauseDepend dependType = *orderedOp.getDoacrossDependType();
1310 bool isDependSource = dependType == omp::ClauseDepend::dependsource;
1311 unsigned numLoops = *orderedOp.getDoacrossNumLoops();
1312 SmallVector<llvm::Value *> vecValues =
1313 moduleTranslation.lookupValues(orderedOp.getDoacrossDependVars());
1314
1315 size_t indexVecValues = 0;
1316 while (indexVecValues < vecValues.size()) {
1317 SmallVector<llvm::Value *> storeValues;
1318 storeValues.reserve(numLoops);
1319 for (unsigned i = 0; i < numLoops; i++) {
1320 storeValues.push_back(vecValues[indexVecValues]);
1321 indexVecValues++;
1322 }
1323 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
1324 findAllocInsertPoints(builder, moduleTranslation);
1325 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
1326 builder.restoreIP(moduleTranslation.getOpenMPBuilder()->createOrderedDepend(
1327 ompLoc, allocaIP, numLoops, storeValues, ".cnt.addr", isDependSource));
1328 }
1329 return success();
1330}
1331
1332/// Converts an OpenMP 'ordered_region' operation into LLVM IR using
1333/// OpenMPIRBuilder.
1334static LogicalResult
1335convertOmpOrderedRegion(Operation &opInst, llvm::IRBuilderBase &builder,
1336 LLVM::ModuleTranslation &moduleTranslation) {
1337 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
1338 auto orderedRegionOp = cast<omp::OrderedRegionOp>(opInst);
1339
1340 if (failed(checkImplementationStatus(opInst)))
1341 return failure();
1342
1343 auto bodyGenCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
1344 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) {
1345 // OrderedOp has only one region associated with it.
1346 auto &region = cast<omp::OrderedRegionOp>(opInst).getRegion();
1347 builder.restoreIP(codeGenIP);
1348 return convertOmpOpRegions(region, "omp.ordered.region", builder,
1349 moduleTranslation)
1350 .takeError();
1351 };
1352
1353 // TODO: Perform finalization actions for variables. This has to be
1354 // called for variables which have destructors/finalizers.
1355 auto finiCB = [&](InsertPointTy codeGenIP) { return llvm::Error::success(); };
1356
1357 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
1358 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
1359 moduleTranslation.getOpenMPBuilder()->createOrderedThreadsSimd(
1360 ompLoc, bodyGenCB, finiCB, !orderedRegionOp.getParLevelSimd());
1361
1362 if (failed(handleError(afterIP, opInst)))
1363 return failure();
1364
1365 builder.restoreIP(*afterIP);
1366 return success();
1367}
1368
1369namespace {
1370/// Contains the arguments for an LLVM store operation
1371struct DeferredStore {
1372 DeferredStore(llvm::Value *value, llvm::Value *address)
1373 : value(value), address(address) {}
1374
1375 llvm::Value *value;
1376 llvm::Value *address;
1377};
1378} // namespace
1379
1380/// Allocate space for privatized reduction variables.
1381/// `deferredStores` contains information to create store operations which needs
1382/// to be inserted after all allocas
1383template <typename T>
1384static LogicalResult
1386 llvm::IRBuilderBase &builder,
1387 LLVM::ModuleTranslation &moduleTranslation,
1388 const llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1390 SmallVectorImpl<llvm::Value *> &privateReductionVariables,
1391 DenseMap<Value, llvm::Value *> &reductionVariableMap,
1392 SmallVectorImpl<DeferredStore> &deferredStores,
1393 llvm::ArrayRef<bool> isByRefs) {
1394 llvm::IRBuilderBase::InsertPointGuard guard(builder);
1395 builder.SetInsertPoint(allocaIP.getNodeParent()->getTerminator());
1396
1397 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
1398 bool useDeviceSharedMem = omp::opInSharedDeviceContext(*op);
1399
1400 // delay creating stores until after all allocas
1401 deferredStores.reserve(op.getNumReductionVars());
1402
1403 for (std::size_t i = 0; i < op.getNumReductionVars(); ++i) {
1404 Region &allocRegion = reductionDecls[i].getAllocRegion();
1405 if (isByRefs[i]) {
1406 if (allocRegion.empty())
1407 continue;
1408
1410 if (failed(inlineConvertOmpRegions(allocRegion, "omp.reduction.alloc",
1411 builder, moduleTranslation, &phis)))
1412 return op.emitError(
1413 "failed to inline `alloc` region of `omp.declare_reduction`");
1414
1415 assert(phis.size() == 1 && "expected one allocation to be yielded");
1416 builder.SetInsertPoint(allocaIP.getNodeParent()->getTerminator());
1417
1418 // Allocate reduction variable (which is a pointer to the real reduction
1419 // variable allocated in the inlined region)
1420 llvm::Type *ptrTy = builder.getPtrTy();
1421 llvm::Type *varTy =
1422 moduleTranslation.convertType(reductionDecls[i].getType());
1423 llvm::Value *var;
1424 if (useDeviceSharedMem) {
1425 var = ompBuilder->createOMPAllocShared(builder, varTy);
1426 } else {
1427 var = builder.CreateAlloca(varTy);
1428 var = builder.CreatePointerBitCastOrAddrSpaceCast(var, ptrTy);
1429 }
1430
1431 llvm::Value *castPhi =
1432 builder.CreatePointerBitCastOrAddrSpaceCast(phis[0], ptrTy);
1433
1434 deferredStores.emplace_back(castPhi, var);
1435
1436 privateReductionVariables[i] = var;
1437 moduleTranslation.mapValue(reductionArgs[i], castPhi);
1438 reductionVariableMap.try_emplace(op.getReductionVars()[i], castPhi);
1439 } else {
1440 assert(allocRegion.empty() &&
1441 "allocaction is implicit for by-val reduction");
1442
1443 llvm::Type *ptrTy = builder.getPtrTy();
1444 llvm::Type *varTy =
1445 moduleTranslation.convertType(reductionDecls[i].getType());
1446 llvm::Value *var;
1447 if (useDeviceSharedMem) {
1448 var = ompBuilder->createOMPAllocShared(builder, varTy);
1449 } else {
1450 var = builder.CreateAlloca(varTy);
1451 var = builder.CreatePointerBitCastOrAddrSpaceCast(var, ptrTy);
1452 }
1453
1454 moduleTranslation.mapValue(reductionArgs[i], var);
1455 privateReductionVariables[i] = var;
1456 reductionVariableMap.try_emplace(op.getReductionVars()[i], var);
1457 }
1458 }
1459
1460 return success();
1461}
1462
1463/// Map input arguments to reduction initialization region
1464template <typename T>
1465static void
1467 llvm::IRBuilderBase &builder,
1469 DenseMap<Value, llvm::Value *> &reductionVariableMap,
1470 unsigned i) {
1471 // map input argument to the initialization region
1472 mlir::omp::DeclareReductionOp &reduction = reductionDecls[i];
1473 Region &initializerRegion = reduction.getInitializerRegion();
1474 Block &entry = initializerRegion.front();
1475
1476 mlir::Value mlirSource = loop.getReductionVars()[i];
1477 llvm::Value *llvmSource = moduleTranslation.lookupValue(mlirSource);
1478 llvm::Value *origVal = llvmSource;
1479 // If a non-pointer value is expected, load the value from the source pointer.
1480 if (!isa<LLVM::LLVMPointerType>(
1481 reduction.getInitializerMoldArg().getType()) &&
1482 isa<LLVM::LLVMPointerType>(mlirSource.getType())) {
1483 origVal =
1484 builder.CreateLoad(moduleTranslation.convertType(
1485 reduction.getInitializerMoldArg().getType()),
1486 llvmSource, "omp_orig");
1487 }
1488 moduleTranslation.mapValue(reduction.getInitializerMoldArg(), origVal);
1489
1490 if (entry.getNumArguments() > 1) {
1491 llvm::Value *allocation =
1492 reductionVariableMap.lookup(loop.getReductionVars()[i]);
1493 moduleTranslation.mapValue(reduction.getInitializerAllocArg(), allocation);
1494 }
1495}
1496
1497static void
1498setInsertPointForPossiblyEmptyBlock(llvm::IRBuilderBase &builder,
1499 llvm::BasicBlock *block = nullptr) {
1500 if (block == nullptr)
1501 block = builder.GetInsertBlock();
1502
1503 if (!block->hasTerminator())
1504 builder.SetInsertPoint(block);
1505 else
1506 builder.SetInsertPoint(block->getTerminator());
1507}
1508
1509/// Inline reductions' `init` regions. This functions assumes that the
1510/// `builder`'s insertion point is where the user wants the `init` regions to be
1511/// inlined; i.e. it does not try to find a proper insertion location for the
1512/// `init` regions. It also leaves the `builder's insertions point in a state
1513/// where the user can continue the code-gen directly afterwards.
1514template <typename OP>
1515static LogicalResult
1516initReductionVars(OP op, ArrayRef<BlockArgument> reductionArgs,
1517 llvm::IRBuilderBase &builder,
1518 LLVM::ModuleTranslation &moduleTranslation,
1519 llvm::BasicBlock *latestAllocaBlock,
1521 SmallVectorImpl<llvm::Value *> &privateReductionVariables,
1522 DenseMap<Value, llvm::Value *> &reductionVariableMap,
1523 llvm::ArrayRef<bool> isByRef,
1524 SmallVectorImpl<DeferredStore> &deferredStores) {
1525 if (op.getNumReductionVars() == 0)
1526 return success();
1527
1528 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
1529 bool useDeviceSharedMem = omp::opInSharedDeviceContext(*op);
1530
1531 llvm::BasicBlock *initBlock = splitBB(builder, true, "omp.reduction.init");
1532 auto allocaIP = latestAllocaBlock->getTerminator()->getIterator();
1533 builder.restoreIP(allocaIP);
1534 SmallVector<llvm::Value *> byRefVars(op.getNumReductionVars());
1535
1536 for (unsigned i = 0; i < op.getNumReductionVars(); ++i) {
1537 if (isByRef[i]) {
1538 if (!reductionDecls[i].getAllocRegion().empty())
1539 continue;
1540
1541 // TODO: remove after all users of by-ref are updated to use the alloc
1542 // region: Allocate reduction variable (which is a pointer to the real
1543 // reduciton variable allocated in the inlined region)
1544 llvm::Type *varTy =
1545 moduleTranslation.convertType(reductionDecls[i].getType());
1546 if (useDeviceSharedMem)
1547 byRefVars[i] = ompBuilder->createOMPAllocShared(builder, varTy);
1548 else
1549 byRefVars[i] = builder.CreateAlloca(varTy);
1550 }
1551 }
1552
1553 setInsertPointForPossiblyEmptyBlock(builder, initBlock);
1554
1555 // store result of the alloc region to the allocated pointer to the real
1556 // reduction variable
1557 for (auto [data, addr] : deferredStores)
1558 builder.CreateStore(data, addr);
1559
1560 // Before the loop, store the initial values of reductions into reduction
1561 // variables. Although this could be done after allocas, we don't want to mess
1562 // up with the alloca insertion point.
1563 for (unsigned i = 0; i < op.getNumReductionVars(); ++i) {
1565
1566 // map block argument to initializer region
1567 mapInitializationArgs(op, moduleTranslation, builder, reductionDecls,
1568 reductionVariableMap, i);
1569
1570 // TODO In some cases (specially on the GPU), the init regions may
1571 // contains stack alloctaions. If the region is inlined in a loop, this is
1572 // problematic. Instead of just inlining the region, handle allocations by
1573 // hoisting fixed length allocations to the function entry and using
1574 // stacksave and restore for variable length ones.
1575 if (failed(inlineConvertOmpRegions(reductionDecls[i].getInitializerRegion(),
1576 "omp.reduction.neutral", builder,
1577 moduleTranslation, &phis)))
1578 return failure();
1579
1580 assert(phis.size() == 1 && "expected one value to be yielded from the "
1581 "reduction neutral element declaration region");
1582
1584
1585 if (isByRef[i]) {
1586 if (!reductionDecls[i].getAllocRegion().empty())
1587 // done in allocReductionVars
1588 continue;
1589
1590 // TODO: this path can be removed once all users of by-ref are updated to
1591 // use an alloc region
1592
1593 // Store the result of the inlined region to the allocated reduction var
1594 // ptr
1595 builder.CreateStore(phis[0], byRefVars[i]);
1596
1597 privateReductionVariables[i] = byRefVars[i];
1598 moduleTranslation.mapValue(reductionArgs[i], phis[0]);
1599 reductionVariableMap.try_emplace(op.getReductionVars()[i], phis[0]);
1600 } else {
1601 // for by-ref case the store is inside of the reduction region
1602 builder.CreateStore(phis[0], privateReductionVariables[i]);
1603 // the rest was handled in allocByValReductionVars
1604 }
1605
1606 // forget the mapping for the initializer region because we might need a
1607 // different mapping if this reduction declaration is re-used for a
1608 // different variable
1609 moduleTranslation.forgetMapping(reductionDecls[i].getInitializerRegion());
1610 }
1611
1612 return success();
1613}
1614
1615/// Collect reduction info
1616template <typename T>
1617static void collectReductionInfo(
1618 T loop, llvm::IRBuilderBase &builder,
1619 LLVM::ModuleTranslation &moduleTranslation,
1622 SmallVectorImpl<OwningAtomicReductionGen> &owningAtomicReductionGens,
1624 const ArrayRef<llvm::Value *> privateReductionVariables,
1626 ArrayRef<bool> isByRef) {
1627 unsigned numReductions = loop.getNumReductionVars();
1628
1629 for (unsigned i = 0; i < numReductions; ++i) {
1630 owningReductionGens.push_back(
1631 makeReductionGen(reductionDecls[i], builder, moduleTranslation));
1632 owningAtomicReductionGens.push_back(
1633 makeAtomicReductionGen(reductionDecls[i], builder, moduleTranslation));
1635 reductionDecls[i], builder, moduleTranslation, isByRef[i]));
1636 }
1637
1638 // Collect the reduction information.
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]);
1646 mlir::Type allocatedType;
1647 reductionDecls[i].getAllocRegion().walk([&](mlir::Operation *op) {
1648 if (auto alloca = mlir::dyn_cast<LLVM::AllocaOp>(op)) {
1649 allocatedType = alloca.getElemType();
1651 }
1652
1654 });
1655
1656 reductionInfos.push_back(
1657 {moduleTranslation.convertType(reductionDecls[i].getType()), variable,
1658 privateReductionVariables[i],
1659 /*EvaluationKind=*/llvm::OpenMPIRBuilder::EvalKind::Scalar,
1661 /*ReductionGenClang=*/nullptr, atomicGen,
1663 allocatedType ? moduleTranslation.convertType(allocatedType) : nullptr,
1664 reductionDecls[i].getByrefElementType()
1665 ? moduleTranslation.convertType(
1666 *reductionDecls[i].getByrefElementType())
1667 : nullptr});
1668 }
1669}
1670
1671/// handling of DeclareReductionOp's cleanup region
1672static LogicalResult
1674 llvm::ArrayRef<llvm::Value *> privateVariables,
1675 LLVM::ModuleTranslation &moduleTranslation,
1676 llvm::IRBuilderBase &builder, StringRef regionName,
1677 bool shouldLoadCleanupRegionArg = true) {
1678 for (auto [i, cleanupRegion] : llvm::enumerate(cleanupRegions)) {
1679 if (cleanupRegion->empty())
1680 continue;
1681
1682 // map the argument to the cleanup region
1683 Block &entry = cleanupRegion->front();
1684
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(
1693 moduleTranslation.convertType(entry.getArgument(0).getType()),
1694 privateVariables[i])
1695 : privateVariables[i];
1696
1697 moduleTranslation.mapValue(entry.getArgument(0), privateVarValue);
1698
1699 if (failed(inlineConvertOmpRegions(*cleanupRegion, regionName, builder,
1700 moduleTranslation)))
1701 return failure();
1702
1703 // clear block argument mapping in case it needs to be re-created with a
1704 // different source for another use of the same reduction decl
1705 moduleTranslation.forgetMapping(*cleanupRegion);
1706 }
1707 return success();
1708}
1709
1710// TODO: not used by ParallelOp
1711template <class OP>
1712static LogicalResult createReductionsAndCleanup(
1713 OP op, llvm::IRBuilderBase &builder,
1714 LLVM::ModuleTranslation &moduleTranslation,
1715 llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1717 ArrayRef<llvm::Value *> privateReductionVariables, ArrayRef<bool> isByRef,
1718 bool isNowait = false, bool isTeamsReduction = false) {
1719 // Process the reductions if required.
1720 if (op.getNumReductionVars() == 0)
1721 return success();
1722
1724 SmallVector<OwningAtomicReductionGen> owningAtomicReductionGens;
1725 SmallVector<OwningDataPtrPtrReductionGen> owningReductionGenRefDataPtrGens;
1727
1728 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
1729
1730 // Create the reduction generators. We need to own them here because
1731 // ReductionInfo only accepts references to the generators.
1732 collectReductionInfo(op, builder, moduleTranslation, reductionDecls,
1733 owningReductionGens, owningAtomicReductionGens,
1734 owningReductionGenRefDataPtrGens,
1735 privateReductionVariables, reductionInfos, isByRef);
1736
1737 // The call to createReductions below expects the block to have a
1738 // terminator. Create an unreachable instruction to serve as terminator
1739 // and remove it later.
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);
1746
1747 if (failed(handleError(contInsertPoint, *op)))
1748 return failure();
1749
1750 if (!contInsertPoint->isValid())
1751 return op->emitOpError() << "failed to convert reductions";
1752
1753 llvm::OpenMPIRBuilder::InsertPointTy afterIP = *contInsertPoint;
1754 if (!isTeamsReduction) {
1755 llvm::OpenMPIRBuilder::InsertPointOrErrorTy barrierIP =
1756 ompBuilder->createBarrier({*contInsertPoint, reductionLoc},
1757 llvm::omp::OMPD_for);
1758
1759 if (failed(handleError(barrierIP, *op)))
1760 return failure();
1761 afterIP = *barrierIP;
1762 }
1763
1764 tempTerminator->eraseFromParent();
1765 builder.restoreIP(afterIP);
1766
1767 // after the construct, deallocate private reduction variables
1768 SmallVector<Region *> reductionRegions;
1769 llvm::transform(reductionDecls, std::back_inserter(reductionRegions),
1770 [](omp::DeclareReductionOp reductionDecl) {
1771 return &reductionDecl.getCleanupRegion();
1772 });
1773 LogicalResult result = inlineOmpRegionCleanup(
1774 reductionRegions, privateReductionVariables, moduleTranslation, builder,
1775 "omp.reduction.cleanup");
1776
1777 bool useDeviceSharedMem = omp::opInSharedDeviceContext(*op);
1778 if (useDeviceSharedMem) {
1779 for (auto [var, reductionDecl] :
1780 llvm::zip_equal(privateReductionVariables, reductionDecls))
1781 ompBuilder->createOMPFreeShared(
1782 builder, var, moduleTranslation.convertType(reductionDecl.getType()));
1783 }
1784
1785 return result;
1786}
1787
1788static ArrayRef<bool> getIsByRef(std::optional<ArrayRef<bool>> attr) {
1789 if (!attr)
1790 return {};
1791 return *attr;
1792}
1793
1794// TODO: not used by omp.parallel
1795template <typename OP>
1797 OP op, ArrayRef<BlockArgument> reductionArgs, llvm::IRBuilderBase &builder,
1798 LLVM::ModuleTranslation &moduleTranslation,
1799 llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1801 SmallVectorImpl<llvm::Value *> &privateReductionVariables,
1802 DenseMap<Value, llvm::Value *> &reductionVariableMap,
1803 llvm::ArrayRef<bool> isByRef) {
1804 if (op.getNumReductionVars() == 0)
1805 return success();
1806
1807 SmallVector<DeferredStore> deferredStores;
1808
1809 if (failed(allocReductionVars(op, reductionArgs, builder, moduleTranslation,
1810 allocaIP, reductionDecls,
1811 privateReductionVariables, reductionVariableMap,
1812 deferredStores, isByRef)))
1813 return failure();
1814
1815 return initReductionVars(op, reductionArgs, builder, moduleTranslation,
1816 allocaIP.getNodeParent(), reductionDecls,
1817 privateReductionVariables, reductionVariableMap,
1818 isByRef, deferredStores);
1819}
1820
1821/// Return the llvm::Value * corresponding to the `privateVar` that
1822/// is being privatized. It isn't always as simple as looking up
1823/// moduleTranslation with privateVar. For instance, in case of
1824/// an allocatable, the descriptor for the allocatable is privatized.
1825/// This descriptor is mapped using an MapInfoOp. So, this function
1826/// will return a pointer to the llvm::Value corresponding to the
1827/// block argument for the mapped descriptor.
1828static llvm::Value *
1829findAssociatedValue(Value privateVar, llvm::IRBuilderBase &builder,
1830 LLVM::ModuleTranslation &moduleTranslation,
1831 llvm::DenseMap<Value, Value> *mappedPrivateVars = nullptr) {
1832 if (mappedPrivateVars == nullptr || !mappedPrivateVars->contains(privateVar))
1833 return moduleTranslation.lookupValue(privateVar);
1834
1835 Value blockArg = (*mappedPrivateVars)[privateVar];
1836 Type privVarType = privateVar.getType();
1837 Type blockArgType = blockArg.getType();
1838 assert(isa<LLVM::LLVMPointerType>(blockArgType) &&
1839 "A block argument corresponding to a mapped var should have "
1840 "!llvm.ptr type");
1841
1842 if (privVarType == blockArgType)
1843 return moduleTranslation.lookupValue(blockArg);
1844
1845 // This typically happens when the privatized type is lowered from
1846 // boxchar<KIND> and gets lowered to !llvm.struct<(ptr, i64)>. That is the
1847 // struct/pair is passed by value. But, mapped values are passed only as
1848 // pointers, so before we privatize, we must load the pointer.
1849 if (!isa<LLVM::LLVMPointerType>(privVarType))
1850 return builder.CreateLoad(moduleTranslation.convertType(privVarType),
1851 moduleTranslation.lookupValue(blockArg));
1852
1853 return moduleTranslation.lookupValue(privateVar);
1854}
1855
1856// Privatizer region arguments may be by-value even when the available LLVM
1857// value is storage for that value, e.g. lowered Fortran boxchar descriptors in
1858// task context structs. Materialize the value expected by the region argument
1859// while preserving the existing pointer mapping for pointer arguments.
1860static llvm::Value *
1861materializeRegionArgValue(llvm::IRBuilderBase &builder,
1862 LLVM::ModuleTranslation &moduleTranslation,
1863 BlockArgument regionArg, llvm::Value *value) {
1864 if (!regionArg)
1865 return value;
1866
1867 llvm::Type *regionArgType =
1868 moduleTranslation.convertType(regionArg.getType());
1869 if (regionArgType->isPointerTy() || !value->getType()->isPointerTy())
1870 return value;
1871
1872 return builder.CreateLoad(regionArgType, value);
1873}
1874
1875/// Initialize a single (first)private variable. You probably want to use
1876/// allocateAndInitPrivateVars instead of this.
1877/// This returns the private variable which has been initialized. This
1878/// variable should be mapped before constructing the body of the Op.
1880initPrivateVar(llvm::IRBuilderBase &builder,
1881 LLVM::ModuleTranslation &moduleTranslation,
1882 omp::PrivateClauseOp &privDecl, llvm::Value *nonPrivateVar,
1883 BlockArgument &blockArg, llvm::Value *llvmPrivateVar,
1884 llvm::BasicBlock *privInitBlock,
1885 llvm::DenseMap<Value, Value> *mappedPrivateVars = nullptr) {
1886 Region &initRegion = privDecl.getInitRegion();
1887 if (initRegion.empty())
1888 return llvmPrivateVar;
1889
1890 assert(nonPrivateVar);
1891 moduleTranslation.mapValue(privDecl.getInitMoldArg(), nonPrivateVar);
1892 moduleTranslation.mapValue(privDecl.getInitPrivateArg(), llvmPrivateVar);
1893
1894 // in-place convert the private initialization region
1896 if (failed(inlineConvertOmpRegions(initRegion, "omp.private.init", builder,
1897 moduleTranslation, &phis)))
1898 return llvm::createStringError(
1899 "failed to inline `init` region of `omp.private`");
1900
1901 assert(phis.size() == 1 && "expected one allocation to be yielded");
1902
1903 // clear init region block argument mapping in case it needs to be
1904 // re-created with a different source for another use of the same
1905 // reduction decl
1906 moduleTranslation.forgetMapping(initRegion);
1907
1908 // Prefer the value yielded from the init region to the allocated private
1909 // variable in case the region is operating on arguments by-value (e.g.
1910 // Fortran character boxes).
1911 return phis[0];
1912}
1913
1914/// Version of initPrivateVar which looks up the nonPrivateVar from mlirPrivVar.
1916 llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation,
1917 omp::PrivateClauseOp &privDecl, Value mlirPrivVar, BlockArgument &blockArg,
1918 llvm::Value *llvmPrivateVar, llvm::BasicBlock *privInitBlock,
1919 llvm::DenseMap<Value, Value> *mappedPrivateVars = nullptr) {
1920 return initPrivateVar(
1921 builder, moduleTranslation, privDecl,
1922 findAssociatedValue(mlirPrivVar, builder, moduleTranslation,
1923 mappedPrivateVars),
1924 blockArg, llvmPrivateVar, privInitBlock, mappedPrivateVars);
1925}
1926
1927static llvm::Error
1928initPrivateVars(llvm::IRBuilderBase &builder,
1929 LLVM::ModuleTranslation &moduleTranslation,
1930 PrivateVarsInfo &privateVarsInfo,
1931 llvm::DenseMap<Value, Value> *mappedPrivateVars = nullptr) {
1932 if (privateVarsInfo.blockArgs.empty())
1933 return llvm::Error::success();
1934
1935 llvm::BasicBlock *privInitBlock = splitBB(builder, true, "omp.private.init");
1936 setInsertPointForPossiblyEmptyBlock(builder, privInitBlock);
1937
1938 for (auto [idx, zip] : llvm::enumerate(llvm::zip_equal(
1939 privateVarsInfo.privatizers, privateVarsInfo.mlirVars,
1940 privateVarsInfo.blockArgs, privateVarsInfo.llvmVars))) {
1941 auto [privDecl, mlirPrivVar, blockArg, llvmPrivateVar] = zip;
1943 builder, moduleTranslation, privDecl, mlirPrivVar, blockArg,
1944 llvmPrivateVar, privInitBlock, mappedPrivateVars);
1945
1946 if (!privVarOrErr)
1947 return privVarOrErr.takeError();
1948
1949 llvmPrivateVar = privVarOrErr.get();
1950 moduleTranslation.mapValue(blockArg, llvmPrivateVar);
1951
1953 }
1954
1955 return llvm::Error::success();
1956}
1957
1958static LogicalResult
1960 llvm::IRBuilderBase &builder,
1961 LLVM::ModuleTranslation &moduleTranslation,
1962 PrivateVarsInfo &privateVarsInfo) {
1963 for (Value allocatorVar : allocatorVars) {
1964 if (privateVarsInfo.convertedAllocators.contains(allocatorVar))
1965 continue;
1966
1967 llvm::Value *allocator = moduleTranslation.lookupValue(allocatorVar);
1968 if (!allocator)
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());
1975 else
1976 return op.emitError(
1977 "OpenMP allocator operand must have integer or pointer type");
1978
1979 privateVarsInfo.convertedAllocators.try_emplace(allocatorVar, allocator);
1980 }
1981 return success();
1982}
1983
1984/// Allocate and initialize delayed private variables. Returns the basic block
1985/// which comes after all of these allocations. llvm::Value * for each of these
1986/// private variables are populated in llvmPrivateVars.
1987template <typename T>
1989 T op, llvm::IRBuilderBase &builder,
1990 LLVM::ModuleTranslation &moduleTranslation,
1991 PrivateVarsInfo &privateVarsInfo,
1992 llvm::OpenMPIRBuilder::InsertPointTy &allocaIP,
1993 llvm::DenseMap<Value, Value> *mappedPrivateVars = nullptr,
1994 std::optional<llvm::OpenMPIRBuilder::InsertPointTy> allocatorIP =
1995 std::nullopt) {
1996 // Allocate private vars
1997 // Save blocks before splits since allocaIP/allocatorIP iterators may follow
1998 // spliced instructions to the new block.
1999 llvm::BasicBlock *allocaBB = allocaIP.getNodeParent();
2000 llvm::Instruction *allocaTerminator = allocaBB->getTerminator();
2001 splitBB(allocaTerminator->getIterator(), true,
2002 allocaTerminator->getStableDebugLoc(), "omp.region.after_alloca");
2003 // Update the allocaTerminator since the alloca block was split above.
2004 allocaTerminator = allocaBB->getTerminator();
2005 // The new terminator is an uncondition branch created by the splitBB above.
2006 assert(allocaTerminator->getNumSuccessors() == 1 &&
2007 "This is an unconditional branch created by splitBB");
2008 allocaIP = allocaTerminator->getIterator();
2009
2010 llvm::Instruction *allocatorTerminator = nullptr;
2011 llvm::BasicBlock *afterAllocatorAllocations = nullptr;
2012 if (allocatorIP) {
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");
2021 }
2022
2023 std::optional<llvm::IRBuilderBase::InsertPointGuard> guard;
2024 if (!allocatorIP)
2025 guard.emplace(builder);
2026 builder.SetInsertPoint(allocaTerminator);
2027
2028 llvm::DataLayout dataLayout = builder.GetInsertBlock()->getDataLayout();
2029 llvm::BasicBlock *afterAllocas = allocaTerminator->getSuccessor(0);
2030
2031 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
2032 bool mightUseDeviceSharedMem = omp::opInSharedDeviceContext(*op);
2033 unsigned int allocaAS =
2034 moduleTranslation.getLLVMModule()->getDataLayout().getAllocaAddrSpace();
2035 unsigned int defaultAS = moduleTranslation.getLLVMModule()
2036 ->getDataLayout()
2037 .getProgramAddressSpace();
2038
2039 SmallVector<int64_t> allocateItemForPrivate(privateVarsInfo.blockArgs.size(),
2040 -1);
2041 ValueRange allocatorVars;
2042 DenseI64ArrayAttr allocateAlignments;
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;
2051 }
2052
2053 for (auto [privateIndex, tuple] : llvm::enumerate(llvm::zip_equal(
2054 privateVarsInfo.privatizers, privateVarsInfo.mlirVars,
2055 privateVarsInfo.blockArgs))) {
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(
2078 moduleTranslation.getLLVMModule()->getContext());
2079 if (!llvm::isUIntN(sizeTy->getBitWidth(), size.getFixedValue()))
2080 return llvm::createStringError(
2081 "OpenMP allocation size cannot be represented by the target size "
2082 "type");
2083 llvm::Value *sizeValue =
2084 llvm::ConstantInt::get(sizeTy, size.getFixedValue());
2085
2086 Value allocatorVar = allocatorVars[allocateIndex];
2087 auto allocator = privateVarsInfo.convertedAllocators.find(allocatorVar);
2088 if (allocator == privateVarsInfo.convertedAllocators.end())
2089 return llvm::createStringError(
2090 "failed to find converted OpenMP allocator operand");
2091 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2092 int64_t alignment =
2093 allocateAlignments ? allocateAlignments[allocateIndex] : 0;
2094 if (alignment != 0) {
2095 // The allocation must be aligned to at least the maximum of the
2096 // requested alignment and the alignment the base language requires
2097 // for the type being allocated.
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");
2108 } else {
2109 llvmPrivateVar = ompBuilder->createOMPAlloc(
2110 ompLoc, sizeValue, allocator->second, "omp.private.alloc");
2111 }
2112 if (!llvmPrivateVar)
2113 return llvm::createStringError(
2114 "failed to create OpenMP private allocation");
2115 privateVarsInfo.allocatorPrivates.push_back(
2116 {llvmPrivateVar, allocator->second});
2117 } else if (mightUseDeviceSharedMem &&
2119 llvmPrivateVar = ompBuilder->createOMPAllocShared(builder, llvmAllocType);
2120 } else {
2121 llvmPrivateVar = builder.CreateAlloca(
2122 llvmAllocType, /*ArraySize=*/nullptr, "omp.private.alloc");
2123 if (allocaAS != defaultAS)
2124 llvmPrivateVar = builder.CreateAddrSpaceCast(
2125 llvmPrivateVar, builder.getPtrTy(defaultAS));
2126 }
2127
2128 privateVarsInfo.llvmVars.push_back(llvmPrivateVar);
2129 }
2130
2131 return afterAllocatorAllocations ? afterAllocatorAllocations : afterAllocas;
2132}
2133
2134/// This can't always be determined statically, but when we can, it is good to
2135/// avoid generating compiler-added barriers which will deadlock the program.
2137 for (mlir::Operation *parent = op->getParentOp(); parent != nullptr;
2138 parent = parent->getParentOp()) {
2139 if (mlir::isa<omp::SingleOp, omp::CriticalOp>(parent))
2140 return true;
2141
2142 // e.g.
2143 // omp.single {
2144 // omp.parallel {
2145 // op
2146 // }
2147 // }
2148 if (mlir::isa<omp::ParallelOp>(parent))
2149 return false;
2150 }
2151 return false;
2152}
2153
2154static LogicalResult copyFirstPrivateVars(
2155 mlir::Operation *op, llvm::IRBuilderBase &builder,
2156 LLVM::ModuleTranslation &moduleTranslation,
2158 ArrayRef<llvm::Value *> llvmPrivateVars,
2159 SmallVectorImpl<omp::PrivateClauseOp> &privateDecls, bool insertBarrier,
2160 llvm::DenseMap<Value, Value> *mappedPrivateVars = nullptr) {
2161 // Apply copy region for firstprivate.
2162 bool needsFirstprivate =
2163 llvm::any_of(privateDecls, [](omp::PrivateClauseOp &privOp) {
2164 return privOp.getDataSharingType() ==
2165 omp::DataSharingClauseType::FirstPrivate;
2166 });
2167
2168 if (!needsFirstprivate)
2169 return success();
2170
2171 llvm::BasicBlock *copyBlock =
2172 splitBB(builder, /*CreateBranch=*/true, "omp.private.copy");
2173 setInsertPointForPossiblyEmptyBlock(builder, copyBlock);
2174
2175 for (auto [decl, moldVar, llvmVar] :
2176 llvm::zip_equal(privateDecls, moldVars, llvmPrivateVars)) {
2177 if (decl.getDataSharingType() != omp::DataSharingClauseType::FirstPrivate)
2178 continue;
2179
2180 // copyRegion implements `lhs = rhs`
2181 Region &copyRegion = decl.getCopyRegion();
2182
2183 llvm::Value *copyMoldVar = materializeRegionArgValue(
2184 builder, moduleTranslation, decl.getCopyMoldArg(), moldVar);
2185 llvm::Value *copyPrivateVar = materializeRegionArgValue(
2186 builder, moduleTranslation, decl.getCopyPrivateArg(), llvmVar);
2187
2188 moduleTranslation.mapValue(decl.getCopyMoldArg(), copyMoldVar);
2189
2190 // map copyRegion lhs arg
2191 moduleTranslation.mapValue(decl.getCopyPrivateArg(), copyPrivateVar);
2192
2193 // in-place convert copy region
2194 if (failed(inlineConvertOmpRegions(copyRegion, "omp.private.copy", builder,
2195 moduleTranslation)))
2196 return decl.emitError("failed to inline `copy` region of `omp.private`");
2197
2199
2200 // ignore unused value yielded from copy region
2201
2202 // clear copy region block argument mapping in case it needs to be
2203 // re-created with different sources for reuse of the same reduction
2204 // decl
2205 moduleTranslation.forgetMapping(copyRegion);
2206 }
2207
2208 if (insertBarrier && !opIsInSingleThread(op)) {
2209 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
2210 llvm::OpenMPIRBuilder::InsertPointOrErrorTy res =
2211 ompBuilder->createBarrier(builder, llvm::omp::OMPD_barrier);
2212 if (failed(handleError(res, *op)))
2213 return failure();
2214 }
2215
2216 return success();
2217}
2218
2219static LogicalResult copyFirstPrivateVars(
2220 mlir::Operation *op, llvm::IRBuilderBase &builder,
2221 LLVM::ModuleTranslation &moduleTranslation,
2222 SmallVectorImpl<mlir::Value> &mlirPrivateVars,
2223 ArrayRef<llvm::Value *> llvmPrivateVars,
2224 SmallVectorImpl<omp::PrivateClauseOp> &privateDecls, bool insertBarrier,
2225 llvm::DenseMap<Value, Value> *mappedPrivateVars = nullptr) {
2226 llvm::SmallVector<llvm::Value *> moldVars(mlirPrivateVars.size());
2227 llvm::transform(mlirPrivateVars, moldVars.begin(), [&](mlir::Value mlirVar) {
2228 // map copyRegion rhs arg
2229 llvm::Value *moldVar = findAssociatedValue(
2230 mlirVar, builder, moduleTranslation, mappedPrivateVars);
2231 assert(moldVar);
2232 return moldVar;
2233 });
2234 return copyFirstPrivateVars(op, builder, moduleTranslation, moldVars,
2235 llvmPrivateVars, privateDecls, insertBarrier,
2236 mappedPrivateVars);
2237}
2238
2239template <typename T>
2240static LogicalResult
2241cleanupPrivateVars(T op, llvm::IRBuilderBase &builder,
2242 LLVM::ModuleTranslation &moduleTranslation, Location loc,
2243 PrivateVarsInfo &privateVarsInfo) {
2244 // private variable deallocation
2245 SmallVector<Region *> privateCleanupRegions;
2246 llvm::transform(privateVarsInfo.privatizers,
2247 std::back_inserter(privateCleanupRegions),
2248 [](omp::PrivateClauseOp privatizer) {
2249 return &privatizer.getDeallocRegion();
2250 });
2251
2252 if (failed(inlineOmpRegionCleanup(privateCleanupRegions,
2253 privateVarsInfo.llvmVars, moduleTranslation,
2254 builder, "omp.private.dealloc",
2255 /*shouldLoadCleanupRegionArg=*/false)))
2256 return mlir::emitError(loc, "failed to inline `dealloc` region of an "
2257 "`omp.private` op in");
2259
2260 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
2261 bool mightUseDeviceSharedMem = omp::opInSharedDeviceContext(*op);
2262 for (auto [privDecl, llvmPrivVar, blockArg] :
2263 llvm::zip_equal(privateVarsInfo.privatizers, privateVarsInfo.llvmVars,
2264 privateVarsInfo.blockArgs)) {
2265 if (mightUseDeviceSharedMem && omp::allocaUsesRequireSharedMem(blockArg)) {
2266 ompBuilder->createOMPFreeShared(
2267 builder, llvmPrivVar,
2268 moduleTranslation.convertType(privDecl.getType()));
2269 }
2270 }
2271
2272 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2273 for (const PrivateVarsInfo::AllocatorPrivateInfo &allocation :
2274 llvm::reverse(privateVarsInfo.allocatorPrivates))
2275 ompBuilder->createOMPFree(ompLoc, allocation.allocatedPtr,
2276 allocation.allocator);
2277
2278 return success();
2279}
2280
2281/// Returns true if the construct contains omp.cancel or omp.cancellation_point
2283 // omp.cancel and omp.cancellation_point must be "closely nested" so they will
2284 // be visible and not inside of function calls. This is enforced by the
2285 // verifier.
2286 return op
2287 ->walk([](Operation *child) {
2288 if (mlir::isa<omp::CancelOp, omp::CancellationPointOp>(child))
2289 return WalkResult::interrupt();
2290 return WalkResult::advance();
2291 })
2292 .wasInterrupted();
2293}
2294
2295// Forward declarations for the task-reduction helpers defined alongside the
2296// omp.taskgroup lowering further down in this file. These are shared by the
2297// `reduction(task, ...)` modifier lowering on the parallel/worksharing
2298// constructs and by the omp.taskgroup / omp.taskloop.context task_reduction
2299// lowering. When \p isModifier is set, `__kmpc_taskred_modifier_init` is
2300// emitted (opening a task-reduction scope) instead of `__kmpc_taskred_init`,
2301// with \p isWorksharing selecting the runtime `is_ws` argument.
2302static llvm::Value *emitTaskReductionInitCall(
2304 ArrayRef<llvm::Value *> origPtrs, StringRef helperNamePrefix,
2305 llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder::InsertPointTy allocaIP,
2306 LLVM::ModuleTranslation &moduleTranslation, bool isModifier = false,
2307 bool isWorksharing = false);
2308static void
2309emitTaskReductionModifierFini(bool isWorksharing, llvm::IRBuilderBase &builder,
2310 LLVM::ModuleTranslation &moduleTranslation);
2311
2312static LogicalResult
2313convertOmpSections(Operation &opInst, llvm::IRBuilderBase &builder,
2314 LLVM::ModuleTranslation &moduleTranslation) {
2315 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2316 using StorableBodyGenCallbackTy =
2317 llvm::OpenMPIRBuilder::StorableBodyGenCallbackTy;
2318
2319 auto sectionsOp = cast<omp::SectionsOp>(opInst);
2320
2321 if (failed(checkImplementationStatus(opInst)))
2322 return failure();
2323
2324 llvm::ArrayRef<bool> isByRef = getIsByRef(sectionsOp.getReductionByref());
2325 assert(isByRef.size() == sectionsOp.getNumReductionVars());
2326
2328 collectReductionDecls(sectionsOp, reductionDecls);
2329 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
2330 findAllocInsertPoints(builder, moduleTranslation);
2331
2332 SmallVector<llvm::Value *> privateReductionVariables(
2333 sectionsOp.getNumReductionVars());
2334 DenseMap<Value, llvm::Value *> reductionVariableMap;
2335
2336 MutableArrayRef<BlockArgument> reductionArgs =
2337 cast<omp::BlockArgOpenMPOpInterface>(opInst).getReductionBlockArgs();
2338
2340 sectionsOp, reductionArgs, builder, moduleTranslation, allocaIP,
2341 reductionDecls, privateReductionVariables, reductionVariableMap,
2342 isByRef)))
2343 return failure();
2344
2345 bool isTaskReductionMod =
2346 sectionsOp.getReductionMod() == omp::ReductionModifier::task &&
2347 sectionsOp.getNumReductionVars() > 0;
2348
2350
2351 for (Operation &op : *sectionsOp.getRegion().begin()) {
2352 auto sectionOp = dyn_cast<omp::SectionOp>(op);
2353 if (!sectionOp) // omp.terminator
2354 continue;
2355
2356 Region &region = sectionOp.getRegion();
2357 auto sectionCB = [&sectionsOp, &region, &builder, &moduleTranslation](
2358 InsertPointTy allocaIP, InsertPointTy codeGenIP,
2359 ArrayRef<llvm::BasicBlock *> deallocBlocks) {
2360 builder.restoreIP(codeGenIP);
2361
2362 // map the omp.section reduction block argument to the omp.sections block
2363 // arguments
2364 // TODO: this assumes that the only block arguments are reduction
2365 // variables
2366 assert(region.getNumArguments() ==
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);
2371 assert(llvmVal);
2372 moduleTranslation.mapValue(sectionArg, llvmVal);
2373 }
2374
2375 return convertOmpOpRegions(region, "omp.section.region", builder,
2376 moduleTranslation)
2377 .takeError();
2378 };
2379 sectionCBs.push_back(sectionCB);
2380 }
2381
2382 // No sections within omp.sections operation - skip generation. This situation
2383 // is only possible if there is only a terminator operation inside the
2384 // sections operation
2385 if (sectionCBs.empty())
2386 return success();
2387
2388 // For `reduction(task, ...)` open a task-reduction scope for the worksharing
2389 // region. Participating explicit tasks accumulate into the per-thread private
2390 // copies, which the worksharing reduction then combines across threads. This
2391 // is emitted only after the empty-sections early return above, so it stays
2392 // balanced with the matching fini emitted after the sections region.
2393 if (isTaskReductionMod &&
2394 !emitTaskReductionInitCall(reductionDecls, privateReductionVariables,
2395 "__omp_taskred_mod_", builder, allocaIP,
2396 moduleTranslation, /*isModifier=*/true,
2397 /*isWorksharing=*/true))
2398 return sectionsOp.emitError(
2399 "failed to emit task reduction modifier initialization");
2400
2401 assert(isa<omp::SectionOp>(*sectionsOp.getRegion().op_begin()));
2402
2403 // TODO: Perform appropriate actions according to the data-sharing
2404 // attribute (shared, private, firstprivate, ...) of variables.
2405 // Currently defaults to shared.
2406 auto privCB = [&](InsertPointTy, InsertPointTy codeGenIP, llvm::Value &,
2407 llvm::Value &vPtr, llvm::Value *&replacementValue)
2408 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
2409 replacementValue = &vPtr;
2410 return codeGenIP;
2411 };
2412
2413 // TODO: Perform finalization actions for variables. This has to be
2414 // called for variables which have destructors/finalizers.
2415 auto finiCB = [&](InsertPointTy codeGenIP) { return llvm::Error::success(); };
2416
2417 allocaIP = findAllocInsertPoints(builder, moduleTranslation);
2418 bool isCancellable = constructIsCancellable(sectionsOp);
2419 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2420 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2421 moduleTranslation.getOpenMPBuilder()->createSections(
2422 ompLoc, allocaIP, sectionCBs, privCB, finiCB, isCancellable,
2423 sectionsOp.getNowait());
2424
2425 if (failed(handleError(afterIP, opInst)))
2426 return failure();
2427
2428 builder.restoreIP(*afterIP);
2429
2430 // Close the task-reduction scope before combining the worksharing copies.
2431 if (isTaskReductionMod)
2432 emitTaskReductionModifierFini(/*isWorksharing=*/true, builder,
2433 moduleTranslation);
2434
2435 // Process the reductions if required.
2437 sectionsOp, builder, moduleTranslation, allocaIP, reductionDecls,
2438 privateReductionVariables, isByRef, sectionsOp.getNowait());
2439}
2440
2441/// Converts an OpenMP scope construct into LLVM IR.
2442static LogicalResult
2443convertOmpScope(omp::ScopeOp &scopeOp, llvm::IRBuilderBase &builder,
2444 LLVM::ModuleTranslation &moduleTranslation) {
2445 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2446 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
2447
2448 if (failed(checkImplementationStatus(*scopeOp)))
2449 return failure();
2450
2451 llvm::ArrayRef<bool> isByRef = getIsByRef(scopeOp.getReductionByref());
2452 assert(isByRef.size() == scopeOp.getNumReductionVars());
2453
2454 PrivateVarsInfo privateVarsInfo(scopeOp);
2455 if (failed(convertAllocatorVars(*scopeOp, scopeOp.getAllocatorVars(), builder,
2456 moduleTranslation, privateVarsInfo)))
2457 return failure();
2458
2460 collectReductionDecls(scopeOp, reductionDecls);
2461 InsertPointTy privateAllocaIP =
2462 findAllocInsertPoints(builder, moduleTranslation);
2463
2464 SmallVector<llvm::Value *> privateReductionVariables(
2465 scopeOp.getNumReductionVars());
2466 DenseMap<Value, llvm::Value *> reductionVariableMap;
2467
2468 MutableArrayRef<BlockArgument> reductionArgs =
2469 cast<omp::BlockArgOpenMPOpInterface>(*scopeOp).getReductionBlockArgs();
2470
2471 if (scopeOp.getAllocateVars().empty()) {
2473 scopeOp, builder, moduleTranslation, privateVarsInfo, privateAllocaIP);
2474 if (failed(handleError(afterAllocas, *scopeOp)))
2475 return failure();
2476 }
2477
2479 scopeOp, reductionArgs, builder, moduleTranslation, privateAllocaIP,
2480 reductionDecls, privateReductionVariables, reductionVariableMap,
2481 isByRef)))
2482 return failure();
2483
2484 auto bodyCB =
2485 [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
2486 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) -> llvm::Error {
2487 if (!scopeOp.getAllocateVars().empty()) {
2488 // Runtime storage must be allocated on each dynamic Scope entry. Ordinary
2489 // private allocas still use the enclosing alloca insertion point.
2491 scopeOp, builder, moduleTranslation, privateVarsInfo, privateAllocaIP,
2492 /*mappedPrivateVars=*/nullptr, codeGenIP);
2493 if (handleError(afterAllocas, *scopeOp).failed())
2494 return llvm::make_error<PreviouslyReportedError>();
2495 builder.SetInsertPoint(afterAllocas.get()->getTerminator());
2496 } else {
2497 builder.restoreIP(codeGenIP);
2498 }
2499
2500 if (handleError(
2501 initPrivateVars(builder, moduleTranslation, privateVarsInfo),
2502 *scopeOp)
2503 .failed())
2504 return llvm::make_error<PreviouslyReportedError>();
2505
2506 if (failed(copyFirstPrivateVars(
2507 scopeOp, builder, moduleTranslation, privateVarsInfo.mlirVars,
2508 privateVarsInfo.llvmVars, privateVarsInfo.privatizers,
2509 scopeOp.getPrivateNeedsBarrier())))
2510 return llvm::make_error<PreviouslyReportedError>();
2511
2512 return convertOmpOpRegions(scopeOp.getRegion(), "omp.scope.region", builder,
2513 moduleTranslation)
2514 .takeError();
2515 };
2516
2517 auto finiCB = [&](InsertPointTy codeGenIP) -> llvm::Error {
2518 InsertPointTy oldIP = builder.saveIP();
2519 builder.restoreIP(codeGenIP);
2520 if (failed(cleanupPrivateVars(scopeOp, builder, moduleTranslation,
2521 scopeOp.getLoc(), privateVarsInfo)))
2522 return llvm::make_error<PreviouslyReportedError>();
2523 builder.restoreIP(oldIP);
2524 return llvm::Error::success();
2525 };
2526
2527 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2528 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2529 ompBuilder->createScope(ompLoc, bodyCB, finiCB, scopeOp.getNowait());
2530
2531 if (failed(handleError(afterIP, *scopeOp)))
2532 return failure();
2533
2534 builder.restoreIP(*afterIP);
2535
2536 // Process the reductions if required.
2538 scopeOp, builder, moduleTranslation, privateAllocaIP, reductionDecls,
2539 privateReductionVariables, isByRef, scopeOp.getNowait(),
2540 /*isTeamsReduction=*/false);
2541}
2542
2543/// Converts an OpenMP single construct into LLVM IR using OpenMPIRBuilder.
2544static LogicalResult
2545convertOmpSingle(omp::SingleOp &singleOp, llvm::IRBuilderBase &builder,
2546 LLVM::ModuleTranslation &moduleTranslation) {
2547 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2548 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2549
2550 if (failed(checkImplementationStatus(*singleOp)))
2551 return failure();
2552
2553 auto bodyCB = [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
2554 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) {
2555 builder.restoreIP(codegenIP);
2556 return convertOmpOpRegions(singleOp.getRegion(), "omp.single.region",
2557 builder, moduleTranslation)
2558 .takeError();
2559 };
2560 auto finiCB = [&](InsertPointTy codeGenIP) { return llvm::Error::success(); };
2561
2562 // Handle copyprivate
2563 Operation::operand_range cpVars = singleOp.getCopyprivateVars();
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(
2572 moduleTranslation.lookupFunction(llvmFuncOp.getName()));
2573 }
2574
2575 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2576 moduleTranslation.getOpenMPBuilder()->createSingle(
2577 ompLoc, bodyCB, finiCB, singleOp.getNowait(), llvmCPVars,
2578 llvmCPFuncs);
2579
2580 if (failed(handleError(afterIP, *singleOp)))
2581 return failure();
2582
2583 builder.restoreIP(*afterIP);
2584 return success();
2585}
2586
2587static omp::DistributeOp
2589 // Early return if we found more than one distribute op or if we can't find
2590 // any distribute op in the teams region.
2591 omp::DistributeOp distOp;
2592 WalkResult walk = teamsOp.getRegion().walk([&](omp::DistributeOp op) {
2593 if (distOp)
2594 return WalkResult::interrupt();
2595 distOp = op;
2596 return WalkResult::skip();
2597 });
2598 if (walk.wasInterrupted() || !distOp)
2599 return {};
2600
2601 auto iface =
2602 llvm::cast<mlir::omp::BlockArgOpenMPOpInterface>(teamsOp.getOperation());
2603 // Check that all uses of the reduction block arg has the same distribute op
2604 // parent.
2606 for (auto ra : iface.getReductionBlockArgs())
2607 for (auto &use : ra.getUses()) {
2608 auto *useOp = use.getOwner();
2609 // Ignore debug uses.
2610 if (mlir::isa<LLVM::DbgDeclareOp, LLVM::DbgValueOp>(useOp)) {
2611 debugUses.push_back(useOp);
2612 continue;
2613 }
2614 if (!distOp->isProperAncestor(useOp))
2615 return {};
2616 }
2617
2618 // If we are going to use distribute reduction then remove any debug uses of
2619 // the reduction parameters in teamsOp. Otherwise they will be left without
2620 // any mapped value in moduleTranslation and will eventually error out.
2621 for (auto *use : debugUses)
2622 use->erase();
2623 return distOp;
2624}
2625
2626// Convert an OpenMP Teams construct to LLVM IR using OpenMPIRBuilder
2627static LogicalResult
2628convertOmpTeams(omp::TeamsOp op, llvm::IRBuilderBase &builder,
2629 LLVM::ModuleTranslation &moduleTranslation) {
2630 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
2631 if (failed(checkImplementationStatus(*op)))
2632 return failure();
2633
2634 DenseMap<Value, llvm::Value *> reductionVariableMap;
2635 unsigned numReductionVars = op.getNumReductionVars();
2637 SmallVector<llvm::Value *> privateReductionVariables(numReductionVars);
2638 llvm::ArrayRef<bool> isByRef;
2639 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
2640 findAllocInsertPoints(builder, moduleTranslation);
2641
2642 // Only do teams reduction if there is no distribute op that captures the
2643 // reduction instead.
2644 bool doTeamsReduction = !getDistributeCapturingTeamsReduction(op);
2645 if (doTeamsReduction) {
2646 isByRef = getIsByRef(op.getReductionByref());
2647
2648 assert(isByRef.size() == op.getNumReductionVars());
2649
2650 MutableArrayRef<BlockArgument> reductionArgs =
2651 llvm::cast<omp::BlockArgOpenMPOpInterface>(*op).getReductionBlockArgs();
2652
2653 collectReductionDecls(op, reductionDecls);
2654
2656 op, reductionArgs, builder, moduleTranslation, allocaIP,
2657 reductionDecls, privateReductionVariables, reductionVariableMap,
2658 isByRef)))
2659 return failure();
2660 }
2661
2662 auto bodyCB = [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
2663 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) {
2665 moduleTranslation, allocaIP, deallocBlocks);
2666 builder.restoreIP(codegenIP);
2667 return convertOmpOpRegions(op.getRegion(), "omp.teams.region", builder,
2668 moduleTranslation)
2669 .takeError();
2670 };
2671
2672 llvm::Value *numTeamsLower = nullptr;
2673 if (Value numTeamsLowerVar = op.getNumTeamsLower())
2674 numTeamsLower = moduleTranslation.lookupValue(numTeamsLowerVar);
2675
2676 llvm::Value *numTeamsUpper = nullptr;
2677 if (!op.getNumTeamsUpperVars().empty())
2678 numTeamsUpper = moduleTranslation.lookupValue(op.getNumTeams(0));
2679
2680 llvm::Value *threadLimit = nullptr;
2681 if (!op.getThreadLimitVars().empty())
2682 threadLimit = moduleTranslation.lookupValue(op.getThreadLimit(0));
2683
2684 llvm::Value *ifExpr = nullptr;
2685 if (Value ifVar = op.getIfExpr())
2686 ifExpr = moduleTranslation.lookupValue(ifVar);
2687
2688 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
2689 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
2690 moduleTranslation.getOpenMPBuilder()->createTeams(
2691 ompLoc, bodyCB, numTeamsLower, numTeamsUpper, threadLimit, ifExpr);
2692
2693 if (failed(handleError(afterIP, *op)))
2694 return failure();
2695
2696 builder.restoreIP(*afterIP);
2697 if (doTeamsReduction) {
2698 // Process the reductions if required.
2700 op, builder, moduleTranslation, allocaIP, reductionDecls,
2701 privateReductionVariables, isByRef,
2702 /*isNoWait*/ false, /*isTeamsReduction*/ true);
2703 }
2704 return success();
2705}
2706
2707static llvm::omp::RTLDependenceKindTy
2708convertDependKind(mlir::omp::ClauseTaskDepend kind) {
2709 switch (kind) {
2710 case mlir::omp::ClauseTaskDepend::taskdependin:
2711 return llvm::omp::RTLDependenceKindTy::DepIn;
2712 // The OpenMP runtime requires that the codegen for 'depend' clause for
2713 // 'out' dependency kind must be the same as codegen for 'depend' clause
2714 // with 'inout' dependency.
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;
2722 }
2723 llvm_unreachable("unhandled depend kind");
2724}
2725
2727 std::optional<ArrayAttr> dependKinds, OperandRange dependVars,
2728 LLVM::ModuleTranslation &moduleTranslation,
2730 if (dependVars.empty())
2731 return;
2732 for (auto dep : llvm::zip(dependVars, dependKinds->getValue())) {
2733 auto kind =
2734 cast<mlir::omp::ClauseTaskDependAttr>(std::get<1>(dep)).getValue();
2735 llvm::omp::RTLDependenceKindTy type = convertDependKind(kind);
2736 llvm::Value *depVal = moduleTranslation.lookupValue(std::get<0>(dep));
2737 llvm::OpenMPIRBuilder::DependData dd(type, depVal->getType(), depVal);
2738 dds.emplace_back(dd);
2739 }
2740}
2741
2742/// Shared implementation of a callback which adds a termiator for the new block
2743/// created for the branch taken when an openmp construct is cancelled. The
2744/// terminator is saved in \p cancelTerminators. This callback is invoked only
2745/// if there is cancellation inside of the taskgroup body.
2746/// The terminator will need to be fixed to branch to the correct block to
2747/// cleanup the construct.
2749 SmallVectorImpl<llvm::UncondBrInst *> &cancelTerminators,
2750 llvm::IRBuilderBase &llvmBuilder, llvm::OpenMPIRBuilder &ompBuilder,
2751 mlir::Operation *op, llvm::omp::Directive cancelDirective) {
2752 auto finiCB = [&](llvm::OpenMPIRBuilder::InsertPointTy ip) -> llvm::Error {
2753 llvm::IRBuilderBase::InsertPointGuard guard(llvmBuilder);
2754
2755 // ip is currently in the block branched to if cancellation occurred.
2756 // We need to create a branch to terminate that block.
2757 llvmBuilder.restoreIP(ip);
2758
2759 // We must still clean up the construct after cancelling it, so we need to
2760 // branch to the block that finalizes the taskgroup.
2761 // That block has not been created yet so use this block as a dummy for now
2762 // and fix this after creating the operation.
2763 cancelTerminators.push_back(llvmBuilder.CreateBr(ip.getNodeParent()));
2764 return llvm::Error::success();
2765 };
2766 // We have to add the cleanup to the OpenMPIRBuilder before the body gets
2767 // created in case the body contains omp.cancel (which will then expect to be
2768 // able to find this cleanup callback).
2769 ompBuilder.pushFinalizationCB(
2770 {finiCB, cancelDirective, constructIsCancellable(op)});
2771}
2772
2773/// If we cancelled the construct, we should branch to the finalization block of
2774/// that construct. OMPIRBuilder structures the CFG such that the cleanup block
2775/// is immediately before the continuation block. Now this finalization has
2776/// been created we can fix the branch.
2777static void
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);
2785}
2786
2787namespace {
2788/// TaskContextStructManager takes care of creating and freeing a structure
2789/// containing information needed by the task body to execute.
2790class TaskContextStructManager {
2791public:
2792 TaskContextStructManager(llvm::IRBuilderBase &builder,
2793 LLVM::ModuleTranslation &moduleTranslation,
2794 MutableArrayRef<omp::PrivateClauseOp> privateDecls)
2795 : builder{builder}, moduleTranslation{moduleTranslation},
2796 privateDecls{privateDecls} {}
2797
2798 /// Creates a heap allocated struct containing space for each private
2799 /// variable. Invariant: privateVarTypes, privateDecls, and the elements of
2800 /// the structure should all have the same order (although privateDecls which
2801 /// do not read from the mold argument are skipped).
2802 void generateTaskContextStruct();
2803
2804 /// Create GEPs to access each member of the structure representing a private
2805 /// variable, adding them to llvmPrivateVars. Null values are added where
2806 /// private decls were skipped so that the ordering continues to match the
2807 /// private decls.
2808 void createGEPsToPrivateVars();
2809
2810 /// Given the address of the structure, return a GEP for each private variable
2811 /// in the structure. Null values are added where private decls were skipped
2812 /// so that the ordering continues to match the private decls.
2813 /// Must be called after generateTaskContextStruct().
2814 SmallVector<llvm::Value *>
2815 createGEPsToPrivateVars(llvm::Value *altStructPtr) const;
2816
2817 /// De-allocate the task context structure.
2818 void freeStructPtr();
2819
2820 MutableArrayRef<llvm::Value *> getLLVMPrivateVarGEPs() {
2821 return llvmPrivateVarGEPs;
2822 }
2823
2824 llvm::Value *getStructPtr() { return structPtr; }
2825
2826private:
2827 llvm::IRBuilderBase &builder;
2828 LLVM::ModuleTranslation &moduleTranslation;
2829 MutableArrayRef<omp::PrivateClauseOp> privateDecls;
2830
2831 /// The type of each member of the structure, in order.
2832 SmallVector<llvm::Type *> privateVarTypes;
2833
2834 /// LLVM values for each private variable, or null if that private variable is
2835 /// not included in the task context structure
2836 SmallVector<llvm::Value *> llvmPrivateVarGEPs;
2837
2838 /// A pointer to the structure containing context for this task.
2839 llvm::Value *structPtr = nullptr;
2840 /// The type of the structure
2841 llvm::Type *structTy = nullptr;
2842};
2843
2844/// IteratorInfo extracts and prepares loop bounds information from an
2845/// mlir::omp::IteratorOp for lowering to LLVM IR.
2846///
2847/// It computes the per-dimension trip counts and the total linearized trip
2848/// count, casted to i64. These are used to build a canonical loop and to
2849/// reconstruct the physical induction variables inside the loop body.
2850class IteratorInfo {
2851private:
2852 llvm::SmallVector<llvm::Value *> lowerBounds;
2853 llvm::SmallVector<llvm::Value *> upperBounds;
2854 llvm::SmallVector<llvm::Value *> steps;
2855 llvm::SmallVector<llvm::Value *> trips;
2856 unsigned dims;
2857 llvm::Value *totalTrips;
2858
2859 llvm::Value *lookUpAsI64(mlir::Value val, const LLVM::ModuleTranslation &mt,
2860 llvm::IRBuilderBase &builder) {
2861 llvm::Value *v = mt.lookupValue(val);
2862 if (!v)
2863 return nullptr;
2864 if (v->getType()->isIntegerTy(64))
2865 return v;
2866 if (v->getType()->isIntegerTy())
2867 return builder.CreateSExtOrTrunc(v, builder.getInt64Ty());
2868 return nullptr;
2869 }
2870
2871public:
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);
2878 steps.resize(dims);
2879 trips.resize(dims);
2880
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);
2886 llvm::Value *st =
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");
2893
2894 lowerBounds[d] = lb;
2895 upperBounds[d] = ub;
2896 steps[d] = st;
2897
2898 // trips = ((ub - lb) / step) + 1 (inclusive ub, assume positive step)
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));
2903 }
2904
2905 totalTrips = llvm::ConstantInt::get(builder.getInt64Ty(), 1);
2906 for (unsigned d = 0; d < dims; ++d)
2907 totalTrips = builder.CreateMul(totalTrips, trips[d]);
2908 }
2909
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; }
2916};
2917
2918} // namespace
2919
2920void TaskContextStructManager::generateTaskContextStruct() {
2921 if (privateDecls.empty())
2922 return;
2923 privateVarTypes.reserve(privateDecls.size());
2924
2925 for (omp::PrivateClauseOp &privOp : privateDecls) {
2926 // Skip private variables which can safely be allocated and initialised
2927 // inside of the task
2928 if (!privOp.readsFromMold())
2929 continue;
2930 Type mlirType = privOp.getType();
2931 privateVarTypes.push_back(moduleTranslation.convertType(mlirType));
2932 }
2933
2934 if (privateVarTypes.empty())
2935 return;
2936
2937 structTy = llvm::StructType::get(moduleTranslation.getLLVMContext(),
2938 privateVarTypes);
2939
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));
2945
2946 // Heap allocate the structure
2947 structPtr = builder.CreateMalloc(intPtrTy, allocSize,
2948 /*ArraySize=*/nullptr, /*MallocF=*/nullptr,
2949 "omp.task.context_ptr");
2950}
2951
2952SmallVector<llvm::Value *> TaskContextStructManager::createGEPsToPrivateVars(
2953 llvm::Value *altStructPtr) const {
2954 SmallVector<llvm::Value *> ret;
2955
2956 // Create GEPs for each struct member
2957 ret.reserve(privateDecls.size());
2958 llvm::Value *zero = builder.getInt32(0);
2959 unsigned i = 0;
2960 for (auto privDecl : privateDecls) {
2961 if (!privDecl.readsFromMold()) {
2962 // Handle this inside of the task so we don't pass unnessecary vars in
2963 ret.push_back(nullptr);
2964 continue;
2965 }
2966 llvm::Value *iVal = builder.getInt32(i);
2967 llvm::Value *gep = builder.CreateGEP(structTy, altStructPtr, {zero, iVal});
2968 ret.push_back(gep);
2969 i += 1;
2970 }
2971 return ret;
2972}
2973
2974void TaskContextStructManager::createGEPsToPrivateVars() {
2975 if (!structPtr)
2976 assert(privateVarTypes.empty());
2977 // Still need to run createGEPsToPrivateVars to populate llvmPrivateVarGEPs
2978 // with null values for skipped private decls
2979
2980 llvmPrivateVarGEPs = createGEPsToPrivateVars(structPtr);
2981}
2982
2983void TaskContextStructManager::freeStructPtr() {
2984 if (!structPtr)
2985 return;
2986
2987 llvm::IRBuilderBase::InsertPointGuard guard{builder};
2988 // Ensure we don't put the call to free() after the terminator
2989 builder.SetInsertPoint(builder.GetInsertBlock()->getTerminator());
2990 builder.CreateFree(structPtr);
2991}
2992
2993static void storeAffinityEntry(llvm::IRBuilderBase &builder,
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");
3001
3002 addr = builder.CreatePtrToInt(addr, kmpTaskAffinityInfoTy->getElementType(0));
3003 len = builder.CreateIntCast(len, kmpTaskAffinityInfoTy->getElementType(1),
3004 /*isSigned=*/false);
3005 llvm::Value *flags = builder.getInt32(0);
3006
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));
3013}
3014
3016 llvm::IRBuilderBase &builder,
3017 LLVM::ModuleTranslation &moduleTranslation,
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");
3022
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");
3026 storeAffinityEntry(builder, *moduleTranslation.getOpenMPBuilder(),
3027 affinityList, builder.getInt64(i), addr, len);
3028 }
3029}
3030
3031static mlir::LogicalResult
3032convertIteratorRegion(llvm::Value *linearIV, IteratorInfo &iterInfo,
3033 mlir::Block &iteratorRegionBlock,
3034 llvm::IRBuilderBase &builder,
3035 LLVM::ModuleTranslation &moduleTranslation) {
3036 llvm::Value *tmp = linearIV;
3037 for (int d = (int)iterInfo.getDims() - 1; d >= 0; --d) {
3038 llvm::Value *trip = iterInfo.getTrips()[d];
3039 // idx_d = tmp % trip_d
3040 llvm::Value *idx = builder.CreateURem(tmp, trip);
3041 // tmp = tmp / trip_d
3042 tmp = builder.CreateUDiv(tmp, trip);
3043
3044 // physIV_d = lb_d + idx_d * step_d
3045 llvm::Value *physIV = builder.CreateAdd(
3046 iterInfo.getLowerBounds()[d],
3047 builder.CreateMul(idx, iterInfo.getSteps()[d]), "omp.it.phys_iv");
3048
3049 moduleTranslation.mapValue(iteratorRegionBlock.getArgument(d), physIV);
3050 }
3051
3052 // Translate the iterator region into the loop body.
3053 moduleTranslation.mapBlock(&iteratorRegionBlock, builder.GetInsertBlock());
3054 if (mlir::failed(moduleTranslation.convertBlock(iteratorRegionBlock,
3055 /*ignoreArguments=*/true,
3056 builder))) {
3057 return mlir::failure();
3058 }
3059 return mlir::success();
3060}
3061
3063 llvm::function_ref<void(llvm::Value *linearIV, mlir::omp::YieldOp yield)>;
3064
3065static mlir::LogicalResult
3066fillIteratorLoop(mlir::omp::IteratorOp itersOp, llvm::IRBuilderBase &builder,
3067 mlir::LLVM::ModuleTranslation &moduleTranslation,
3068 IteratorInfo &iterInfo, llvm::StringRef loopName,
3069 IteratorStoreEntryTy genStoreEntry) {
3070 mlir::Region &itersRegion = itersOp.getRegion();
3071 mlir::Block &iteratorRegionBlock = itersRegion.front();
3072
3073 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
3074
3075 auto bodyGen = [&](llvm::OpenMPIRBuilder::InsertPointTy bodyIP,
3076 llvm::Value *linearIV) -> llvm::Error {
3077 llvm::IRBuilderBase::InsertPointGuard guard(builder);
3078 builder.restoreIP(bodyIP);
3079
3080 if (failed(convertIteratorRegion(linearIV, iterInfo, iteratorRegionBlock,
3081 builder, moduleTranslation))) {
3082 return llvm::make_error<llvm::StringError>(
3083 "failed to convert iterator region", llvm::inconvertibleErrorCode());
3084 }
3085
3086 auto yield =
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");
3090
3091 genStoreEntry(linearIV, yield);
3092
3093 // Iterator-region block/value mappings are temporary for this conversion,
3094 // clear them to avoid stale entries in ModuleTranslation.
3095 moduleTranslation.forgetMapping(itersRegion);
3096
3097 return llvm::Error::success();
3098 };
3099
3100 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
3101 moduleTranslation.getOpenMPBuilder()->createIteratorLoop(
3102 loc, iterInfo.getTotalTrips(), bodyGen, loopName);
3103 if (failed(handleError(afterIP, *itersOp)))
3104 return failure();
3105
3106 builder.restoreIP(*afterIP);
3107
3108 return mlir::success();
3109}
3110
3111static mlir::LogicalResult
3112buildAffinityData(mlir::omp::TaskOp &taskOp, llvm::IRBuilderBase &builder,
3113 mlir::LLVM::ModuleTranslation &moduleTranslation,
3114 llvm::OpenMPIRBuilder::AffinityData &ad) {
3115
3116 if (taskOp.getAffinityVars().empty() && taskOp.getIterated().empty()) {
3117 ad.Count = nullptr;
3118 ad.Info = nullptr;
3119 return mlir::success();
3120 }
3121
3123 llvm::StructType *kmpTaskAffinityInfoTy =
3124 moduleTranslation.getOpenMPBuilder()->getKmpTaskAffinityInfoTy();
3125
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))
3129 builder.restoreIP(findAllocInsertPoints(builder, moduleTranslation));
3130 return builder.CreateAlloca(kmpTaskAffinityInfoTy, count,
3131 "omp.affinity_list");
3132 };
3133
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());
3139 ad.Info =
3140 builder.CreatePointerBitCastOrAddrSpaceCast(info, builder.getPtrTy(0));
3141 return ad;
3142 };
3143
3144 if (!taskOp.getAffinityVars().empty()) {
3145 llvm::Value *count = llvm::ConstantInt::get(
3146 builder.getInt64Ty(), taskOp.getAffinityVars().size());
3147 llvm::Value *list = allocateAffinityList(count);
3148 fillAffinityLocators(taskOp.getAffinityVars(), builder, moduleTranslation,
3149 list);
3150 ads.emplace_back(createAffinity(count, list));
3151 }
3152
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());
3159 if (failed(fillIteratorLoop(
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");
3165 llvm::Value *addr =
3166 moduleTranslation.lookupValue(entryOp.getAddr());
3167 llvm::Value *len =
3168 moduleTranslation.lookupValue(entryOp.getLen());
3169 storeAffinityEntry(builder,
3170 *moduleTranslation.getOpenMPBuilder(),
3171 affList, linearIV, addr, len);
3172 })))
3173 return llvm::failure();
3174 ads.emplace_back(createAffinity(iterInfo.getTotalTrips(), affList));
3175 }
3176 }
3177
3178 llvm::Value *totalAffinityCount = builder.getInt32(0);
3179 for (const auto &affinity : ads)
3180 totalAffinityCount = builder.CreateAdd(
3181 totalAffinityCount,
3182 builder.CreateIntCast(affinity.Count, builder.getInt32Ty(),
3183 /*isSigned=*/false));
3184
3185 llvm::Value *affinityInfo = ads.front().Info;
3186 if (ads.size() > 1) {
3187 llvm::StructType *kmpTaskAffinityInfoTy =
3188 moduleTranslation.getOpenMPBuilder()->getKmpTaskAffinityInfoTy();
3189 llvm::Value *affinityInfoElemSize = builder.getInt64(
3190 moduleTranslation.getLLVMModule()->getDataLayout().getTypeAllocSize(
3191 kmpTaskAffinityInfoTy));
3192
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(), /*isSigned=*/false);
3198 llvm::Value *affinityCountInt64 = builder.CreateIntCast(
3199 affinityCount, builder.getInt64Ty(), /*isSigned=*/false);
3200 llvm::Value *affinityInfoSize =
3201 builder.CreateMul(affinityCountInt64, affinityInfoElemSize);
3202
3203 llvm::Value *packedAffinityInfoIndex = builder.CreateIntCast(
3204 packedAffinityInfoOffset, kmpTaskAffinityInfoTy->getElementType(0),
3205 /*isSigned=*/false);
3206 packedAffinityInfoIndex = builder.CreateInBoundsGEP(
3207 kmpTaskAffinityInfoTy, packedAffinityInfo, packedAffinityInfoIndex);
3208
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);
3215
3216 packedAffinityInfoOffset =
3217 builder.CreateAdd(packedAffinityInfoOffset, affinityCount);
3218 }
3219
3220 affinityInfo = packedAffinityInfo;
3221 }
3222
3223 ad.Count = totalAffinityCount;
3224 ad.Info = affinityInfo;
3225
3226 return mlir::success();
3227}
3228
3229// Allocates a single kmp_dep_info array sized to hold both locator
3230// (non-iterated) and iterated entries, fills the locator entries first, then
3231// runs an iterator loop for each iterator modifier object.
3232static mlir::LogicalResult
3233buildDependData(OperandRange dependVars, std::optional<ArrayAttr> dependKinds,
3234 OperandRange dependIterated,
3235 std::optional<ArrayAttr> dependIteratedKinds,
3236 llvm::IRBuilderBase &builder,
3237 mlir::LLVM::ModuleTranslation &moduleTranslation,
3238 llvm::OpenMPIRBuilder::DependenciesInfo &taskDeps) {
3239 if (dependIterated.empty()) {
3240 buildDependDataLocator(dependKinds, dependVars, moduleTranslation,
3241 taskDeps.Deps);
3242 return mlir::success();
3243 }
3244
3245 llvm::OpenMPIRBuilder &ompBuilder = *moduleTranslation.getOpenMPBuilder();
3246 llvm::Type *dependInfoTy = ompBuilder.DependInfo;
3247 unsigned numLocator = dependVars.size();
3248
3249 // Compute total count: locator deps + sum of iterator trip counts.
3250 llvm::Value *totalCount =
3251 llvm::ConstantInt::get(builder.getInt64Ty(), numLocator);
3252
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);
3258 totalCount =
3259 builder.CreateAdd(totalCount, iterInfos.back().getTotalTrips());
3260 }
3261
3262 // Heap-allocate the kmp_depend_info array so we don't risk
3263 // dynamic-sized alloca outside the entry block (e.g. inside loops).
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 /*MallocF=*/nullptr, ".dep.arr.addr");
3271
3272 // Fill non-iterated entries at indices [0, numLocator).
3273 if (numLocator > 0) {
3275 buildDependDataLocator(dependKinds, dependVars, moduleTranslation, dds);
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);
3281 }
3282 }
3283
3284 // Fill iterated entries starting at index numLocator.
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 =
3291 convertDependKind(kindAttr.getValue());
3292
3293 auto itersOp = dependIterated[i].getDefiningOp<mlir::omp::IteratorOp>();
3294 if (failed(fillIteratorLoop(
3295 itersOp, builder, moduleTranslation, iterInfo, "dep_iterator",
3296 [&](llvm::Value *linearIV, mlir::omp::YieldOp yield) {
3297 llvm::Value *addr =
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(
3303 builder, entry,
3304 llvm::OpenMPIRBuilder::DependData{rtlKind, addr->getType(),
3305 addr});
3306 })))
3307 return mlir::failure();
3308
3309 // Advance offset by the trip count of this iterator.
3310 offset = builder.CreateAdd(offset, iterInfo.getTotalTrips());
3311 }
3312
3313 taskDeps.DepArray = depArray;
3314 taskDeps.NumDeps = builder.CreateTrunc(totalCount, builder.getInt32Ty());
3315 return mlir::success();
3316}
3317
3318/// Converts an OpenMP task construct into LLVM IR using OpenMPIRBuilder.
3319static LogicalResult
3320convertOmpTaskOp(omp::TaskOp taskOp, llvm::IRBuilderBase &builder,
3321 LLVM::ModuleTranslation &moduleTranslation) {
3322 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
3323 if (failed(checkImplementationStatus(*taskOp)))
3324 return failure();
3325
3326 PrivateVarsInfo privateVarsInfo(taskOp);
3327 TaskContextStructManager taskStructMgr{builder, moduleTranslation,
3328 privateVarsInfo.privatizers};
3329
3330 // Allocate and copy private variables before creating the task. This avoids
3331 // accessing invalid memory if (after this scope ends) the private variables
3332 // are initialized from host variables or if the variables are copied into
3333 // from host variables (firstprivate). The insertion point is just before
3334 // where the code for creating and scheduling the task will go. That puts this
3335 // code outside of the outlined task region, which is what we want because
3336 // this way the initialization and copy regions are executed immediately while
3337 // the host variable data are still live.
3339 InsertPointTy allocaIP =
3340 findAllocInsertPoints(builder, moduleTranslation, &deallocBlocks);
3341
3342 // Not using splitBB() because that requires the current block to have a
3343 // terminator.
3344 assert(builder.GetInsertPoint() == builder.GetInsertBlock()->end());
3345 llvm::BasicBlock *taskStartBlock = llvm::BasicBlock::Create(
3346 builder.getContext(), "omp.task.start",
3347 /*Parent=*/builder.GetInsertBlock()->getParent());
3348 llvm::Instruction *branchToTaskStartBlock = builder.CreateBr(taskStartBlock);
3349 builder.SetInsertPoint(branchToTaskStartBlock);
3350
3351 // Now do this again to make the initialization and copy blocks
3352 llvm::BasicBlock *copyBlock =
3353 splitBB(builder, /*CreateBranch=*/true, "omp.private.copy");
3354 llvm::BasicBlock *initBlock =
3355 splitBB(builder, /*CreateBranch=*/true, "omp.private.init");
3356
3357 // Now the control flow graph should look like
3358 // starter_block:
3359 // <---- where we started when convertOmpTaskOp was called
3360 // br %omp.private.init
3361 // omp.private.init:
3362 // br %omp.private.copy
3363 // omp.private.copy:
3364 // br %omp.task.start
3365 // omp.task.start:
3366 // <---- where we want the insertion point to be when we call createTask()
3367
3368 // Save the alloca insertion point on ModuleTranslation stack for use in
3369 // nested regions.
3371 moduleTranslation, allocaIP, deallocBlocks);
3372
3373 // Allocate and initialize private variables
3374 builder.SetInsertPoint(initBlock->getTerminator());
3375
3376 // Create task variable structure
3377 taskStructMgr.generateTaskContextStruct();
3378 // GEPs so that we can initialize the variables. Don't use these GEPs inside
3379 // of the body otherwise it will be the GEP not the struct which is fowarded
3380 // to the outlined function. GEPs forwarded in this way are passed in a
3381 // stack-allocated (by OpenMPIRBuilder) structure which is not safe for tasks
3382 // which may not be executed until after the current stack frame goes out of
3383 // scope.
3384 taskStructMgr.createGEPsToPrivateVars();
3385
3386 for (auto [privDecl, mlirPrivVar, blockArg, llvmPrivateVarAlloc] :
3387 llvm::zip_equal(privateVarsInfo.privatizers, privateVarsInfo.mlirVars,
3388 privateVarsInfo.blockArgs,
3389 taskStructMgr.getLLVMPrivateVarGEPs())) {
3390 // To be handled inside the task.
3391 if (!privDecl.readsFromMold())
3392 continue;
3393 assert(llvmPrivateVarAlloc &&
3394 "reads from mold so shouldn't have been skipped");
3395
3396 llvm::Expected<llvm::Value *> privateVarOrErr =
3397 initPrivateVar(builder, moduleTranslation, privDecl, mlirPrivVar,
3398 blockArg, llvmPrivateVarAlloc, initBlock);
3399 if (!privateVarOrErr)
3400 return handleError(privateVarOrErr, *taskOp.getOperation());
3401
3403
3404 // TODO: this is a bit of a hack for Fortran character boxes.
3405 // Character boxes are passed by value into the init region and then the
3406 // initialized character box is yielded by value. Here we need to store the
3407 // yielded value into the private allocation, and load the private
3408 // allocation to match the type expected by region block arguments.
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);
3413 // Load it so we have the value pointed to by the GEP
3414 llvmPrivateVar = builder.CreateLoad(privateVarOrErr.get()->getType(),
3415 llvmPrivateVarAlloc);
3416 }
3417 assert(llvmPrivateVar->getType() ==
3418 moduleTranslation.convertType(blockArg.getType()));
3419
3420 // Mapping blockArg -> llvmPrivateVarAlloc is done inside the body callback
3421 // so that OpenMPIRBuilder doesn't try to pass each GEP address through a
3422 // stack allocated structure.
3423 }
3424
3425 // firstprivate copy region
3426 setInsertPointForPossiblyEmptyBlock(builder, copyBlock);
3427 if (failed(copyFirstPrivateVars(
3428 taskOp, builder, moduleTranslation, privateVarsInfo.mlirVars,
3429 taskStructMgr.getLLVMPrivateVarGEPs(), privateVarsInfo.privatizers,
3430 taskOp.getPrivateNeedsBarrier())))
3431 return llvm::failure();
3432
3433 llvm::OpenMPIRBuilder::AffinityData ad;
3434 if (failed(buildAffinityData(taskOp, builder, moduleTranslation, ad)))
3435 return llvm::failure();
3436
3437 // Resolve and validate in_reduction declarations. Byref in_reduction has
3438 // already been rejected by checkImplementationStatus; the helper rejects the
3439 // remaining richer declare_reduction shapes (two-argument initializer,
3440 // cleanup region, missing combiner). This is pure MLIR symbol-table work and
3441 // emits no IR. The matching task_reduction descriptor is registered by an
3442 // enclosing taskgroup; here we only look the per-task storage up at runtime.
3445 taskOp.getOperation(), taskOp.getInReductionSyms(), "omp.task",
3446 "in_reduction", inRedDecls)))
3447 return failure();
3448 SmallVector<llvm::Value *> inRedOrigPtrs;
3449 inRedOrigPtrs.reserve(inRedDecls.size());
3450 for (Value v : taskOp.getInReductionVars())
3451 inRedOrigPtrs.push_back(moduleTranslation.lookupValue(v));
3452
3453 // Set up for call to createTask()
3454 builder.SetInsertPoint(taskStartBlock);
3455
3456 auto bodyCB =
3457 [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
3458 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) -> llvm::Error {
3459 // Save the alloca insertion point on ModuleTranslation stack for use in
3460 // nested regions.
3462 moduleTranslation, allocaIP, deallocBlocks);
3463
3464 // translate the body of the task:
3465 builder.restoreIP(codegenIP);
3466
3467 llvm::BasicBlock *privInitBlock = nullptr;
3468 privateVarsInfo.llvmVars.resize(privateVarsInfo.blockArgs.size());
3469 for (auto [i, zip] : llvm::enumerate(llvm::zip_equal(
3470 privateVarsInfo.blockArgs, privateVarsInfo.privatizers,
3471 privateVarsInfo.mlirVars))) {
3472 auto [blockArg, privDecl, mlirPrivVar] = zip;
3473 // This is handled before the task executes
3474 if (privDecl.readsFromMold())
3475 continue;
3476
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, /*ArraySize=*/nullptr, "omp.private.alloc");
3483
3484 llvm::Expected<llvm::Value *> privateVarOrError =
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();
3491 }
3492
3493 taskStructMgr.createGEPsToPrivateVars();
3494 for (auto [i, llvmPrivVar] :
3495 llvm::enumerate(taskStructMgr.getLLVMPrivateVarGEPs())) {
3496 if (!llvmPrivVar) {
3497 assert(privateVarsInfo.llvmVars[i] &&
3498 "This is added in the loop above");
3499 continue;
3500 }
3501 privateVarsInfo.llvmVars[i] = llvmPrivVar;
3502 }
3503
3504 // Find and map the addresses of each variable within the task context
3505 // structure
3506 for (auto [blockArg, llvmPrivateVar, privateDecl] :
3507 llvm::zip_equal(privateVarsInfo.blockArgs, privateVarsInfo.llvmVars,
3508 privateVarsInfo.privatizers)) {
3509 // This was handled above.
3510 if (!privateDecl.readsFromMold())
3511 continue;
3512 // Fix broken pass-by-value case for Fortran character boxes
3513 if (!mlir::isa<LLVM::LLVMPointerType>(blockArg.getType())) {
3514 llvmPrivateVar = builder.CreateLoad(
3515 moduleTranslation.convertType(blockArg.getType()), llvmPrivateVar);
3516 }
3517 assert(llvmPrivateVar->getType() ==
3518 moduleTranslation.convertType(blockArg.getType()));
3519 moduleTranslation.mapValue(blockArg, llvmPrivateVar);
3520 }
3521
3522 // Map in_reduction block arguments to the per-task private storage returned
3523 // by __kmpc_task_reduction_get_th_data. This call must be emitted inside
3524 // the to-be-outlined task body so that it returns the *executing* thread's
3525 // gtid (not the encountering thread's). The descriptor is NULL: the runtime
3526 // walks up enclosing taskgroups to find the matching task_reduction
3527 // registration for `origPtr`. The original pointers are auto-captured into
3528 // the task shareds aggregate by CodeExtractor during
3529 // OpenMPIRBuilder::finalize.
3530 if (!inRedDecls.empty()) {
3531 auto iface = cast<omp::BlockArgOpenMPOpInterface>(taskOp.getOperation());
3532 llvm::OpenMPIRBuilder &ompB = *moduleTranslation.getOpenMPBuilder();
3533 llvm::Module *m = moduleTranslation.getLLVMModule();
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);
3540 // Align OpenMPIRBuilder's internal IRBuilder with `builder` so the gtid
3541 // call lands inside the to-be-outlined task body.
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);
3548 ArrayRef<BlockArgument> inRedBlockArgs = iface.getInReductionBlockArgs();
3549 for (auto [blockArg, origPtr] :
3550 llvm::zip_equal(inRedBlockArgs, inRedOrigPtrs)) {
3551 // __kmpc_task_reduction_get_th_data takes and returns a generic,
3552 // default-address-space `ptr`. Normalize a non-default-address-space
3553 // original pointer to the generic address space before the call, and
3554 // cast the returned private pointer back to the block argument's
3555 // address space when it differs (mirrors the taskloop reduction
3556 // remapping in convertOmpTaskloopContextOp).
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);
3569 }
3570 }
3571
3572 auto continuationBlockOrError = convertOmpOpRegions(
3573 taskOp.getRegion(), "omp.task.region", builder, moduleTranslation);
3574 if (failed(handleError(continuationBlockOrError, *taskOp)))
3575 return llvm::make_error<PreviouslyReportedError>();
3576
3577 builder.SetInsertPoint(continuationBlockOrError.get()->getTerminator());
3578
3579 if (failed(cleanupPrivateVars(taskOp, builder, moduleTranslation,
3580 taskOp.getLoc(), privateVarsInfo)))
3581 return llvm::make_error<PreviouslyReportedError>();
3582
3583 // Free heap allocated task context structure at the end of the task.
3584 taskStructMgr.freeStructPtr();
3585
3586 return llvm::Error::success();
3587 };
3588
3589 llvm::OpenMPIRBuilder &ompBuilder = *moduleTranslation.getOpenMPBuilder();
3590 SmallVector<llvm::UncondBrInst *> cancelTerminators;
3591 // The directive to match here is OMPD_taskgroup because it is the taskgroup
3592 // which is canceled. This is handled here because it is the task's cleanup
3593 // block which should be branched to.
3594 pushCancelFinalizationCB(cancelTerminators, builder, ompBuilder, taskOp,
3595 llvm::omp::Directive::OMPD_taskgroup);
3596
3597 llvm::OpenMPIRBuilder::DependenciesInfo dependencies;
3598 if (failed(buildDependData(taskOp.getDependVars(), taskOp.getDependKinds(),
3599 taskOp.getDependIterated(),
3600 taskOp.getDependIteratedKinds(), builder,
3601 moduleTranslation, dependencies)))
3602 return failure();
3603
3604 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
3605 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
3606 moduleTranslation.getOpenMPBuilder()->createTask(
3607 ompLoc, allocaIP, deallocBlocks, bodyCB, !taskOp.getUntied(),
3608 moduleTranslation.lookupValue(taskOp.getFinal()),
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);
3614
3615 if (failed(handleError(afterIP, *taskOp)))
3616 return failure();
3617
3618 // Set the correct branch target for task cancellation
3619 popCancelFinalizationCB(cancelTerminators, ompBuilder,
3620 afterIP->getNodeParent());
3621
3622 builder.restoreIP(*afterIP);
3623
3624 if (dependencies.DepArray)
3625 builder.CreateFree(dependencies.DepArray);
3626
3627 return success();
3628}
3629
3630/// The correct entry point is convertOmpTaskloopContextOp. This gets called
3631/// whilst lowering the body of the taskloop context (i.e. the task function).
3632static LogicalResult
3633convertOmpTaskloopWrapperOp(omp::TaskloopWrapperOp loopWrapperOp,
3634 llvm::IRBuilderBase &builder,
3635 LLVM::ModuleTranslation &moduleTranslation) {
3636 mlir::Operation &opInst = *loopWrapperOp.getOperation();
3637 if (failed(checkImplementationStatus(opInst)))
3638 return failure();
3639
3640 // Recurse into the loop body.
3641 auto continuationBlockOrError = convertOmpOpRegions(
3642 loopWrapperOp.getRegion(), "omp.taskloop.wrapper.region", builder,
3643 moduleTranslation);
3644
3645 if (failed(handleError(continuationBlockOrError, opInst)))
3646 return failure();
3647
3648 builder.SetInsertPoint(continuationBlockOrError.get());
3649 return success();
3650}
3651
3652/// Look up the given value in the mapping, and if it's not there, translate its
3653/// defining operation at the current builder insertion point. Only pure,
3654/// regionless operations are supported because the same operation will later be
3655/// translated again when the taskloop body itself is lowered.
3656static llvm::Expected<llvm::Value *>
3658 LLVM::ModuleTranslation &moduleTranslation,
3659 llvm::IRBuilderBase &builder) {
3660 if (llvm::Value *mapped = moduleTranslation.lookupValue(value))
3661 return mapped;
3662
3663 Operation *defOp = value.getDefiningOp();
3664 if (!defOp)
3665 return llvm::make_error<llvm::StringError>(
3666 "value is a block argument and is not mapped",
3667 llvm::inconvertibleErrorCode());
3668 if (defOp->getNumRegions() != 0 || !isPure(defOp))
3669 return llvm::make_error<llvm::StringError>(
3670 "unsupported op defining taskloop loop bound",
3671 llvm::inconvertibleErrorCode());
3672
3673 SmallVector<Value> mappingsToRemove;
3674 mappingsToRemove.reserve(defOp->getNumOperands() + defOp->getNumResults());
3675 for (Value operand : defOp->getOperands()) {
3676 if (moduleTranslation.lookupValue(operand))
3677 continue;
3678
3679 llvm::Expected<llvm::Value *> operandOrError =
3680 lookupOrTranslatePureValue(operand, moduleTranslation, builder);
3681 if (!operandOrError)
3682 return operandOrError.takeError();
3683 moduleTranslation.mapValue(operand, *operandOrError);
3684 mappingsToRemove.push_back(operand);
3685 }
3686
3687 if (failed(moduleTranslation.convertOperation(*defOp, builder)))
3688 return llvm::make_error<llvm::StringError>(
3689 "failed to convert op defining taskloop loop bound",
3690 llvm::inconvertibleErrorCode());
3691
3692 llvm::Value *result = moduleTranslation.lookupValue(value);
3693 assert(result && "expected conversion of loop bound op to produce a value");
3694
3695 for (Value resultValue : defOp->getResults()) {
3696 if (moduleTranslation.lookupValue(resultValue))
3697 mappingsToRemove.push_back(resultValue);
3698 }
3699 for (Value mappedValue : mappingsToRemove)
3700 moduleTranslation.forgetMapping(mappedValue);
3701
3702 return result;
3703}
3704
3705static llvm::Error
3706computeTaskloopBounds(omp::LoopNestOp loopOp, llvm::IRBuilderBase &builder,
3707 LLVM::ModuleTranslation &moduleTranslation,
3708 llvm::Value *&lbVal, llvm::Value *&ubVal,
3709 llvm::Value *&stepVal) {
3710 Operation::operand_range lowerBounds = loopOp.getLoopLowerBounds();
3711 Operation::operand_range upperBounds = loopOp.getLoopUpperBounds();
3712 Operation::operand_range steps = loopOp.getLoopSteps();
3713
3714 llvm::Expected<llvm::Value *> firstLbOrErr =
3715 lookupOrTranslatePureValue(lowerBounds[0], moduleTranslation, builder);
3716 if (!firstLbOrErr)
3717 return firstLbOrErr.takeError();
3718
3719 llvm::Type *boundType = (*firstLbOrErr)->getType();
3720 ubVal = builder.getIntN(boundType->getIntegerBitWidth(), 1);
3721 if (loopOp.getCollapseNumLoops() > 1) {
3722 // In cases where Collapse is used with Taskloop, the upper bound of the
3723 // iteration space needs to be recalculated to cater for the collapsed loop.
3724 // The Collapsed Loop UpperBound is the product of all collapsed
3725 // loop's tripcount.
3726 // The LowerBound for collapsed loops is always 1. When the loops are
3727 // collapsed, it will reset the bounds and introduce processing to ensure
3728 // the index's are presented as expected. As this happens after creating
3729 // Taskloop, these bounds need predicting. Example:
3730 // !$omp taskloop collapse(2)
3731 // do i = 1, 10
3732 // do j = 1, 5
3733 // ..
3734 // end do
3735 // end do
3736 // This loop above has a total of 50 iterations, so the lb will be 1, and
3737 // the ub will be 50. collapseLoops in OMPIRBuilder then handles ensuring
3738 // that i and j are properly presented when used in the loop.
3739 for (uint64_t i = 0; i < loopOp.getCollapseNumLoops(); i++) {
3741 i == 0 ? std::move(firstLbOrErr)
3742 : lookupOrTranslatePureValue(lowerBounds[i], moduleTranslation,
3743 builder);
3744 if (!lbOrErr)
3745 return lbOrErr.takeError();
3747 upperBounds[i], moduleTranslation, builder);
3748 if (!ubOrErr)
3749 return ubOrErr.takeError();
3751 lookupOrTranslatePureValue(steps[i], moduleTranslation, builder);
3752 if (!stepOrErr)
3753 return stepOrErr.takeError();
3754
3755 llvm::Value *loopLb = *lbOrErr;
3756 llvm::Value *loopUb = *ubOrErr;
3757 llvm::Value *loopStep = *stepOrErr;
3758 // In some cases, such as where the ub is less than the lb so the loop
3759 // steps down, the calculation for the loopTripCount is swapped. To ensure
3760 // the correct value is found, calculate both UB - LB and LB - UB then
3761 // select which value to use depending on how the loop has been
3762 // configured.
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());
3774 // For loops that have a step value not equal to 1, we need to adjust the
3775 // trip count to ensure the correct number of iterations for the loop is
3776 // captured.
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(
3786 loopTripCountRem,
3787 builder.getIntN(loopTripCountRem->getType()->getIntegerBitWidth(),
3788 0));
3789 loopTripCount =
3790 builder.CreateAdd(loopTripCountDivStep,
3791 builder.CreateZExtOrTrunc(
3792 needsRoundUp, loopTripCountDivStep->getType()));
3793 ubVal = builder.CreateMul(ubVal, loopTripCount);
3794 }
3795 lbVal = builder.getIntN(boundType->getIntegerBitWidth(), 1);
3796 stepVal = builder.getIntN(boundType->getIntegerBitWidth(), 1);
3797 } else {
3799 lookupOrTranslatePureValue(upperBounds[0], moduleTranslation, builder);
3800 if (!ubOrErr)
3801 return ubOrErr.takeError();
3803 lookupOrTranslatePureValue(steps[0], moduleTranslation, builder);
3804 if (!stepOrErr)
3805 return stepOrErr.takeError();
3806 lbVal = *firstLbOrErr;
3807 ubVal = *ubOrErr;
3808 stepVal = *stepOrErr;
3809 }
3810
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();
3815}
3816
3817// Converts an OpenMP taskloop construct into LLVM IR using OpenMPIRBuilder.
3818static LogicalResult
3819convertOmpTaskloopContextOp(omp::TaskloopContextOp contextOp,
3820 llvm::IRBuilderBase &builder,
3821 LLVM::ModuleTranslation &moduleTranslation) {
3822 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
3823 mlir::Operation &opInst = *contextOp.getOperation();
3824 omp::TaskloopWrapperOp loopWrapperOp = contextOp.getLoopOp();
3825 if (failed(checkImplementationStatus(opInst)))
3826 return failure();
3827
3828 // It stores the pointer of allocated firstprivate copies,
3829 // which can be used later for freeing the allocated space.
3830 SmallVector<llvm::Value *> llvmFirstPrivateVars;
3831 PrivateVarsInfo privateVarsInfo(contextOp);
3832 TaskContextStructManager taskStructMgr{builder, moduleTranslation,
3833 privateVarsInfo.privatizers};
3834
3836 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
3837 findAllocInsertPoints(builder, moduleTranslation, &deallocBlocks);
3838
3839 assert(builder.GetInsertPoint() == builder.GetInsertBlock()->end());
3840 llvm::BasicBlock *taskloopStartBlock = llvm::BasicBlock::Create(
3841 builder.getContext(), "omp.taskloop.wrapper.start",
3842 /*Parent=*/builder.GetInsertBlock()->getParent());
3843 llvm::Instruction *branchToTaskloopStartBlock =
3844 builder.CreateBr(taskloopStartBlock);
3845 builder.SetInsertPoint(branchToTaskloopStartBlock);
3846
3847 llvm::BasicBlock *copyBlock =
3848 splitBB(builder, /*CreateBranch=*/true, "omp.private.copy");
3849 llvm::BasicBlock *initBlock =
3850 splitBB(builder, /*CreateBranch=*/true, "omp.private.init");
3851
3853 moduleTranslation, allocaIP, deallocBlocks);
3854
3855 // Allocate and initialize private variables
3856 builder.SetInsertPoint(initBlock->getTerminator());
3857
3858 // TODO: don't allocate if the loop has zero iterations.
3859 taskStructMgr.generateTaskContextStruct();
3860 taskStructMgr.createGEPsToPrivateVars();
3861
3862 llvmFirstPrivateVars.resize(privateVarsInfo.blockArgs.size());
3863
3864 for (auto [i, zip] : llvm::enumerate(llvm::zip_equal(
3865 privateVarsInfo.privatizers, privateVarsInfo.mlirVars,
3866 privateVarsInfo.blockArgs, taskStructMgr.getLLVMPrivateVarGEPs()))) {
3867 auto [privDecl, mlirPrivVar, blockArg, llvmPrivateVarAlloc] = zip;
3868 // To be handled inside the taskloop.
3869 if (!privDecl.readsFromMold())
3870 continue;
3871 assert(llvmPrivateVarAlloc &&
3872 "reads from mold so shouldn't have been skipped");
3873
3874 llvm::Expected<llvm::Value *> privateVarOrErr =
3875 initPrivateVar(builder, moduleTranslation, privDecl, mlirPrivVar,
3876 blockArg, llvmPrivateVarAlloc, initBlock);
3877 if (!privateVarOrErr)
3878 return handleError(privateVarOrErr, *contextOp.getOperation());
3879
3880 llvmFirstPrivateVars[i] = privateVarOrErr.get();
3881
3882 llvm::IRBuilderBase::InsertPointGuard guard(builder);
3883 builder.SetInsertPoint(builder.GetInsertBlock()->getTerminator());
3884
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);
3889 // Load it so we have the value pointed to by the GEP
3890 llvmPrivateVar = builder.CreateLoad(privateVarOrErr.get()->getType(),
3891 llvmPrivateVarAlloc);
3892 }
3893 assert(llvmPrivateVar->getType() ==
3894 moduleTranslation.convertType(blockArg.getType()));
3895 }
3896
3897 // firstprivate copy region
3898 setInsertPointForPossiblyEmptyBlock(builder, copyBlock);
3899 if (failed(copyFirstPrivateVars(
3900 contextOp, builder, moduleTranslation, privateVarsInfo.mlirVars,
3901 taskStructMgr.getLLVMPrivateVarGEPs(), privateVarsInfo.privatizers,
3902 contextOp.getPrivateNeedsBarrier())))
3903 return llvm::failure();
3904
3905 // Resolve and validate reduction / in_reduction declarations up front.
3906 // This is pure MLIR symbol-table work and does not emit IR, so do it
3907 // before moving the builder to the taskloop start block. Richer
3908 // declare_reduction shapes (byref) have been rejected already by
3909 // checkImplementationStatus; the rest (two-argument initializer, cleanup
3910 // region, missing combiner) are rejected by the helper.
3913 contextOp.getOperation(), contextOp.getReductionSyms(),
3914 "omp.taskloop.context", "reduction", redDecls)))
3915 return failure();
3918 contextOp.getOperation(), contextOp.getInReductionSyms(),
3919 "omp.taskloop.context", "in_reduction", inRedDecls)))
3920 return failure();
3921
3922 // The op verifier rejects nogroup + reduction, so no check is needed here.
3923
3924 SmallVector<llvm::Value *> redOrigPtrs;
3925 redOrigPtrs.reserve(redDecls.size());
3926 for (Value v : contextOp.getReductionVars())
3927 redOrigPtrs.push_back(moduleTranslation.lookupValue(v));
3928 SmallVector<llvm::Value *> inRedOrigPtrs;
3929 inRedOrigPtrs.reserve(inRedDecls.size());
3930 for (Value v : contextOp.getInReductionVars())
3931 inRedOrigPtrs.push_back(moduleTranslation.lookupValue(v));
3932
3933 // Set up insertion point for emitting the implicit-taskgroup reduction
3934 // setup (if any) and for the subsequent call to createTaskloop().
3935 builder.SetInsertPoint(taskloopStartBlock);
3936
3937 llvm::OpenMPIRBuilder &ompBuilderRef = *moduleTranslation.getOpenMPBuilder();
3938 llvm::Module *module = moduleTranslation.getLLVMModule();
3939
3940 // If we have task_reduction items, we must emit our own implicit
3941 // __kmpc_taskgroup so that the descriptor returned by __kmpc_taskred_init
3942 // is associated with that taskgroup. We then force NoGroup=true so that
3943 // OpenMPIRBuilder::createTaskloop does not emit a second taskgroup.
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);
3952 // Align OpenMPIRBuilder's internal IRBuilder with `builder` so the
3953 // gtid call lands at our insertion point.
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});
3959
3960 redDesc = emitTaskReductionInitCall(redDecls, redOrigPtrs,
3961 "__omp_taskloop_taskred_", builder,
3962 allocaIP, moduleTranslation);
3963 if (!redDesc)
3964 return failure();
3965 }
3966
3967 auto loopOp = cast<omp::LoopNestOp>(loopWrapperOp.getWrappedLoop());
3968 llvm::Value *lbVal = nullptr;
3969 llvm::Value *ubVal = nullptr;
3970 llvm::Value *stepVal = nullptr;
3971 if (llvm::Error err = computeTaskloopBounds(
3972 loopOp, builder, moduleTranslation, lbVal, ubVal, stepVal))
3973 return handleError(std::move(err), opInst);
3974
3975 auto bodyCB =
3976 [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
3977 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) -> llvm::Error {
3978 // Save the alloca insertion point on ModuleTranslation stack for use in
3979 // nested regions.
3981 moduleTranslation, allocaIP, deallocBlocks);
3982
3983 // translate the body of the taskloop:
3984 builder.restoreIP(codegenIP);
3985
3986 llvm::BasicBlock *privInitBlock = nullptr;
3987 privateVarsInfo.llvmVars.resize(privateVarsInfo.blockArgs.size());
3988 for (auto [i, zip] : llvm::enumerate(llvm::zip_equal(
3989 privateVarsInfo.blockArgs, privateVarsInfo.privatizers,
3990 privateVarsInfo.mlirVars))) {
3991 auto [blockArg, privDecl, mlirPrivVar] = zip;
3992 // This is handled before the task executes
3993 if (privDecl.readsFromMold())
3994 continue;
3995
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, /*ArraySize=*/nullptr, "omp.private.alloc");
4002
4003 llvm::Expected<llvm::Value *> privateVarOrError =
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();
4010 }
4011
4012 taskStructMgr.createGEPsToPrivateVars();
4013 for (auto [i, llvmPrivVar] :
4014 llvm::enumerate(taskStructMgr.getLLVMPrivateVarGEPs())) {
4015 if (!llvmPrivVar) {
4016 assert(privateVarsInfo.llvmVars[i] &&
4017 "This is added in the loop above");
4018 continue;
4019 }
4020 privateVarsInfo.llvmVars[i] = llvmPrivVar;
4021 }
4022
4023 // Find and map the addresses of each variable within the taskloop context
4024 // structure
4025 for (auto [blockArg, llvmPrivateVar, privateDecl] :
4026 llvm::zip_equal(privateVarsInfo.blockArgs, privateVarsInfo.llvmVars,
4027 privateVarsInfo.privatizers)) {
4028 // This was handled above.
4029 if (!privateDecl.readsFromMold())
4030 continue;
4031 // Fix broken pass-by-value case for Fortran character boxes
4032 if (!mlir::isa<LLVM::LLVMPointerType>(blockArg.getType())) {
4033 llvmPrivateVar = builder.CreateLoad(
4034 moduleTranslation.convertType(blockArg.getType()), llvmPrivateVar);
4035 }
4036 assert(llvmPrivateVar->getType() ==
4037 moduleTranslation.convertType(blockArg.getType()));
4038 moduleTranslation.mapValue(blockArg, llvmPrivateVar);
4039 }
4040
4041 // Map reduction and in_reduction block arguments to the per-task private
4042 // storage returned by __kmpc_task_reduction_get_th_data. This call must
4043 // be emitted inside the to-be-outlined task body so that it returns the
4044 // *executing* thread's gtid (not the encountering thread's). The
4045 // taskgroup descriptor `redDesc` is computed in the outer scope and is
4046 // auto-captured into the task shareds aggregate by CodeExtractor during
4047 // OpenMPIRBuilder::finalize. For in_reduction the descriptor is NULL:
4048 // the runtime walks up enclosing taskgroups to find the matching
4049 // task_reduction registration for `origPtr`.
4050 if (!redDecls.empty() || !inRedDecls.empty()) {
4051 auto iface =
4052 cast<omp::BlockArgOpenMPOpInterface>(contextOp.getOperation());
4053 llvm::OpenMPIRBuilder &ompB = *moduleTranslation.getOpenMPBuilder();
4054 llvm::Module *m = moduleTranslation.getLLVMModule();
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);
4061 // Align OpenMPIRBuilder's internal IRBuilder with `builder` so the
4062 // gtid call lands inside the to-be-outlined task body.
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);
4068
4069 // Emit one __kmpc_task_reduction_get_th_data lookup for a reduction /
4070 // in_reduction item and map its block argument to the per-task private
4071 // storage the runtime returns. The runtime entry point takes (and
4072 // returns) a generic, default-address-space `ptr`, so normalize a
4073 // non-default-address-space original pointer to the generic address
4074 // space before the call (mirroring the descriptor setup in
4075 // emitTaskReductionInitCall), and cast the returned private pointer back
4076 // to the block argument's address space when that differs.
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);
4084 llvm::Value *priv =
4085 builder.CreateCall(getThData, {bodyGtid, desc, origPtr}, name);
4086 if (auto *argPtrTy = llvm::dyn_cast<llvm::PointerType>(
4087 moduleTranslation.convertType(blockArg.getType()));
4088 argPtrTy && argPtrTy->getAddressSpace() != 0)
4089 priv = builder.CreateAddrSpaceCast(priv, argPtrTy);
4090 moduleTranslation.mapValue(blockArg, priv);
4091 };
4092
4093 ArrayRef<BlockArgument> redBlockArgs = iface.getReductionBlockArgs();
4094 for (auto [blockArg, origPtr] :
4095 llvm::zip_equal(redBlockArgs, redOrigPtrs))
4096 remapReductionArg(blockArg, redDesc, origPtr, "omp.taskred.priv");
4097 ArrayRef<BlockArgument> inRedBlockArgs = iface.getInReductionBlockArgs();
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");
4102 }
4103
4104 // Lower the contents of the taskloop context region: this is the body of
4105 // the generated task, not the loop.
4106 auto continuationBlockOrError = convertOmpOpRegions(
4107 contextOp.getRegion(), "omp.taskloop.context.region", builder,
4108 moduleTranslation);
4109
4110 if (failed(handleError(continuationBlockOrError, opInst)))
4111 return llvm::make_error<PreviouslyReportedError>();
4112
4113 builder.SetInsertPoint(continuationBlockOrError.get()->getTerminator());
4114
4115 // This is freeing the private variables as mapped inside of the task: these
4116 // will be per-task private copies possibly after task duplication. This is
4117 // handled transparently by how these are passed to the structure passed
4118 // into the outlined function. When the task is duplicated, that structure
4119 // is duplicated too.
4120 if (failed(cleanupPrivateVars(contextOp, builder, moduleTranslation,
4121 contextOp.getLoc(), privateVarsInfo)))
4122 return llvm::make_error<PreviouslyReportedError>();
4123 // Similarly, the task context structure freed inside the task is the
4124 // per-task copy after task duplication.
4125 taskStructMgr.freeStructPtr();
4126
4127 return llvm::Error::success();
4128 };
4129
4130 // Taskloop divides into an appropriate number of tasks by repeatedly
4131 // duplicating the original task. Each time this is done, the task context
4132 // structure must be duplicated too.
4133 auto taskDupCB = [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
4134 llvm::Value *destPtr, llvm::Value *srcPtr)
4136 llvm::IRBuilderBase::InsertPointGuard guard(builder);
4137 builder.restoreIP(codegenIP);
4138
4139 llvm::Type *ptrTy =
4140 builder.getPtrTy(srcPtr->getType()->getPointerAddressSpace());
4141 llvm::Value *src =
4142 builder.CreateLoad(ptrTy, srcPtr, "omp.taskloop.context.src");
4143
4144 TaskContextStructManager &srcStructMgr = taskStructMgr;
4145 TaskContextStructManager destStructMgr(builder, moduleTranslation,
4146 privateVarsInfo.privatizers);
4147 destStructMgr.generateTaskContextStruct();
4148 llvm::Value *dest = destStructMgr.getStructPtr();
4149 dest->setName("omp.taskloop.context.dest");
4150 builder.CreateStore(dest, destPtr);
4151
4153 srcStructMgr.createGEPsToPrivateVars(src);
4155 destStructMgr.createGEPsToPrivateVars(dest);
4156
4157 // Inline init regions.
4158 for (auto [privDecl, mold, blockArg, llvmPrivateVarAlloc] :
4159 llvm::zip_equal(privateVarsInfo.privatizers, srcGEPs,
4160 privateVarsInfo.blockArgs, destGEPs)) {
4161 // To be handled inside task body.
4162 if (!privDecl.readsFromMold())
4163 continue;
4164 assert(llvmPrivateVarAlloc &&
4165 "reads from mold so shouldn't have been skipped");
4166
4167 llvm::Value *moldArg = materializeRegionArgValue(
4168 builder, moduleTranslation, privDecl.getInitMoldArg(), mold);
4170 builder, moduleTranslation, privDecl, moldArg, blockArg,
4171 llvmPrivateVarAlloc, builder.GetInsertBlock());
4172 if (!privateVarOrErr)
4173 return privateVarOrErr.takeError();
4174
4176
4177 // TODO: this is a bit of a hack for Fortran character boxes.
4178 // Character boxes are passed by value into the init region and then the
4179 // initialized character box is yielded by value. Here we need to store
4180 // the yielded value into the private allocation, and load the private
4181 // allocation to match the type expected by region block arguments.
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);
4186 // Load it so we have the value pointed to by the GEP
4187 llvmPrivateVar = builder.CreateLoad(privateVarOrErr.get()->getType(),
4188 llvmPrivateVarAlloc);
4189 }
4190 assert(llvmPrivateVar->getType() ==
4191 moduleTranslation.convertType(blockArg.getType()));
4192
4193 // Mapping blockArg -> llvmPrivateVarAlloc is done inside the body
4194 // callback so that OpenMPIRBuilder doesn't try to pass each GEP address
4195 // through a stack allocated structure.
4196 }
4197
4198 if (failed(copyFirstPrivateVars(contextOp.getOperation(), builder,
4199 moduleTranslation, srcGEPs, destGEPs,
4200 privateVarsInfo.privatizers,
4201 contextOp.getPrivateNeedsBarrier())))
4202 return llvm::make_error<PreviouslyReportedError>();
4203
4204 return builder.saveIP();
4205 };
4206
4207 auto loopInfo = [&]() -> llvm::Expected<llvm::CanonicalLoopInfo *> {
4208 llvm::CanonicalLoopInfo *loopInfo = findCurrentLoopInfo(moduleTranslation);
4209 return loopInfo;
4210 };
4211
4212 llvm::Value *ifCond = nullptr;
4213 llvm::Value *grainsize = nullptr;
4214 int sched = 0; // default
4215 mlir::Value grainsizeVal = contextOp.getGrainsize();
4216 mlir::Value numTasksVal = contextOp.getNumTasks();
4217 if (Value ifVar = contextOp.getIfExpr())
4218 ifCond = moduleTranslation.lookupValue(ifVar);
4219 if (grainsizeVal) {
4220 grainsize = moduleTranslation.lookupValue(grainsizeVal);
4221 sched = 1; // grainsize
4222 } else if (numTasksVal) {
4223 grainsize = moduleTranslation.lookupValue(numTasksVal);
4224 sched = 2; // num_tasks
4225 }
4226
4227 llvm::OpenMPIRBuilder::TaskDupCallbackTy taskDupOrNull = nullptr;
4228 if (taskStructMgr.getStructPtr())
4229 taskDupOrNull = taskDupCB;
4230
4231 llvm::OpenMPIRBuilder &ompBuilder = *moduleTranslation.getOpenMPBuilder();
4232 SmallVector<llvm::UncondBrInst *> cancelTerminators;
4233 // The directive to match here is OMPD_taskgroup because it is the
4234 // taskgroup which is canceled. This is handled here because it is the
4235 // task's cleanup block which should be branched to. It doesn't depend upon
4236 // nogroup because even in that case the taskloop might still be inside an
4237 // explicit taskgroup.
4238 pushCancelFinalizationCB(cancelTerminators, builder, ompBuilder, contextOp,
4239 llvm::omp::Directive::OMPD_taskgroup);
4240
4241 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4242 bool effectiveNoGroup = contextOp.getNogroup() || implicitTaskgroup;
4243 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
4244 moduleTranslation.getOpenMPBuilder()->createTaskloop(
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);
4253
4254 if (failed(handleError(afterIP, opInst)))
4255 return failure();
4256
4257 popCancelFinalizationCB(cancelTerminators, ompBuilder,
4258 afterIP->getNodeParent());
4259
4260 builder.restoreIP(*afterIP);
4261
4262 // Close the implicit taskgroup we opened for task_reduction. The end call
4263 // must execute on the encountering thread, so use the outer-scope gtid.
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);
4270 // Align OpenMPIRBuilder's internal IRBuilder with `builder` so the
4271 // gtid call lands at our insertion point.
4272 ompBuilder.updateToLocation(endLoc);
4273 llvm::Value *outerGtid = ompBuilder.getOrCreateThreadID(ident);
4274 llvm::FunctionCallee endTgFn = ompBuilder.getOrCreateRuntimeFunction(
4275 *moduleTranslation.getLLVMModule(),
4276 llvm::omp::OMPRTL___kmpc_end_taskgroup);
4277 builder.CreateCall(endTgFn, {ident, outerGtid});
4278 }
4279 return success();
4280}
4281
4282/// Build an outlined init helper for a task_reduction declare_reduction op.
4283/// Signature: void(ptr %priv, ptr %orig). For non-byref reductions, the init
4284/// region's mold argument is mapped following the same rule as the regular
4285/// reduction path (`mapInitializationArgs`): a non-pointer mold loads the
4286/// value from %orig, while a pointer-typed mold receives %orig directly. The
4287/// yielded value is stored into %priv.
4288static llvm::Function *
4289emitTaskReductionInitFn(omp::DeclareReductionOp decl, StringRef baseName,
4290 LLVM::ModuleTranslation &moduleTranslation) {
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");
4303
4304 llvm::BasicBlock *entry = llvm::BasicBlock::Create(ctx, "entry", fn);
4305 llvm::IRBuilder<> b(entry);
4306
4307 // Map the initializer's mold argument the same way the regular reduction
4308 // path does in `mapInitializationArgs`: only load the original value when a
4309 // non-pointer mold is expected. For a pointer-typed mold the storage pointer
4310 // (%orig) is passed through directly, so a mold-yielding initializer lowers
4311 // to `store ptr %orig, ptr %priv` rather than emitting a spurious load.
4312 Value moldArg = decl.getInitializerMoldArg();
4313 llvm::Value *origVal = fn->getArg(1);
4314 if (!isa<LLVM::LLVMPointerType>(moldArg.getType()))
4315 origVal = b.CreateLoad(moduleTranslation.convertType(moldArg.getType()),
4316 fn->getArg(1), "omp.orig");
4317 moduleTranslation.mapValue(moldArg, origVal);
4319 if (failed(inlineConvertOmpRegions(decl.getInitializerRegion(),
4320 "omp.taskred.init", b, moduleTranslation,
4321 &phis))) {
4322 fn->eraseFromParent();
4323 return nullptr;
4324 }
4325 assert(phis.size() == 1 &&
4326 "expected one value yielded from reduction initializer");
4327 b.CreateStore(phis[0], fn->getArg(0));
4328 b.CreateRetVoid();
4329
4330 moduleTranslation.forgetMapping(decl.getInitializerRegion());
4331 return fn;
4332}
4333
4334/// Build an outlined combiner helper for a task_reduction declare_reduction op.
4335/// Signature: void(ptr %lhs, ptr %rhs). For non-byref reductions, the values
4336/// at *%lhs and *%rhs are loaded, fed into the combiner region, and the
4337/// yielded scalar is stored back into *%lhs.
4338static llvm::Function *
4339emitTaskReductionCombFn(omp::DeclareReductionOp decl, StringRef baseName,
4340 LLVM::ModuleTranslation &moduleTranslation) {
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");
4353
4354 llvm::BasicBlock *entry = llvm::BasicBlock::Create(ctx, "entry", fn);
4355 llvm::IRBuilder<> b(entry);
4356
4357 llvm::Type *elemTy = moduleTranslation.convertType(decl.getType());
4358 Block &combBlock = decl.getReductionRegion().front();
4359 assert(combBlock.getNumArguments() == 2 &&
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");
4363 moduleTranslation.mapValue(combBlock.getArgument(0), lhsVal);
4364 moduleTranslation.mapValue(combBlock.getArgument(1), rhsVal);
4365
4367 if (failed(inlineConvertOmpRegions(decl.getReductionRegion(),
4368 "omp.taskred.comb", b, moduleTranslation,
4369 &phis))) {
4370 fn->eraseFromParent();
4371 return nullptr;
4372 }
4373 assert(phis.size() == 1 &&
4374 "expected one value yielded from reduction combiner");
4375 b.CreateStore(phis[0], fn->getArg(0));
4376 b.CreateRetVoid();
4377
4378 moduleTranslation.forgetMapping(decl.getReductionRegion());
4379 return fn;
4380}
4381
4382/// Emit the per-taskgroup task_reduction descriptor array and the
4383/// `__kmpc_taskred_init` runtime call. \p origPtrs holds the LLVM values for
4384/// the original (shared) variables, one per declaration in \p redDecls.
4385/// `builder` must be set to the point at which the descriptor stores and the
4386/// init call should be emitted; the descriptor array itself is allocated at
4387/// \p allocaIP. \p helperNamePrefix is used to disambiguate the generated
4388/// init/combiner helper symbol names between taskgroup and taskloop callers.
4389///
4390/// When \p isModifier is false, emits `__kmpc_taskred_init` and returns the
4391/// `ptr` value it produces (the taskgroup reduction handle). When \p isModifier
4392/// is true, emits `__kmpc_taskred_modifier_init` instead to open a
4393/// task-reduction scope for a parallel or worksharing construct, passing
4394/// \p isWorksharing as the runtime `is_ws` argument. Returns null on failure.
4395///
4396/// Only the non-byref form is handled here. Byref reductions have already
4397/// been rejected by `checkImplementationStatus`.
4398static llvm::Value *emitTaskReductionInitCall(
4400 ArrayRef<llvm::Value *> origPtrs, StringRef helperNamePrefix,
4401 llvm::IRBuilderBase &builder, llvm::OpenMPIRBuilder::InsertPointTy allocaIP,
4402 LLVM::ModuleTranslation &moduleTranslation, bool isModifier,
4403 bool isWorksharing) {
4404 assert(redDecls.size() == origPtrs.size() &&
4405 "expected one orig pointer per reduction decl");
4406 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
4407 llvm::Module *llvmModule = moduleTranslation.getLLVMModule();
4408 llvm::LLVMContext &ctx = llvmModule->getContext();
4409 const llvm::DataLayout &dl = llvmModule->getDataLayout();
4410
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(/*AddrSpace=*/0));
4415
4416 // Identified `kmp_taskred_input_t` struct, matching the layout used by
4417 // Clang's CGOpenMPRuntime::emitTaskReductionInit.
4418 llvm::StructType *redInputTy =
4419 llvm::StructType::getTypeByName(ctx, "kmp_taskred_input_t");
4420 if (!redInputTy)
4421 redInputTy = llvm::StructType::create(
4422 ctx, {ptrTy, ptrTy, sizeTy, ptrTy, ptrTy, ptrTy, i32Ty},
4423 "kmp_taskred_input_t");
4424
4425 unsigned n = redDecls.size();
4426 llvm::ArrayType *arrTy = llvm::ArrayType::get(redInputTy, n);
4427
4428 // Allocate the descriptor array in the enclosing function's alloca block.
4429 llvm::AllocaInst *arrAlloca;
4430 {
4431 llvm::IRBuilderBase::InsertPointGuard guard(builder);
4432 builder.restoreIP(allocaIP);
4433 arrAlloca =
4434 builder.CreateAlloca(arrTy, /*ArraySize=*/nullptr, ".taskred.input");
4435 }
4436
4437 // Fill each descriptor entry at the current builder insertion point.
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();
4447
4448 std::string baseName =
4449 (llvm::Twine(helperNamePrefix) + decl.getSymName()).str();
4450 llvm::Function *initFn =
4451 emitTaskReductionInitFn(decl, baseName, moduleTranslation);
4452 llvm::Function *combFn =
4453 emitTaskReductionCombFn(decl, baseName, moduleTranslation);
4454 if (!initFn || !combFn)
4455 return nullptr;
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);
4462 };
4463 storeField(0, orig); // reduce_shar
4464 storeField(1, orig); // reduce_orig
4465 storeField(2, llvm::ConstantInt::get(sizeTy, size)); // reduce_size
4466 storeField(3, initFn); // reduce_init
4467 storeField(4, llvm::ConstantPointerNull::get(ptrTy)); // reduce_fini
4468 storeField(5, combFn); // reduce_comb
4469 storeField(6, llvm::ConstantInt::get(i32Ty, 0)); // flags
4470 }
4471
4472 // Emit the runtime call that registers the task reduction data.
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);
4480 if (isModifier) {
4481 // __kmpc_taskred_modifier_init(loc, gtid, is_ws, num, &arr) opens a
4482 // task-reduction scope for the enclosing parallel/worksharing region.
4483 llvm::FunctionCallee modInit = ompBuilder->getOrCreateRuntimeFunction(
4484 *llvmModule, llvm::omp::OMPRTL___kmpc_taskred_modifier_init);
4485 return builder.CreateCall(modInit,
4486 {ident, gtid,
4487 builder.getInt32(isWorksharing ? 1 : 0),
4488 builder.getInt32(n), arrAlloca},
4489 ".taskred.desc");
4490 }
4491 // __kmpc_taskred_init(gtid, num, &arr).
4492 llvm::FunctionCallee taskredInit = ompBuilder->getOrCreateRuntimeFunction(
4493 *llvmModule, llvm::omp::OMPRTL___kmpc_taskred_init);
4494 return builder.CreateCall(taskredInit, {gtid, builder.getInt32(n), arrAlloca},
4495 ".taskred.desc");
4496}
4497
4498/// Emits `__kmpc_task_reduction_modifier_fini(loc, gtid, is_ws)` at the current
4499/// builder insertion point, closing the task-reduction scope opened by the
4500/// `task` reduction modifier on a parallel or worksharing construct.
4501static void
4502emitTaskReductionModifierFini(bool isWorksharing, llvm::IRBuilderBase &builder,
4503 LLVM::ModuleTranslation &moduleTranslation) {
4504 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
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)});
4517}
4518
4519/// Converts an OpenMP taskgroup construct into LLVM IR using OpenMPIRBuilder.
4520static LogicalResult
4521convertOmpTaskgroupOp(omp::TaskgroupOp tgOp, llvm::IRBuilderBase &builder,
4522 LLVM::ModuleTranslation &moduleTranslation) {
4523 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
4524 if (failed(checkImplementationStatus(*tgOp)))
4525 return failure();
4526
4527 // Resolve and validate task_reduction declarations up front. We only handle
4528 // declare_reduction ops shaped like a non-byref scalar reduction in this
4529 // first cut; richer shapes (two-argument initializer, cleanup region,
4530 // missing combiner) require additional infrastructure.
4532 if (auto syms = tgOp.getTaskReductionSyms()) {
4533 redDecls.reserve(syms->size());
4534 for (auto sym : syms->getAsRange<SymbolRefAttr>()) {
4536 tgOp, sym);
4537 if (!decl)
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 "
4549 "combiner region");
4550 redDecls.push_back(decl);
4551 }
4552 }
4553
4554 auto bodyCB =
4555 [&](InsertPointTy allocaIP, InsertPointTy codegenIP,
4556 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) -> llvm::Error {
4557 builder.restoreIP(codegenIP);
4558
4559 if (!redDecls.empty()) {
4561 origPtrs.reserve(redDecls.size());
4562 for (Value v : tgOp.getTaskReductionVars())
4563 origPtrs.push_back(moduleTranslation.lookupValue(v));
4564 if (!emitTaskReductionInitCall(redDecls, origPtrs, "__omp_taskred_",
4565 builder, allocaIP, moduleTranslation))
4566 return llvm::createStringError(
4567 llvm::inconvertibleErrorCode(),
4568 "failed to emit task_reduction initialization for omp.taskgroup");
4569 }
4570
4571 // Inside the taskgroup body, each task_reduction block argument refers to
4572 // the same shared/original storage that the runtime now knows about via
4573 // the descriptor array. Inner tasks that declare in_reduction look up
4574 // per-task private copies through the runtime; the taskgroup body itself
4575 // uses the original variable.
4576 for (auto [i, blockArg] :
4577 llvm::enumerate(tgOp.getRegion().getArguments())) {
4578 llvm::Value *orig =
4579 moduleTranslation.lookupValue(tgOp.getTaskReductionVars()[i]);
4580 moduleTranslation.mapValue(blockArg, orig);
4581 }
4582
4583 return convertOmpOpRegions(tgOp.getRegion(), "omp.taskgroup.region",
4584 builder, moduleTranslation)
4585 .takeError();
4586 };
4587
4589 InsertPointTy allocaIP =
4590 findAllocInsertPoints(builder, moduleTranslation, &deallocBlocks);
4591 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4592 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
4593 moduleTranslation.getOpenMPBuilder()->createTaskgroup(
4594 ompLoc, allocaIP, deallocBlocks, bodyCB);
4595
4596 if (failed(handleError(afterIP, *tgOp)))
4597 return failure();
4598
4599 builder.restoreIP(*afterIP);
4600 return success();
4601}
4602
4603static LogicalResult
4604convertOmpInteropInitOp(omp::InteropInitOp initOp, llvm::IRBuilderBase &builder,
4605 LLVM::ModuleTranslation &moduleTranslation) {
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";
4611
4612 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
4613 llvm::Value *interopVar =
4614 moduleTranslation.lookupValue(initOp.getInteropVar());
4615 llvm::Value *device = initOp.getDevice()
4616 ? moduleTranslation.lookupValue(initOp.getDevice())
4617 : nullptr;
4618
4619 // TODO: Handle depend clauses when supported.
4620 llvm::Value *numDeps = llvm::ConstantInt::get(builder.getInt32Ty(), 0);
4621 llvm::Value *depArray = llvm::ConstantPointerNull::get(builder.getPtrTy());
4622 bool hasNowait = initOp.getNowait();
4623
4624 // A single `init` clause may list both `target` and `targetsync`, but the
4625 // runtime init call takes a single interop-type. Collapse the set to one
4626 // value, matching Clang: if `target` is present use Target, otherwise
4627 // TargetSync. The offload runtime object model supports only one type per
4628 // object; representing both would require a runtime change.
4629 bool hasTarget = false, hasTargetSync = false;
4630 for (mlir::Attribute typeAttr : initOp.getInteropTypes()) {
4631 switch (cast<omp::InteropTypeAttr>(typeAttr).getValue()) {
4632 case omp::InteropType::target:
4633 hasTarget = true;
4634 break;
4635 case omp::InteropType::targetsync:
4636 hasTargetSync = true;
4637 break;
4638 }
4639 }
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);
4645 return success();
4646}
4647
4648static LogicalResult
4649convertOmpInteropDestroyOp(omp::InteropDestroyOp destroyOp,
4650 llvm::IRBuilderBase &builder,
4651 LLVM::ModuleTranslation &moduleTranslation) {
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";
4658
4659 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
4660 llvm::Value *interopVar =
4661 moduleTranslation.lookupValue(destroyOp.getInteropVar());
4662 llvm::Value *device =
4663 destroyOp.getDevice()
4664 ? moduleTranslation.lookupValue(destroyOp.getDevice())
4665 : nullptr;
4666
4667 llvm::Value *numDeps = llvm::ConstantInt::get(builder.getInt32Ty(), 0);
4668 llvm::Value *depArray = llvm::ConstantPointerNull::get(builder.getPtrTy());
4669 bool hasNowait = destroyOp.getNowait();
4670
4671 ompBuilder->createOMPInteropDestroy(builder, interopVar, device, numDeps,
4672 depArray, hasNowait);
4673 return success();
4674}
4675
4676static LogicalResult
4677convertOmpInteropUseOp(omp::InteropUseOp useOp, llvm::IRBuilderBase &builder,
4678 LLVM::ModuleTranslation &moduleTranslation) {
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";
4684
4685 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
4686 llvm::Value *interopVar =
4687 moduleTranslation.lookupValue(useOp.getInteropVar());
4688 llvm::Value *device = useOp.getDevice()
4689 ? moduleTranslation.lookupValue(useOp.getDevice())
4690 : nullptr;
4691
4692 llvm::Value *numDeps = llvm::ConstantInt::get(builder.getInt32Ty(), 0);
4693 llvm::Value *depArray = llvm::ConstantPointerNull::get(builder.getPtrTy());
4694 bool hasNowait = useOp.getNowait();
4695
4696 ompBuilder->createOMPInteropUse(builder, interopVar, device, numDeps,
4697 depArray, hasNowait);
4698 return success();
4699}
4700
4701static LogicalResult
4702convertOmpTaskwaitOp(omp::TaskwaitOp twOp, llvm::IRBuilderBase &builder,
4703 LLVM::ModuleTranslation &moduleTranslation) {
4704 if (failed(checkImplementationStatus(*twOp)))
4705 return failure();
4706
4707 llvm::OpenMPIRBuilder::DependenciesInfo dds;
4708 if (failed(buildDependData(
4709 twOp.getDependVars(), twOp.getDependKinds(), twOp.getDependIterated(),
4710 twOp.getDependIteratedKinds(), builder, moduleTranslation, dds))) {
4711 return failure();
4712 }
4713
4714 moduleTranslation.getOpenMPBuilder()->createTaskwait(builder, dds,
4715 twOp.getNowait());
4716 if (dds.DepArray) {
4717 builder.CreateFree(dds.DepArray);
4718 }
4719
4720 return success();
4721}
4722
4723/// Converts an OpenMP workshare loop into LLVM IR using OpenMPIRBuilder.
4724static LogicalResult
4725convertOmpWsloop(Operation &opInst, llvm::IRBuilderBase &builder,
4726 LLVM::ModuleTranslation &moduleTranslation) {
4727 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
4728 auto wsloopOp = cast<omp::WsloopOp>(opInst);
4729 if (failed(checkImplementationStatus(opInst)))
4730 return failure();
4731
4732 auto loopOp = cast<omp::LoopNestOp>(wsloopOp.getWrappedLoop());
4733 llvm::ArrayRef<bool> isByRef = getIsByRef(wsloopOp.getReductionByref());
4734 assert(isByRef.size() == wsloopOp.getNumReductionVars());
4735
4736 // Static is the default.
4737 auto schedule =
4738 wsloopOp.getScheduleKind().value_or(omp::ClauseScheduleKind::Static);
4739
4740 // Find the loop configuration.
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);
4748 }
4749
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);
4760 }
4761 }
4762
4763 PrivateVarsInfo privateVarsInfo(wsloopOp);
4764
4766 collectReductionDecls(wsloopOp, reductionDecls);
4767
4768 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
4769 findAllocInsertPoints(builder, moduleTranslation);
4770
4771 SmallVector<llvm::Value *> privateReductionVariables(
4772 wsloopOp.getNumReductionVars());
4773
4775 wsloopOp, builder, moduleTranslation, privateVarsInfo, allocaIP);
4776 if (handleError(afterAllocas, opInst).failed())
4777 return failure();
4778
4779 DenseMap<Value, llvm::Value *> reductionVariableMap;
4780
4781 MutableArrayRef<BlockArgument> reductionArgs =
4782 cast<omp::BlockArgOpenMPOpInterface>(opInst).getReductionBlockArgs();
4783
4784 SmallVector<DeferredStore> deferredStores;
4785
4786 if (failed(allocReductionVars(wsloopOp, reductionArgs, builder,
4787 moduleTranslation, allocaIP, reductionDecls,
4788 privateReductionVariables, reductionVariableMap,
4789 deferredStores, isByRef)))
4790 return failure();
4791
4792 if (handleError(initPrivateVars(builder, moduleTranslation, privateVarsInfo),
4793 opInst)
4794 .failed())
4795 return failure();
4796
4797 if (failed(copyFirstPrivateVars(
4798 wsloopOp, builder, moduleTranslation, privateVarsInfo.mlirVars,
4799 privateVarsInfo.llvmVars, privateVarsInfo.privatizers,
4800 wsloopOp.getPrivateNeedsBarrier())))
4801 return failure();
4802
4803 assert(afterAllocas.get()->getSinglePredecessor());
4804 if (failed(initReductionVars(wsloopOp, reductionArgs, builder,
4805 moduleTranslation,
4806 afterAllocas.get()->getSinglePredecessor(),
4807 reductionDecls, privateReductionVariables,
4808 reductionVariableMap, isByRef, deferredStores)))
4809 return failure();
4810
4811 // For `reduction(task, ...)` open a task-reduction scope for the worksharing
4812 // loop. Participating explicit tasks accumulate into the per-thread private
4813 // copies, which the worksharing reduction then combines across threads.
4814 bool isTaskReductionMod =
4815 wsloopOp.getReductionMod() == omp::ReductionModifier::task &&
4816 wsloopOp.getNumReductionVars() > 0;
4817 if (isTaskReductionMod &&
4818 !emitTaskReductionInitCall(reductionDecls, privateReductionVariables,
4819 "__omp_taskred_mod_", builder, allocaIP,
4820 moduleTranslation, /*isModifier=*/true,
4821 /*isWorksharing=*/true))
4822 return wsloopOp.emitError(
4823 "failed to emit task reduction modifier initialization");
4824
4825 // TODO: Handle doacross loops when the ordered clause has a parameter.
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();
4830
4831 // The only legal way for the direct parent to be omp.distribute is that this
4832 // represents 'distribute parallel do'. Otherwise, this is a regular
4833 // worksharing loop.
4834 llvm::omp::WorksharingLoopType workshareLoopType =
4835 llvm::isa_and_present<omp::DistributeOp>(opInst.getParentOp())
4836 ? llvm::omp::WorksharingLoopType::DistributeForStaticLoop
4837 : llvm::omp::WorksharingLoopType::ForStaticLoop;
4838
4839 SmallVector<llvm::UncondBrInst *> cancelTerminators;
4840 pushCancelFinalizationCB(cancelTerminators, builder, *ompBuilder, wsloopOp,
4841 llvm::omp::Directive::OMPD_for);
4842
4843 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
4844
4845 // Initialize linear variables and linear step
4846 LinearClauseProcessor linearClauseProcessor;
4847
4848 if (!wsloopOp.getLinearVars().empty()) {
4849 auto linearVarTypes = wsloopOp.getLinearVarTypes().value();
4850 for (mlir::Attribute linearVarType : linearVarTypes)
4851 linearClauseProcessor.registerType(moduleTranslation, linearVarType);
4852
4853 for (auto [idx, linearVar] : llvm::enumerate(wsloopOp.getLinearVars()))
4854 linearClauseProcessor.createLinearVar(
4855 builder, moduleTranslation, moduleTranslation.lookupValue(linearVar),
4856 idx);
4857 for (mlir::Value linearStep : wsloopOp.getLinearStepVars())
4858 linearClauseProcessor.initLinearStep(moduleTranslation, linearStep);
4859 }
4860
4862 wsloopOp.getRegion(), "omp.wsloop.region", builder, moduleTranslation);
4863
4864 if (failed(handleError(regionBlock, opInst)))
4865 return failure();
4866
4867 llvm::CanonicalLoopInfo *loopInfo = findCurrentLoopInfo(moduleTranslation);
4868
4869 // Emit Initialization and Update IR for linear variables
4870 if (!wsloopOp.getLinearVars().empty()) {
4871 linearClauseProcessor.initLinearVar(builder, moduleTranslation,
4872 loopInfo->getPreheader());
4873 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterBarrierIP =
4874 moduleTranslation.getOpenMPBuilder()->createBarrier(
4875 builder, llvm::omp::OMPD_barrier);
4876 if (failed(handleError(afterBarrierIP, *loopOp)))
4877 return failure();
4878 builder.restoreIP(*afterBarrierIP);
4879 linearClauseProcessor.updateLinearVar(builder, loopInfo->getBody(),
4880 loopInfo->getIndVar());
4881 linearClauseProcessor.splitLinearFiniBB(builder, loopInfo->getExit());
4882 }
4883
4884 builder.SetInsertPoint((*regionBlock)->begin());
4885
4886 // Check if we can generate no-loop kernel
4887 bool noLoopMode = false;
4888 omp::TargetOp targetOp = wsloopOp->getParentOfType<mlir::omp::TargetOp>();
4889 if (targetOp &&
4890 targetOp.getKernelType() == omp::TargetExecMode::spmd_no_loop) {
4891 Operation *targetCapturedOp =
4892 cast<omp::ComposableOpInterface>(*targetOp).findCapturedOp();
4893 // We need this check because, without it, noLoopMode would be set to true
4894 // for every omp.wsloop nested inside a no-loop SPMD target region, even if
4895 // that loop is not the top-level SPMD one.
4896 if (loopOp == targetCapturedOp)
4897 noLoopMode = true;
4898 }
4899
4900 for (size_t index = 0; index < wsloopOp.getLinearVars().size(); index++)
4901 linearClauseProcessor.rewriteInPlace(builder, loopInfo->getBody(),
4902 loopInfo->getLatch(), index);
4903
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);
4911
4912 if (failed(handleError(wsloopIP, opInst)))
4913 return failure();
4914
4915 // Save the continuation block before linear var finalization, which may
4916 // invalidate the wsloopIP iterator.
4917 llvm::BasicBlock *wsloopContinuationBB = wsloopIP->getNodeParent();
4918
4919 // Emit finalization and in-place rewrites for linear vars.
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());
4927 if (failed(handleError(afterBarrierIP, *loopOp)))
4928 return failure();
4929
4930 builder.restoreIP(oldIP);
4931 }
4932
4933 // Set the correct branch target for task cancellation
4934 popCancelFinalizationCB(cancelTerminators, *ompBuilder, wsloopContinuationBB);
4935
4936 // Close the task-reduction scope before the worksharing reduction combine.
4937 if (isTaskReductionMod)
4938 emitTaskReductionModifierFini(/*isWorksharing=*/true, builder,
4939 moduleTranslation);
4940
4941 // Process the reductions if required.
4942 if (failed(createReductionsAndCleanup(
4943 wsloopOp, builder, moduleTranslation, allocaIP, reductionDecls,
4944 privateReductionVariables, isByRef, wsloopOp.getNowait(),
4945 /*isTeamsReduction=*/false)))
4946 return failure();
4947
4948 return cleanupPrivateVars(wsloopOp, builder, moduleTranslation,
4949 wsloopOp.getLoc(), privateVarsInfo);
4950}
4951
4952/// Converts the OpenMP parallel operation to LLVM IR.
4953static LogicalResult
4954convertOmpParallel(omp::ParallelOp opInst, llvm::IRBuilderBase &builder,
4955 LLVM::ModuleTranslation &moduleTranslation) {
4956 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
4957 ArrayRef<bool> isByRef = getIsByRef(opInst.getReductionByref());
4958 assert(isByRef.size() == opInst.getNumReductionVars());
4959 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
4960 bool isCancellable = constructIsCancellable(opInst);
4961
4962 if (failed(checkImplementationStatus(*opInst)))
4963 return failure();
4964
4965 PrivateVarsInfo privateVarsInfo(opInst);
4966 if (failed(convertAllocatorVars(*opInst, opInst.getAllocatorVars(), builder,
4967 moduleTranslation, privateVarsInfo)))
4968 return failure();
4969
4970 // Collect reduction declarations
4972 collectReductionDecls(opInst, reductionDecls);
4973 SmallVector<llvm::Value *> privateReductionVariables(
4974 opInst.getNumReductionVars());
4975 SmallVector<DeferredStore> deferredStores;
4976 // Only open a task-reduction scope when the `task` modifier is present and
4977 // there are reduction variables to combine; otherwise the matching fini in
4978 // the reduction-combine path (guarded by getNumReductionVars() > 0) would be
4979 // skipped, leaving the modifier init unbalanced.
4980 bool isTaskReductionMod =
4981 opInst.getReductionMod() == omp::ReductionModifier::task &&
4982 opInst.getNumReductionVars() > 0;
4983
4984 auto bodyGenCB =
4985 [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
4986 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) -> llvm::Error {
4988 opInst, builder, moduleTranslation, privateVarsInfo, allocaIP);
4989 if (handleError(afterAllocas, *opInst).failed())
4990 return llvm::make_error<PreviouslyReportedError>();
4991
4992 // Allocate reduction vars
4993 DenseMap<Value, llvm::Value *> reductionVariableMap;
4994
4995 MutableArrayRef<BlockArgument> reductionArgs =
4996 cast<omp::BlockArgOpenMPOpInterface>(*opInst).getReductionBlockArgs();
4997
4998 allocaIP = allocaIP.getNodeParent()->getTerminator()->getIterator();
4999
5000 if (failed(allocReductionVars(
5001 opInst, reductionArgs, builder, moduleTranslation, allocaIP,
5002 reductionDecls, privateReductionVariables, reductionVariableMap,
5003 deferredStores, isByRef)))
5004 return llvm::make_error<PreviouslyReportedError>();
5005
5006 assert(afterAllocas.get()->getSinglePredecessor());
5007 builder.restoreIP(codeGenIP);
5008
5009 if (handleError(
5010 initPrivateVars(builder, moduleTranslation, privateVarsInfo),
5011 *opInst)
5012 .failed())
5013 return llvm::make_error<PreviouslyReportedError>();
5014
5015 if (failed(copyFirstPrivateVars(
5016 opInst, builder, moduleTranslation, privateVarsInfo.mlirVars,
5017 privateVarsInfo.llvmVars, privateVarsInfo.privatizers,
5018 opInst.getPrivateNeedsBarrier())))
5019 return llvm::make_error<PreviouslyReportedError>();
5020
5021 if (failed(
5022 initReductionVars(opInst, reductionArgs, builder, moduleTranslation,
5023 afterAllocas.get()->getSinglePredecessor(),
5024 reductionDecls, privateReductionVariables,
5025 reductionVariableMap, isByRef, deferredStores)))
5026 return llvm::make_error<PreviouslyReportedError>();
5027
5028 // For `reduction(task, ...)` open a task-reduction scope so participating
5029 // explicit tasks accumulate into the per-thread private copies; the
5030 // parallel reduction then combines those copies across the team.
5031 if (isTaskReductionMod &&
5032 !emitTaskReductionInitCall(reductionDecls, privateReductionVariables,
5033 "__omp_taskred_mod_", builder, allocaIP,
5034 moduleTranslation, /*isModifier=*/true,
5035 /*isWorksharing=*/false))
5036 return llvm::createStringError(
5037 "failed to emit task reduction modifier initialization");
5038
5039 // Save the alloca insertion point on ModuleTranslation stack for use in
5040 // nested regions.
5042 moduleTranslation, allocaIP, deallocBlocks);
5043
5044 // ParallelOp has only one region associated with it.
5046 opInst.getRegion(), "omp.par.region", builder, moduleTranslation);
5047 if (!regionBlock)
5048 return regionBlock.takeError();
5049
5050 // Process the reductions if required.
5051 if (opInst.getNumReductionVars() > 0) {
5052 // Collect reduction info
5054 SmallVector<OwningAtomicReductionGen> owningAtomicReductionGens;
5056 owningReductionGenRefDataPtrGens;
5058 collectReductionInfo(opInst, builder, moduleTranslation, reductionDecls,
5059 owningReductionGens, owningAtomicReductionGens,
5060 owningReductionGenRefDataPtrGens,
5061 privateReductionVariables, reductionInfos, isByRef);
5062
5063 // Move to region cont block
5064 builder.SetInsertPoint((*regionBlock)->getTerminator());
5065
5066 // Close the task-reduction scope before the per-thread reduction
5067 // contributions are combined across the team.
5068 if (isTaskReductionMod)
5069 emitTaskReductionModifierFini(/*isWorksharing=*/false, builder,
5070 moduleTranslation);
5071
5072 // Generate reductions from info
5073 llvm::UnreachableInst *tempTerminator = builder.CreateUnreachable();
5074 builder.SetInsertPoint(tempTerminator);
5075
5076 llvm::OpenMPIRBuilder::InsertPointOrErrorTy contInsertPoint =
5077 ompBuilder->createReductions(builder, allocaIP, reductionInfos,
5078 isByRef,
5079 /*IsNoWait=*/false,
5080 /*IsTeamsReduction=*/false);
5081 if (!contInsertPoint)
5082 return contInsertPoint.takeError();
5083
5084 if (!contInsertPoint->isValid())
5085 return llvm::make_error<PreviouslyReportedError>();
5086
5087 tempTerminator->eraseFromParent();
5088 builder.restoreIP(*contInsertPoint);
5089 }
5090
5091 return llvm::Error::success();
5092 };
5093
5094 auto privCB = [](InsertPointTy allocaIP, InsertPointTy codeGenIP,
5095 llvm::Value &, llvm::Value &val, llvm::Value *&replVal) {
5096 // tell OpenMPIRBuilder not to do anything. We handled Privatisation in
5097 // bodyGenCB.
5098 replVal = &val;
5099 return codeGenIP;
5100 };
5101
5102 // TODO: Perform finalization actions for variables. This has to be
5103 // called for variables which have destructors/finalizers.
5104 auto finiCB = [&](InsertPointTy codeGenIP) -> llvm::Error {
5105 InsertPointTy oldIP = builder.saveIP();
5106 builder.restoreIP(codeGenIP);
5107
5108 // if the reduction has a cleanup region, inline it here to finalize the
5109 // reduction variables
5110 SmallVector<Region *> reductionCleanupRegions;
5111 llvm::transform(reductionDecls, std::back_inserter(reductionCleanupRegions),
5112 [](omp::DeclareReductionOp reductionDecl) {
5113 return &reductionDecl.getCleanupRegion();
5114 });
5115 if (failed(inlineOmpRegionCleanup(
5116 reductionCleanupRegions, privateReductionVariables,
5117 moduleTranslation, builder, "omp.reduction.cleanup")))
5118 return llvm::createStringError(
5119 "failed to inline `cleanup` region of `omp.declare_reduction`");
5120
5121 if (failed(cleanupPrivateVars(opInst, builder, moduleTranslation,
5122 opInst.getLoc(), privateVarsInfo)))
5123 return llvm::make_error<PreviouslyReportedError>();
5124
5125 // If we could be performing cancellation, add the cancellation barrier on
5126 // the way out of the outlined region.
5127 if (isCancellable) {
5128 auto IPOrErr = ompBuilder->createBarrier(
5129 llvm::OpenMPIRBuilder::LocationDescription(builder),
5130 llvm::omp::Directive::OMPD_unknown,
5131 /* ForceSimpleCall */ false,
5132 /* CheckCancelFlag */ false);
5133 if (!IPOrErr)
5134 return IPOrErr.takeError();
5135 }
5136
5137 builder.restoreIP(oldIP);
5138 return llvm::Error::success();
5139 };
5140
5141 llvm::Value *ifCond = nullptr;
5142 if (auto ifVar = opInst.getIfExpr())
5143 ifCond = moduleTranslation.lookupValue(ifVar);
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())
5149 pbKind = getProcBindKind(*bind);
5150
5152 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5153 findAllocInsertPoints(builder, moduleTranslation, &deallocBlocks);
5154 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5155
5156 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
5157 ompBuilder->createParallel(ompLoc, allocaIP, deallocBlocks, bodyGenCB,
5158 privCB, finiCB, ifCond, numThreads, pbKind,
5159 isCancellable);
5160
5161 if (failed(handleError(afterIP, *opInst)))
5162 return failure();
5163
5164 builder.restoreIP(*afterIP);
5165 return success();
5166}
5167
5168/// Convert Order attribute to llvm::omp::OrderKind.
5169static llvm::omp::OrderKind
5170convertOrderKind(std::optional<omp::ClauseOrderKind> o) {
5171 if (!o)
5172 return llvm::omp::OrderKind::OMP_ORDER_unknown;
5173 switch (*o) {
5174 case omp::ClauseOrderKind::Concurrent:
5175 return llvm::omp::OrderKind::OMP_ORDER_concurrent;
5176 }
5177 llvm_unreachable("Unknown ClauseOrderKind kind");
5178}
5179
5180/// Converts an OpenMP simd loop into LLVM IR using OpenMPIRBuilder.
5181static LogicalResult
5182convertOmpSimd(Operation &opInst, llvm::IRBuilderBase &builder,
5183 LLVM::ModuleTranslation &moduleTranslation) {
5184 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5185 auto simdOp = cast<omp::SimdOp>(opInst);
5186
5187 if (failed(checkImplementationStatus(opInst)))
5188 return failure();
5189
5190 PrivateVarsInfo privateVarsInfo(simdOp);
5191
5192 MutableArrayRef<BlockArgument> reductionArgs =
5193 cast<omp::BlockArgOpenMPOpInterface>(opInst).getReductionBlockArgs();
5194 DenseMap<Value, llvm::Value *> reductionVariableMap;
5195 SmallVector<llvm::Value *> privateReductionVariables(
5196 simdOp.getNumReductionVars());
5197 SmallVector<DeferredStore> deferredStores;
5199 collectReductionDecls(simdOp, reductionDecls);
5200 llvm::ArrayRef<bool> isByRef = getIsByRef(simdOp.getReductionByref());
5201 assert(isByRef.size() == simdOp.getNumReductionVars());
5202
5203 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5204 findAllocInsertPoints(builder, moduleTranslation);
5205
5207 simdOp, builder, moduleTranslation, privateVarsInfo, allocaIP);
5208 if (handleError(afterAllocas, opInst).failed())
5209 return failure();
5210
5211 // Initialize linear variables and linear step
5212 LinearClauseProcessor linearClauseProcessor;
5213 if (linearClauseProcessor.initLinearIV(simdOp).failed())
5214 return failure();
5215
5216 if (!simdOp.getLinearVars().empty()) {
5217 auto linearVarTypes = simdOp.getLinearVarTypes().value();
5218 for (mlir::Attribute linearVarType : linearVarTypes)
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(
5223 privateVarsInfo.mlirVars, privateVarsInfo.llvmVars)) {
5224 // If the linear variable is implicit, reuse the already
5225 // existing llvm::Value
5226 if (linearVar == mlirPrivVar) {
5227 isImplicit = true;
5228 linearClauseProcessor.createLinearVar(builder, moduleTranslation,
5229 llvmPrivateVar, idx);
5230 break;
5231 }
5232 }
5233
5234 if (!isImplicit)
5235 linearClauseProcessor.createLinearVar(
5236 builder, moduleTranslation,
5237 moduleTranslation.lookupValue(linearVar), idx);
5238 }
5239 for (mlir::Value linearStep : simdOp.getLinearStepVars())
5240 linearClauseProcessor.initLinearStep(moduleTranslation, linearStep);
5241 }
5242
5243 if (failed(allocReductionVars(simdOp, reductionArgs, builder,
5244 moduleTranslation, allocaIP, reductionDecls,
5245 privateReductionVariables, reductionVariableMap,
5246 deferredStores, isByRef)))
5247 return failure();
5248
5249 if (handleError(initPrivateVars(builder, moduleTranslation, privateVarsInfo),
5250 opInst)
5251 .failed())
5252 return failure();
5253
5254 // No call to copyFirstPrivateVars because FIRSTPRIVATE is not allowed for
5255 // SIMD.
5256
5257 assert(afterAllocas.get()->getSinglePredecessor());
5258 if (failed(initReductionVars(simdOp, reductionArgs, builder,
5259 moduleTranslation,
5260 afterAllocas.get()->getSinglePredecessor(),
5261 reductionDecls, privateReductionVariables,
5262 reductionVariableMap, isByRef, deferredStores)))
5263 return failure();
5264
5265 llvm::ConstantInt *simdlen = nullptr;
5266 if (std::optional<uint64_t> simdlenVar = simdOp.getSimdlen())
5267 simdlen = builder.getInt64(simdlenVar.value());
5268
5269 llvm::ConstantInt *safelen = nullptr;
5270 if (std::optional<uint64_t> safelenVar = simdOp.getSafelen())
5271 safelen = builder.getInt64(safelenVar.value());
5272
5273 llvm::MapVector<llvm::Value *, llvm::Value *> alignedVars;
5274 llvm::omp::OrderKind order = convertOrderKind(simdOp.getOrder());
5275
5276 llvm::BasicBlock *sourceBlock = builder.GetInsertBlock();
5277 std::optional<ArrayAttr> alignmentValues = simdOp.getAlignments();
5278 mlir::OperandRange operands = simdOp.getAlignedVars();
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();
5283
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");
5288
5289 // Check if the alignment value is not a power of 2. If so, skip emitting
5290 // alignment.
5291 if (!intAttr.getValue().isPowerOf2())
5292 continue;
5293
5294 auto curInsert = builder.saveIP();
5295 builder.SetInsertPoint(sourceBlock);
5296 llvmVal = builder.CreateLoad(ty, llvmVal);
5297 builder.restoreIP(curInsert);
5298 alignedVars[llvmVal] = alignment;
5299 }
5300
5302 simdOp.getRegion(), "omp.simd.region", builder, moduleTranslation);
5303
5304 if (failed(handleError(regionBlock, opInst)))
5305 return failure();
5306
5307 llvm::CanonicalLoopInfo *loopInfo = findCurrentLoopInfo(moduleTranslation);
5308 // Emit Initialization for linear variables
5309 if (simdOp.getLinearVars().size()) {
5310 linearClauseProcessor.initLinearVar(builder, moduleTranslation,
5311 loopInfo->getPreheader());
5312
5313 linearClauseProcessor.updateLinearVar(builder, loopInfo->getBody(),
5314 loopInfo->getIndVar());
5315 }
5316 builder.SetInsertPoint((*regionBlock)->begin());
5317
5318 for (size_t index = 0; index < simdOp.getLinearVars().size(); index++)
5319 linearClauseProcessor.rewriteInPlace(builder, loopInfo->getBody(),
5320 loopInfo->getLatch(), index);
5321
5322 ompBuilder->applySimd(loopInfo, alignedVars,
5323 simdOp.getIfExpr()
5324 ? moduleTranslation.lookupValue(simdOp.getIfExpr())
5325 : nullptr,
5326 order, simdlen, safelen);
5327
5328 linearClauseProcessor.updateLinearIV(builder, moduleTranslation);
5329 linearClauseProcessor.emitStoresForLinearVar(builder);
5330
5331 // We now need to reduce the per-simd-lane reduction variable into the
5332 // original variable. This works a bit differently to other reductions (e.g.
5333 // wsloop) because we don't need to call into the OpenMP runtime to handle
5334 // threads: everything happened in this one thread.
5335 for (auto [i, tuple] : llvm::enumerate(
5336 llvm::zip(reductionDecls, isByRef, simdOp.getReductionVars(),
5337 privateReductionVariables))) {
5338 auto [decl, byRef, reductionVar, privateReductionVar] = tuple;
5339
5340 OwningReductionGen gen = makeReductionGen(decl, builder, moduleTranslation);
5341 llvm::Value *originalVariable = moduleTranslation.lookupValue(reductionVar);
5342 llvm::Type *reductionType = moduleTranslation.convertType(decl.getType());
5343
5344 // We have one less load for by-ref case because that load is now inside of
5345 // the reduction region.
5346 llvm::Value *redValue = originalVariable;
5347 if (!byRef)
5348 redValue =
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;
5353
5354 auto res = gen(builder.saveIP(), redValue, privateRedValue, reduced);
5355 if (failed(handleError(res, opInst)))
5356 return failure();
5357 builder.restoreIP(res.get());
5358
5359 // For by-ref case, the store is inside of the reduction region.
5360 if (!byRef)
5361 builder.CreateStore(reduced, originalVariable);
5362 }
5363
5364 // After the construct, deallocate private reduction variables.
5365 SmallVector<Region *> reductionRegions;
5366 llvm::transform(reductionDecls, std::back_inserter(reductionRegions),
5367 [](omp::DeclareReductionOp reductionDecl) {
5368 return &reductionDecl.getCleanupRegion();
5369 });
5370 if (failed(inlineOmpRegionCleanup(reductionRegions, privateReductionVariables,
5371 moduleTranslation, builder,
5372 "omp.reduction.cleanup")))
5373 return failure();
5374
5375 return cleanupPrivateVars(simdOp, builder, moduleTranslation, simdOp.getLoc(),
5376 privateVarsInfo);
5377}
5378
5379/// Converts an OpenMP loop nest into LLVM IR using OpenMPIRBuilder.
5380static LogicalResult
5381convertOmpLoopNest(Operation &opInst, llvm::IRBuilderBase &builder,
5382 LLVM::ModuleTranslation &moduleTranslation) {
5383 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5384 auto loopOp = cast<omp::LoopNestOp>(opInst);
5385
5386 if (failed(checkImplementationStatus(opInst)))
5387 return failure();
5388
5389 // Set up the source location value for OpenMP runtime.
5390 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5391
5392 // Generator of the canonical loop body.
5395 auto bodyGen = [&](llvm::OpenMPIRBuilder::InsertPointTy ip,
5396 llvm::Value *iv) -> llvm::Error {
5397 // Make sure further conversions know about the induction variable.
5398 moduleTranslation.mapValue(
5399 loopOp.getRegion().front().getArgument(loopInfos.size()), iv);
5400
5401 // Capture the body insertion point for use in nested loops. BodyIP of the
5402 // CanonicalLoopInfo always points to the beginning of the entry block of
5403 // the body.
5404 bodyInsertPoints.push_back(ip);
5405
5406 if (loopInfos.size() != loopOp.getNumLoops() - 1)
5407 return llvm::Error::success();
5408
5409 // Convert the body of the loop.
5410 builder.restoreIP(ip);
5412 loopOp.getRegion(), "omp.loop_nest.region", builder, moduleTranslation);
5413 if (!regionBlock)
5414 return regionBlock.takeError();
5415
5416 builder.SetInsertPoint((*regionBlock)->begin());
5417 return llvm::Error::success();
5418 };
5419
5420 // Delegate actual loop construction to the OpenMP IRBuilder.
5421 // TODO: this currently assumes omp.loop_nest is semantically similar to SCF
5422 // loop, i.e. it has a positive step, uses signed integer semantics.
5423 // Reconsider this code when the nested loop operation clearly supports more
5424 // cases.
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]);
5431
5432 // Make sure loop trip count are emitted in the preheader of the outermost
5433 // loop at the latest so that they are all available for the new collapsed
5434 // loop will be created below.
5435 llvm::OpenMPIRBuilder::LocationDescription loc = ompLoc;
5436 llvm::OpenMPIRBuilder::InsertPointTy computeIP = ompLoc.IP;
5437 if (i != 0) {
5438 loc = llvm::OpenMPIRBuilder::LocationDescription(bodyInsertPoints.back(),
5439 ompLoc.DL);
5440 computeIP = loopInfos.front()->getPreheaderIP();
5441 }
5442
5444 ompBuilder->createCanonicalLoop(
5445 loc, bodyGen, lowerBound, upperBound, step,
5446 /*IsSigned=*/true, loopOp.getLoopInclusive(), computeIP);
5447
5448 if (failed(handleError(loopResult, *loopOp)))
5449 return failure();
5450
5451 loopInfos.push_back(*loopResult);
5452 }
5453
5454 llvm::OpenMPIRBuilder::InsertPointTy afterIP =
5455 loopInfos.front()->getAfterIP();
5456
5457 // Do tiling.
5458 if (const auto &tiles = loopOp.getTileSizes()) {
5459 llvm::Type *ivType = loopInfos.front()->getIndVarType();
5461
5462 for (auto tile : tiles.value()) {
5463 llvm::Value *tileVal = llvm::ConstantInt::get(ivType, tile);
5464 tileSizes.push_back(tileVal);
5465 }
5466
5467 std::vector<llvm::CanonicalLoopInfo *> newLoops =
5468 ompBuilder->tileLoops(ompLoc.DL, loopInfos, tileSizes);
5469
5470 // Update afterIP to get the correct insertion point after
5471 // tiling.
5472 llvm::BasicBlock *afterBB = newLoops.front()->getAfter();
5473 llvm::BasicBlock *afterAfterBB = afterBB->getSingleSuccessor();
5474 afterIP = afterAfterBB->begin();
5475
5476 // Update the loop infos.
5477 loopInfos.clear();
5478 for (const auto &newLoop : newLoops)
5479 loopInfos.push_back(newLoop);
5480 } // Tiling done.
5481
5482 // Do collapse.
5483 const auto &numCollapse = loopOp.getCollapseNumLoops();
5485 loopInfos.begin(), loopInfos.begin() + (numCollapse));
5486
5487 auto newTopLoopInfo =
5488 ompBuilder->collapseLoops(ompLoc.DL, collapseLoopInfos, {});
5489
5490 assert(newTopLoopInfo && "New top loop information is missing");
5491 moduleTranslation.stackWalk<OpenMPLoopInfoStackFrame>(
5492 [&](OpenMPLoopInfoStackFrame &frame) {
5493 frame.loopInfo = newTopLoopInfo;
5494 return WalkResult::interrupt();
5495 });
5496
5497 // Continue building IR after the loop. Note that the LoopInfo returned by
5498 // `collapseLoops` points inside the outermost loop and is intended for
5499 // potential further loop transformations. Use the insertion point stored
5500 // before collapsing loops instead.
5501 builder.restoreIP(afterIP);
5502 return success();
5503}
5504
5505/// Convert an omp.canonical_loop to LLVM-IR
5506static LogicalResult
5507convertOmpCanonicalLoopOp(omp::CanonicalLoopOp op, llvm::IRBuilderBase &builder,
5508 LLVM::ModuleTranslation &moduleTranslation) {
5509 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5510
5511 llvm::OpenMPIRBuilder::LocationDescription loopLoc(builder);
5512 Value loopIV = op.getInductionVar();
5513 Value loopTC = op.getTripCount();
5514
5515 llvm::Value *llvmTC = moduleTranslation.lookupValue(loopTC);
5516
5518 ompBuilder->createCanonicalLoop(
5519 loopLoc,
5520 [&](llvm::OpenMPIRBuilder::InsertPointTy ip, llvm::Value *llvmIV) {
5521 // Register the mapping of MLIR induction variable to LLVM-IR
5522 // induction variable
5523 moduleTranslation.mapValue(loopIV, llvmIV);
5524
5525 builder.restoreIP(ip);
5527 convertOmpOpRegions(op.getRegion(), "omp.loop.region", builder,
5528 moduleTranslation);
5529
5530 return bodyGenStatus.takeError();
5531 },
5532 llvmTC, "omp.loop");
5533 if (!llvmOrError)
5534 return op.emitError(llvm::toString(llvmOrError.takeError()));
5535
5536 llvm::CanonicalLoopInfo *llvmCLI = *llvmOrError;
5537 llvm::IRBuilderBase::InsertPoint afterIP = llvmCLI->getAfterIP();
5538 builder.restoreIP(afterIP);
5539
5540 // Register the mapping of MLIR loop to LLVM-IR OpenMPIRBuilder loop
5541 if (Value cli = op.getCli())
5542 moduleTranslation.mapOmpLoop(cli, llvmCLI);
5543
5544 return success();
5545}
5546
5547/// Apply a `#pragma omp unroll` / "!$omp unroll" transformation using the
5548/// OpenMPIRBuilder.
5549static LogicalResult
5550applyUnrollHeuristic(omp::UnrollHeuristicOp op, llvm::IRBuilderBase &builder,
5551 LLVM::ModuleTranslation &moduleTranslation) {
5552 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5553
5554 Value applyee = op.getApplyee();
5555 assert(applyee && "Loop to apply unrolling on required");
5556
5557 llvm::CanonicalLoopInfo *consBuilderCLI =
5558 moduleTranslation.lookupOMPLoop(applyee);
5559 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5560 ompBuilder->unrollLoopHeuristic(loc.DL, consBuilderCLI);
5561
5562 moduleTranslation.invalidateOmpLoop(applyee);
5563 return success();
5564}
5565
5566/// Apply a `#pragma omp unroll full` / `!$omp unroll full` transformation
5567/// using the OpenMPIRBuilder.
5568static LogicalResult
5569applyUnrollFull(omp::UnrollFullOp op, llvm::IRBuilderBase &builder,
5570 LLVM::ModuleTranslation &moduleTranslation) {
5571 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5572
5573 Value applyee = op.getApplyee();
5574 assert(applyee && "Loop to apply unrolling on required");
5575
5576 llvm::CanonicalLoopInfo *consBuilderCLI =
5577 moduleTranslation.lookupOMPLoop(applyee);
5578 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5579 ompBuilder->unrollLoopFull(loc.DL, consBuilderCLI);
5580
5581 moduleTranslation.invalidateOmpLoop(applyee);
5582 return success();
5583}
5584
5585/// Apply a `#pragma omp unroll partial` / `!$omp unroll partial`
5586/// transformation using the OpenMPIRBuilder.
5587static LogicalResult
5588applyUnrollPartial(omp::UnrollPartialOp op, llvm::IRBuilderBase &builder,
5589 LLVM::ModuleTranslation &moduleTranslation) {
5590 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5591
5592 Value applyee = op.getApplyee();
5593 assert(applyee && "Loop to apply unrolling on required");
5594
5595 llvm::CanonicalLoopInfo *consBuilderCLI =
5596 moduleTranslation.lookupOMPLoop(applyee);
5597 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5598
5599 // No generatee is supported yet, so the unrolled loop's CanonicalLoopInfo is
5600 // not requested and unrolling is deferred to LLVM's LoopUnroll pass.
5601 int32_t factor = static_cast<int32_t>(op.getUnrollFactor());
5602 ompBuilder->unrollLoopPartial(loc.DL, consBuilderCLI, factor,
5603 /*UnrolledCLI=*/nullptr);
5604
5605 moduleTranslation.invalidateOmpLoop(applyee);
5606 return success();
5607}
5608
5609/// Apply a `#pragma omp tile` / `!$omp tile` transformation using the
5610/// OpenMPIRBuilder.
5611static LogicalResult applyTile(omp::TileOp op, llvm::IRBuilderBase &builder,
5612 LLVM::ModuleTranslation &moduleTranslation) {
5613 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5614 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5615
5617 SmallVector<llvm::Value *> translatedSizes;
5618
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);
5624 }
5625
5626 for (Value applyee : op.getApplyees()) {
5627 llvm::CanonicalLoopInfo *consBuilderCLI =
5628 moduleTranslation.lookupOMPLoop(applyee);
5629 assert(applyee && "Canonical loop must already been translated");
5630 translatedLoops.push_back(consBuilderCLI);
5631 }
5632
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))
5638 moduleTranslation.mapOmpLoop(mlirLoop, genLoop);
5639 }
5640
5641 // CLIs can only be consumed once
5642 for (Value applyee : op.getApplyees())
5643 moduleTranslation.invalidateOmpLoop(applyee);
5644
5645 return success();
5646}
5647
5648/// Apply a `#pragma omp fuse` / `!$omp fuse` transformation using the
5649/// OpenMPIRBuilder.
5650static LogicalResult applyFuse(omp::FuseOp op, llvm::IRBuilderBase &builder,
5651 LLVM::ModuleTranslation &moduleTranslation) {
5652 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5653 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
5654
5655 // Select what CLIs are going to be fused
5656 SmallVector<llvm::CanonicalLoopInfo *> beforeFuse, toFuse, afterFuse;
5657 for (size_t i = 0; i < op.getApplyees().size(); i++) {
5658 Value applyee = op.getApplyees()[i];
5659 llvm::CanonicalLoopInfo *consBuilderCLI =
5660 moduleTranslation.lookupOMPLoop(applyee);
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);
5667 else
5668 toFuse.push_back(consBuilderCLI);
5669 }
5670 assert(
5671 (op.getGeneratees().empty() ||
5672 beforeFuse.size() + afterFuse.size() + 1 == op.getGeneratees().size()) &&
5673 "Wrong number of generatees");
5674
5675 // do the fuse
5676 auto generatedLoop = ompBuilder->fuseLoops(loc.DL, toFuse);
5677 if (!op.getGeneratees().empty()) {
5678 size_t i = 0;
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]);
5684 }
5685
5686 // CLIs can only be consumed once
5687 for (Value applyee : op.getApplyees())
5688 moduleTranslation.invalidateOmpLoop(applyee);
5689
5690 return success();
5691}
5692
5693/// Convert an Atomic Ordering attribute to llvm::AtomicOrdering.
5694static llvm::AtomicOrdering
5695convertAtomicOrdering(std::optional<omp::ClauseMemoryOrderKind> ao) {
5696 if (!ao)
5697 return llvm::AtomicOrdering::Monotonic; // Default Memory Ordering
5698
5699 switch (*ao) {
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;
5710 }
5711 llvm_unreachable("Unknown ClauseMemoryOrderKind kind");
5712}
5713
5714/// Compute the cmpxchg failure ordering for an atomic compare op: use the
5715/// `fail` clause ordering when present (the verifier guarantees it is a valid
5716/// cmpxchg failure ordering), otherwise the strongest failure ordering derived
5717/// from the success ordering (which matches the OpenMPIRBuilder default).
5718static llvm::AtomicOrdering
5719getAtomicCompareFailureOrdering(omp::AtomicCompareOp atomicCompareOp,
5720 llvm::AtomicOrdering atomicOrdering) {
5721 if (atomicCompareOp.getFailMemoryOrder())
5722 return convertAtomicOrdering(atomicCompareOp.getFailMemoryOrder());
5723 return llvm::AtomicCmpXchgInst::getStrongestFailureOrdering(atomicOrdering);
5724}
5725
5726/// Convert omp.atomic.read operation to LLVM IR.
5727static LogicalResult
5728convertOmpAtomicRead(Operation &opInst, llvm::IRBuilderBase &builder,
5729 LLVM::ModuleTranslation &moduleTranslation) {
5730 auto readOp = cast<omp::AtomicReadOp>(opInst);
5731 if (failed(checkImplementationStatus(opInst)))
5732 return failure();
5733
5734 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5735 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5736 findAllocInsertPoints(builder, moduleTranslation);
5737
5738 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5739
5740 llvm::AtomicOrdering AO = convertAtomicOrdering(readOp.getMemoryOrder());
5741 llvm::Value *x = moduleTranslation.lookupValue(readOp.getX());
5742 llvm::Value *v = moduleTranslation.lookupValue(readOp.getV());
5743
5744 llvm::Type *elementType =
5745 moduleTranslation.convertType(readOp.getElementType());
5746
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));
5750 return success();
5751}
5752
5753/// Converts an omp.atomic.write operation to LLVM IR.
5754static LogicalResult
5755convertOmpAtomicWrite(Operation &opInst, llvm::IRBuilderBase &builder,
5756 LLVM::ModuleTranslation &moduleTranslation) {
5757 auto writeOp = cast<omp::AtomicWriteOp>(opInst);
5758 if (failed(checkImplementationStatus(opInst)))
5759 return failure();
5760
5761 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5762 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
5763 findAllocInsertPoints(builder, moduleTranslation);
5764
5765 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
5766 llvm::AtomicOrdering ao = convertAtomicOrdering(writeOp.getMemoryOrder());
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, /*isSigned=*/false,
5771 /*isVolatile=*/false};
5772 builder.restoreIP(
5773 ompBuilder->createAtomicWrite(ompLoc, x, expr, ao, allocaIP));
5774 return success();
5775}
5776
5777/// Converts an LLVM dialect binary operation to the corresponding enum value
5778/// for `atomicrmw` supported binary operation.
5779static llvm::AtomicRMWInst::BinOp convertBinOpToAtomic(Operation &op) {
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);
5791}
5792
5793static void extractAtomicControlFlags(omp::AtomicUpdateOp atomicUpdateOp,
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();
5806 }
5807}
5808
5809/// Converts an OpenMP atomic update operation using OpenMPIRBuilder.
5810static LogicalResult
5811convertOmpAtomicUpdate(omp::AtomicUpdateOp &opInst,
5812 llvm::IRBuilderBase &builder,
5813 LLVM::ModuleTranslation &moduleTranslation) {
5814 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
5815 if (failed(checkImplementationStatus(*opInst)))
5816 return failure();
5817
5818 // Convert values and types.
5819 auto &innerOpList = opInst.getRegion().front().getOperations();
5820 bool isXBinopExpr{false};
5821 llvm::AtomicRMWInst::BinOp binop;
5822 mlir::Value mlirExpr;
5823 llvm::Value *llvmExpr = nullptr;
5824 llvm::Value *llvmX = nullptr;
5825 llvm::Type *llvmXElementType = nullptr;
5826 if (innerOpList.size() == 2) {
5827 // The two operations here are the update and the terminator.
5828 // Since we can identify the update operation, there is a possibility
5829 // that we can generate the atomicrmw instruction.
5830 mlir::Operation &innerOp = *opInst.getRegion().front().begin();
5831 if (!llvm::is_contained(innerOp.getOperands(),
5832 opInst.getRegion().getArgument(0))) {
5833 return opInst.emitError("no atomic update operation with region argument"
5834 " as operand found inside atomic.update region");
5835 }
5836 binop = convertBinOpToAtomic(innerOp);
5837 isXBinopExpr = innerOp.getOperand(0) == opInst.getRegion().getArgument(0);
5838 mlirExpr = (isXBinopExpr ? innerOp.getOperand(1) : innerOp.getOperand(0));
5839 llvmExpr = moduleTranslation.lookupValue(mlirExpr);
5840 } else {
5841 // Since the update region includes more than one operation
5842 // we will resort to generating a cmpxchg loop.
5843 binop = llvm::AtomicRMWInst::BinOp::BAD_BINOP;
5844 }
5845 llvmX = moduleTranslation.lookupValue(opInst.getX());
5846 llvmXElementType = moduleTranslation.convertType(
5847 opInst.getRegion().getArgument(0).getType());
5848 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicX = {llvmX, llvmXElementType,
5849 /*isSigned=*/false,
5850 /*isVolatile=*/false};
5851
5852 llvm::AtomicOrdering atomicOrdering =
5853 convertAtomicOrdering(opInst.getMemoryOrder());
5854
5855 // Generate update code.
5856 auto updateFn =
5857 [&opInst, &moduleTranslation](
5858 llvm::Value *atomicx,
5859 llvm::IRBuilder<> &builder) -> llvm::Expected<llvm::Value *> {
5860 Block &bb = *opInst.getRegion().begin();
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>();
5865
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 "
5869 "argument");
5870 return moduleTranslation.lookupValue(yieldop.getResults()[0]);
5871 };
5872
5873 bool isIgnoreDenormalMode;
5874 bool isFineGrainedMemory;
5875 bool isRemoteMemory;
5876 extractAtomicControlFlags(opInst, isIgnoreDenormalMode, isFineGrainedMemory,
5877 isRemoteMemory);
5878 // Handle ambiguous alloca, if any.
5879 auto allocaIP = findAllocInsertPoints(builder, moduleTranslation);
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);
5886
5887 if (failed(handleError(afterIP, *opInst)))
5888 return failure();
5889
5890 builder.restoreIP(*afterIP);
5891 return success();
5892}
5893
5894/// Helper to extract the OMPAtomicCompareOp from an integer comparison
5895/// predicate. Returns std::nullopt for unsupported predicates.
5896static std::optional<llvm::omp::OMPAtomicCompareOp>
5897convertICmpPredicateToAtomicCompareOp(LLVM::ICmpPredicate predicate) {
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;
5907 default:
5908 return std::nullopt;
5909 }
5910}
5911
5912/// Helper to extract the OMPAtomicCompareOp from a floating-point comparison
5913/// predicate. Returns std::nullopt for unsupported predicates.
5914static std::optional<llvm::omp::OMPAtomicCompareOp>
5915convertFCmpPredicateToAtomicCompareOp(LLVM::FCmpPredicate predicate) {
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;
5926 default:
5927 return std::nullopt;
5928 }
5929}
5930
5931/// Result of matching the decomposed complex equality pattern inside an atomic
5932/// compare region.
5934 bool isComplex = false;
5935 bool isNE = false; // `or` of the field compares => NE (unsupported).
5936 mlir::Value eAggregate; // The complex expected value (`e`).
5937 bool isXBinopExpr = false; // True if x is the first fcmp operand.
5938};
5939
5940/// Detect a decomposed complex equality comparison in an atomic compare region:
5941/// %re_x = llvm.extractvalue %xval[0]
5942/// %re_e = llvm.extractvalue %eStruct[0]
5943/// %cmp_re = llvm.fcmp "oeq" %re_x, %re_e
5944/// %im_x = llvm.extractvalue %xval[1]
5945/// %im_e = llvm.extractvalue %eStruct[1]
5946/// %cmp_im = llvm.fcmp "oeq" %im_x, %im_e
5947/// %cmp = llvm.and %cmp_re, %cmp_im (llvm.or would be NE)
5948/// It is recognised by an and/or whose operands are both fcmps operating on
5949/// extractvalues, one chain rooted at the block argument (x) and the other at
5950/// the expected complex value (e).
5953 auto traceToAggregate = [](mlir::Value v) -> mlir::Value {
5954 if (auto extractOp = v.getDefiningOp<LLVM::ExtractValueOp>())
5955 return extractOp.getContainer();
5956 return nullptr;
5957 };
5958 for (Operation &op : block.getOperations()) {
5959 if (!isa<LLVM::AndOp, LLVM::OrOp>(op))
5960 continue;
5961 auto lhsFcmp = op.getOperand(0).getDefiningOp<LLVM::FCmpOp>();
5962 auto rhsFcmp = op.getOperand(1).getDefiningOp<LLVM::FCmpOp>();
5963 if (!lhsFcmp || !rhsFcmp)
5964 continue;
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)
5970 continue;
5971 mlir::Value eAggregate = lhsXIsOp0 ? lhsAgg1 : lhsAgg0;
5972 if (!eAggregate)
5973 continue;
5974 result.isComplex = true;
5975 result.isNE = isa<LLVM::OrOp>(op);
5976 result.eAggregate = eAggregate;
5977 result.isXBinopExpr = lhsXIsOp0;
5978 break;
5979 }
5980 return result;
5981}
5982
5983/// Emit an IEEE-754-correct `cmpxchg` for a complex (struct-typed) atomic
5984/// compare with `fcmp oeq`. The old value of X is returned (as the complex
5985/// struct type) in \p oldComplex and the success flag (i1) in \p cmpOk.
5986/// \p failOrdering is the memory ordering used when the compare-exchange does
5987/// not store; it must be a valid cmpxchg failure ordering.
5988static void emitComplexAtomicCmpXchg(llvm::IRBuilderBase &builder,
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);
6003
6004 // Spill D to obtain its integer bit pattern for the swap value.
6005 llvm::AllocaInst *dAlloca =
6006 builder.CreateAlloca(complexTy, nullptr, "cmplx.d");
6007 dAlloca->setAlignment(maxAlign);
6008 builder.CreateAlignedStore(dVal, dAlloca, maxAlign);
6009 llvm::Value *dInt =
6010 builder.CreateAlignedLoad(intTy, dAlloca, maxAlign, "cmplx.d.int");
6011
6012 // Load X atomically and reinterpret as complex. Use the failure ordering: on
6013 // a failed component comparison we branch around the cmpxchg, so this load is
6014 // the only memory op on that path.
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");
6024
6025 // Component-wise IEEE-754 equality: `fcmp oeq` yields false for NaN (so a
6026 // NaN component correctly makes the compare fail) and true for +0.0 vs -0.0
6027 // (so a zero-sign difference does not spuriously fail the compare).
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");
6035
6036 // When the components compare equal, attempt the swap using X's own loaded
6037 // bit pattern as the comparand; otherwise the compare fails and X is left
6038 // unchanged (the captured old value is the value just loaded).
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);
6046
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);
6054
6055 // Merge the swap and no-swap paths.
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);
6063
6064 // Reinterpret the old integer value as the complex struct via memory.
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,
6070 "cmplx.old.val");
6071 cmpOk = okPHI;
6072}
6073
6074/// Holds the extracted comparison pattern information from an atomic compare
6075/// region.
6077 llvm::omp::OMPAtomicCompareOp compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6078 llvm::Value *eVal = nullptr;
6079 llvm::Value *dVal = nullptr;
6080 bool isXBinopExpr = false;
6081 bool isSigned = false;
6082};
6083/// Extract comparison predicate, expected value (e), desired value (d), and
6084/// related flags from an atomic compare region block by scanning for
6085/// icmp/fcmp/select/min/max operations.
6086static LogicalResult extractAtomicComparePattern(
6087 Block &block,
6088 llvm::function_ref<llvm::Value *(mlir::Value)> materializeValue,
6089 omp::AtomicCompareOp atomicCompareOp, AtomicComparePatternInfo &info) {
6090 // Complex equality is a decomposed per-field pattern (extractvalue + fcmp +
6091 // and) rather than a single scalar compare. Detect it first so the scalar
6092 // icmp/fcmp handling below does not mistake a real/imaginary field for the
6093 // whole expected value.
6095 cplx.isComplex) {
6096 if (cplx.isNE)
6097 return atomicCompareOp.emitError(
6098 "unsupported comparison predicate (NE) for complex atomic compare");
6099 info.compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6100 info.isXBinopExpr = cplx.isXBinopExpr;
6101 info.eVal = materializeValue(cplx.eAggregate);
6102 for (Operation &op : block.getOperations()) {
6103 if (auto selectOp = dyn_cast<LLVM::SelectOp>(op)) {
6104 info.dVal = materializeValue(selectOp.getTrueValue());
6105 break;
6106 }
6107 }
6108 return success();
6109 }
6110
6111 for (Operation &op : block.getOperations()) {
6112 // Pre-filter: skip icmps that don't involve the block argument
6113 // (e.g., truthiness extractions from logical-to-integer conversion).
6114 if (auto icmpOp = dyn_cast<LLVM::ICmpOp>(op);
6115 icmpOp && icmpOp.getOperand(0) != block.getArgument(0) &&
6116 icmpOp.getOperand(1) != block.getArgument(0))
6117 continue;
6118
6119 LogicalResult result =
6121 .Case<LLVM::ICmpOp>([&](LLVM::ICmpOp icmpOp) -> LogicalResult {
6122 auto maybeOp =
6123 convertICmpPredicateToAtomicCompareOp(icmpOp.getPredicate());
6124 if (!maybeOp)
6125 return atomicCompareOp.emitError(
6126 "unsupported comparison predicate in atomic compare");
6127 info.compareOp = *maybeOp;
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);
6133 info.isXBinopExpr =
6134 (icmpOp.getOperand(0) == block.getArgument(0));
6135 mlir::Value eOperand = info.isXBinopExpr ? icmpOp.getOperand(1)
6136 : icmpOp.getOperand(0);
6137 info.eVal = materializeValue(eOperand);
6138 return success();
6139 })
6140 .Case<LLVM::FCmpOp>([&](LLVM::FCmpOp fcmpOp) -> LogicalResult {
6141 auto maybeOp =
6142 convertFCmpPredicateToAtomicCompareOp(fcmpOp.getPredicate());
6143 if (!maybeOp)
6144 return atomicCompareOp.emitError(
6145 "unsupported comparison predicate in atomic compare");
6146 info.compareOp = *maybeOp;
6147 info.isXBinopExpr =
6148 (fcmpOp.getOperand(0) == block.getArgument(0));
6149 mlir::Value eOperand = info.isXBinopExpr ? fcmpOp.getOperand(1)
6150 : fcmpOp.getOperand(0);
6151 info.eVal = materializeValue(eOperand);
6152 return success();
6153 })
6154 .Case<LLVM::SelectOp>([&](LLVM::SelectOp selectOp) {
6155 if (!info.dVal)
6156 info.dVal = materializeValue(selectOp.getTrueValue());
6157 return success();
6158 })
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 *) {
6164 // Canonicalized min/max ops (arith or LLVM intrinsic form).
6165 // max(x,e) came from slt/ult/olt -> OMPAtomicCompareOp::MIN
6166 // min(x,e) came from sgt/ugt/ogt -> OMPAtomicCompareOp::MAX
6167 // (OMPIRBuilder inverts: MIN->atomicrmw max, MAX->atomicrmw min)
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);
6175 info.isXBinopExpr = (op.getOperand(0) == block.getArgument(0));
6176 mlir::Value eOperand =
6177 info.isXBinopExpr ? op.getOperand(1) : op.getOperand(0);
6178 info.eVal = materializeValue(eOperand);
6179 info.dVal = info.eVal;
6180 return success();
6181 })
6182 .Default([](Operation *) { return success(); });
6183
6184 if (failed(result))
6185 return result;
6186 }
6187 return success();
6188}
6189
6190static LogicalResult
6191convertOmpAtomicCapture(omp::AtomicCaptureOp atomicCaptureOp,
6192 llvm::IRBuilderBase &builder,
6193 LLVM::ModuleTranslation &moduleTranslation) {
6194 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
6195 if (failed(checkImplementationStatus(*atomicCaptureOp)))
6196 return failure();
6197
6198 omp::AtomicUpdateOp atomicUpdateOp = atomicCaptureOp.getAtomicUpdateOp();
6199 omp::AtomicWriteOp atomicWriteOp = atomicCaptureOp.getAtomicWriteOp();
6200 omp::AtomicCompareOp atomicCompareOp = atomicCaptureOp.getAtomicCompareOp();
6201
6202 // If the capture contains an atomic.compare, delegate to
6203 // createAtomicCompare with the capture variable (V) set.
6204 if (atomicCompareOp) {
6205 omp::AtomicReadOp atomicReadOp = atomicCaptureOp.getAtomicReadOp();
6206 assert(atomicReadOp && "expected atomic.read in capture+compare");
6207
6208 Region &region = atomicCompareOp.getRegion();
6209 Block &block = region.front();
6210
6211 llvm::Type *llvmXElementType =
6212 moduleTranslation.convertType(block.getArgument(0).getType());
6213 llvm::Value *llvmX = moduleTranslation.lookupValue(atomicCompareOp.getX());
6214 llvm::Value *llvmV = moduleTranslation.lookupValue(atomicReadOp.getV());
6215
6216 bool isSigned = false;
6217 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicX = {
6218 llvmX, llvmXElementType, isSigned, /*IsVolatile=*/false};
6219 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicV = {
6220 llvmV, llvmXElementType, /*isSigned=*/false, /*IsVolatile=*/false};
6221 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicR = {nullptr, nullptr, false,
6222 false};
6223
6224 llvm::AtomicOrdering atomicOrdering =
6225 convertAtomicOrdering(atomicCaptureOp.getMemoryOrder());
6226
6227 // Pre-translate non-pattern operations inside the compare region.
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);
6235 };
6236 for (Operation &op : block.without_terminator()) {
6237 if (isAtomicComparePatternOp(op))
6238 continue;
6239 bool allOperandsMapped =
6240 llvm::all_of(op.getOperands(), [&](mlir::Value v) {
6241 return moduleTranslation.lookupValue(v) != nullptr;
6242 });
6243 if (!allOperandsMapped)
6244 continue;
6245 if (failed(moduleTranslation.convertOperation(op, builder)))
6246 return atomicCompareOp.emitError(
6247 "failed to translate operation inside atomic compare region");
6248 }
6249
6250 auto materializeValue = [&](mlir::Value val) -> llvm::Value * {
6251 if (llvm::Value *existing = moduleTranslation.lookupValue(val))
6252 return existing;
6253 if (auto loadOp = val.getDefiningOp<LLVM::LoadOp>()) {
6254 if (loadOp->getParentRegion() == &region) {
6255 llvm::Value *loadAddr =
6256 moduleTranslation.lookupValue(loadOp.getAddr());
6257 if (!loadAddr)
6258 return nullptr;
6259 llvm::Type *loadType =
6260 moduleTranslation.convertType(loadOp.getResult().getType());
6261 return builder.CreateLoad(loadType, loadAddr);
6262 }
6263 }
6264 return nullptr;
6265 };
6266
6267 // Extract comparison predicate, eVal, and dVal from the region.
6268 AtomicComparePatternInfo patternInfo;
6269 if (failed(extractAtomicComparePattern(block, materializeValue,
6270 atomicCompareOp, patternInfo)))
6271 return failure();
6272
6273 llvm::omp::OMPAtomicCompareOp compareOp = patternInfo.compareOp;
6274 llvm::Value *eVal = patternInfo.eVal;
6275 llvm::Value *dVal = patternInfo.dVal;
6276 bool isXBinopExpr = patternInfo.isXBinopExpr;
6277 isSigned = patternInfo.isSigned;
6278
6279 if (!eVal)
6280 return atomicCompareOp.emitError(
6281 "failed to extract expected value (e) from atomic compare region");
6282 if (!dVal) {
6283 auto yieldOp = cast<omp::YieldOp>(block.getTerminator());
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]);
6288 }
6289
6290 llvmAtomicX.IsSigned = isSigned;
6291
6292 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6293 bool isReadFirst = isa<omp::AtomicReadOp>(atomicCaptureOp.getFirstOp());
6294 bool isPostfixCapture = !isReadFirst;
6295 bool isFailOnly = atomicCaptureOp.getFailOnly();
6296
6297 // Complex equality capture: x is struct-typed, which the OMPIRBuilder
6298 // cannot handle, so emit an IEEE-754-correct cmpxchg (as in the non-capture
6299 // complex path) and reconstruct the captured value from its result. The
6300 // helper compares components with `fcmp oeq` so `-0.0 == +0.0` and `NaN`
6301 // are handled as in the scalar float path. Complex only supports the ==
6302 // comparison.
6303 if (llvmXElementType->isStructTy()) {
6304 llvm::Value *oldComplex = nullptr;
6305 llvm::Value *cmpOk = nullptr;
6306 llvm::AtomicOrdering failOrdering =
6307 getAtomicCompareFailureOrdering(atomicCompareOp, atomicOrdering);
6308 emitComplexAtomicCmpXchg(builder, llvmX, llvmXElementType, eVal, dVal,
6309 atomicOrdering, failOrdering,
6310 atomicCompareOp.getWeak(), oldComplex, cmpOk);
6311
6312 if (isFailOnly) {
6313 // v is written only when the compare fails (cmpOk == false).
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) {
6328 // v gets the new value of x: d on success, old x otherwise.
6329 llvm::Value *newComplex = builder.CreateSelect(cmpOk, dVal, oldComplex);
6330 builder.CreateStore(newComplex, llvmAtomicV.Var,
6331 llvmAtomicV.IsVolatile);
6332 } else {
6333 // Prefix: v gets the old value of x.
6334 builder.CreateStore(oldComplex, llvmAtomicV.Var,
6335 llvmAtomicV.IsVolatile);
6336 }
6337
6338 // Emit flush after atomic compare if needed (release/acq_rel/seq_cst).
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);
6344 }
6345 return success();
6346 }
6347
6348 // Min/max (<, >) comparisons lower to an atomicrmw. The OMPIRBuilder has no
6349 // notion of a failed compare for an atomicrmw, so the fail-only capture
6350 // form (v written only when the compare fails) has no valid mapping and is
6351 // Min/max (<, >) comparisons lower to an atomicrmw. The OMPIRBuilder has no
6352 // notion of a failed compare for an atomicrmw (it asserts on IsFailOnly),
6353 // so for min/max the fail-only capture is reconstructed manually below.
6354 bool isMinMax = compareOp != llvm::omp::OMPAtomicCompareOp::EQ;
6355
6356 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicVForCall = llvmAtomicV;
6357 // Capture into V is reconstructed manually below for:
6358 // * postfix capture (v gets the new value of x): equality from the
6359 // cmpxchg result, min/max from the atomicrmw result;
6360 // * min/max fail-only capture (the OMPIRBuilder cannot express it).
6361 // Bypass V in the OMPIRBuilder for those cases so it does not also emit its
6362 // own (for min/max, incorrect or unsupported) capture store.
6363 bool minMaxManualCapture = isMinMax && (isPostfixCapture || isFailOnly);
6364 bool eqPostfixManualCapture = !isMinMax && isPostfixCapture && !isFailOnly;
6365 if (minMaxManualCapture || eqPostfixManualCapture)
6366 llvmAtomicVForCall = {nullptr, nullptr, false, false};
6367
6368 // The OMPIRBuilder only understands IsFailOnly for the equality (cmpxchg)
6369 // path; for min/max it would assert. Min/max fail-only is handled here.
6370 bool builderFailOnly = isFailOnly && !isMinMax;
6371
6372 // IsPostfixUpdate selects which value the OMPIRBuilder captures into V:
6373 // * min/max prefix and equality prefix: a direct store of the old value
6374 // (IsPostfixUpdate=true).
6375 // * equality fail-only: a conditional store (IsPostfixUpdate=false).
6376 // Manually-reconstructed captures bypass V above.
6377 bool isPostfixUpdate = !builderFailOnly;
6378
6379 bool isWeak = atomicCompareOp.getWeak();
6380 bool savedHandleFPNegZero = ompBuilder->setHandleFPNegZero(true);
6381 llvm::AtomicOrdering failureOrdering =
6382 getAtomicCompareFailureOrdering(atomicCompareOp, atomicOrdering);
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);
6389
6390 if (failed(handleError(afterIP, *atomicCaptureOp)))
6391 return failure();
6392
6393 builder.restoreIP(*afterIP);
6394
6395 // Min/max capture is reconstructed from the atomicrmw the OMPIRBuilder
6396 // emits (its result is the old value of x). V was bypassed above.
6397 // * postfix: v gets the new value min/max(old, e);
6398 // * fail-only: v gets the old value, but only when the compare failed
6399 // (i.e. the atomicrmw did not change x).
6400 // (Prefix min/max captures the old value directly through V, so nothing
6401 // extra is needed there.)
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)) {
6407 rmw = r;
6408 break;
6409 }
6410 }
6411 assert(rmw && "expected atomicrmw for min/max compare capture");
6412 llvm::Value *oldVal = rmw;
6413 llvm::Value *rhs = rmw->getValOperand();
6414
6415 if (isFailOnly) {
6416 // The compare "failed" (the else branch runs) exactly when the
6417 // atomicrmw did not change x. Recompute the original update condition
6418 // on the old value and negate it. v is stored only in that case.
6419 llvm::CmpInst::Predicate updatePred;
6420 switch (rmw->getOperation()) {
6421 case llvm::AtomicRMWInst::Min:
6422 updatePred = llvm::CmpInst::ICMP_SGT;
6423 break;
6424 case llvm::AtomicRMWInst::Max:
6425 updatePred = llvm::CmpInst::ICMP_SLT;
6426 break;
6427 case llvm::AtomicRMWInst::UMin:
6428 updatePred = llvm::CmpInst::ICMP_UGT;
6429 break;
6430 case llvm::AtomicRMWInst::UMax:
6431 updatePred = llvm::CmpInst::ICMP_ULT;
6432 break;
6433 case llvm::AtomicRMWInst::FMin:
6434 updatePred = llvm::CmpInst::FCMP_OGT;
6435 break;
6436 case llvm::AtomicRMWInst::FMax:
6437 updatePred = llvm::CmpInst::FCMP_OLT;
6438 break;
6439 default:
6440 llvm_unreachable(
6441 "unexpected atomicrmw op for min/max compare capture");
6442 }
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);
6455 } else {
6456 llvm::Intrinsic::ID id;
6457 switch (rmw->getOperation()) {
6458 case llvm::AtomicRMWInst::Min:
6459 id = llvm::Intrinsic::smin;
6460 break;
6461 case llvm::AtomicRMWInst::Max:
6462 id = llvm::Intrinsic::smax;
6463 break;
6464 case llvm::AtomicRMWInst::UMin:
6465 id = llvm::Intrinsic::umin;
6466 break;
6467 case llvm::AtomicRMWInst::UMax:
6468 id = llvm::Intrinsic::umax;
6469 break;
6470 case llvm::AtomicRMWInst::FMin:
6471 id = llvm::Intrinsic::minnum;
6472 break;
6473 case llvm::AtomicRMWInst::FMax:
6474 id = llvm::Intrinsic::maxnum;
6475 break;
6476 default:
6477 llvm_unreachable(
6478 "unexpected atomicrmw op for min/max compare capture");
6479 }
6480 llvm::Value *newVal = builder.CreateBinaryIntrinsic(id, oldVal, rhs);
6481 builder.CreateStore(newVal, llvmAtomicV.Var, llvmAtomicV.IsVolatile);
6482 }
6483 }
6484
6485 // Equality postfix: v = select(success, D, old) — reconstructs the new
6486 // value of x from the cmpxchg result.
6487 if (!isMinMax && isPostfixCapture && !isFailOnly) {
6488 llvm::BasicBlock *curBB = builder.GetInsertBlock();
6489 llvm::Value *oldVal = nullptr;
6490 llvm::Value *successVal = nullptr;
6491
6492 // Integer path (and non-HandleFPNegZero FP path): a single cmpxchg
6493 // lives in the current block.
6494 for (auto &inst : llvm::reverse(*curBB)) {
6495 if (isa<llvm::AtomicCmpXchgInst>(&inst)) {
6496 oldVal = builder.CreateExtractValue(&inst, /*Idxs=*/0);
6497 successVal = builder.CreateExtractValue(&inst, /*Idxs=*/1);
6498 break;
6499 }
6500 }
6501
6502 // FP HandleFPNegZero path: the OMPIRBuilder emits a multi-block
6503 // structure (NaN / ±0.0 handling) with cmpxchg in predecessor
6504 // blocks. Results are merged via PHI nodes in the current (exit)
6505 // block: an i1 PHI for success and a bitcast of an integer PHI
6506 // for the old FP value.
6507 if (!oldVal) {
6508 for (auto &inst : *curBB) {
6509 auto *phi = dyn_cast<llvm::PHINode>(&inst);
6510 if (!phi)
6511 break;
6512 if (phi->getType()->isIntegerTy(1))
6513 successVal = phi;
6514 }
6515 for (auto &inst : *curBB) {
6516 if (auto *bc = dyn_cast<llvm::BitCastInst>(&inst)) {
6517 oldVal = bc;
6518 break;
6519 }
6520 }
6521 }
6522
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);
6527 }
6528
6529 return success();
6530 }
6531
6532 mlir::Value mlirExpr;
6533 bool isXBinopExpr = false, isPostfixUpdate = false;
6534 llvm::AtomicRMWInst::BinOp binop = llvm::AtomicRMWInst::BinOp::BAD_BINOP;
6535
6536 assert((atomicUpdateOp || atomicWriteOp) &&
6537 "internal op must be an atomic.update or atomic.write op");
6538
6539 if (atomicWriteOp) {
6540 isPostfixUpdate = true;
6541 mlirExpr = atomicWriteOp.getExpr();
6542 } else {
6543 isPostfixUpdate = atomicCaptureOp.getSecondOp() ==
6544 atomicCaptureOp.getAtomicUpdateOp().getOperation();
6545 auto &innerOpList = atomicUpdateOp.getRegion().front().getOperations();
6546 // Find the binary update operation that uses the region argument
6547 // and get the expression to update
6548 if (innerOpList.size() == 2) {
6549 mlir::Operation &innerOp = *atomicUpdateOp.getRegion().front().begin();
6550 if (!llvm::is_contained(innerOp.getOperands(),
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");
6555 }
6556 binop = convertBinOpToAtomic(innerOp);
6557 isXBinopExpr =
6558 innerOp.getOperand(0) == atomicUpdateOp.getRegion().getArgument(0);
6559 mlirExpr = (isXBinopExpr ? innerOp.getOperand(1) : innerOp.getOperand(0));
6560 } else {
6561 binop = llvm::AtomicRMWInst::BinOp::BAD_BINOP;
6562 }
6563 }
6564
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,
6573 /*isSigned=*/false,
6574 /*isVolatile=*/false};
6575 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicV = {llvmV, llvmXElementType,
6576 /*isSigned=*/false,
6577 /*isVolatile=*/false};
6578
6579 llvm::AtomicOrdering atomicOrdering =
6580 convertAtomicOrdering(atomicCaptureOp.getMemoryOrder());
6581
6582 auto updateFn =
6583 [&](llvm::Value *atomicx,
6584 llvm::IRBuilder<> &builder) -> llvm::Expected<llvm::Value *> {
6585 if (atomicWriteOp)
6586 return moduleTranslation.lookupValue(atomicWriteOp.getExpr());
6587 Block &bb = *atomicUpdateOp.getRegion().begin();
6588 moduleTranslation.mapValue(*atomicUpdateOp.getRegion().args_begin(),
6589 atomicx);
6590 moduleTranslation.mapBlock(&bb, builder.GetInsertBlock());
6591 if (failed(moduleTranslation.convertBlock(bb, true, builder)))
6592 return llvm::make_error<PreviouslyReportedError>();
6593
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 "
6597 "argument");
6598 return moduleTranslation.lookupValue(yieldop.getResults()[0]);
6599 };
6600
6601 bool isIgnoreDenormalMode;
6602 bool isFineGrainedMemory;
6603 bool isRemoteMemory;
6604 extractAtomicControlFlags(atomicUpdateOp, isIgnoreDenormalMode,
6605 isFineGrainedMemory, isRemoteMemory);
6606 // Handle ambiguous alloca, if any.
6607 auto allocaIP = findAllocInsertPoints(builder, moduleTranslation);
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);
6614
6615 if (failed(handleError(afterIP, *atomicCaptureOp)))
6616 return failure();
6617
6618 builder.restoreIP(*afterIP);
6619 return success();
6620}
6621
6622/// Converts an omp.atomic.compare operation to LLVM IR.
6623///
6624/// if (x == e) x = d
6625/// The region contains a comparison + select pattern:
6626/// ^bb0(%xval: T):
6627/// %cmp = llvm.icmp/fcmp <pred> %xval, %e : T
6628/// %sel = llvm.select %cmp, %d, %xval : i1, T
6629/// omp.yield(%sel : T)
6630///
6631/// From MLIR extract:
6632/// 1) comparison operator
6633/// 2) expected value (e)
6634/// 3) desired value (d)
6635/// These are passed to OpenMPIRBuilder::createAtomicCompare which generates
6636/// the actual cmpxchg / atomicrmw instruction.
6637///
6638static LogicalResult
6639convertOmpAtomicCompare(omp::AtomicCompareOp atomicCompareOp,
6640 llvm::IRBuilderBase &builder,
6641 LLVM::ModuleTranslation &moduleTranslation) {
6642 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
6643 if (failed(checkImplementationStatus(*atomicCompareOp)))
6644 return failure();
6645
6646 Region &region = atomicCompareOp.getRegion();
6647 Block &block = region.front();
6648
6649 // Determine element type from the region block argument
6650 llvm::Type *llvmXElementType =
6651 moduleTranslation.convertType(block.getArgument(0).getType());
6652 if (!llvmXElementType)
6653 return atomicCompareOp.emitError(
6654 "unable to determine element type for atomic compare");
6655
6656 llvm::Value *llvmX = moduleTranslation.lookupValue(atomicCompareOp.getX());
6657
6658 // IsSigned is determined from the comparison predicate in the region.
6659 // Signed ICmp predicates (slt/sgt) set this to true; unsigned (ult/ugt)
6660 // leave it false. For EQ and float comparisons, signedness is irrelevant.
6661 bool isSigned = false;
6662 llvm::OpenMPIRBuilder::AtomicOpValue llvmAtomicX = {llvmX, llvmXElementType,
6663 isSigned,
6664 /*IsVolatile=*/false};
6665
6666 llvm::AtomicOrdering atomicOrdering =
6667 convertAtomicOrdering(atomicCompareOp.getMemoryOrder());
6668
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);
6676 };
6677
6678 // Pre-translate operations inside the region that compute e and d (e.g.,
6679 // GEP, loads for dereferencing Fortran pointers) but are not part of the
6680 // atomic compare-and-swap pattern (icmp/fcmp, select, and/or).
6681 //
6682 // 1) Validity: The OpenMP spec requires e and d to be evaluated before the
6683 // atomic operation, so emitting their computation here is correct.
6684 // 2) Memory effects: These ops only depend on values defined outside the
6685 // region. They cannot observe the block argument (%xval), which is the
6686 // value loaded atomically by cmpxchg and does not exist yet.
6687 // 3) Invariant enforcement: The `allOperandsMapped` check below skips any
6688 // op whose operands include the unmapped block argument, guaranteeing
6689 // only region-external-dependent ops are pre-translated.
6690 for (Operation &op : block.without_terminator()) {
6691 // Skip operations that form the atomic compare pattern — these are
6692 // not emitted as individual instructions but are analyzed below to
6693 // extract the comparison predicate, expected value (e), and desired
6694 // value (d) for generating a single cmpxchg/atomicrmw.
6695 if (isAtomicComparePatternOp(op))
6696 continue;
6697
6698 // Avoid translating ops that depend on the unmapped block argument.
6699 bool allOperandsMapped = llvm::all_of(op.getOperands(), [&](mlir::Value v) {
6700 return moduleTranslation.lookupValue(v) != nullptr;
6701 });
6702 if (!allOperandsMapped)
6703 continue;
6704
6705 if (failed(moduleTranslation.convertOperation(op, builder)))
6706 return atomicCompareOp.emitError(
6707 "failed to translate operation inside atomic compare region");
6708 }
6709
6710 // Look up a value that may have been pre-translated or defined outside the
6711 // region.
6712 auto materializeValue = [&](mlir::Value val) -> llvm::Value * {
6713 // Check if the value is already mapped (pre-translated or defined outside).
6714 if (llvm::Value *existing = moduleTranslation.lookupValue(val))
6715 return existing;
6716 // Fallback for a single LoadOp whose address is mapped but whose result
6717 // was not pre-translated.
6718 if (auto loadOp = val.getDefiningOp<LLVM::LoadOp>()) {
6719 if (loadOp->getParentRegion() == &region) {
6720 llvm::Value *loadAddr = moduleTranslation.lookupValue(loadOp.getAddr());
6721 if (!loadAddr)
6722 return nullptr;
6723 llvm::Type *loadType =
6724 moduleTranslation.convertType(loadOp.getResult().getType());
6725 return builder.CreateLoad(loadType, loadAddr);
6726 }
6727 }
6728 return nullptr;
6729 };
6730
6731 // Walk the region to extract comparison predicate, eVal, and dVal.
6732 // if (x == eVal) x = dVal
6733 llvm::omp::OMPAtomicCompareOp compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6734 llvm::Value *eVal = nullptr;
6735 llvm::Value *dVal = nullptr;
6736 bool isXBinopExpr = false;
6737
6738 // Check for a decomposed complex comparison pattern (extractvalue + fcmp +
6739 // and/or of the real/imaginary fields).
6741 bool isComplexPattern = cplx.isComplex;
6742 if (isComplexPattern) {
6743 if (cplx.isNE)
6744 // OrOp corresponds to NE, which is not a valid atomic compare op.
6745 return atomicCompareOp.emitError(
6746 "unsupported comparison predicate (NE) for complex atomic compare");
6747 compareOp = llvm::omp::OMPAtomicCompareOp::EQ;
6748 isXBinopExpr = cplx.isXBinopExpr;
6749 eVal = materializeValue(cplx.eAggregate);
6750 }
6751
6752 if (isComplexPattern) {
6753 // dVal from SelectOp or YieldOp.
6754 for (Operation &op : block.getOperations()) {
6755 if (auto selectOp = dyn_cast<LLVM::SelectOp>(op)) {
6756 dVal = materializeValue(selectOp.getTrueValue());
6757 break;
6758 }
6759 }
6760 if (!dVal) {
6761 auto yieldOp = cast<omp::YieldOp>(block.getTerminator());
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]);
6766 }
6767
6768 llvm::Value *oldComplex = nullptr;
6769 llvm::Value *cmpOk = nullptr;
6770 llvm::AtomicOrdering failOrdering =
6771 getAtomicCompareFailureOrdering(atomicCompareOp, atomicOrdering);
6772 emitComplexAtomicCmpXchg(builder, llvmX, llvmXElementType, eVal, dVal,
6773 atomicOrdering, failOrdering,
6774 atomicCompareOp.getWeak(), oldComplex, cmpOk);
6775 (void)oldComplex;
6776 (void)cmpOk;
6777
6778 // Emit flush after atomic compare if needed (for release, acq_rel,
6779 // seq_cst orderings).
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);
6785 }
6786 return success();
6787 } else {
6788 AtomicComparePatternInfo patternInfo;
6789 if (failed(extractAtomicComparePattern(block, materializeValue,
6790 atomicCompareOp, patternInfo)))
6791 return failure();
6792 compareOp = patternInfo.compareOp;
6793 eVal = patternInfo.eVal;
6794 dVal = patternInfo.dVal;
6795 isXBinopExpr = patternInfo.isXBinopExpr;
6796 isSigned = patternInfo.isSigned;
6797 }
6798
6799 if (!eVal)
6800 return atomicCompareOp.emitError(
6801 "failed to extract expected value (e) from atomic compare region");
6802 if (!dVal) {
6803 // Fall back to the yield operand.
6804 auto yieldOp = cast<omp::YieldOp>(block.getTerminator());
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]);
6809 }
6810
6811 llvmAtomicX.IsSigned = isSigned;
6812
6813 llvm::OpenMPIRBuilder::AtomicOpValue vOpVal = {nullptr, nullptr, false,
6814 false};
6815 llvm::OpenMPIRBuilder::AtomicOpValue rOpVal = {nullptr, nullptr, false,
6816 false};
6817 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6818
6819 bool isWeak = atomicCompareOp.getWeak();
6820
6821 bool savedHandleFPNegZero = ompBuilder->setHandleFPNegZero(true);
6822 llvm::AtomicOrdering failureOrdering =
6823 getAtomicCompareFailureOrdering(atomicCompareOp, atomicOrdering);
6824 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6825 ompBuilder->createAtomicCompare(
6826 ompLoc, llvmAtomicX, vOpVal, rOpVal, eVal, dVal, atomicOrdering,
6827 compareOp, isXBinopExpr, /*IsPostfixUpdate=*/false,
6828 /*IsFailOnly=*/false, failureOrdering, isWeak);
6829 ompBuilder->setHandleFPNegZero(savedHandleFPNegZero);
6830
6831 if (failed(handleError(afterIP, *atomicCompareOp)))
6832 return failure();
6833
6834 builder.restoreIP(*afterIP);
6835 return success();
6836}
6837
6838static llvm::omp::Directive convertCancellationConstructType(
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;
6849 }
6850 llvm_unreachable("Unhandled cancellation construct type");
6851}
6852
6853static LogicalResult
6854convertOmpCancel(omp::CancelOp op, llvm::IRBuilderBase &builder,
6855 LLVM::ModuleTranslation &moduleTranslation) {
6856 if (failed(checkImplementationStatus(*op.getOperation())))
6857 return failure();
6858
6859 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6860 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
6861
6862 llvm::Value *ifCond = nullptr;
6863 if (Value ifVar = op.getIfExpr())
6864 ifCond = moduleTranslation.lookupValue(ifVar);
6865
6866 llvm::omp::Directive cancelledDirective =
6867 convertCancellationConstructType(op.getCancelDirective());
6868
6869 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6870 ompBuilder->createCancel(ompLoc, ifCond, cancelledDirective);
6871
6872 if (failed(handleError(afterIP, *op.getOperation())))
6873 return failure();
6874
6875 builder.restoreIP(afterIP.get());
6876
6877 return success();
6878}
6879
6880static LogicalResult
6881convertOmpCancellationPoint(omp::CancellationPointOp op,
6882 llvm::IRBuilderBase &builder,
6883 LLVM::ModuleTranslation &moduleTranslation) {
6884 if (failed(checkImplementationStatus(*op.getOperation())))
6885 return failure();
6886
6887 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6888 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
6889
6890 llvm::omp::Directive cancelledDirective =
6891 convertCancellationConstructType(op.getCancelDirective());
6892
6893 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
6894 ompBuilder->createCancellationPoint(ompLoc, cancelledDirective);
6895
6896 if (failed(handleError(afterIP, *op.getOperation())))
6897 return failure();
6898
6899 builder.restoreIP(afterIP.get());
6900
6901 return success();
6902}
6903
6904/// Converts an OpenMP Threadprivate operation into LLVM IR using
6905/// OpenMPIRBuilder.
6906static LogicalResult
6907convertOmpThreadprivate(Operation &opInst, llvm::IRBuilderBase &builder,
6908 LLVM::ModuleTranslation &moduleTranslation) {
6909 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
6910 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
6911 auto threadprivateOp = cast<omp::ThreadprivateOp>(opInst);
6912
6913 if (failed(checkImplementationStatus(opInst)))
6914 return failure();
6915
6916 Value symAddr = threadprivateOp.getSymAddr();
6917 auto *symOp = symAddr.getDefiningOp();
6918
6919 if (auto asCast = dyn_cast<LLVM::AddrSpaceCastOp>(symOp))
6920 symOp = asCast.getOperand().getDefiningOp();
6921
6922 if (!isa<LLVM::AddressOfOp>(symOp))
6923 return opInst.emitError("Addressing symbol not found");
6924 LLVM::AddressOfOp addressOfOp = dyn_cast<LLVM::AddressOfOp>(symOp);
6925
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(
6932 type);
6933 llvm::ConstantInt *size = builder.getInt64(typeSize.getFixedValue());
6934 llvm::Value *callInst = ompBuilder->createCachedThreadPrivate(
6935 ompLoc, globalValue, size, global.getSymName() + ".cache");
6936 moduleTranslation.mapValue(opInst.getResult(0), callInst);
6937
6938 return success();
6939}
6940
6941static llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseKind
6942convertToDeviceClauseKind(mlir::omp::DeclareTargetDeviceType deviceClause) {
6943 switch (deviceClause) {
6944 case mlir::omp::DeclareTargetDeviceType::host:
6945 return llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseHost;
6946 break;
6947 case mlir::omp::DeclareTargetDeviceType::nohost:
6948 return llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseNoHost;
6949 break;
6950 case mlir::omp::DeclareTargetDeviceType::any:
6951 return llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseAny;
6952 break;
6953 }
6954 llvm_unreachable("unhandled device clause");
6955}
6956
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;
6969 }
6970 llvm_unreachable("unhandled capture clause");
6971}
6972
6974 Operation *op = value.getDefiningOp();
6975 if (auto addrCast = dyn_cast_if_present<LLVM::AddrSpaceCastOp>(op))
6976 op = addrCast->getOperand(0).getDefiningOp();
6977 if (auto addressOfOp = dyn_cast_if_present<LLVM::AddressOfOp>(op)) {
6978 auto modOp = addressOfOp->getParentOfType<mlir::ModuleOp>();
6979 return modOp.lookupSymbol(addressOfOp.getGlobalName());
6980 }
6981 return nullptr;
6982}
6983
6985 while (Operation *op = value.getDefiningOp()) {
6986 if (auto addrCast = dyn_cast_if_present<LLVM::AddrSpaceCastOp>(op))
6987 value = addrCast.getOperand();
6988 // Traces through hlfir.declare, fir.declare to reach the base address and
6989 // use for type lookup.
6990 else if (op->getName().getIdentifier() &&
6991 (op->getName().getIdentifier().str() == "hlfir.declare" ||
6992 op->getName().getIdentifier().str() == "fir.declare")) {
6993 if (op->getNumOperands() > 0)
6994 value = op->getOperand(0);
6995 else
6996 break;
6997 } else {
6998 break;
6999 }
7000 }
7001 return value;
7002}
7003
7004// Determine the LLVM type whose storage size should be allocated for an
7005// OpenMP allocate directive list item. Opaque pointers lose element type, so
7006// trace through declare wrappers to the underlying global or stack allocation.
7007static llvm::Type *
7009 LLVM::ModuleTranslation &moduleTranslation) {
7010 llvm::Type *llvmVarTy = moduleTranslation.convertType(var.getType());
7011 if (!llvmVarTy->isPointerTy())
7012 return llvmVarTy;
7013
7014 if (Operation *globalOp = getGlobalOpFromValue(baseVar))
7015 if (auto gop = dyn_cast<LLVM::GlobalOp>(globalOp))
7016 return moduleTranslation.convertType(gop.getGlobalType());
7017
7018 if (auto allocaOp =
7019 dyn_cast_if_present<LLVM::AllocaOp>(baseVar.getDefiningOp()))
7020 return moduleTranslation.convertType(allocaOp.getElemType());
7021
7022 if (llvm::Value *baseLlvm = moduleTranslation.lookupValue(baseVar))
7023 if (auto *allocaInst = dyn_cast<llvm::AllocaInst>(baseLlvm))
7024 return allocaInst->getAllocatedType();
7025
7026 return llvmVarTy;
7027}
7028
7029// For dynamically-sized stack allocations, compute the allocation size from
7030// the alloca's element count at runtime.
7031static std::optional<llvm::Value *> getDynamicAllocatedSize(
7032 Value var, Value baseVar, LLVM::ModuleTranslation &moduleTranslation,
7033 llvm::IRBuilderBase &builder, const llvm::DataLayout &dataLayout) {
7034 if (auto allocaOp =
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));
7044 }
7045 }
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())) {
7050 uint64_t elemSize =
7051 dataLayout.getTypeAllocSize(allocaInst->getAllocatedType())
7052 .getFixedValue();
7053 return builder.CreateMul(allocaInst->getArraySize(),
7054 builder.getInt64(elemSize));
7055 }
7056 }
7057 }
7058 return std::nullopt;
7059}
7060
7061static llvm::SmallString<64>
7062getDeclareTargetRefPtrSuffix(LLVM::GlobalOp globalOp,
7063 llvm::OpenMPIRBuilder &ompBuilder,
7064 llvm::vfs::FileSystem &vfs) {
7065 llvm::SmallString<64> suffix;
7066 llvm::raw_svector_ostream os(suffix);
7067 if (globalOp.getVisibility() == mlir::SymbolTable::Visibility::Private) {
7068 auto loc = globalOp->getLoc()->findInstanceOf<FileLineColLoc>();
7069 auto fileInfoCallBack = [&loc]() {
7070 return std::pair<std::string, uint64_t>(
7071 llvm::StringRef(loc.getFilename()), loc.getLine());
7072 };
7073
7074 os << llvm::format(
7075 "_%x",
7076 ompBuilder.getTargetEntryUniqueInfo(fileInfoCallBack, vfs).FileID);
7077 }
7078 os << "_decl_tgt_ref_ptr";
7079
7080 return suffix;
7081}
7082
7083static bool isDeclareTargetLink(Value value) {
7084 if (auto declareTargetGlobal =
7085 dyn_cast_if_present<omp::DeclareTargetInterface>(
7086 getGlobalOpFromValue(value))) {
7087 omp::DeclareTargetAttr declareTargetAttr =
7088 declareTargetGlobal.getDeclareTarget();
7089 if (declareTargetAttr && declareTargetAttr.getCaptureClause() ==
7090 omp::DeclareTargetCaptureClause::link)
7091 return true;
7092 }
7093 return false;
7094}
7095
7096static bool isDeclareTargetTo(Value value) {
7097 if (auto declareTargetGlobal =
7098 dyn_cast_if_present<omp::DeclareTargetInterface>(
7099 getGlobalOpFromValue(value))) {
7100 omp::DeclareTargetAttr declareTargetAttr =
7101 declareTargetGlobal.getDeclareTarget();
7102 if (declareTargetAttr && (declareTargetAttr.getCaptureClause() ==
7103 omp::DeclareTargetCaptureClause::to ||
7104 declareTargetAttr.getCaptureClause() ==
7105 omp::DeclareTargetCaptureClause::enter))
7106 return true;
7107 }
7108 return false;
7109}
7110
7111// Returns the reference pointer generated by the lowering of the declare
7112// target operation in cases where the link clause is used or the to clause is
7113// used in USM mode.
7114static llvm::Value *
7116 LLVM::ModuleTranslation &moduleTranslation) {
7117 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
7118 if (auto gOp =
7119 dyn_cast_or_null<LLVM::GlobalOp>(getGlobalOpFromValue(value))) {
7120 // In this case, we must utilise the reference pointer generated by
7121 // the declare target operation, similar to Clang
7122 if (isDeclareTargetLink(value) ||
7123 (isDeclareTargetTo(value) &&
7124 ompBuilder->Config.hasRequiresUnifiedSharedMemory())) {
7126 gOp, *ompBuilder, moduleTranslation.getFileSystem());
7127
7128 if (gOp.getSymName().contains(suffix))
7129 return moduleTranslation.getLLVMModule()->getNamedValue(
7130 gOp.getSymName());
7131
7132 return moduleTranslation.getLLVMModule()->getNamedValue(
7133 (gOp.getSymName().str() + suffix.str()).str());
7134 }
7135 }
7136 return nullptr;
7137}
7138
7139namespace {
7140// Append customMappers information to existing MapInfosTy
7141struct MapInfosTy : llvm::OpenMPIRBuilder::MapInfosTy {
7142 SmallVector<Operation *, 4> Mappers;
7143
7144 /// Append arrays in \a CurInfo.
7145 void append(MapInfosTy &curInfo) {
7146 Mappers.append(curInfo.Mappers.begin(), curInfo.Mappers.end());
7147 llvm::OpenMPIRBuilder::MapInfosTy::append(curInfo);
7148 }
7149};
7150// A small helper structure to contain data gathered
7151// for map lowering and coalese it into one area and
7152// avoiding extra computations such as searches in the
7153// llvm module for lowered mapped variables or checking
7154// if something is declare target (and retrieving the
7155// value) more than neccessary.
7156struct MapInfoData : MapInfosTy {
7157 llvm::SmallVector<bool, 4> IsDeclareTarget;
7158 llvm::SmallVector<bool, 4> IsAMember;
7159 // Identify if mapping was added by mapClause or use_device clauses.
7160 llvm::SmallVector<bool, 4> IsAMapping;
7161 llvm::SmallVector<mlir::Operation *, 4> MapClause;
7162 llvm::SmallVector<llvm::Value *, 4> OriginalValue;
7163 // Stripped off array/pointer to get the underlying
7164 // element type
7165 llvm::SmallVector<llvm::Type *, 4> BaseType;
7166
7167 /// Append arrays in \a CurInfo.
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);
7176 }
7177};
7178
7179enum class TargetDirectiveEnumTy : uint32_t {
7180 None = 0,
7181 Target = 1,
7182 TargetData = 2,
7183 TargetEnterData = 3,
7184 TargetExitData = 4,
7185 TargetUpdate = 5
7186};
7187
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;
7193 })
7194 .Case([&](omp::TargetExitDataOp) {
7195 return TargetDirectiveEnumTy::TargetExitData;
7196 })
7197 .Case([&](omp::TargetUpdateOp) {
7198 return TargetDirectiveEnumTy::TargetUpdate;
7199 })
7200 .Case([&](omp::TargetOp) { return TargetDirectiveEnumTy::Target; })
7201 .Default([&](Operation *op) { return TargetDirectiveEnumTy::None; });
7202}
7203
7204} // namespace
7205
7206static uint64_t getArrayElementSizeInBits(LLVM::LLVMArrayType arrTy,
7207 DataLayout &dl) {
7208 if (auto nestedArrTy = llvm::dyn_cast_if_present<LLVM::LLVMArrayType>(
7209 arrTy.getElementType()))
7210 return getArrayElementSizeInBits(nestedArrTy, dl);
7211 return dl.getTypeSizeInBits(arrTy.getElementType());
7212}
7213
7214// The intent is to verify if the mapped data being passed is a
7215// pointer -> pointee that requires special handling in certain cases,
7216// e.g. applying the OMP_MAP_PTR_AND_OBJ map type.
7217//
7218// There may be a better way to verify this, but unfortunately with
7219// opaque pointers we lose the ability to easily check if something is
7220// a pointer whilst maintaining access to the underlying type.
7221static bool checkIfPointerMap(omp::MapInfoOp mapOp) {
7222 // If we have a varPtrPtr field assigned then the underlying type is a pointer
7223 if (mapOp.getVarPtrPtr())
7224 return true;
7225
7226 // If the map data is declare target with a link clause, then it's represented
7227 // as a pointer when we lower it to LLVM-IR even if at the MLIR level it has
7228 // no relation to pointers.
7229 if (isDeclareTargetLink(mapOp.getVarPtr()))
7230 return true;
7231
7232 return false;
7233}
7234
7235// A privatizeable attach map is a pointer/descriptor that is privatized and
7236// passed directly as a kernel argument (target_param) rather than undergoing
7237// the standard attach/parent mapping. These are handled specially in a couple
7238// of places in map lowering.
7239static bool isPrivatizeableAttachMap(omp::ClauseMapFlags mapType) {
7240 return bitEnumContainsAll(mapType, omp::ClauseMapFlags::priv |
7241 omp::ClauseMapFlags::target_param |
7242 omp::ClauseMapFlags::attach);
7243}
7244
7245// This function calculates the size to be offloaded for a specified type, given
7246// its associated map clause (which can contain bounds information which affects
7247// the total size), this size is calculated based on the underlying element type
7248// e.g. given a 1-D array of ints, we will calculate the size from the integer
7249// type * number of elements in the array. This size can be used in other
7250// calculations but is ultimately used as an argument to the OpenMP runtimes
7251// kernel argument structure which is generated through the combinedInfo data
7252// structures.
7253// This function is somewhat equivalent to Clang's getExprTypeSize inside of
7254// CGOpenMPRuntime.cpp.
7255static llvm::Value *getSizeInBytes(DataLayout &dl, const mlir::Type &type,
7256 Operation *clauseOp,
7257 llvm::Value *basePointer,
7258 llvm::Type *baseType,
7259 llvm::IRBuilderBase &builder,
7260 LLVM::ModuleTranslation &moduleTranslation) {
7261 if (auto memberClause =
7262 mlir::dyn_cast_if_present<mlir::omp::MapInfoOp>(clauseOp)) {
7263 // This calculates the size to transfer based on bounds and the underlying
7264 // element type, provided bounds have been specified (Fortran
7265 // pointers/allocatables/target and arrays that have sections specified fall
7266 // into this as well)
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())) {
7272 // The below calculation for the size to be mapped calculated from the
7273 // map.info's bounds is: (elemCount * [UB - LB] + 1), later we
7274 // multiply by the underlying element types byte size to get the full
7275 // size to be offloaded based on the bounds
7276 elementCount = builder.CreateMul(
7277 elementCount,
7278 builder.CreateAdd(
7279 builder.CreateSub(
7280 moduleTranslation.lookupValue(boundOp.getUpperBound()),
7281 moduleTranslation.lookupValue(boundOp.getLowerBound())),
7282 builder.getInt64(1)));
7283 }
7284 }
7285
7286 // utilising getTypeSizeInBits instead of getTypeSize as getTypeSize gives
7287 // the size in inconsistent byte or bit format.
7288 uint64_t underlyingTypeSzInBits = dl.getTypeSizeInBits(type);
7289 if (auto arrTy = llvm::dyn_cast_if_present<LLVM::LLVMArrayType>(type))
7290 underlyingTypeSzInBits = getArrayElementSizeInBits(arrTy, dl);
7291
7292 // The size in bytes x number of elements, the sizeInBytes stored is
7293 // the underyling types size, e.g. if ptr<i32>, it'll be the i32's
7294 // size, so we do some on the fly runtime math to get the size in
7295 // bytes from the extent (ub - lb) * sizeInBytes. NOTE: This may need
7296 // some adjustment for members with more complex types.
7297 llvm::Value *sizeCalc = builder.CreateMul(
7298 elementCount, builder.getInt64(underlyingTypeSzInBits / 8),
7299 "element_count");
7300
7301 // This is a part of a "complicated" bit of size calculation logic that is
7302 // in place to handle a couple of scenarios, one specific to Fortran and
7303 // the other a more general OpenMP issue. The other piece of the
7304 // calculation can be found as the final size calculation within the
7305 // processIndividualMap function. Ideally we would move it here, but due
7306 // to the complexity of calculating the final base address of some
7307 // constructs (required for a nullary check), it's left as the final step.
7308 // So, in the below 2 cases, the nullary check is in processIndividualMap
7309 // and the size equality check is here. The cases this modifications help
7310 // cover are:
7311 //
7312 // 1) If an argument has a null base pointer, then the size must be set to
7313 // 0 to avoid the runtime exploding/complaining about an illegal
7314 // pointer map. The size returning non-zero is feasible in certain
7315 // cases if for example someone has specified there own bounds/range.
7316 // 2) We wish to support a very specific OpenMP Fortran edge-case where a
7317 // size zero array can be legally presence checked and found to be on
7318 // device when it has been mapped. In these rare occasions the
7319 // allocatable/pointer will have a size of 1 allocated for the
7320 // underlying data, but this wall not be represented within the size of
7321 // the descriptor, so we get a non-nullary pointer and a size of 0,
7322 // allowing us to specify a size of 1 in these cases registering it on
7323 // the device mapping table as present.
7324 //
7325 // The default fall through case is just returning the size calculation
7326 // above, if we are not nullary and the size we calculate is non-zero,
7327 // which is basically any pointer type that is allocated in someway
7328 // (providing you are not running on a rare system that allows malloc's of
7329 // size 0 with whatever caveats that may come with).
7330 //
7331 // Later in the nullary check in processIndividualMap it just devolves to
7332 // selecting a size of 0 if we are nullary, if we are not, we will return
7333 // either 1 or the calculated size, depending on the outcome of this
7334 // select.
7335 if (checkIfPointerMap(memberClause)) {
7336 return builder.CreateSelect(
7337 builder.CreateICmpEQ(sizeCalc, builder.getInt64(0)),
7338 builder.getInt64(1), sizeCalc);
7339 }
7340
7341 return sizeCalc;
7342 }
7343 }
7344
7345 return builder.getInt64(dl.getTypeSizeInBits(type) / 8);
7346}
7347
7348// Convert the MLIR map flag set to the runtime map flag set for embedding
7349// in LLVM-IR. This is important as the two bit-flag lists do not correspond
7350// 1-to-1 as there's flags the runtime doesn't care about and vice versa.
7351// Certain flags are discarded here such as RefPtee and co.
7352static llvm::omp::OpenMPOffloadMappingFlags
7353convertClauseMapFlags(omp::ClauseMapFlags mlirFlags) {
7354 const bool hasExplicitMap =
7355 (mlirFlags & ~omp::ClauseMapFlags::is_device_ptr) !=
7356 omp::ClauseMapFlags::none;
7357
7358 llvm::omp::OpenMPOffloadMappingFlags mapType =
7359 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE;
7360
7361 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::to))
7362 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TO;
7363
7364 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::from))
7365 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_FROM;
7366
7367 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::always))
7368 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ALWAYS;
7369
7370 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::del))
7371 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_DELETE;
7372
7373 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::return_param))
7374 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
7375
7376 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::priv))
7377 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_PRIVATE;
7378
7379 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::literal))
7380 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL;
7381
7382 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::implicit))
7383 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_IMPLICIT;
7384
7385 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::close))
7386 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_CLOSE;
7387
7388 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::present))
7389 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_PRESENT;
7390
7391 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::ompx_hold))
7392 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_OMPX_HOLD;
7393
7394 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::attach))
7395 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH;
7396
7397 if (bitEnumContainsAll(mlirFlags, omp::ClauseMapFlags::target_param))
7398 mapType |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TARGET_PARAM;
7399
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;
7404 }
7405
7406 return mapType;
7407}
7408
7410 MapInfoData &mapData, SmallVectorImpl<Value> &mapVars,
7411 LLVM::ModuleTranslation &moduleTranslation, DataLayout &dl,
7412 llvm::IRBuilderBase &builder, ArrayRef<Value> useDevPtrOperands = {},
7413 ArrayRef<Value> useDevAddrOperands = {},
7414 ArrayRef<Value> hasDevAddrOperands = {}) {
7415
7416 auto checkRefPtrOrPteeMapWithAttach = [](omp::ClauseMapFlags mapType) {
7417 bool hasRefType =
7418 bitEnumContainsAll(mapType, omp::ClauseMapFlags::ref_ptr) ||
7419 bitEnumContainsAll(mapType, omp::ClauseMapFlags::ref_ptee);
7420 return hasRefType &&
7421 bitEnumContainsAll(mapType, omp::ClauseMapFlags::attach);
7422 };
7423
7424 auto checkIsAMember = [](const auto &mapVars, auto mapOp) {
7425 // Check if this is a member mapping and correctly assign that it is, if
7426 // it is a member of a larger object.
7427 // TODO: Need better handling of members, and distinguishing of members
7428 // that are implicitly allocated on device vs explicitly passed in as
7429 // arguments.
7430 // TODO: May require some further additions to support nested record
7431 // types, i.e. member maps that can have member maps.
7432 for (Value mapValue : mapVars) {
7433 auto map = cast<omp::MapInfoOp>(mapValue.getDefiningOp());
7434 for (auto member : map.getMembers())
7435 if (member == mapOp)
7436 return true;
7437 }
7438 return false;
7439 };
7440
7441 // Process MapOperands
7442 for (Value mapValue : mapVars) {
7443 auto mapOp = cast<omp::MapInfoOp>(mapValue.getDefiningOp());
7444 bool isAttachStyleMap =
7445 checkRefPtrOrPteeMapWithAttach(mapOp.getMapType()) ||
7446 isPrivatizeableAttachMap(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());
7454
7455 if (llvm::Value *refPtr =
7456 getRefPtrIfDeclareTarget(offloadPtr, moduleTranslation)) {
7457 mapData.IsDeclareTarget.push_back(true);
7458 mapData.BasePointers.push_back(refPtr);
7459 } else if (isDeclareTargetTo(offloadPtr)) {
7460 mapData.IsDeclareTarget.push_back(true);
7461 mapData.BasePointers.push_back(mapData.OriginalValue.back());
7462 } else { // regular mapped variable
7463 mapData.IsDeclareTarget.push_back(false);
7464 mapData.BasePointers.push_back(mapData.OriginalValue.back());
7465 }
7466
7467 // In every situation we currently have if we have a varPtrPtr present
7468 // we wish to utilise it's type for the base type, main cases are
7469 // currently Fortran descriptor base address maps and attach maps.
7470 mapData.BaseType.push_back(moduleTranslation.convertType(
7471 mapOp.getVarPtrPtr() ? mapOp.getVarPtrPtrType().value()
7472 : mapOp.getVarPtrType()));
7473
7474 // For the attach map cases, it's a little odd, as we effectively have to
7475 // utilise the base address (including all bounds offsets) for the pointer
7476 // field, the pointer address for the base address field, and the pointer
7477 // not the data (base addresses) size. So we end up with a mix of base
7478 // types and sizes we wish to insert here.
7479 mlir::Type sizeType = (isAttachStyleMap || !mapOp.getVarPtrPtr())
7480 ? mapOp.getVarPtrType()
7481 : mapOp.getVarPtrPtrType().value();
7482 mapData.Sizes.push_back(getSizeInBytes(
7483 dl, sizeType, isAttachStyleMap ? nullptr : mapOp,
7484 mapData.Pointers.back(), moduleTranslation.convertType(sizeType),
7485 builder, moduleTranslation));
7486 mapData.MapClause.push_back(mapOp.getOperation());
7487 mapData.Types.push_back(convertClauseMapFlags(mapOp.getMapType()));
7488 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
7489 mapData.HasAttachPtr.push_back(false);
7490 mapData.Names.push_back(LLVM::createMappingInformation(
7491 mapOp.getLoc(), *moduleTranslation.getOpenMPBuilder()));
7492 mapData.DevicePointers.push_back(llvm::OpenMPIRBuilder::DeviceInfoTy::None);
7493 if (mapOp.getMapperId())
7494 mapData.Mappers.push_back(
7496 mapOp, mapOp.getMapperIdAttr()));
7497 else
7498 mapData.Mappers.push_back(nullptr);
7499 mapData.IsAMapping.push_back(true);
7500 mapData.IsAMember.push_back(checkIsAMember(mapVars, mapOp));
7501 }
7502
7503 auto findMapInfo = [&mapData](llvm::Value *val,
7504 llvm::OpenMPIRBuilder::DeviceInfoTy devInfoTy,
7505 size_t memberCount) {
7506 unsigned index = 0;
7507 bool found = false;
7508 for (llvm::Value *basePtr : mapData.OriginalValue) {
7509 auto mapOp = cast<omp::MapInfoOp>(mapData.MapClause[index]);
7510 // TODO: Currently we define an equivalent mapping as
7511 // the same base pointer and an equivalent member count, but
7512 // that is a loose definition. We may have to extend to check
7513 // for other fields (varPtrPtr/individual members being mapped).
7514 // Note: Attach maps are not the same as a normal data transfer
7515 // they specify to the runtime to perform an attach map and they
7516 // (at least at the moment) are never something we would aim to
7517 // return in a use_dev_* clause, so they are skipped in terms of
7518 // duplicate maps.
7519 bool isAttachMap =
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()) {
7525 found = true;
7526 mapData.Types[index] |=
7527 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
7528 mapData.DevicePointers[index] = devInfoTy;
7529 }
7530 index++;
7531 }
7532 return found;
7533 };
7534
7535 // Process useDevPtr(Addr)Operands
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());
7540 Value offloadPtr =
7541 mapOp.getVarPtrPtr() ? mapOp.getVarPtrPtr() : mapOp.getVarPtr();
7542 llvm::Value *origValue = moduleTranslation.lookupValue(offloadPtr);
7543
7544 // Check if map info is already present for this entry.
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);
7558 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
7559 mapData.HasAttachPtr.push_back(false);
7560 mapData.Names.push_back(LLVM::createMappingInformation(
7561 mapOp.getLoc(), *moduleTranslation.getOpenMPBuilder()));
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));
7566 }
7567 }
7568 };
7569
7570 addDevInfos(useDevAddrOperands, llvm::OpenMPIRBuilder::DeviceInfoTy::Address);
7571 addDevInfos(useDevPtrOperands, llvm::OpenMPIRBuilder::DeviceInfoTy::Pointer);
7572
7573 for (Value mapValue : hasDevAddrOperands) {
7574 auto mapOp = cast<omp::MapInfoOp>(mapValue.getDefiningOp());
7575 Value offloadPtr =
7576 mapOp.getVarPtrPtr() ? mapOp.getVarPtrPtr() : mapOp.getVarPtr();
7577 llvm::Value *origValue = moduleTranslation.lookupValue(offloadPtr);
7578 auto mapType = convertClauseMapFlags(mapOp.getMapType());
7579 auto mapTypeAlways = llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ALWAYS;
7580 bool isDevicePtr =
7581 (mapOp.getMapType() & omp::ClauseMapFlags::is_device_ptr) !=
7582 omp::ClauseMapFlags::none;
7583
7584 mapData.OriginalValue.push_back(origValue);
7585 mapData.BasePointers.push_back(origValue);
7586 mapData.Pointers.push_back(origValue);
7587 mapData.IsDeclareTarget.push_back(false);
7588
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)));
7593
7594 mapData.MapClause.push_back(mapOp.getOperation());
7595 if (llvm::to_underlying(mapType & mapTypeAlways)) {
7596 // Descriptors are mapped with the ALWAYS flag, since they can get
7597 // rematerialized, so the address of the decriptor for a given object
7598 // may change from one place to another.
7599 mapData.Types.push_back(mapType);
7600 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
7601 mapData.HasAttachPtr.push_back(false);
7602 // Technically it's possible for a non-descriptor mapping to have
7603 // both has-device-addr and ALWAYS, so lookup the mapper in case it
7604 // exists.
7605 if (mapOp.getMapperId()) {
7606 mapData.Mappers.push_back(
7608 mapOp, mapOp.getMapperIdAttr()));
7609 } else {
7610 mapData.Mappers.push_back(nullptr);
7611 }
7612 } else {
7613 // For is_device_ptr we need the map type to propagate so the runtime
7614 // can materialize the device-side copy of the pointer container.
7615 mapData.Types.push_back(
7616 isDevicePtr ? mapType
7617 : llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL);
7618 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
7619 mapData.HasAttachPtr.push_back(false);
7620 mapData.Mappers.push_back(nullptr);
7621 }
7622 mapData.Names.push_back(LLVM::createMappingInformation(
7623 mapOp.getLoc(), *moduleTranslation.getOpenMPBuilder()));
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));
7629 }
7630}
7631
7632static int getMapDataMemberIdx(MapInfoData &mapData, omp::MapInfoOp memberOp) {
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);
7637}
7638
7640 omp::MapInfoOp mapInfo, bool first = true) {
7641 ArrayAttr indexAttr = mapInfo.getMembersIndexAttr();
7642 llvm::SmallVector<size_t> occludedChildren;
7643 llvm::sort(
7644 indices.begin(), indices.end(), [&](const size_t a, const size_t b) {
7645 // Bail early if we are asked to look at the same index. If we do not
7646 // bail early, we can end up mistakenly adding indices to
7647 // occludedChildren. This can occur with some types of libc++ hardening.
7648 if (a == b)
7649 return false;
7650
7651 auto memberIndicesA = cast<ArrayAttr>(indexAttr[a]);
7652 auto memberIndicesB = cast<ArrayAttr>(indexAttr[b]);
7653
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();
7657
7658 if (aIndex == bIndex)
7659 continue;
7660
7661 if (aIndex < bIndex)
7662 return first;
7663
7664 if (aIndex > bIndex)
7665 return !first;
7666 }
7667
7668 // Iterated up until the end of the smallest member and
7669 // they were found to be equal up to that point, so select
7670 // the member with the lowest index count, so the "parent"
7671 bool memberAParent = memberIndicesA.size() < memberIndicesB.size();
7672 if (memberAParent)
7673 occludedChildren.push_back(b);
7674 else
7675 occludedChildren.push_back(a);
7676 return memberAParent;
7677 });
7678
7679 for (auto v : occludedChildren)
7680 indices.erase(std::remove(indices.begin(), indices.end(), v),
7681 indices.end());
7682}
7683
7684static omp::MapInfoOp getFirstOrLastMappedMemberPtr(omp::MapInfoOp mapInfo,
7685 bool first) {
7686 ArrayAttr indexAttr = mapInfo.getMembersIndexAttr();
7687 // Only 1 member has been mapped, we can return it.
7688 if (indexAttr.size() == 1)
7689 return cast<omp::MapInfoOp>(mapInfo.getMembers()[0].getDefiningOp());
7690 llvm::SmallVector<size_t> indices(indexAttr.size());
7691 std::iota(indices.begin(), indices.end(), 0);
7692 sortMapIndices(indices, mapInfo, first);
7693 return llvm::cast<omp::MapInfoOp>(
7694 mapInfo.getMembers()[indices.front()].getDefiningOp());
7695}
7696
7697/// This function calculates the array/pointer offset for map data provided
7698/// with bounds operations, e.g. when provided something like the following:
7699///
7700/// Fortran
7701/// map(tofrom: array(2:5, 3:2))
7702///
7703/// We must calculate the initial pointer offset to pass across, this function
7704/// performs this using bounds.
7705///
7706/// TODO/WARNING: This only supports Fortran's column major indexing currently
7707/// as is noted in the note below and comments in the function, we must extend
7708/// this function when we add a C++ frontend.
7709/// NOTE: which while specified in row-major order it currently needs to be
7710/// flipped for Fortran's column order array allocation and access (as
7711/// opposed to C++'s row-major, hence the backwards processing where order is
7712/// important). This is likely important to keep in mind for the future when
7713/// we incorporate a C++ frontend, both frontends will need to agree on the
7714/// ordering of generated bounds operations (one may have to flip them) to
7715/// make the below lowering frontend agnostic. The offload size
7716/// calcualtion may also have to be adjusted for C++.
7717static std::vector<llvm::Value *>
7719 llvm::IRBuilderBase &builder, bool isArrayTy,
7720 OperandRange bounds) {
7721 std::vector<llvm::Value *> idx;
7722 // There's no bounds to calculate an offset from, we can safely
7723 // ignore and return no indices.
7724 if (bounds.empty())
7725 return idx;
7726
7727 // If we have an array type, then we have its type so can treat it as a
7728 // normal GEP instruction where the bounds operations are simply indexes
7729 // into the array. We currently do reverse order of the bounds, which
7730 // I believe leans more towards Fortran's column-major in memory.
7731 if (isArrayTy) {
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()));
7737 }
7738 }
7739 } else {
7740 // If we do not have an array type, but we have bounds, then we're dealing
7741 // with a pointer that's being treated like an array and we have the
7742 // underlying type e.g. an i32, or f64 etc, e.g. a fortran descriptor base
7743 // address (pointer pointing to the actual data) so we must caclulate the
7744 // offset using a single index which the following loop attempts to
7745 // compute using the standard column-major algorithm e.g for a 3D array:
7746 //
7747 // ((((c_idx * b_len) + b_idx) * a_len) + a_idx)
7748 //
7749 // It is of note that it's doing column-major rather than row-major at the
7750 // moment, but having a way for the frontend to indicate which major format
7751 // to use or standardizing/canonicalizing the order of the bounds to compute
7752 // the offset may be useful in the future when there's other frontends with
7753 // different formats.
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))
7758 idx.emplace_back(
7759 moduleTranslation.lookupValue(boundOp.getLowerBound()));
7760 else
7761 idx.back() = builder.CreateAdd(
7762 builder.CreateMul(idx.back(), moduleTranslation.lookupValue(
7763 boundOp.getExtent())),
7764 moduleTranslation.lookupValue(boundOp.getLowerBound()));
7765 }
7766 }
7767 }
7768
7769 return idx;
7770}
7771
7773 llvm::transform(values, std::back_inserter(ints), [](Attribute value) {
7774 return cast<IntegerAttr>(value).getInt();
7775 });
7776}
7777
7778// Gathers members that are overlapping in the parent, excluding members that
7779// themselves overlap, keeping the top-most (closest to parents level) map.
7780static void
7782 omp::MapInfoOp parentOp) {
7783 // No members mapped, no overlaps.
7784 if (parentOp.getMembers().empty())
7785 return;
7786
7787 // Single member, we can insert and return early.
7788 if (parentOp.getMembers().size() == 1) {
7789 overlapMapDataIdxs.push_back(0);
7790 return;
7791 }
7792
7793 ArrayAttr indexAttr = parentOp.getMembersIndexAttr();
7794 size_t numMembers = indexAttr.size();
7795
7796 // Pre-convert all member indices to integer arrays for efficient comparison.
7797 llvm::SmallVector<llvm::SmallVector<int64_t>> memberIndices(numMembers);
7798 for (auto [i, indicesAttr] : llvm::enumerate(indexAttr))
7799 getAsIntegers(cast<ArrayAttr>(indicesAttr), memberIndices[i]);
7800
7801 // For each member, check if it's superseded by another (shorter prefix)
7802 // member. If member j's indices are a prefix of member i's indices, then
7803 // i is a child of j and should be skipped. e.g. if member [0] is mapped,
7804 // we skip members [0,1], [0,2], etc.
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) {
7809 if (i == j)
7810 continue;
7811 const auto &jIndices = memberIndices[j];
7812 // If j's indices are a strict prefix of i's indices, skip i
7813 if (jIndices.size() < iIndices.size() &&
7814 std::equal(jIndices.begin(), jIndices.end(), iIndices.begin())) {
7815 skipIndices.insert(i);
7816 break; // No need to check other potential parents
7817 }
7818 }
7819 }
7820
7821 // Collect indices of members that are not superseded by a parent.
7822 for (size_t i = 0; i < numMembers; ++i)
7823 if (!skipIndices.contains(i))
7824 overlapMapDataIdxs.push_back(i);
7825}
7826
7827/// This function handles the insertion of a single item of map data from
7828/// MapInfoData into the OMPIRBuilder's MapInfo list. Utilising this function
7829/// means the map being inserted can be treated as a non-parent map entity,
7830/// if the memberOfFlag is set then the map being inserted is treated as
7831/// a member map of a larger entity. The insertion into the MapInfo list of
7832/// the OMPIRBuilder can vary based on a number of factors, such as if it's
7833/// a ref_ptr or ref_ptee map, if it's a member of a record, what construct
7834/// the map belongs to and the various map type bit flags that are set for
7835/// the map.
7836static void
7837processIndividualMap(llvm::IRBuilderBase &builder,
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]);
7846
7847 bool isPtrTy = checkIfPointerMap(mapInfoOp);
7848 bool isAttachMap = ((convertClauseMapFlags(mapInfoOp.getMapType()) &
7849 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH) ==
7850 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH);
7851
7852 // Declare target variables are not passed to the kernel, and for the moment
7853 // attach maps are not passed to the kernel. However, it is possible to create
7854 // attach maps that transfer data and thus can be kernel arguments, but our
7855 // existing frontend does not do this.
7856 if (isTargetParam &&
7857 (targetDirective == TargetDirectiveEnumTy::Target &&
7858 !mapData.IsDeclareTarget[mapDataIdx]) &&
7859 !isAttachMap)
7860 mapFlag |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_TARGET_PARAM;
7861
7862 if (mapInfoOp.getMapCaptureType() == omp::VariableCaptureKind::ByCopy &&
7863 !isPtrTy)
7864 mapFlag |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_LITERAL;
7865
7866 // If we have a pointer and it's part of a MEMBER_OF mapping we do not apply
7867 // MEMBER_OF, as the runtime currently has a work-around that utilises
7868 // MEMBER_OF to prevent reference updating in certain scenarios instead of
7869 // target_param. However, this causes a noticeable issue in cases where we
7870 // map some data (Fortran descriptor primarily at the moment), alter it on
7871 // the host, and then expect it to not be updated in a subsequent implicit map
7872 // (such as an implicit map on a target).
7873 if (memberOfFlag != llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE) {
7874 if (!isPtrTy && !isAttachMap)
7875 ompBuilder.setCorrectMemberOfFlag(mapFlag, memberOfFlag);
7876
7877 // The return parameter should be the over-riding parent in cases where we
7878 // have a return parameter that is echoed to all members, the main case of
7879 // this currently is with fortran descriptors. It may need more finessing
7880 // for C/C++ in the future or descriptors that are members of derived
7881 // types.
7882 mapFlag &= ~llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
7883 }
7884
7885 // We apply MAP_PTR_AND_OBJ when within a declare mapper object as it enforces
7886 // MEMBER_OF mappings on maps that are passed the initial nesting depth, which
7887 // includes pointed to data and attach members, both of which are technically
7888 // not part of the main object. This has the side effect of causing early
7889 // map-backs in certain cases where an implicit declare mapper has been
7890 // emitted for a target region. Applying MAP_PTR_AND_OBJ in these situations
7891 // circumvents this.
7892 if (isPtrTy && !isAttachMap && mapData.IsDeclareTarget[mapDataIdx])
7893 mapFlag |= llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_PTR_AND_OBJ;
7894
7895 // if we're provided a mapDataParentIdx, then the data being mapped is
7896 // part of a larger object (in a parent <-> member mapping) and in this
7897 // case our BasePointer should be the parent. Except in the edge case
7898 // where we are mapping pointee data, where we try staying close to
7899 // what Clang currently does and utilise the regular base pointer of the
7900 // data.
7901 bool isRefPtee =
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);
7908
7909 if (!mapInfoOp->getParentOfType<omp::DeclareMapperOp>() &&
7910 mapDataParentIdx >= 0 && !(isRefPtee || (isRefPtrPtee && isPtrTy))) {
7911 combinedInfo.BasePointers.emplace_back(
7912 mapData.BasePointers[mapDataParentIdx]);
7913 } else {
7914 combinedInfo.BasePointers.emplace_back(mapData.BasePointers[mapDataIdx]);
7915 }
7916
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);
7925 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
7926 combinedInfo.HasAttachPtr.emplace_back(false);
7927 // A privatized attach map needs storage for the pointer or descriptor even
7928 // when its pointee is null. Its size describes that storage, not the pointee,
7929 // and is needed by the runtime for corresponding-pointer-initialization.
7930 combinedInfo.Sizes.emplace_back(
7931 isPtrTy && !isPrivatizeableAttachMap(mapInfoOp.getMapType())
7932 ? builder.CreateSelect(
7933 builder.CreateIsNull(mapData.Pointers[mapDataIdx]),
7934 builder.getInt64(0), mapData.Sizes[mapDataIdx])
7935 : mapData.Sizes[mapDataIdx]);
7936}
7937
7938// This creates two insertions into the MapInfosTy data structure for the
7939// "parent" of a set of members, (usually a container e.g.
7940// class/structure/derived type) when subsequent members have also been
7941// explicitly mapped on the same map clause. Certain types, such as Fortran
7942// descriptors are mapped like this as well, however, the members are
7943// implicit as far as a user is concerned, but we must explicitly map them
7944// internally.
7945//
7946// This function also returns the memberOfFlag for this particular parent,
7947// which is utilised in subsequent member mappings (by modifying there map type
7948// with it) to indicate that a member is part of this parent and should be
7949// treated by the runtime as such. Important to achieve the correct mapping.
7950//
7951// This function borrows a lot from Clang's emitCombinedEntry function
7952// inside of CGOpenMPRuntime.cpp
7954 LLVM::ModuleTranslation &moduleTranslation, llvm::IRBuilderBase &builder,
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");
7962 auto parentClause =
7963 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
7964 auto *parentMapper = mapData.Mappers[mapDataIndex];
7965
7966 // Map the first segment of the parent. If a user-defined mapper is attached,
7967 // include the parent's to/from-style bits (and common modifiers) in this
7968 // base entry so the mapper receives correct copy semantics via its 'type'
7969 // parameter. Also keep TARGET_PARAM when required for kernel arguments.
7970 MapFlags baseFlag = (targetDirective == TargetDirectiveEnumTy::Target &&
7971 !mapData.IsDeclareTarget[mapDataIndex])
7972 ? MapFlags::OMP_MAP_TARGET_PARAM
7973 : MapFlags::OMP_MAP_NONE;
7974
7975 if (parentMapper) {
7976 // Preserve relevant map-type bits from the parent clause. These include
7977 // the copy direction (TO/FROM), as well as commonly used modifiers that
7978 // should be visible to the mapper for correct behaviour.
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);
7986 } else {
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);
7993 }
7994
7995 combinedInfo.Types.emplace_back(baseFlag);
7996 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
7997 combinedInfo.HasAttachPtr.emplace_back(false);
7998 combinedInfo.DevicePointers.emplace_back(
7999 mapData.DevicePointers[mapDataIndex]);
8000 // Only attach the mapper to the base entry when we are mapping the whole
8001 // parent. Combined/segment entries must not carry a mapper; otherwise the
8002 // mapper can be invoked with a partial size, which is undefined behaviour.
8003 combinedInfo.Mappers.emplace_back(
8004 parentMapper && !parentClause.getPartialMap() ? parentMapper : nullptr);
8005 combinedInfo.Names.emplace_back(LLVM::createMappingInformation(
8006 mapData.MapClause[mapDataIndex]->getLoc(), ompBuilder));
8007 combinedInfo.BasePointers.emplace_back(mapData.BasePointers[mapDataIndex]);
8008
8009 // Calculate size of the parent object being mapped based on the
8010 // addresses at runtime, highAddr - lowAddr = size. This of course
8011 // doesn't factor in allocated data like pointers, hence the further
8012 // processing of members specified by users, or in the case of
8013 // Fortran pointers and allocatables, the mapping of the pointed to
8014 // data by the descriptor (which itself, is a structure containing
8015 // runtime information on the dynamically allocated data).
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]);
8025 } else {
8026 auto mapOp = dyn_cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
8027 int firstMemberIdx = getMapDataMemberIdx(
8028 mapData, getFirstOrLastMappedMemberPtr(mapOp, true));
8029 lowAddr = builder.CreatePointerCast(mapData.BasePointers[firstMemberIdx],
8030 builder.getPtrTy());
8031
8032 int lastMemberIdx = getMapDataMemberIdx(
8033 mapData, getFirstOrLastMappedMemberPtr(mapOp, false));
8034 auto lastMemberMapInfo =
8035 cast<omp::MapInfoOp>(mapData.MapClause[lastMemberIdx]);
8036
8037 // NOTE: Currently, for RefPtee the BaseType is set to the varPtrPtr field,
8038 // which is the pointer datas type and not the member within the structure
8039 // that it's part of, so we have to make sure we use the member type in this
8040 // case when calculating the parents size offsets.
8041 // TODO: May be good to extend MapInfoData to support tracking of both
8042 // VarPtr/VarPtrPtr BaseType's to better distinguish what's being used more
8043 // consistently.
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];
8049 if (isRefPteeMap)
8050 castType =
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]);
8057 }
8058
8059 llvm::Value *size = builder.CreateIntCast(
8060 builder.CreatePtrDiff(builder.getInt8Ty(), highAddr, lowAddr),
8061 builder.getInt64Ty(),
8062 /*isSigned=*/false);
8063 combinedInfo.Sizes.push_back(size);
8064
8065 // This creates the initial MEMBER_OF mapping that consists of
8066 // the parent/top level container (same as above effectively, except
8067 // with a fixed initial compile time size and separate maptype which
8068 // indicates the true mape type (tofrom etc.). This parent mapping is
8069 // only relevant if the structure in its totality is being mapped,
8070 // otherwise the above suffices.
8071 if (!parentClause.getPartialMap()) {
8072 // TODO: This will need to be expanded to include the whole host of logic
8073 // for the map flags that Clang currently supports (e.g. it should do some
8074 // further case specific flag modifications). For the moment, it handles
8075 // what we support as expected.
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);
8080
8081 llvm::SmallVector<size_t> overlapIdxs;
8082 // Find all of the members that "overlap", i.e. occlude other members that
8083 // were mapped alongside the parent, e.g. member [0], occludes [0,1] and
8084 // [0,2], but not [1,0].
8085 getOverlappedMembers(overlapIdxs, parentClause);
8086
8087 // When we only have one overlap we skip the case that tries to segment the
8088 // mapping as best it can without creating holes, as the calculation is more
8089 // likely to have more overhead than anything we gain from mapping a smaller
8090 // chunk of data. This can be seen in cases where we are mapping Fortran
8091 // descriptors which are a special case of record type mapping.
8092 //
8093 // The cases for close and update are unique edge cases where the segmenting
8094 // does not play well with the runtime currently.
8095 if (targetDirective == TargetDirectiveEnumTy::TargetUpdate || hasMapClose ||
8096 overlapIdxs.size() == 1) {
8097 combinedInfo.Types.emplace_back(mapFlag);
8098 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
8099 combinedInfo.HasAttachPtr.emplace_back(false);
8100 combinedInfo.DevicePointers.emplace_back(
8101 mapData.DevicePointers[mapDataIndex]);
8102 combinedInfo.Names.emplace_back(LLVM::createMappingInformation(
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);
8109 } else {
8110 // We need to make sure the overlapped members are sorted in order of
8111 // lowest address to highest address.
8112 sortMapIndices(overlapIdxs, parentClause);
8113
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());
8120
8121 // Currently, the return parameter should be the over-riding parent in
8122 // cases where we have a return parameter that is echoed to all members,
8123 // the main case of this currently is with fortran descriptors. It may
8124 // need more finessing for C/C++ in the future or descriptors that are
8125 // members of derived types.
8126 mapFlag &= ~llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_RETURN_PARAM;
8127
8128 // TODO: We may want to skip arrays/array sections in this as Clang does.
8129 // It appears to be an optimisation rather than a necessity though,
8130 // but this requires further investigation. However, we would have to make
8131 // sure to not exclude maps with bounds that ARE pointers, as these are
8132 // processed as separate components, i.e. pointer + data.
8133 for (auto v : overlapIdxs) {
8134 auto mapDataOverlapIdx = getMapDataMemberIdx(
8135 mapData,
8136 cast<omp::MapInfoOp>(parentClause.getMembers()[v].getDefiningOp()));
8137 auto isPtrMap = checkIfPointerMap(
8138 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataOverlapIdx]));
8139 combinedInfo.Types.emplace_back(mapFlag);
8140 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
8141 combinedInfo.HasAttachPtr.emplace_back(false);
8142 combinedInfo.DevicePointers.emplace_back(
8143 llvm::OpenMPIRBuilder::DeviceInfoTy::None);
8144 combinedInfo.Names.emplace_back(LLVM::createMappingInformation(
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],
8153 lowAddr),
8154 builder.getInt64Ty(), /*isSigned=*/true);
8155 // In certain cases, we'll generate a size of 0 if we're not careful
8156 // (e.g. if lowAddr happens to be the first member), which isn't
8157 // correct, even if the runtimes is sometimes fine with it so, in these
8158 // scenarios we select the types size instead.
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);
8168 }
8169
8170 combinedInfo.Types.emplace_back(mapFlag);
8171 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
8172 combinedInfo.HasAttachPtr.emplace_back(false);
8173 combinedInfo.DevicePointers.emplace_back(
8174 llvm::OpenMPIRBuilder::DeviceInfoTy::None);
8175 combinedInfo.Names.emplace_back(LLVM::createMappingInformation(
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));
8184 }
8185 }
8186}
8187
8189 llvm::IRBuilderBase &builder,
8190 llvm::OpenMPIRBuilder &ompBuilder,
8191 DataLayout &dl, MapInfosTy &combinedInfo,
8192 MapInfoData &mapData, uint64_t mapDataIndex,
8193 TargetDirectiveEnumTy targetDirective) {
8194 assert(!ompBuilder.Config.isTargetDevice() &&
8195 "function only supported for host device codegen");
8196
8197 auto parentClause =
8198 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
8199
8200 // If we have a partial map (no parent referenced in the map clauses of the
8201 // directive, only members) and only a single member, we do not need to bind
8202 // the map of the member to the parent, we can pass the member separately.
8203 if (parentClause.getMembers().size() == 1 && parentClause.getPartialMap()) {
8204 auto memberClause = llvm::cast<omp::MapInfoOp>(
8205 parentClause.getMembers()[0].getDefiningOp());
8206 int memberDataIdx = getMapDataMemberIdx(mapData, memberClause);
8207 // Note: Clang treats arrays with explicit bounds that fall into this
8208 // category as a parent with map case, however, it seems this isn't a
8209 // requirement, and processing them as an individual map is fine. So,
8210 // we will handle them as individual maps for the moment, as it's
8211 // difficult for us to check this as we always require bounds to be
8212 // specified currently and it's also marginally more optimal (single
8213 // map rather than two). The difference may come from the fact that
8214 // Clang maps array without bounds as pointers (which we do not
8215 // currently do), whereas we treat them as arrays in all cases
8216 // currently.
8218 builder, ompBuilder, mapData, memberDataIdx, combinedInfo,
8219 targetDirective,
8220 /*MemberOfFlag=*/llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE,
8221 /*isTargetParam=*/true, mapDataIndex);
8222 return;
8223 }
8224
8225 auto collectMapInfoIdxs =
8226 [&](llvm::SmallVectorImpl<int64_t> &mapsAndInfoIdx) {
8227 auto parentClause =
8228 llvm::cast<omp::MapInfoOp>(mapData.MapClause[mapDataIndex]);
8229 mapsAndInfoIdx.push_back(getMapDataMemberIdx(mapData, parentClause));
8230 for (auto member : parentClause.getMembers())
8231 mapsAndInfoIdx.push_back(getMapDataMemberIdx(
8232 mapData, llvm::cast<omp::MapInfoOp>(member.getDefiningOp())));
8233 };
8234
8235 llvm::SmallVector<int64_t> mapInfoIdx;
8236 collectMapInfoIdxs(mapInfoIdx);
8237
8238 llvm::omp::OpenMPOffloadMappingFlags memberOfFlag =
8239 ompBuilder.getMemberOfFlag(combinedInfo.Types.size());
8240
8241 // The first index is the parent map, the rest are its members. The parent
8242 // normally undergoes the standard parent-with-members mapping, contributing
8243 // the MEMBER_OF flag that binds each member to it. The one exception is a
8244 // privatizeable attach map (a privatized pointer/descriptor passed directly
8245 // as a kernel argument): here the parent is emitted as an individual map
8246 // instead, for the time being, as it's used only in pointer/allocatable to
8247 // array cases for the moment. This only ever applies to the parent, so it is
8248 // checked once here rather than inside the loop below.
8249 bool parentIsPrivatizeableAttach =
8250 isPrivatizeableAttachMap(parentClause.getMapType());
8251 for (auto [i, idx] : llvm::enumerate(mapInfoIdx)) {
8252 bool emitParentMap = i == 0 && !parentIsPrivatizeableAttach;
8253 if (emitParentMap) {
8254 mapParentWithMembers(moduleTranslation, builder, ompBuilder, dl,
8255 combinedInfo, mapData, idx, memberOfFlag,
8256 targetDirective);
8257 } else {
8259 builder, ompBuilder, mapData, idx, combinedInfo, targetDirective,
8260 parentIsPrivatizeableAttach
8261 ? llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_NONE
8262 : memberOfFlag,
8263 /*isTargetParam=*/false, mapDataIndex);
8264 }
8265 }
8266}
8267
8268// This is a variation on Clang's GenerateOpenMPCapturedVars, which
8269// generates different operation (e.g. load/store) combinations for
8270// arguments to the kernel, based on map capture kinds which are then
8271// utilised in the combinedInfo in place of the original Map value.
8272static void
8273createAlteredByCaptureMap(MapInfoData &mapData,
8274 LLVM::ModuleTranslation &moduleTranslation,
8275 llvm::IRBuilderBase &builder) {
8276 assert(!moduleTranslation.getOpenMPBuilder()->Config.isTargetDevice() &&
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]);
8280 bool isAttachMap =
8281 ((convertClauseMapFlags(mapOp.getMapType()) &
8282 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH) ==
8283 llvm::omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH);
8284
8285 // If it's declare target, skip it, it's handled separately. However, if
8286 // it's declare target, and an attach map, we want to calculate the exact
8287 // address offset so that we attach correctly.
8288 if (!mapData.IsDeclareTarget[i] ||
8289 (mapData.IsDeclareTarget[i] && isAttachMap)) {
8290 omp::VariableCaptureKind captureKind = mapOp.getMapCaptureType();
8291 bool isPtrTy = checkIfPointerMap(mapOp);
8292
8293 // Currently handles array sectioning lowerbound case, but more
8294 // logic may be required in the future. Clang invokes EmitLValue,
8295 // which has specialised logic for special Clang types such as user
8296 // defines, so it is possible we will have to extend this for
8297 // structures or other complex types. As the general idea is that this
8298 // function mimics some of the logic from Clang that we require for
8299 // kernel argument passing from host -> device.
8300 switch (captureKind) {
8301 case omp::VariableCaptureKind::ByRef: {
8302 llvm::Value *newV = mapData.Pointers[i];
8303 std::vector<llvm::Value *> offsetIdx = calculateBoundsOffset(
8304 moduleTranslation, builder, mapData.BaseType[i]->isArrayTy(),
8305 mapOp.getBounds());
8306 if (isPtrTy)
8307 newV = builder.CreateLoad(builder.getPtrTy(), newV);
8308
8309 if (!offsetIdx.empty())
8310 newV = builder.CreateInBoundsGEP(mapData.BaseType[i], newV, offsetIdx,
8311 "array_offset");
8312 mapData.Pointers[i] = newV;
8313 } break;
8314 case omp::VariableCaptureKind::ByCopy: {
8315 llvm::Type *type = mapData.BaseType[i];
8316 llvm::Value *newV;
8317 if (mapData.Pointers[i]->getType()->isPointerTy())
8318 newV = builder.CreateLoad(type, mapData.Pointers[i]);
8319 else
8320 newV = mapData.Pointers[i];
8321
8322 if (!isPtrTy) {
8323 auto curInsert = builder.saveIP();
8324 llvm::DebugLoc DbgLoc = builder.getCurrentDebugLocation();
8325 builder.restoreIP(findAllocInsertPoints(builder, moduleTranslation));
8326 auto *memTempAlloc =
8327 builder.CreateAlloca(builder.getPtrTy(), nullptr, ".casted");
8328 builder.SetCurrentDebugLocation(DbgLoc);
8329 builder.restoreIP(curInsert);
8330
8331 builder.CreateStore(newV, memTempAlloc);
8332 newV = builder.CreateLoad(builder.getPtrTy(), memTempAlloc);
8333 }
8334
8335 mapData.Pointers[i] = newV;
8336 mapData.BasePointers[i] = newV;
8337 } break;
8338 case omp::VariableCaptureKind::This:
8339 case omp::VariableCaptureKind::VLAType:
8340 mapData.MapClause[i]->emitOpError("Unhandled capture kind");
8341 break;
8342 }
8343 }
8344 }
8345}
8346
8347// Generate all map related information and fill the combinedInfo.
8348static void genMapInfos(llvm::IRBuilderBase &builder,
8349 LLVM::ModuleTranslation &moduleTranslation,
8350 DataLayout &dl, MapInfosTy &combinedInfo,
8351 MapInfoData &mapData,
8352 TargetDirectiveEnumTy targetDirective) {
8353 assert(!moduleTranslation.getOpenMPBuilder()->Config.isTargetDevice() &&
8354 "function only supported for host device codegen");
8355 // We wish to modify some of the methods in which arguments are
8356 // passed based on their capture type by the target region, this can
8357 // involve generating new loads and stores, which changes the
8358 // MLIR value to LLVM value mapping, however, we only wish to do this
8359 // locally for the current function/target and also avoid altering
8360 // ModuleTranslation, so we remap the base pointer or pointer stored
8361 // in the map infos corresponding MapInfoData, which is later accessed
8362 // by genMapInfos and createTarget to help generate the kernel and
8363 // kernel arg structure. It primarily becomes relevant in cases like
8364 // bycopy, or byref range'd arrays. In the default case, we simply
8365 // pass thee pointer byref as both basePointer and pointer.
8366 createAlteredByCaptureMap(mapData, moduleTranslation, builder);
8367
8368 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
8369
8370 // We operate under the assumption that all vectors that are
8371 // required in MapInfoData are of equal lengths (either filled with
8372 // default constructed data or appropiate information) so we can
8373 // utilise the size from any component of MapInfoData, if we can't
8374 // something is missing from the initial MapInfoData construction.
8375 for (size_t i = 0; i < mapData.MapClause.size(); ++i) {
8376 if (mapData.IsAMember[i])
8377 continue;
8378
8379 auto mapInfoOp = dyn_cast<omp::MapInfoOp>(mapData.MapClause[i]);
8380 if (!mapInfoOp.getMembers().empty()) {
8381 processMapWithMembersOf(moduleTranslation, builder, *ompBuilder, dl,
8382 combinedInfo, mapData, i, targetDirective);
8383 continue;
8384 }
8385
8386 processIndividualMap(builder, *ompBuilder, mapData, i, combinedInfo,
8387 targetDirective);
8388 }
8389}
8390
8391static llvm::Expected<llvm::Function *>
8392emitUserDefinedMapper(Operation *declMapperOp, llvm::IRBuilderBase &builder,
8393 LLVM::ModuleTranslation &moduleTranslation,
8394 llvm::StringRef mapperFuncName,
8395 TargetDirectiveEnumTy targetDirective);
8396
8397static llvm::Expected<llvm::Function *>
8398getOrCreateUserDefinedMapperFunc(Operation *op, llvm::IRBuilderBase &builder,
8399 LLVM::ModuleTranslation &moduleTranslation,
8400 TargetDirectiveEnumTy targetDirective) {
8401 assert(!moduleTranslation.getOpenMPBuilder()->Config.isTargetDevice() &&
8402 "function only supported for host device codegen");
8403 auto declMapperOp = cast<omp::DeclareMapperOp>(op);
8404 std::string mapperFuncName =
8405 moduleTranslation.getOpenMPBuilder()->createPlatformSpecificName(
8406 {"omp_mapper", declMapperOp.getSymName()});
8407
8408 if (auto *lookupFunc = moduleTranslation.lookupFunction(mapperFuncName))
8409 return lookupFunc;
8410
8411 // Recursive types can cause re-entrant mapper emission. The mapper function
8412 // is created by OpenMPIRBuilder before the callbacks run, so it may already
8413 // exist in the LLVM module even though it is not yet registered in the
8414 // ModuleTranslation mapping table. Reuse and register it to break the
8415 // recursion.
8416 if (llvm::Function *existingFunc =
8417 moduleTranslation.getLLVMModule()->getFunction(mapperFuncName)) {
8418 moduleTranslation.mapFunction(mapperFuncName, existingFunc);
8419 return existingFunc;
8420 }
8421
8422 return emitUserDefinedMapper(declMapperOp, builder, moduleTranslation,
8423 mapperFuncName, targetDirective);
8424}
8425
8426static llvm::Expected<llvm::Function *>
8427emitUserDefinedMapper(Operation *op, llvm::IRBuilderBase &builder,
8428 LLVM::ModuleTranslation &moduleTranslation,
8429 llvm::StringRef mapperFuncName,
8430 TargetDirectiveEnumTy targetDirective) {
8431 assert(!moduleTranslation.getOpenMPBuilder()->Config.isTargetDevice() &&
8432 "function only supported for host device codegen");
8433 auto declMapperOp = cast<omp::DeclareMapperOp>(op);
8434 auto declMapperInfoOp = declMapperOp.getDeclareMapperInfo();
8435 if (failed(checkImplementationStatus(*declMapperInfoOp)))
8436 return llvm::make_error<PreviouslyReportedError>();
8437
8438 DataLayout dl = DataLayout(declMapperOp->getParentOfType<ModuleOp>());
8439 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
8440 llvm::Type *varType = moduleTranslation.convertType(declMapperOp.getType());
8441 SmallVector<Value> mapVars = declMapperInfoOp.getMapVars();
8442
8443 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
8444
8445 // Fill up the arrays with all the mapped variables.
8446 MapInfosTy combinedInfo;
8447 auto genMapInfoCB =
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(),
8455 /*ignoreArguments=*/true,
8456 builder)))
8457 return llvm::make_error<PreviouslyReportedError>();
8458 MapInfoData mapData;
8459 collectMapDataFromMapOperands(mapData, mapVars, moduleTranslation, dl,
8460 builder);
8461 genMapInfos(builder, moduleTranslation, dl, combinedInfo, mapData,
8462 targetDirective);
8463
8464 // Drop the mapping that is no longer necessary so that the same region
8465 // can be processed multiple times.
8466 moduleTranslation.forgetMapping(declMapperOp.getRegion());
8467 return combinedInfo;
8468 };
8469
8470 auto customMapperCB = [&](unsigned i) -> llvm::Expected<llvm::Function *> {
8471 if (!combinedInfo.Mappers[i])
8472 return nullptr;
8473 return getOrCreateUserDefinedMapperFunc(combinedInfo.Mappers[i], builder,
8474 moduleTranslation, targetDirective);
8475 };
8476
8477 llvm::Expected<llvm::Function *> newFn = ompBuilder->emitUserDefinedMapper(
8478 genMapInfoCB, varType, mapperFuncName, customMapperCB,
8479 /*PreserveMemberOfFlags=*/true);
8480 if (!newFn)
8481 return newFn.takeError();
8482 if ([[maybe_unused]] llvm::Function *mappedFunc =
8483 moduleTranslation.lookupFunction(mapperFuncName)) {
8484 assert(mappedFunc == *newFn &&
8485 "mapper function mapping disagrees with emitted function");
8486 } else {
8487 moduleTranslation.mapFunction(mapperFuncName, *newFn);
8488 }
8489 return *newFn;
8490}
8491
8492static llvm::Value *getSourceLocIdentFromOp(llvm::IRBuilderBase &builder,
8493 llvm::OpenMPIRBuilder &ompBuilder,
8494 Operation *op) {
8495 auto fileLoc = op->getLoc()->findInstanceOf<FileLineColLoc>();
8496 if (!fileLoc)
8497 return nullptr;
8498 uint32_t strSize;
8499 llvm::Function *parentFn = builder.GetInsertBlock()->getParent();
8500 llvm::StringRef fnName = parentFn ? parentFn->getName() : "";
8501 llvm::Constant *srcStr = LLVM::createSourceLocStrFromLocation(
8502 fileLoc, ompBuilder, fnName, strSize);
8503 return ompBuilder.getOrCreateIdent(srcStr, strSize);
8504}
8505
8506static LogicalResult
8507convertOmpTargetData(Operation *op, llvm::IRBuilderBase &builder,
8508 LLVM::ModuleTranslation &moduleTranslation) {
8509 llvm::Value *ifCond = nullptr;
8510 llvm::Value *deviceID = builder.getInt64(llvm::omp::OMP_DEVICEID_UNDEF);
8511 SmallVector<Value> mapVars;
8512 SmallVector<Value> useDevicePtrVars;
8513 SmallVector<Value> useDeviceAddrVars;
8514 llvm::omp::RuntimeFunction RTLFn;
8515 DataLayout DL = DataLayout(op->getParentOfType<ModuleOp>());
8516 TargetDirectiveEnumTy targetDirective = getTargetDirectiveEnumTyFromOp(op);
8517
8518 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
8519 llvm::OpenMPIRBuilder::TargetDataInfo info(
8520 /*RequiresDevicePointerInfo=*/true,
8521 /*SeparateBeginEndCalls=*/true);
8522
8523 if (ompBuilder->Config.isTargetDevice())
8524 return op->emitOpError() << "not allowed in a target device";
8525
8526 bool isOffloadEntry = !ompBuilder->Config.TargetTriples.empty();
8527
8528 auto getDeviceID = [&](mlir::Value dev) -> llvm::Value * {
8529 llvm::Value *v = moduleTranslation.lookupValue(dev);
8530 return builder.CreateIntCast(v, builder.getInt64Ty(), /*isSigned=*/true);
8531 };
8532
8533 LogicalResult result =
8535 .Case([&](omp::TargetDataOp dataOp) {
8536 if (failed(checkImplementationStatus(*dataOp)))
8537 return failure();
8538
8539 if (auto ifVar = dataOp.getIfExpr())
8540 ifCond = moduleTranslation.lookupValue(ifVar);
8541
8542 if (mlir::Value devId = dataOp.getDevice())
8543 deviceID = getDeviceID(devId);
8544
8545 mapVars = dataOp.getMapVars();
8546 useDevicePtrVars = dataOp.getUseDevicePtrVars();
8547 useDeviceAddrVars = dataOp.getUseDeviceAddrVars();
8548 return success();
8549 })
8550 .Case([&](omp::TargetEnterDataOp enterDataOp) -> LogicalResult {
8551 if (failed(checkImplementationStatus(*enterDataOp)))
8552 return failure();
8553
8554 if (auto ifVar = enterDataOp.getIfExpr())
8555 ifCond = moduleTranslation.lookupValue(ifVar);
8556
8557 if (mlir::Value devId = enterDataOp.getDevice())
8558 deviceID = getDeviceID(devId);
8559
8560 RTLFn =
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();
8566 return success();
8567 })
8568 .Case([&](omp::TargetExitDataOp exitDataOp) -> LogicalResult {
8569 if (failed(checkImplementationStatus(*exitDataOp)))
8570 return failure();
8571
8572 if (auto ifVar = exitDataOp.getIfExpr())
8573 ifCond = moduleTranslation.lookupValue(ifVar);
8574
8575 if (mlir::Value devId = exitDataOp.getDevice())
8576 deviceID = getDeviceID(devId);
8577
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();
8583 return success();
8584 })
8585 .Case([&](omp::TargetUpdateOp updateDataOp) -> LogicalResult {
8586 if (failed(checkImplementationStatus(*updateDataOp)))
8587 return failure();
8588
8589 if (auto ifVar = updateDataOp.getIfExpr())
8590 ifCond = moduleTranslation.lookupValue(ifVar);
8591
8592 if (mlir::Value devId = updateDataOp.getDevice())
8593 deviceID = getDeviceID(devId);
8594
8595 RTLFn =
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();
8601 return success();
8602 })
8603 .DefaultUnreachable("unexpected operation");
8604
8605 if (failed(result))
8606 return failure();
8607 // Pretend we have IF(false) if we're not doing offload.
8608 if (!isOffloadEntry)
8609 ifCond = builder.getFalse();
8610
8611 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
8612 MapInfoData mapData;
8613 collectMapDataFromMapOperands(mapData, mapVars, moduleTranslation, DL,
8614 builder, useDevicePtrVars, useDeviceAddrVars);
8615
8616 // Fill up the arrays with all the mapped variables.
8617 MapInfosTy combinedInfo;
8618 auto genMapInfoCB = [&](InsertPointTy codeGenIP) -> MapInfosTy & {
8619 builder.restoreIP(codeGenIP);
8620 genMapInfos(builder, moduleTranslation, DL, combinedInfo, mapData,
8621 targetDirective);
8622 return combinedInfo;
8623 };
8624
8625 // Define a lambda to apply mappings between use_device_addr and
8626 // use_device_ptr base pointers, and their associated block arguments.
8627 auto mapUseDevice =
8628 [&moduleTranslation](
8629 llvm::OpenMPIRBuilder::DeviceInfoTy type,
8631 llvm::SmallVectorImpl<Value> &useDeviceVars, MapInfoData &mapInfoData,
8632 llvm::function_ref<llvm::Value *(llvm::Value *)> mapper = nullptr) {
8633 for (auto [arg, useDevVar] :
8634 llvm::zip_equal(blockArgs, useDeviceVars)) {
8635
8636 auto getMapBasePtr = [](omp::MapInfoOp mapInfoOp) {
8637 return mapInfoOp.getVarPtrPtr() ? mapInfoOp.getVarPtrPtr()
8638 : mapInfoOp.getVarPtr();
8639 };
8640
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)
8648 continue;
8649
8650 if (llvm::Value *devPtrInfoMap =
8651 mapper ? mapper(basePointer) : basePointer) {
8652 moduleTranslation.mapValue(arg, devPtrInfoMap);
8653 break;
8654 }
8655 }
8656 }
8657 };
8658
8659 using BodyGenTy = llvm::OpenMPIRBuilder::BodyGenTy;
8660 auto bodyGenCB = [&](InsertPointTy codeGenIP, BodyGenTy bodyGenType)
8661 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
8662 // We must always restoreIP regardless of doing anything the caller
8663 // does not restore it, leading to incorrect (no) branch generation.
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 &region = cast<omp::TargetDataOp>(op).getRegion();
8669 switch (bodyGenType) {
8670 case BodyGenTy::Priv:
8671 // Check if any device ptr/addr info is available
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)
8678 return nullptr;
8679 return builder.CreateLoad(
8680 builder.getPtrTy(),
8681 info.DevicePtrInfoMap[basePointer].second);
8682 });
8683 mapUseDevice(llvm::OpenMPIRBuilder::DeviceInfoTy::Pointer,
8684 blockArgIface.getUseDevicePtrBlockArgs(), useDevicePtrVars,
8685 mapData, [&](llvm::Value *basePointer) {
8686 return info.DevicePtrInfoMap[basePointer].second;
8687 });
8688
8689 if (failed(inlineConvertOmpRegions(region, "omp.data.region", builder,
8690 moduleTranslation)))
8691 return llvm::make_error<PreviouslyReportedError>();
8692 }
8693 break;
8694 case BodyGenTy::DupNoPriv:
8695 if (info.DevicePtrInfoMap.empty()) {
8696 // For host device we still need to do the mapping for codegen,
8697 // otherwise it may try to lookup a missing value.
8698 mapUseDevice(llvm::OpenMPIRBuilder::DeviceInfoTy::Address,
8699 blockArgIface.getUseDeviceAddrBlockArgs(),
8700 useDeviceAddrVars, mapData);
8701 mapUseDevice(llvm::OpenMPIRBuilder::DeviceInfoTy::Pointer,
8702 blockArgIface.getUseDevicePtrBlockArgs(), useDevicePtrVars,
8703 mapData);
8704 }
8705 break;
8706 case BodyGenTy::NoPriv:
8707 // If device info is available then region has already been generated
8708 if (info.DevicePtrInfoMap.empty()) {
8709 if (failed(inlineConvertOmpRegions(region, "omp.data.region", builder,
8710 moduleTranslation)))
8711 return llvm::make_error<PreviouslyReportedError>();
8712 }
8713 break;
8714 }
8715 return builder.saveIP();
8716 };
8717
8718 auto customMapperCB =
8719 [&](unsigned int i) -> llvm::Expected<llvm::Function *> {
8720 if (!combinedInfo.Mappers[i])
8721 return nullptr;
8722 info.HasMapper = true;
8723 return getOrCreateUserDefinedMapperFunc(combinedInfo.Mappers[i], builder,
8724 moduleTranslation, targetDirective);
8725 };
8726
8727 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
8729 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
8730 findAllocInsertPoints(builder, moduleTranslation, &deallocBlocks);
8731
8732 // Pass the region's source location to the runtime, taken from the op's own
8733 // location; only offloading entries emit the mapper calls that consume it.
8734 // No need to also guard on !isTargetDevice here: this function bails out
8735 // earlier for the target device, so isOffloadEntry alone is sufficient.
8736 llvm::Value *srcLocOverride =
8737 isOffloadEntry ? getSourceLocIdentFromOp(builder, *ompBuilder, op)
8738 : nullptr;
8739
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,
8745 /*MapperFunc=*/nullptr, bodyGenCB,
8746 /*DeviceAddrCB=*/nullptr, srcLocOverride);
8747 return ompBuilder->createTargetData(
8748 ompLoc, allocaIP, builder.saveIP(), deallocBlocks, deviceID, ifCond,
8749 info, genMapInfoCB, customMapperCB, &RTLFn,
8750 /*BodyGenCB=*/nullptr,
8751 /*DeviceAddrCB=*/nullptr, srcLocOverride);
8752 }();
8753
8754 if (failed(handleError(afterIP, *op)))
8755 return failure();
8756
8757 builder.restoreIP(*afterIP);
8758 return success();
8759}
8760
8761static LogicalResult
8762convertOmpDistribute(Operation &opInst, llvm::IRBuilderBase &builder,
8763 LLVM::ModuleTranslation &moduleTranslation) {
8764 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
8765 auto distributeOp = cast<omp::DistributeOp>(opInst);
8766 if (failed(checkImplementationStatus(opInst)))
8767 return failure();
8768
8769 /// Process teams op reduction in distribute if the reduction is contained in
8770 /// this specific distribute op.
8771 omp::TeamsOp teamsOp = opInst.getParentOfType<omp::TeamsOp>();
8772 bool doDistributeReduction =
8773 teamsOp && getDistributeCapturingTeamsReduction(teamsOp) == distributeOp;
8774
8775 DenseMap<Value, llvm::Value *> reductionVariableMap;
8776 unsigned numReductionVars = teamsOp ? teamsOp.getNumReductionVars() : 0;
8778 SmallVector<llvm::Value *> privateReductionVariables(numReductionVars);
8779 llvm::ArrayRef<bool> isByRef;
8780
8781 if (doDistributeReduction) {
8782 isByRef = getIsByRef(teamsOp.getReductionByref());
8783 assert(isByRef.size() == teamsOp.getNumReductionVars());
8784
8785 collectReductionDecls(teamsOp, reductionDecls);
8786 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
8787 findAllocInsertPoints(builder, moduleTranslation);
8788
8789 MutableArrayRef<BlockArgument> reductionArgs =
8790 llvm::cast<omp::BlockArgOpenMPOpInterface>(*teamsOp)
8791 .getReductionBlockArgs();
8792
8794 teamsOp, reductionArgs, builder, moduleTranslation, allocaIP,
8795 reductionDecls, privateReductionVariables, reductionVariableMap,
8796 isByRef)))
8797 return failure();
8798 }
8799
8800 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
8801 auto bodyGenCB =
8802 [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
8803 llvm::ArrayRef<llvm::BasicBlock *> deallocBlocks) -> llvm::Error {
8804 // DistributeOp has only one region associated with it.
8805 builder.restoreIP(codeGenIP);
8806 PrivateVarsInfo privVarsInfo(distributeOp);
8807
8809 distributeOp, builder, moduleTranslation, privVarsInfo, allocaIP);
8810 if (handleError(afterAllocas, opInst).failed())
8811 return llvm::make_error<PreviouslyReportedError>();
8812
8813 // Save the alloca insertion point on ModuleTranslation stack for use in
8814 // nested regions. This must be after allocatePrivateVars, which splits
8815 // the alloca block and updates allocaIP.
8817 moduleTranslation, allocaIP, deallocBlocks);
8818
8819 if (handleError(initPrivateVars(builder, moduleTranslation, privVarsInfo),
8820 opInst)
8821 .failed())
8822 return llvm::make_error<PreviouslyReportedError>();
8823
8824 if (failed(copyFirstPrivateVars(
8825 distributeOp, builder, moduleTranslation, privVarsInfo.mlirVars,
8826 privVarsInfo.llvmVars, privVarsInfo.privatizers,
8827 distributeOp.getPrivateNeedsBarrier())))
8828 return llvm::make_error<PreviouslyReportedError>();
8829
8830 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
8831 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
8833 convertOmpOpRegions(distributeOp.getRegion(), "omp.distribute.region",
8834 builder, moduleTranslation);
8835 if (!regionBlock)
8836 return regionBlock.takeError();
8837 builder.SetInsertPoint((*regionBlock)->begin());
8838
8839 // Skip applying a workshare loop below when translating 'distribute
8840 // parallel do' (it's been already handled by this point while translating
8841 // the nested omp.wsloop).
8842 if (!isa_and_present<omp::WsloopOp>(distributeOp.getNestedWrapper())) {
8843 // TODO: Add support for clauses which are valid for DISTRIBUTE
8844 // constructs. Static schedule is the default.
8845 bool hasDistSchedule = distributeOp.getDistScheduleStatic();
8846 auto schedule = hasDistSchedule ? omp::ClauseScheduleKind::Distribute
8847 : omp::ClauseScheduleKind::Static;
8848 // dist_schedule clauses are ordered - otherise this should be false
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 =
8858 findCurrentLoopInfo(moduleTranslation);
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);
8866
8867 if (!wsloopIP)
8868 return wsloopIP.takeError();
8869 }
8870 if (failed(cleanupPrivateVars(distributeOp, builder, moduleTranslation,
8871 distributeOp.getLoc(), privVarsInfo)))
8872 return llvm::make_error<PreviouslyReportedError>();
8873
8874 return llvm::Error::success();
8875 };
8876
8878 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
8879 findAllocInsertPoints(builder, moduleTranslation, &deallocBlocks);
8880 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
8881 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
8882 ompBuilder->createDistribute(ompLoc, allocaIP, deallocBlocks, bodyGenCB);
8883
8884 if (failed(handleError(afterIP, opInst)))
8885 return failure();
8886
8887 builder.restoreIP(*afterIP);
8888
8889 if (doDistributeReduction) {
8890 // Process the reductions if required.
8892 teamsOp, builder, moduleTranslation, allocaIP, reductionDecls,
8893 privateReductionVariables, isByRef,
8894 /*isNoWait*/ false, /*isTeamsReduction*/ true);
8895 }
8896 return success();
8897}
8898
8899/// Lowers the FlagsAttr which is applied to the module when offloading. This
8900/// attribute contains OpenMP RTL globals that can be passed as flags to the
8901/// frontend, otherwise they are set to default
8902static LogicalResult
8903convertFlagsAttr(Operation *op, mlir::omp::FlagsAttr attribute,
8904 LLVM::ModuleTranslation &moduleTranslation) {
8905 auto offloadMod = dyn_cast<omp::OffloadModuleInterface>(op);
8906 if (!offloadMod)
8907 return op->emitOpError() << "omp flags attached to non offload module op";
8908
8909 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
8910
8911 if (offloadMod.getIsTargetDevice())
8912 ompBuilder->M.addModuleFlag(llvm::Module::Max, "openmp-device",
8913 attribute.getOpenmpDeviceVersion());
8914
8915 // The flags below are only intended to be emitted for GPU offload targets.
8916 if (!offloadMod.getIsGPU())
8917 return success();
8918
8919 if (attribute.getNoGpuLib())
8920 return success();
8921
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");
8932 return success();
8933}
8934
8935static void getTargetEntryUniqueInfo(llvm::TargetRegionEntryInfo &targetInfo,
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");
8942
8943 auto fileInfoCallBack = [&fileLoc]() {
8944 return std::pair<std::string, uint64_t>(
8945 llvm::StringRef(fileLoc.getFilename()), fileLoc.getLine());
8946 };
8947
8948 targetInfo =
8949 ompBuilder.getTargetEntryUniqueInfo(fileInfoCallBack, vfs, parentName);
8950}
8951
8952// The createDeviceArgumentAccessor function generates
8953// instructions for retrieving (acessing) kernel
8954// arguments inside of the device kernel for use by
8955// the kernel. This enables different semantics such as
8956// the creation of temporary copies of data allowing
8957// semantics like read-only/no host write back kernel
8958// arguments.
8959//
8960// This currently implements a very light version of Clang's
8961// EmitParmDecl's handling of direct argument handling as well
8962// as a portion of the argument access generation based on
8963// capture types found at the end of emitOutlinedFunctionPrologue
8964// in Clang. The indirect path handling of EmitParmDecl's may be
8965// required for future work, but a direct 1-to-1 copy doesn't seem
8966// possible as the logic is rather scattered throughout Clang's
8967// lowering and perhaps we wish to deviate slightly.
8968//
8969// \param mapData - A container containing vectors of information
8970// corresponding to the input argument, which should have a
8971// corresponding entry in the MapInfoData containers
8972// OrigialValue's.
8973// \param arg - This is the generated kernel function argument that
8974// corresponds to the passed in input argument. We generated different
8975// accesses of this Argument, based on capture type and other Input
8976// related information.
8977// \param input - This is the host side value that will be passed to
8978// the kernel i.e. the kernel input, we rewrite all uses of this within
8979// the kernel (as we generate the kernel body based on the target's region
8980// which maintians references to the original input) to the retVal argument
8981// apon exit of this function inside of the OMPIRBuilder. This interlinks
8982// the kernel argument to future uses of it in the function providing
8983// appropriate "glue" instructions inbetween.
8984// \param retVal - This is the value that all uses of input inside of the
8985// kernel will be re-written to, the goal of this function is to generate
8986// an appropriate location for the kernel argument to be accessed from,
8987// e.g. ByRef will result in a temporary allocation location and then
8988// a store of the kernel argument into this allocated memory which
8989// will then be loaded from, ByCopy will use the allocated memory
8990// directly.
8991static llvm::IRBuilderBase::InsertPoint createDeviceArgumentAccessor(
8992 omp::TargetOp targetOp, MapInfoData &mapData, llvm::Argument &arg,
8993 llvm::Value *input, llvm::Value *&retVal, llvm::IRBuilderBase &builder,
8994 llvm::OpenMPIRBuilder &ompBuilder,
8995 LLVM::ModuleTranslation &moduleTranslation,
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);
9002
9003 omp::VariableCaptureKind capture = omp::VariableCaptureKind::ByRef;
9004 LLVM::TypeToLLVMIRTranslator typeToLLVMIRTranslator(
9005 ompBuilder.M.getContext());
9006 unsigned alignmentValue = 0;
9007 BlockArgument mlirArg;
9009 cast<omp::BlockArgOpenMPOpInterface>(*targetOp).getBlockArgsPairs(
9010 blockArgsPairs);
9011 // Find the associated MapInfoData entry for the current input
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();
9016 // Get information of alignment of mapped object
9017 alignmentValue = typeToLLVMIRTranslator.getPreferredAlignment(
9018 mapOp.getVarPtrType(), ompBuilder.M.getDataLayout());
9019
9020 // Find the corresponding entry block argument, which can be associated to
9021 // a map, use_device* or has_device* clause.
9022 for (auto &[val, arg] : blockArgsPairs) {
9023 if (mapOp.getResult() == val) {
9024 mlirArg = arg;
9025 break;
9026 }
9027 }
9028 assert(mlirArg && "expected to find entry block argument for map clause");
9029 break;
9030 }
9031 }
9032
9033 unsigned int allocaAS = ompBuilder.M.getDataLayout().getAllocaAddrSpace();
9034 unsigned int defaultAS =
9035 ompBuilder.M.getDataLayout().getProgramAddressSpace();
9036
9037 // Create the allocation for the argument.
9038 llvm::Value *v = nullptr;
9039 if (omp::opInSharedDeviceContext(*targetOp) &&
9041 // Use the beginning of the codeGenIP rather than the usual allocation point
9042 // for shared memory allocations because otherwise these would be done prior
9043 // to the target initialization call. Also, the exit block (where the
9044 // deallocation is placed) is only executed if the initialization call
9045 // succeeds.
9046 builder.SetInsertPoint(codeGenIP.getNodeParent()->getFirstInsertionPt());
9047 v = ompBuilder.createOMPAllocShared(builder, arg.getType());
9048
9049 // Create deallocations in all provided deallocation points and then restore
9050 // the insertion point to right after the new allocations.
9051 llvm::IRBuilderBase::InsertPointGuard guard(builder);
9052 for (auto deallocIP : deallocIPs) {
9053 builder.SetInsertPoint(deallocIP);
9054 ompBuilder.createOMPFreeShared(builder, v, arg.getType());
9055 }
9056 } else {
9057 // Use the current point, which was previously set to allocaIP.
9058 v = builder.CreateAlloca(arg.getType(), allocaAS);
9059
9060 if (allocaAS != defaultAS && arg.getType()->isPointerTy())
9061 v = builder.CreateAddrSpaceCast(v, builder.getPtrTy(defaultAS));
9062 }
9063
9064 builder.CreateStore(&arg, v);
9065
9066 builder.restoreIP(codeGenIP);
9067
9068 switch (capture) {
9069 case omp::VariableCaptureKind::ByCopy: {
9070 retVal = v;
9071 break;
9072 }
9073 case omp::VariableCaptureKind::ByRef: {
9074 llvm::LoadInst *loadInst = builder.CreateAlignedLoad(
9075 v->getType(), v,
9076 ompBuilder.M.getDataLayout().getPrefTypeAlign(v->getType()));
9077 // CreateAlignedLoad function creates similar LLVM IR:
9078 // %res = load ptr, ptr %input, align 8
9079 // This LLVM IR does not contain information about alignment
9080 // of the loaded value. We need to add !align metadata to unblock
9081 // optimizer. The existence of the !align metadata on the instruction
9082 // tells the optimizer that the value loaded is known to be aligned to
9083 // a boundary specified by the integer value in the metadata node.
9084 // Example:
9085 // %res = load ptr, ptr %input, align 8, !align !align_md_node
9086 // ^ ^
9087 // | |
9088 // alignment of %input address |
9089 // |
9090 // alignment of %res object
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()),
9098 alignmentValue))));
9099 }
9100 retVal = loadInst;
9101
9102 break;
9103 }
9104 case omp::VariableCaptureKind::This:
9105 case omp::VariableCaptureKind::VLAType:
9106 // TODO: Consider returning error to use standard reporting for
9107 // unimplemented features.
9108 assert(false && "Currently unsupported capture kind");
9109 break;
9110 }
9111
9112 return builder.saveIP();
9113}
9114
9115/// Follow uses of `host_eval`-defined block arguments of the given `omp.target`
9116/// operation and populate output variables with their corresponding host value
9117/// (i.e. operand evaluated outside of the target region), based on their uses
9118/// inside of the target region.
9119///
9120/// Loop bounds and steps are only optionally populated, if output vectors are
9121/// provided.
9122static void
9123extractHostEvalClauses(omp::TargetOp targetOp, Value &numThreads,
9124 Value &numTeamsLower, Value &numTeamsUpper,
9125 Value &threadLimit,
9126 llvm::SmallVectorImpl<Value> *lowerBounds = nullptr,
9127 llvm::SmallVectorImpl<Value> *upperBounds = nullptr,
9128 llvm::SmallVectorImpl<Value> *steps = nullptr) {
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);
9133
9134 for (Operation *user : blockArg.getUsers()) {
9136 .Case([&](omp::TeamsOp teamsOp) {
9137 if (teamsOp.getNumTeamsLower() == blockArg)
9138 numTeamsLower = hostEvalVar;
9139 else if (llvm::is_contained(teamsOp.getNumTeamsUpperVars(),
9140 blockArg))
9141 numTeamsUpper = hostEvalVar;
9142 else if (!teamsOp.getThreadLimitVars().empty() &&
9143 teamsOp.getThreadLimit(0) == blockArg)
9144 threadLimit = hostEvalVar;
9145 else
9146 llvm_unreachable("unsupported host_eval use");
9147 })
9148 .Case([&](omp::ParallelOp parallelOp) {
9149 if (!parallelOp.getNumThreadsVars().empty() &&
9150 parallelOp.getNumThreads(0) == blockArg)
9151 numThreads = hostEvalVar;
9152 else
9153 llvm_unreachable("unsupported host_eval use");
9154 })
9155 .Case([&](omp::LoopNestOp loopOp) {
9156 auto processBounds =
9157 [&](OperandRange opBounds,
9158 llvm::SmallVectorImpl<Value> *outBounds) -> bool {
9159 bool found = false;
9160 for (auto [i, lb] : llvm::enumerate(opBounds)) {
9161 if (lb == blockArg) {
9162 found = true;
9163 if (outBounds)
9164 (*outBounds)[i] = hostEvalVar;
9165 }
9166 }
9167 return found;
9168 };
9169 bool found =
9170 processBounds(loopOp.getLoopLowerBounds(), lowerBounds);
9171 found = processBounds(loopOp.getLoopUpperBounds(), upperBounds) ||
9172 found;
9173 found = processBounds(loopOp.getLoopSteps(), steps) || found;
9174 (void)found;
9175 assert(found && "unsupported host_eval use");
9176 })
9177 .DefaultUnreachable("unsupported host_eval use");
9178 }
9179 }
9180}
9181
9182/// If \p op is of the given type parameter, return it casted to that type.
9183/// Otherwise, if its immediate parent operation (or some other higher-level
9184/// parent, if \p immediateParent is false) is of that type, return that parent
9185/// casted to the given type.
9186///
9187/// If \p op is \c null or neither it or its parent(s) are of the specified
9188/// type, return a \c null operation.
9189template <typename OpTy>
9190static OpTy castOrGetParentOfType(Operation *op, bool immediateParent = false) {
9191 if (!op)
9192 return OpTy();
9193
9194 if (OpTy casted = dyn_cast<OpTy>(op))
9195 return casted;
9196
9197 if (immediateParent)
9198 return dyn_cast_if_present<OpTy>(op->getParentOp());
9199
9200 return op->getParentOfType<OpTy>();
9201}
9202
9203/// If the given \p value is defined by an \c llvm.mlir.constant operation and
9204/// it is of an integer type, return its value.
9205static std::optional<int64_t> extractConstInteger(Value value) {
9206 if (!value)
9207 return std::nullopt;
9208
9209 if (auto constOp = value.getDefiningOp<LLVM::ConstantOp>())
9210 if (auto constAttr = dyn_cast<IntegerAttr>(constOp.getValue()))
9211 return constAttr.getInt();
9212
9213 return std::nullopt;
9214}
9215
9216static uint64_t getTypeByteSize(mlir::Type type, const DataLayout &dl) {
9217 uint64_t sizeInBits = dl.getTypeSizeInBits(type);
9218 uint64_t sizeInBytes = sizeInBits / 8;
9219 return sizeInBytes;
9220}
9221
9222template <typename OpTy>
9223static uint64_t getReductionDataSize(OpTy &op) {
9224 if (op.getNumReductionVars() > 0) {
9226 collectReductionDecls(op, reductions);
9227
9229 members.reserve(reductions.size());
9230 for (omp::DeclareReductionOp &red : reductions) {
9231 // For by-ref reductions, use the actual element type rather than the
9232 // pointer type so that the buffer size matches the access pattern in
9233 // the copy/reduce callbacks generated by OMPIRBuilder.
9234 if (red.getByrefElementType())
9235 members.push_back(*red.getByrefElementType());
9236 else
9237 members.push_back(red.getType());
9238 }
9239 Operation *opp = op.getOperation();
9240 auto structType = mlir::LLVM::LLVMStructType::getLiteral(
9241 opp->getContext(), members, /*isPacked=*/false);
9242 DataLayout dl = DataLayout(opp->getParentOfType<ModuleOp>());
9243 return getTypeByteSize(structType, dl);
9244 }
9245 return 0;
9246}
9247
9248/// Populate default `MinTeams`, `MaxTeams` and `MaxThreads` to their default
9249/// values as stated by the corresponding clauses, if constant.
9250///
9251/// These default values must be set before the creation of the outlined LLVM
9252/// function for the target region, so that they can be used to initialize the
9253/// corresponding global `ConfigurationEnvironmentTy` structure.
9254static void
9255initTargetDefaultAttrs(omp::TargetOp targetOp, Operation *capturedOp,
9256 llvm::OpenMPIRBuilder::TargetKernelDefaultAttrs &attrs,
9257 bool isTargetDevice, bool isGPU) {
9258 // TODO: Handle constant 'if' clauses.
9259
9260 Value numThreads, numTeamsLower, numTeamsUpper, threadLimit;
9261 if (!isTargetDevice) {
9262 extractHostEvalClauses(targetOp, numThreads, numTeamsLower, numTeamsUpper,
9263 threadLimit);
9264 } else {
9265 // In the target device, values for these clauses are not passed as
9266 // host_eval, but instead evaluated prior to entry to the region. This
9267 // ensures values are mapped and available inside of the target region.
9268 if (auto teamsOp = castOrGetParentOfType<omp::TeamsOp>(capturedOp)) {
9269 numTeamsLower = teamsOp.getNumTeamsLower();
9270 // Handle num_teams upper bounds (only first value for now)
9271 if (!teamsOp.getNumTeamsUpperVars().empty())
9272 numTeamsUpper = teamsOp.getNumTeams(0);
9273 if (!teamsOp.getThreadLimitVars().empty())
9274 threadLimit = teamsOp.getThreadLimit(0);
9275 }
9276
9277 if (auto parallelOp = castOrGetParentOfType<omp::ParallelOp>(capturedOp)) {
9278 if (!parallelOp.getNumThreadsVars().empty())
9279 numThreads = parallelOp.getNumThreads(0);
9280 }
9281 }
9282
9283 // Handle clauses impacting the number of teams.
9284
9285 int32_t minTeamsVal = 1, maxTeamsVal = -1;
9286 if (castOrGetParentOfType<omp::TeamsOp>(capturedOp)) {
9287 // TODO: Use `hostNumTeamsLower` to initialize `minTeamsVal`. For now,
9288 // match clang and set min and max to the same value.
9289 if (numTeamsUpper) {
9290 if (auto val = extractConstInteger(numTeamsUpper))
9291 minTeamsVal = maxTeamsVal = *val;
9292 } else {
9293 minTeamsVal = maxTeamsVal = 0;
9294 }
9295 } else if (castOrGetParentOfType<omp::ParallelOp>(capturedOp,
9296 /*immediateParent=*/true) ||
9298 /*immediateParent=*/true)) {
9299 minTeamsVal = maxTeamsVal = 1;
9300 } else {
9301 minTeamsVal = maxTeamsVal = -1;
9302 }
9303
9304 // Handle clauses impacting the number of threads.
9305
9306 auto setMaxValueFromClause = [](Value clauseValue, int32_t &result) {
9307 if (!clauseValue)
9308 return;
9309
9310 if (auto val = extractConstInteger(clauseValue))
9311 result = *val;
9312
9313 // Found an applicable clause, so it's not undefined. Mark as unknown
9314 // because it's not constant.
9315 if (result < 0)
9316 result = 0;
9317 };
9318
9319 // Extract 'thread_limit' clause from 'target' and 'teams' directives.
9320 int32_t targetThreadLimitVal = -1, teamsThreadLimitVal = -1;
9321 if (!targetOp.getThreadLimitVars().empty())
9322 setMaxValueFromClause(targetOp.getThreadLimit(0), targetThreadLimitVal);
9323 setMaxValueFromClause(threadLimit, teamsThreadLimitVal);
9324
9325 // Extract 'max_threads' clause from 'parallel' or set to 1 if it's SIMD.
9326 int32_t maxThreadsVal = -1;
9328 setMaxValueFromClause(numThreads, maxThreadsVal);
9329 else if (castOrGetParentOfType<omp::SimdOp>(capturedOp,
9330 /*immediateParent=*/true))
9331 maxThreadsVal = 1;
9332
9333 // For max values, < 0 means unset, == 0 means set but unknown. Select the
9334 // minimum value between 'max_threads' and 'thread_limit' clauses that were
9335 // set.
9336 int32_t combinedMaxThreadsVal = targetThreadLimitVal;
9337 if (combinedMaxThreadsVal < 0 ||
9338 (teamsThreadLimitVal >= 0 && teamsThreadLimitVal < combinedMaxThreadsVal))
9339 combinedMaxThreadsVal = teamsThreadLimitVal;
9340
9341 if (combinedMaxThreadsVal < 0 ||
9342 (maxThreadsVal >= 0 && maxThreadsVal < combinedMaxThreadsVal))
9343 combinedMaxThreadsVal = maxThreadsVal;
9344
9345 int32_t reductionDataSize = 0;
9346 if (isGPU && capturedOp) {
9347 if (auto teamsOp = castOrGetParentOfType<omp::TeamsOp>(capturedOp))
9348 reductionDataSize = getReductionDataSize(teamsOp);
9349 }
9350
9351 // Update kernel bounds structure for the `OpenMPIRBuilder` to use.
9352 // Use the kernel_type attribute set by the frontend instead of analyzing IR.
9353 omp::TargetExecMode execMode = targetOp.getKernelType();
9354 switch (execMode) {
9355 case omp::TargetExecMode::bare:
9356 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_BARE;
9357 break;
9358 case omp::TargetExecMode::generic:
9359 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_GENERIC;
9360 break;
9361 case omp::TargetExecMode::spmd:
9362 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_SPMD;
9363 break;
9364 case omp::TargetExecMode::spmd_no_loop:
9365 attrs.ExecFlags = llvm::omp::OMP_TGT_EXEC_MODE_SPMD_NO_LOOP;
9366 break;
9367 }
9368 attrs.MinTeams.front() = minTeamsVal;
9369 attrs.MaxTeams.front() = maxTeamsVal;
9370 attrs.MinThreads.front() = 1;
9371 attrs.MaxThreads.front() = combinedMaxThreadsVal;
9372 attrs.ReductionDataSize = reductionDataSize;
9373}
9374
9375/// Gather LLVM runtime values for all clauses evaluated in the host that are
9376/// passed to the kernel invocation.
9377///
9378/// This function must be called only when compiling for the host. Also, it will
9379/// only provide correct results if it's called after the body of \c targetOp
9380/// has been fully generated.
9381static void
9382initTargetRuntimeAttrs(llvm::IRBuilderBase &builder,
9383 LLVM::ModuleTranslation &moduleTranslation,
9384 omp::TargetOp targetOp, Operation *capturedOp,
9385 llvm::OpenMPIRBuilder::TargetKernelRuntimeAttrs &attrs) {
9386 omp::LoopNestOp loopOp = castOrGetParentOfType<omp::LoopNestOp>(capturedOp);
9387 unsigned numLoops = loopOp ? loopOp.getNumLoops() : 0;
9388
9389 Value numThreads, numTeamsLower, numTeamsUpper, teamsThreadLimit;
9390 llvm::SmallVector<Value> lowerBounds(numLoops), upperBounds(numLoops),
9391 steps(numLoops);
9392 extractHostEvalClauses(targetOp, numThreads, numTeamsLower, numTeamsUpper,
9393 teamsThreadLimit, &lowerBounds, &upperBounds, &steps);
9394
9395 // TODO: Handle constant 'if' clauses.
9396 if (!targetOp.getThreadLimitVars().empty()) {
9397 Value targetThreadLimit = targetOp.getThreadLimit(0);
9398 attrs.TargetThreadLimit.front() =
9399 moduleTranslation.lookupValue(targetThreadLimit);
9400 }
9401
9402 // The __kmpc_push_num_teams_51 function expects int32 as the arguments. So,
9403 // truncate or sign extend lower and upper num_teams bounds as well as
9404 // thread_limit to match int32 ABI requirements for the OpenMP runtime.
9405 if (numTeamsLower)
9406 attrs.MinTeams.front() = builder.CreateSExtOrTrunc(
9407 moduleTranslation.lookupValue(numTeamsLower), builder.getInt32Ty());
9408
9409 if (numTeamsUpper)
9410 attrs.MaxTeams.front() = builder.CreateSExtOrTrunc(
9411 moduleTranslation.lookupValue(numTeamsUpper), builder.getInt32Ty());
9412
9413 if (teamsThreadLimit)
9414 attrs.TeamsThreadLimit.front() = builder.CreateSExtOrTrunc(
9415 moduleTranslation.lookupValue(teamsThreadLimit), builder.getInt32Ty());
9416
9417 if (numThreads)
9418 attrs.MaxThreads.front() = moduleTranslation.lookupValue(numThreads);
9419
9420 if (targetOp.hasHostEvalTripCount()) {
9421 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
9422 attrs.LoopTripCount = nullptr;
9423
9424 // To calculate the trip count, we multiply together the trip counts of
9425 // every collapsed canonical loop. We don't need to create the loop nests
9426 // here, since we're only interested in the trip count.
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);
9432
9433 if (!lowerBound || !upperBound || !step) {
9434 attrs.LoopTripCount = nullptr;
9435 break;
9436 }
9437
9438 llvm::OpenMPIRBuilder::LocationDescription loc(builder);
9439 llvm::Value *tripCount = ompBuilder->calculateCanonicalLoopTripCount(
9440 loc, lowerBound, upperBound, step, /*IsSigned=*/true,
9441 loopOp.getLoopInclusive());
9442
9443 if (!attrs.LoopTripCount) {
9444 attrs.LoopTripCount = tripCount;
9445 continue;
9446 }
9447
9448 // TODO: Enable UndefinedSanitizer to diagnose an overflow here.
9449 attrs.LoopTripCount = builder.CreateMul(attrs.LoopTripCount, tripCount,
9450 {}, /*HasNUW=*/true);
9451 }
9452 }
9453
9454 attrs.DeviceID = builder.getInt64(llvm::omp::OMP_DEVICEID_UNDEF);
9455 if (mlir::Value devId = targetOp.getDevice()) {
9456 attrs.DeviceID = moduleTranslation.lookupValue(devId);
9457 attrs.DeviceID =
9458 builder.CreateSExtOrTrunc(attrs.DeviceID, builder.getInt64Ty());
9459 }
9460}
9461
9462static llvm::omp::OMPDynGroupprivateFallbackType
9463getDynGroupprivateFallbackType(omp::FallbackModifierAttr fallbackAttr) {
9464 omp::FallbackModifier fb = fallbackAttr ? fallbackAttr.getValue()
9465 : omp::FallbackModifier::default_mem;
9466 switch (fb) {
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;
9473 }
9474
9475 llvm_unreachable("unexpected dyn_groupprivate fallback type");
9476}
9477
9478static LogicalResult
9479convertOmpTarget(Operation &opInst, llvm::IRBuilderBase &builder,
9480 LLVM::ModuleTranslation &moduleTranslation) {
9481 auto targetOp = cast<omp::TargetOp>(opInst);
9482
9483 // The current debug location already has the DISubprogram for the outlined
9484 // function that will be created for the target op. We save it here so that
9485 // we can set it on the outlined function.
9486 llvm::DebugLoc outlinedFnLoc = builder.getCurrentDebugLocation();
9487 if (failed(checkImplementationStatus(opInst)))
9488 return failure();
9489
9490 // During the handling of target op, we will generate instructions in the
9491 // parent function like call to the oulined function or branch to a new
9492 // BasicBlock. We set the debug location here to parent function so that those
9493 // get the correct debug locations. For outlined functions, the normal MLIR op
9494 // conversion will automatically pick the correct location.
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()));
9503
9504 // OMPIRBuilder emits runtime calls into the outlined function before bodyCB
9505 // below gets a chance to attach the subprogram to it, so it needs the
9506 // outlined function's location handed to it separately. Only pass it under
9507 // the same condition that decides whether the subprogram is attached at all:
9508 // a location may not be attached to an instruction in a function that has no
9509 // subprogram.
9510 llvm::DebugLoc outlinedFnDbgLoc;
9511 if (outlinedFnLoc && parentLLVMFn->getSubprogram())
9512 outlinedFnDbgLoc = outlinedFnLoc;
9513
9514 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
9515 bool isTargetDevice = ompBuilder->Config.isTargetDevice();
9516 bool isGPU = ompBuilder->Config.isGPU();
9517
9518 auto parentFn = opInst.getParentOfType<LLVM::LLVMFuncOp>();
9519 auto argIface = cast<omp::BlockArgOpenMPOpInterface>(opInst);
9520 auto &targetRegion = targetOp.getRegion();
9521 // Holds the private vars that have been mapped along with the block
9522 // argument that corresponds to the MapInfoOp corresponding to the private
9523 // var in question. So, for instance:
9524 //
9525 // %10 = omp.map.info var_ptr(%6#0 : !fir.ref<!fir.box<!fir.heap<i32>>>, ..)
9526 // omp.target map_entries(%10 -> %arg0) private(@box.privatizer %6#0-> %arg1)
9527 //
9528 // Then, %10 has been created so that the descriptor can be used by the
9529 // privatizer @box.privatizer on the device side. Here we'd record {%6#0,
9530 // %arg0} in the mappedPrivateVars map.
9531 llvm::DenseMap<Value, Value> mappedPrivateVars;
9532 DataLayout dl = DataLayout(opInst.getParentOfType<ModuleOp>());
9533 SmallVector<Value> mapVars = targetOp.getMapVars();
9534 SmallVector<Value> hdaVars = targetOp.getHasDeviceAddrVars();
9535 ArrayRef<BlockArgument> mapBlockArgs = argIface.getMapBlockArgs();
9536 ArrayRef<BlockArgument> hdaBlockArgs = argIface.getHasDeviceAddrBlockArgs();
9537 llvm::Function *llvmOutlinedFn = nullptr;
9538 TargetDirectiveEnumTy targetDirective =
9539 getTargetDirectiveEnumTyFromOp(&opInst);
9540
9541 // TODO: It can also be false if a compile-time constant `false` IF clause is
9542 // specified.
9543 bool isOffloadEntry =
9544 isTargetDevice || !ompBuilder->Config.TargetTriples.empty();
9545
9546 // Resolve in_reduction clauses on omp.target for the host. From the target
9547 // device's perspective an in_reduction list item behaves as a regular
9548 // map(tofrom) variable, so no special handling is needed there; only the
9549 // host redirects the mapped value to the per-task reduction-private storage
9550 // returned by __kmpc_task_reduction_get_th_data (emitted inside the
9551 // to-be-outlined target task body). This applies to both offloading and
9552 // non-offloading host modules.
9553 //
9554 // The target body has no dedicated in_reduction block argument: each
9555 // in_reduction variable is accessed through its map_entries block argument.
9556 // So each in_reduction variable must also be captured by a matching
9557 // map_entries entry (guaranteed by the verifier); without one the outlined
9558 // body would reference a value defined in the host function. Record, for each
9559 // in_reduction variable, the position of that map entry so the corresponding
9560 // map block argument can be redirected inside the body. The in_reduction
9561 // operand itself is used as the `orig` argument of the runtime lookup.
9562 SmallVector<llvm::Value *> inRedOrigPtrs;
9563 SmallVector<unsigned> inRedMapArgIdx;
9564 if (!targetOp.getInReductionVars().empty() && !isTargetDevice) {
9565 inRedOrigPtrs.reserve(targetOp.getInReductionVars().size());
9566 inRedMapArgIdx.reserve(targetOp.getInReductionVars().size());
9567 for (Value v : targetOp.getInReductionVars()) {
9568 // Select the map_entries entry that captures this in_reduction operand.
9569 // The verifier guarantees at least one match exists; more than one
9570 // matching entry is a lowering ambiguity (the redirect cannot pick which
9571 // map argument to rebind).
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())
9576 continue;
9577 if (matchIdx)
9578 return targetOp.emitError()
9579 << "in_reduction variable on omp.target has multiple matching "
9580 "map_entries entries; the redirect target is ambiguous";
9581 matchIdx = idx;
9582 }
9583 // The verifier requires a capturing map entry for every in_reduction
9584 // operand, so a match must exist here.
9585 assert(matchIdx &&
9586 "TargetOp verifier guarantees a matching map_entries entry for "
9587 "each in_reduction variable");
9588 inRedMapArgIdx.push_back(*matchIdx);
9589 // The runtime `orig` pointer is the in_reduction operand itself, the
9590 // reduction variable the enclosing taskgroup registered.
9591 inRedOrigPtrs.push_back(moduleTranslation.lookupValue(v));
9592 }
9593 }
9594
9595 // For some private variables, the MapsForPrivatizedVariablesPass
9596 // creates MapInfoOp instances. Go through the private variables and
9597 // the mapped variables so that during codegeneration we are able
9598 // to quickly look up the corresponding map variable, if any for each
9599 // private variable.
9600 if (!targetOp.getPrivateVars().empty() && !targetOp.getMapVars().empty()) {
9601 OperandRange privateVars = targetOp.getPrivateVars();
9602 std::optional<ArrayAttr> privateSyms = targetOp.getPrivateSyms();
9603 std::optional<DenseI64ArrayAttr> privateMapIndices =
9604 targetOp.getPrivateMapsAttr();
9605
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);
9610
9611 SymbolRefAttr privatizerName = llvm::cast<SymbolRefAttr>(privSym);
9612 omp::PrivateClauseOp privatizer =
9613 findPrivatizer(targetOp, privatizerName);
9614
9615 if (!privatizer.needsMap())
9616 continue;
9617
9618 mlir::Value mappedValue =
9619 targetOp.getMappedValueForPrivateVar(privVarIdx);
9620 assert(mappedValue && "Expected to find mapped value for a privatized "
9621 "variable that needs mapping");
9622
9623 // The MapInfoOp defining the map var isn't really needed later.
9624 // So, we don't store it in any datastructure. Instead, we just
9625 // do some sanity checks on it right now.
9626 auto mapInfoOp = mappedValue.getDefiningOp<omp::MapInfoOp>();
9627 [[maybe_unused]] Type varType = mapInfoOp.getVarPtrType();
9628
9629 // Check #1: Check that the type of the private variable matches
9630 // the type of the variable being mapped.
9631 if (!isa<LLVM::LLVMPointerType>(privVar.getType()))
9632 assert(
9633 varType == privVar.getType() &&
9634 "Type of private var doesn't match the type of the mapped value");
9635
9636 // Ok, only 1 sanity check for now.
9637 // Record the block argument corresponding to this mapvar.
9638 mappedPrivateVars.insert(
9639 {privVar,
9640 targetRegion.getArgument(argIface.getMapBlockArgsStart() +
9641 (*privateMapIndices)[privVarIdx])});
9642 }
9643 }
9644
9645 using InsertPointTy = llvm::OpenMPIRBuilder::InsertPointTy;
9646 auto bodyCB = [&](InsertPointTy allocaIP, InsertPointTy codeGenIP,
9647 ArrayRef<llvm::BasicBlock *> deallocBlocks)
9648 -> llvm::OpenMPIRBuilder::InsertPointOrErrorTy {
9649 llvm::IRBuilderBase::InsertPointGuard guard(builder);
9650 builder.SetCurrentDebugLocation(llvm::DebugLoc());
9651 // Forward target-cpu and target-features function attributes from the
9652 // original function to the new outlined function.
9653 llvm::Function *llvmParentFn =
9654 moduleTranslation.lookupFunction(parentFn.getName());
9655 llvmOutlinedFn = codeGenIP.getNodeParent()->getParent();
9656 assert(llvmParentFn && llvmOutlinedFn &&
9657 "Both parent and outlined functions must exist at this point");
9658
9659 if (outlinedFnLoc && llvmParentFn->getSubprogram())
9660 llvmOutlinedFn->setSubprogram(outlinedFnLoc->getScope()->getSubprogram());
9661
9662 if (auto attr = llvmParentFn->getFnAttribute("target-cpu");
9663 attr.isStringAttribute())
9664 llvmOutlinedFn->addFnAttr(attr);
9665
9666 if (auto attr = llvmParentFn->getFnAttribute("target-features");
9667 attr.isStringAttribute())
9668 llvmOutlinedFn->addFnAttr(attr);
9669
9670 for (auto [idx, arg] : llvm::enumerate(mapBlockArgs)) {
9671 // in_reduction list items on omp.target are accessed through their
9672 // map_entries block argument, which is redirected below to the per-task
9673 // reduction-private storage returned by the runtime. Skip the default
9674 // host-value mapping for those block arguments so the write-once
9675 // mapValue mapping is free to be set to the private pointer.
9676 if (llvm::is_contained(inRedMapArgIdx, idx))
9677 continue;
9678 auto mapInfoOp = cast<omp::MapInfoOp>(mapVars[idx].getDefiningOp());
9679 llvm::Value *mapOpValue =
9680 moduleTranslation.lookupValue(mapInfoOp.getVarPtr());
9681 moduleTranslation.mapValue(arg, mapOpValue);
9682 }
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);
9688 }
9689
9690 // Do privatization after moduleTranslation has already recorded
9691 // mapped values.
9692 PrivateVarsInfo privateVarsInfo(targetOp);
9693
9695 allocatePrivateVars(targetOp, builder, moduleTranslation,
9696 privateVarsInfo, allocaIP, &mappedPrivateVars);
9697
9698 if (failed(handleError(afterAllocas, *targetOp)))
9699 return llvm::make_error<PreviouslyReportedError>();
9700
9701 builder.restoreIP(codeGenIP);
9702 if (handleError(initPrivateVars(builder, moduleTranslation, privateVarsInfo,
9703 &mappedPrivateVars),
9704 *targetOp)
9705 .failed())
9706 return llvm::make_error<PreviouslyReportedError>();
9707
9708 if (failed(copyFirstPrivateVars(
9709 targetOp, builder, moduleTranslation, privateVarsInfo.mlirVars,
9710 privateVarsInfo.llvmVars, privateVarsInfo.privatizers,
9711 targetOp.getPrivateNeedsBarrier(), &mappedPrivateVars)))
9712 return llvm::make_error<PreviouslyReportedError>();
9713
9714 // The target body accesses each in_reduction variable through its
9715 // map_entries block argument. Redirect that block argument to the per-task
9716 // private storage returned by __kmpc_task_reduction_get_th_data so the body
9717 // accumulates into the reduction-private copy rather than the mapped
9718 // original. The lookup must run inside the target task body so the gtid
9719 // corresponds to the executing thread. The descriptor argument is NULL: the
9720 // runtime walks enclosing taskgroups to locate the matching task_reduction
9721 // registration for `origPtr`. Mirrors the in_reduction handling on
9722 // omp.taskloop.context.
9723 if (!inRedOrigPtrs.empty()) {
9724 // Collect, per item, the type the private pointer must have (the map
9725 // block argument's type), and, through the callback, rebind the map block
9726 // argument that stands in for each in_reduction list item to the per-task
9727 // reduction-private storage the runtime returns.
9728 SmallVector<llvm::Type *> inRedResultPtrTys;
9729 inRedResultPtrTys.reserve(inRedMapArgIdx.size());
9730 for (unsigned mapArgIdx : inRedMapArgIdx)
9731 inRedResultPtrTys.push_back(
9732 moduleTranslation.convertType(mapBlockArgs[mapArgIdx].getType()));
9733
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]],
9740 priv);
9741 });
9742 builder.restoreIP(redIP);
9743 }
9744
9746 moduleTranslation, allocaIP, deallocBlocks);
9748 targetRegion, "omp.target", builder, moduleTranslation);
9749
9750 if (failed(handleError(exitBlock, *targetOp)))
9751 return llvm::make_error<PreviouslyReportedError>();
9752
9753 builder.SetInsertPoint(exitBlock.get()->getTerminator());
9754
9755 if (failed(cleanupPrivateVars(targetOp, builder, moduleTranslation,
9756 targetOp.getLoc(), privateVarsInfo)))
9757 return llvm::make_error<PreviouslyReportedError>();
9758
9759 return builder.saveIP();
9760 };
9761
9762 StringRef parentName = parentFn.getName();
9763
9764 llvm::TargetRegionEntryInfo entryInfo;
9765
9766 getTargetEntryUniqueInfo(entryInfo, targetOp,
9767 *moduleTranslation.getOpenMPBuilder(),
9768 moduleTranslation.getFileSystem(), parentName);
9769
9770 MapInfoData mapData;
9771 collectMapDataFromMapOperands(mapData, mapVars, moduleTranslation, dl,
9772 builder, /*useDevPtrOperands=*/{},
9773 /*useDevAddrOperands=*/{}, hdaVars);
9774
9775 MapInfosTy combinedInfos;
9776 auto genMapInfoCB =
9777 [&](llvm::OpenMPIRBuilder::InsertPointTy codeGenIP) -> MapInfosTy & {
9778 builder.restoreIP(codeGenIP);
9779 genMapInfos(builder, moduleTranslation, dl, combinedInfos, mapData,
9780 targetDirective);
9781
9782 // Append a null entry for the implicit dyn_ptr argument so the argument
9783 // count sent to the runtime already includes it.
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);
9793 // TODO: set HasAttachPtr from Flang for pointee-storage entries.
9794 combinedInfos.HasAttachPtr.push_back(false);
9795 if (!combinedInfos.Names.empty())
9796 combinedInfos.Names.push_back(nullPtr);
9797 combinedInfos.Mappers.push_back(nullptr);
9798
9799 return combinedInfos;
9800 };
9801
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());
9809 // We just return the unaltered argument for the host function
9810 // for now, some alterations may be required in the future to
9811 // keep host fallback functions working identically to the device
9812 // version (e.g. pass ByCopy values should be treated as such on
9813 // host and device, currently not always the case)
9814 if (!isTargetDevice) {
9815 retVal = cast<llvm::Value>(&arg);
9816 return codeGenIP;
9817 }
9818
9819 return createDeviceArgumentAccessor(targetOp, mapData, arg, input, retVal,
9820 builder, *ompBuilder, moduleTranslation,
9821 allocaIP, codeGenIP, deallocIPs);
9822 };
9823
9824 llvm::OpenMPIRBuilder::TargetKernelRuntimeAttrs runtimeAttrs;
9825 llvm::OpenMPIRBuilder::TargetKernelDefaultAttrs defaultAttrs;
9826 Operation *targetCapturedOp =
9827 cast<omp::ComposableOpInterface>(*targetOp).findCapturedOp();
9828 initTargetDefaultAttrs(targetOp, targetCapturedOp, defaultAttrs,
9829 isTargetDevice, isGPU);
9830
9831 // Collect host-evaluated values needed to properly launch the kernel from the
9832 // host.
9833 if (!isTargetDevice)
9834 initTargetRuntimeAttrs(builder, moduleTranslation, targetOp,
9835 targetCapturedOp, runtimeAttrs);
9836
9837 // Pass host-evaluated values as parameters to the kernel / host fallback,
9838 // except if they are constants. In any case, map the MLIR block argument to
9839 // the corresponding LLVM values.
9841 SmallVector<Value> hostEvalVars = targetOp.getHostEvalVars();
9842 ArrayRef<BlockArgument> hostEvalBlockArgs = argIface.getHostEvalBlockArgs();
9843 for (auto [arg, var] : llvm::zip_equal(hostEvalBlockArgs, hostEvalVars)) {
9844 llvm::Value *value = moduleTranslation.lookupValue(var);
9845 moduleTranslation.mapValue(arg, value);
9846
9847 if (!llvm::isa<llvm::Constant>(value))
9848 kernelInput.push_back(value);
9849 }
9850
9851 for (size_t i = 0, e = mapData.OriginalValue.size(); i != e; ++i) {
9852 // 1) Declare target arguments are not passed to kernels as arguments.
9853 // 2) Attach maps are not passed in as arguments to kernels, except for
9854 // private attach maps used for corresponding-pointer initialization.
9855 // 3) Children of record objects are not passed in as arguments.
9856 // TODO: We currently do not handle cases where a member is explicitly
9857 // passed in as an argument, this will likley need to be handled in
9858 // the near future, rather than using IsAMember, it may be better to
9859 // test if the relevant BlockArg is used within the target region and
9860 // then use that as a basis for exclusion in the kernel inputs.
9861 using MapFlags = llvm::omp::OpenMPOffloadMappingFlags;
9862 bool isAttachMap = (mapData.Types[i] & MapFlags::OMP_MAP_ATTACH) ==
9863 MapFlags::OMP_MAP_ATTACH;
9864 bool isPrivateTargetParam =
9865 (mapData.Types[i] &
9866 (MapFlags::OMP_MAP_PRIVATE | MapFlags::OMP_MAP_TARGET_PARAM)) ==
9867 (MapFlags::OMP_MAP_PRIVATE | MapFlags::OMP_MAP_TARGET_PARAM);
9868
9869 if (!mapData.IsDeclareTarget[i] && !mapData.IsAMember[i] &&
9870 (!isAttachMap || (isAttachMap && isPrivateTargetParam)))
9871 kernelInput.push_back(mapData.OriginalValue[i]);
9872 }
9873
9875 llvm::OpenMPIRBuilder::InsertPointTy allocaIP =
9876 findAllocInsertPoints(builder, moduleTranslation, &deallocBlocks);
9877
9878 llvm::OpenMPIRBuilder::DependenciesInfo dds;
9879 if (failed(buildDependData(
9880 targetOp.getDependVars(), targetOp.getDependKinds(),
9881 targetOp.getDependIterated(), targetOp.getDependIteratedKinds(),
9882 builder, moduleTranslation, dds)))
9883 return failure();
9884
9885 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
9886
9887 llvm::OpenMPIRBuilder::TargetDataInfo info(
9888 /*RequiresDevicePointerInfo=*/false,
9889 /*SeparateBeginEndCalls=*/true);
9890
9891 auto customMapperCB =
9892 [&](unsigned int i) -> llvm::Expected<llvm::Function *> {
9893 if (!combinedInfos.Mappers[i])
9894 return nullptr;
9895 info.HasMapper = true;
9896 return getOrCreateUserDefinedMapperFunc(combinedInfos.Mappers[i], builder,
9897 moduleTranslation, targetDirective);
9898 };
9899
9900 llvm::Value *ifCond = nullptr;
9901 if (Value targetIfCond = targetOp.getIfExpr())
9902 ifCond = moduleTranslation.lookupValue(targetIfCond);
9903
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(),
9909 /*isSigned=*/false);
9910 }
9911
9912 llvm::omp::OMPDynGroupprivateFallbackType fallbackType =
9913 getDynGroupprivateFallbackType(targetOp.getDynGroupprivateFallbackAttr());
9914
9915 // Pass the target region's source location to the runtime, taken from the
9916 // op's own location. Restricted to the host offload path that actually emits
9917 // the kernel launch, to avoid creating an unused identifier on the device.
9918 llvm::Value *rtLocOverride =
9919 (!isTargetDevice && isOffloadEntry)
9920 ? getSourceLocIdentFromOp(builder, *ompBuilder, targetOp)
9921 : nullptr;
9922
9923 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
9924 moduleTranslation.getOpenMPBuilder()->createTarget(
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,
9929 rtLocOverride);
9930
9931 if (failed(handleError(afterIP, opInst)))
9932 return failure();
9933
9934 builder.restoreIP(*afterIP);
9935
9936 if (dds.DepArray)
9937 builder.CreateFree(dds.DepArray);
9938
9939 return success();
9940}
9941
9942static LogicalResult
9943convertDeclareTargetAttr(Operation *op, mlir::omp::DeclareTargetAttr attribute,
9944 llvm::OpenMPIRBuilder *ompBuilder,
9945 LLVM::ModuleTranslation &moduleTranslation) {
9946 // Amend omp.declare_target by deleting the IR of the outlined functions
9947 // created for target regions. They cannot be filtered out from MLIR earlier
9948 // because the omp.target operation inside must be translated to LLVM, but
9949 // the wrapper functions themselves must not remain at the end of the
9950 // process. We know that functions where omp.declare_target does not match
9951 // omp.is_target_device at this stage can only be wrapper functions because
9952 // those that aren't are removed earlier as an MLIR transformation pass.
9953 if (FunctionOpInterface funcOp = dyn_cast<FunctionOpInterface>(op)) {
9954 if (auto offloadMod = dyn_cast<omp::OffloadModuleInterface>(
9955 op->getParentOfType<ModuleOp>().getOperation())) {
9956 if (!offloadMod.getIsTargetDevice())
9957 return success();
9958
9959 omp::DeclareTargetDeviceType declareType = attribute.getDeviceType();
9960
9961 if (declareType == omp::DeclareTargetDeviceType::host) {
9962 llvm::Function *llvmFunc =
9963 moduleTranslation.lookupFunction(funcOp.getName());
9964 llvmFunc->dropAllReferences();
9965 llvmFunc->eraseFromParent();
9966
9967 // Invalidate the builder's current insertion point, as it now points to
9968 // a deleted block.
9969 ompBuilder->Builder.ClearInsertionPoint();
9970 ompBuilder->Builder.SetCurrentDebugLocation(llvm::DebugLoc());
9971 } else if (llvm::Function *llvmFunc =
9972 moduleTranslation.lookupFunction(funcOp.getName())) {
9973 // Device-side declare target functions are externally visible by
9974 // default so they can be referenced from other device translation
9975 // units. That also prevents the offload LTO from internalizing and
9976 // deleting them when they end up unused in the final device image.
9977 // Such dead functions can still reference internal LDS and trigger
9978 // spurious "local memory global used by non-kernel function" backend
9979 // warnings. Marking them hidden keeps the symbol usable within the
9980 // device image's linkage unit while letting LTO drop it when nothing
9981 // references it; symbols that must stay reachable (e.g. via an offload
9982 // entry that takes their address) are kept alive by that reference.
9983 if (!llvmFunc->isDeclaration() && llvmFunc->hasExternalLinkage() &&
9984 llvmFunc->getVisibility() == llvm::GlobalValue::DefaultVisibility)
9985 llvmFunc->setVisibility(llvm::GlobalValue::HiddenVisibility);
9986 }
9987 }
9988 return success();
9989 }
9990
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);
9995 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
9996 bool isDeclaration = gOp.isDeclaration();
9997 bool isExternallyVisible =
9998 gOp.getVisibility() != mlir::SymbolTable::Visibility::Private;
9999 auto loc = op->getLoc()->findInstanceOf<FileLineColLoc>();
10000 llvm::StringRef mangledName = gOp.getSymName();
10001 mlir::omp::DeclareTargetCaptureClause captureClause =
10002 attribute.getCaptureClause();
10003 auto captureClauseKind = convertToCaptureClauseKind(captureClause);
10004 auto deviceClause = convertToDeviceClauseKind(attribute.getDeviceType());
10005 llvm::StringRef entryMangledName = mangledName;
10006 llvm::Constant *entryAddr = llvm::cast<llvm::Constant>(gVal);
10007 std::function<llvm::GlobalValue::LinkageTypes()> variableLinkage;
10008 llvm::SmallString<128> entryNameStorage;
10009 bool requiresUSM = ompBuilder->Config.hasRequiresUnifiedSharedMemory();
10010 bool isToOrEnter =
10011 captureClause == omp::DeclareTargetCaptureClause::to ||
10012 captureClause == omp::DeclareTargetCaptureClause::enter;
10013 bool isHostOnly =
10014 attribute.getDeviceType() == omp::DeclareTargetDeviceType::host;
10015
10016 // A to/enter declare-target variable needs a device-resident,
10017 // name-resolvable copy and a host offloading entry. A local-linkage
10018 // global provides neither, so we promote it to external.
10019 if (isToOrEnter && !isHostOnly && !requiresUSM &&
10020 gVar->hasLocalLinkage()) {
10021 gVar->setLinkage(llvm::GlobalValue::ExternalLinkage);
10022 isExternallyVisible = true;
10023
10024 // Clear the stale dso_local flag so it is referenced like a
10025 // module-scope declare target global.
10026 if (ompBuilder->Config.isTargetDevice())
10027 gVar->setDSOLocal(false);
10028 }
10029
10030 if (isToOrEnter &&
10031 deviceClause ==
10032 llvm::OffloadEntriesInfoManager::OMPTargetDeviceClauseAny &&
10033 !requiresUSM && !isDeclaration &&
10034 (gVal->hasLocalLinkage() || gVal->hasHiddenVisibility())) {
10035 // Keep the original symbol as-is for target code, but create a visible
10036 // alias for the offload entry so libomptarget can associate the host
10037 // global with the actual device global.
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);
10043 } else {
10044 entryAddr = llvm::GlobalAlias::create(
10045 gVal->getValueType(), gVal->getAddressSpace(),
10046 llvm::GlobalValue::WeakAnyLinkage, entryMangledName, entryAddr,
10047 llvmModule);
10048 llvm::cast<llvm::GlobalAlias>(entryAddr)->setVisibility(
10049 llvm::GlobalValue::DefaultVisibility);
10050 }
10051 variableLinkage = [] { return llvm::GlobalValue::WeakAnyLinkage; };
10052 }
10053 // unused for MLIR at the moment, required in Clang for book
10054 // keeping
10055 std::vector<llvm::GlobalVariable *> generatedRefs;
10056
10057 std::vector<llvm::Triple> targetTriple;
10058 auto targetTripleAttr = dyn_cast_or_null<mlir::StringAttr>(
10059 op->getParentOfType<mlir::ModuleOp>()->getDiscardableAttr(
10060 LLVM::LLVMDialect::getTargetTripleAttrName()));
10061 if (targetTripleAttr)
10062 targetTriple.emplace_back(targetTripleAttr.data());
10063
10064 auto fileInfoCallBack = [&loc]() {
10065 std::string filename = "";
10066 std::uint64_t lineNo = 0;
10067
10068 if (loc) {
10069 filename = loc.getFilename().str();
10070 lineNo = loc.getLine();
10071 }
10072
10073 return std::pair<std::string, std::uint64_t>(llvm::StringRef(filename),
10074 lineNo);
10075 };
10076
10077 llvm::vfs::FileSystem &vfs = moduleTranslation.getFileSystem();
10078 ompBuilder->registerTargetGlobalVariable(
10079 captureClauseKind, deviceClause, isDeclaration, isExternallyVisible,
10080 ompBuilder->getTargetEntryUniqueInfo(fileInfoCallBack, vfs),
10081 entryMangledName, generatedRefs, /*OpenMPSimd*/ false, targetTriple,
10082 /*GlobalInitializer*/ nullptr, variableLinkage, gVal->getType(),
10083 entryAddr);
10084
10085 if (ompBuilder->Config.isTargetDevice() &&
10086 (captureClause == omp::DeclareTargetCaptureClause::link ||
10087 requiresUSM)) {
10088 // For USM and link we generate a global reference pointer in the
10089 // default address space (e.g address space 0), as opposed to the
10090 // globals original type and address space.
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, /*OpenMPSimd*/ false, targetTriple,
10096 ptrTy, /*GlobalInitializer*/ nullptr,
10097 /*VariableLinkage*/ nullptr);
10098
10099 // For indirectly-accessed global pointers, we rely on "internal"
10100 // linkage to optimize out the unneeded full-variable storage later,
10101 // since we can't prevent the LLVM dialect from generating globals
10102 // without also breaking target lowering. However, We can only do
10103 // this for definiions, as global variable declarations must have
10104 // external or weak linkage.
10105 if (refPtr) {
10106 if (!gVar->isDeclaration())
10107 gVar->setLinkage(llvm::GlobalValue::InternalLinkage);
10108
10109 // Register the (original global, reference pointer) pair so that the
10110 // OpenMPIRBuilder can rewrite uses of the original global during
10111 // finalization.
10112 if (auto *newGV =
10113 dyn_cast<llvm::GlobalValue>(refPtr->stripPointerCasts()))
10114 ompBuilder->registerDeclareTargetGlobalReplacement(gVal, newGV);
10115 }
10116 }
10117
10118 // Mark 'device_type(host) enter(...)' variables as external in the device
10119 // since they're not supposed to have their own copy. This will cause
10120 // linker errors if accesses are attempted from the target device.
10121 if (ompBuilder->Config.isTargetDevice() && isHostOnly && isToOrEnter) {
10122 gVar->setLinkage(llvm::GlobalValue::ExternalLinkage);
10123 gVar->setInitializer(nullptr);
10124 }
10125 }
10126 }
10127
10128 return success();
10129}
10130
10131namespace {
10132
10133/// Implementation of the dialect interface that converts operations belonging
10134/// to the OpenMP dialect to LLVM IR.
10135class OpenMPDialectLLVMIRTranslationInterface
10136 : public LLVMTranslationDialectInterface {
10137public:
10138 using LLVMTranslationDialectInterface::LLVMTranslationDialectInterface;
10139
10140 /// Translates the given operation to LLVM IR using the provided IR builder
10141 /// and saving the state in `moduleTranslation`.
10142 LogicalResult
10143 convertOperation(Operation *op, llvm::IRBuilderBase &builder,
10144 LLVM::ModuleTranslation &moduleTranslation) const final;
10145
10146 /// Given an OpenMP MLIR attribute, create the corresponding LLVM-IR,
10147 /// runtime calls, or operation amendments
10148 LogicalResult
10149 amendOperation(Operation *op, ArrayRef<llvm::Instruction *> instructions,
10150 NamedAttribute attribute,
10151 LLVM::ModuleTranslation &moduleTranslation) const final;
10152
10153 /// Records the LLVM alloc pointer produced for an OMP ALLOCATE variable so
10154 /// that the paired omp.allocate_free op can generate the matching
10155 /// __kmpc_free call.
10156 void registerAllocatedPtr(Value var, llvm::Value *ptr) const {
10157 ompAllocatedPtrs[var] = ptr;
10158 }
10159
10160 /// Returns the LLVM alloc pointer previously registered for var, or
10161 /// nullptr if no allocation was recorded.
10162 llvm::Value *lookupAllocatedPtr(Value var) const {
10163 auto it = ompAllocatedPtrs.find(var);
10164 return it != ompAllocatedPtrs.end() ? it->second : nullptr;
10165 }
10166
10167private:
10168 /// Maps each MLIR variable value that appeared in an omp.allocate_dir op to
10169 /// the LLVM pointer returned by the corresponding __kmpc_alloc call. The
10170 /// paired omp.allocate_free op looks up these pointers to emit __kmpc_free.
10171 mutable DenseMap<Value, llvm::Value *> ompAllocatedPtrs;
10172};
10173
10174} // namespace
10175
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)>>(
10181 attribute.getName())
10182 .Case("omp.is_target_device",
10183 [&](Attribute attr) {
10184 if (auto deviceAttr = dyn_cast<BoolAttr>(attr)) {
10185 llvm::OpenMPIRBuilderConfig &config =
10186 moduleTranslation.getOpenMPBuilder()->Config;
10187 config.setIsTargetDevice(deviceAttr.getValue());
10188 return success();
10189 }
10190 return failure();
10191 })
10192 .Case("omp.is_gpu",
10193 [&](Attribute attr) {
10194 if (auto gpuAttr = dyn_cast<BoolAttr>(attr)) {
10195 llvm::OpenMPIRBuilderConfig &config =
10196 moduleTranslation.getOpenMPBuilder()->Config;
10197 config.setIsGPU(gpuAttr.getValue());
10198 return success();
10199 }
10200 return failure();
10201 })
10202 .Case("omp.host_ir_filepath",
10203 [&](Attribute attr) {
10204 if (auto filepathAttr = dyn_cast<StringAttr>(attr)) {
10205 llvm::OpenMPIRBuilder *ompBuilder =
10206 moduleTranslation.getOpenMPBuilder();
10207 ompBuilder->loadOffloadInfoMetadata(
10208 moduleTranslation.getFileSystem(), filepathAttr.getValue());
10209 return success();
10210 }
10211 return failure();
10212 })
10213 .Case("omp.flags",
10214 [&](Attribute attr) {
10215 if (auto rtlAttr = dyn_cast<omp::FlagsAttr>(attr))
10216 return convertFlagsAttr(op, rtlAttr, moduleTranslation);
10217 return failure();
10218 })
10219 .Case("omp.version",
10220 [&](Attribute attr) {
10221 if (auto versionAttr = dyn_cast<omp::VersionAttr>(attr)) {
10222 llvm::OpenMPIRBuilder *ompBuilder =
10223 moduleTranslation.getOpenMPBuilder();
10224 ompBuilder->M.addModuleFlag(llvm::Module::Max, "openmp",
10225 versionAttr.getVersion());
10226 return success();
10227 }
10228 return failure();
10229 })
10230 .Case("omp.declare_target",
10231 [&](Attribute attr) {
10232 if (auto declareTargetAttr =
10233 dyn_cast<omp::DeclareTargetAttr>(attr)) {
10234 llvm::OpenMPIRBuilder *ompBuilder =
10235 moduleTranslation.getOpenMPBuilder();
10236 return convertDeclareTargetAttr(op, declareTargetAttr,
10237 ompBuilder, moduleTranslation);
10238 }
10239 return failure();
10240 })
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 =
10247 moduleTranslation.getOpenMPBuilder()->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));
10256 return success();
10257 }
10258 return failure();
10259 })
10260 .Case("omp.target_triples",
10261 [&](Attribute attr) {
10262 if (auto triplesAttr = dyn_cast<ArrayAttr>(attr)) {
10263 llvm::OpenMPIRBuilderConfig &config =
10264 moduleTranslation.getOpenMPBuilder()->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());
10270 else
10271 return failure();
10272 }
10273 return success();
10274 }
10275 return failure();
10276 })
10277 .Case("omp.integer_wrap_around",
10278 [&](Attribute attr) {
10279 if (auto wrapAttr = dyn_cast<omp::IntegerWrapAroundAttr>(attr)) {
10280 llvm::OpenMPIRBuilderConfig &config =
10281 moduleTranslation.getOpenMPBuilder()->Config;
10282 config.setNoSignedWrap(!wrapAttr.getIntegerWrapAround());
10283 return success();
10284 }
10285 return failure();
10286 })
10287 .Default([](Attribute) {
10288 // Fall through for omp attributes that do not require lowering.
10289 return success();
10290 })(attribute.getValue());
10291
10292 return failure();
10293}
10294
10295// Returns true if the operation is not inside a TargetOp, it is part of a
10296// function and that function is not declare target.
10297static bool isHostDeviceOp(Operation *op) {
10298 // Assumes no reverse offloading
10299 if (op->getParentOfType<omp::TargetOp>())
10300 return false;
10301
10302 if (auto parentFn = op->getParentOfType<LLVM::LLVMFuncOp>()) {
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)
10310 return false;
10311 }
10312
10313 return true;
10314 }
10315
10316 return false;
10317}
10318
10319static llvm::Function *getOmpTargetAlloc(llvm::IRBuilderBase &builder,
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());
10328 return func;
10329}
10330
10331template <typename T>
10332static llvm::Value *
10333getAllocationSize(llvm::IRBuilderBase &builder,
10334 LLVM::ModuleTranslation &moduleTranslation, T op) {
10335 llvm::DataLayout dataLayout =
10336 moduleTranslation.getLLVMModule()->getDataLayout();
10337 llvm::Type *llvmHeapTy =
10338 moduleTranslation.convertType(op.getMemElemTypeAttr().getValue());
10339
10340 auto alignment = op.getMemAlignment();
10341 llvm::TypeSize typeSize = llvm::alignTo(
10342 dataLayout.getTypeStoreSize(llvmHeapTy),
10343 alignment ? *alignment : dataLayout.getABITypeAlign(llvmHeapTy).value());
10344
10345 llvm::Value *allocSize = builder.getInt64(typeSize.getFixedValue());
10346 return builder.CreateMul(
10347 allocSize,
10348 builder.CreateIntCast(moduleTranslation.lookupValue(op.getMemArraySize()),
10349 builder.getInt64Ty(),
10350 /*isSigned=*/false));
10351}
10352
10353template <>
10354llvm::Value *getAllocationSize(llvm::IRBuilderBase &builder,
10355 LLVM::ModuleTranslation &moduleTranslation,
10356 omp::TargetAllocMemOp op) {
10357 llvm::DataLayout dataLayout =
10358 moduleTranslation.getLLVMModule()->getDataLayout();
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(
10364 allocSize,
10365 builder.CreateIntCast(moduleTranslation.lookupValue(typeParam),
10366 builder.getInt64Ty(),
10367 /*isSigned=*/false));
10368 }
10369 return allocSize;
10370}
10371
10372static LogicalResult
10373convertTargetAllocMemOp(Operation &opInst, llvm::IRBuilderBase &builder,
10374 LLVM::ModuleTranslation &moduleTranslation) {
10375 auto allocMemOp = cast<omp::TargetAllocMemOp>(opInst);
10376 if (!allocMemOp)
10377 return failure();
10378
10379 // Get "omp_target_alloc" function
10380 llvm::Module *llvmModule = moduleTranslation.getLLVMModule();
10381 llvm::Function *ompTargetAllocFunc = getOmpTargetAlloc(builder, llvmModule);
10382 // Get the corresponding device value in llvm
10383 mlir::Value deviceNum = allocMemOp.getDevice();
10384 llvm::Value *llvmDeviceNum = moduleTranslation.lookupValue(deviceNum);
10385 // Get the allocation size.
10386 llvm::Value *allocSize =
10387 getAllocationSize(builder, moduleTranslation, allocMemOp);
10388 // Create call to "omp_target_alloc" with the args as translated llvm values.
10389 llvm::CallInst *call =
10390 builder.CreateCall(ompTargetAllocFunc, {allocSize, llvmDeviceNum});
10391 llvm::Value *resultI64 = builder.CreatePtrToInt(call, builder.getInt64Ty());
10392
10393 // Map the result
10394 moduleTranslation.mapValue(allocMemOp.getResult(), resultI64);
10395 return success();
10396}
10397
10398static LogicalResult
10399convertAllocSharedMemOp(omp::AllocSharedMemOp allocMemOp,
10400 llvm::IRBuilderBase &builder,
10401 LLVM::ModuleTranslation &moduleTranslation) {
10402 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
10403 llvm::Value *size = getAllocationSize(builder, moduleTranslation, allocMemOp);
10404 moduleTranslation.mapValue(allocMemOp.getResult(),
10405 ompBuilder->createOMPAllocShared(builder, size));
10406 return success();
10407}
10408
10409static LogicalResult
10410convertAllocateDirOp(Operation &opInst, llvm::IRBuilderBase &builder,
10411 LLVM::ModuleTranslation &moduleTranslation,
10412 const OpenMPDialectLLVMIRTranslationInterface &ompIface) {
10413 auto allocateDirOp = cast<omp::AllocateDirOp>(opInst);
10414 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
10415
10416 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
10417 llvm::Module *llvmModule = moduleTranslation.getLLVMModule();
10418 llvm::DataLayout dataLayout = llvmModule->getDataLayout();
10419 SmallVector<Value> vars = allocateDirOp.getVarList();
10420 std::optional<int64_t> alignAttr = allocateDirOp.getAlign();
10421
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());
10430 } else {
10431 allocator = llvm::ConstantPointerNull::get(builder.getPtrTy());
10432 }
10433
10434 for (Value var : vars) {
10435 Value baseVar = getBaseValueForTypeLookup(var);
10436 llvm::Type *typeToInspect =
10437 getAllocatedLlvmTypeForVariable(var, baseVar, moduleTranslation);
10438
10439 llvm::Value *size;
10440 if (std::optional<llvm::Value *> dynamicSize = getDynamicAllocatedSize(
10441 var, baseVar, moduleTranslation, builder, dataLayout)) {
10442 size = *dynamicSize;
10443 } else if (typeToInspect->isArrayTy()) {
10444 size = builder.getInt64(
10445 dataLayout.getTypeAllocSize(typeToInspect).getFixedValue());
10446 } else {
10447 size = builder.getInt64(
10448 dataLayout.getTypeAllocSize(typeToInspect).getFixedValue());
10449 }
10450
10451 uint64_t alignValue =
10452 alignAttr ? alignAttr.value()
10453 : dataLayout.getABITypeAlign(typeToInspect).value();
10454 llvm::Value *alignConst = builder.getInt64(alignValue);
10455 // Align the size: ((size + align - 1) / align) * align
10456 size = builder.CreateAdd(size, builder.getInt64(alignValue - 1), "", true);
10457 size = builder.CreateUDiv(size, alignConst);
10458 size = builder.CreateMul(size, alignConst, "", true);
10459
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,
10466 allocName);
10467 } else {
10468 allocCall =
10469 ompBuilder->createOMPAlloc(ompLoc, size, allocator, allocName);
10470 }
10471 // Record the alloc pointer keyed by the MLIR variable value.
10472 ompIface.registerAllocatedPtr(var, allocCall);
10473
10474 if (llvm::Value *baseLlvm = moduleTranslation.lookupValue(baseVar)) {
10475 llvm::Value *boundPtr = builder.CreatePointerBitCastOrAddrSpaceCast(
10476 allocCall, baseLlvm->getType());
10477 moduleTranslation.remapAllValuesWith(baseLlvm, boundPtr);
10478 } else if (llvm::Value *varLlvm = moduleTranslation.lookupValue(var)) {
10479 llvm::Value *boundPtr = builder.CreatePointerBitCastOrAddrSpaceCast(
10480 allocCall, varLlvm->getType());
10481 moduleTranslation.remapAllValuesWith(varLlvm, boundPtr);
10482 }
10483 }
10484
10485 return success();
10486}
10487
10488static LogicalResult
10489convertAllocateFreeOp(Operation &opInst, llvm::IRBuilderBase &builder,
10490 LLVM::ModuleTranslation &moduleTranslation,
10491 const OpenMPDialectLLVMIRTranslationInterface &ompIface) {
10492 auto freeOp = cast<omp::AllocateFreeOp>(opInst);
10493 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
10494 llvm::OpenMPIRBuilder::LocationDescription ompLoc(builder);
10495
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());
10504 } else {
10505 allocator = llvm::ConstantPointerNull::get(builder.getPtrTy());
10506 }
10507
10508 // Emit __kmpc_free for each variable in reverse allocation order.
10509 SmallVector<Value> vars = freeOp.getVarList();
10510 for (Value var : llvm::reverse(vars)) {
10511 llvm::Value *allocPtr = ompIface.lookupAllocatedPtr(var);
10512 if (!allocPtr)
10513 return opInst.emitError("omp.allocate_free: no allocation recorded");
10514 ompBuilder->createOMPFree(ompLoc, allocPtr, allocator, "");
10515 }
10516
10517 return success();
10518}
10519
10520static llvm::Function *getOmpTargetFree(llvm::IRBuilderBase &builder,
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());
10529 return func;
10530}
10531
10532static LogicalResult
10533convertTargetFreeMemOp(Operation &opInst, llvm::IRBuilderBase &builder,
10534 LLVM::ModuleTranslation &moduleTranslation) {
10535 auto freeMemOp = cast<omp::TargetFreeMemOp>(opInst);
10536 if (!freeMemOp)
10537 return failure();
10538
10539 // Get "omp_target_free" function
10540 llvm::Module *llvmModule = moduleTranslation.getLLVMModule();
10541 llvm::Function *ompTragetFreeFunc = getOmpTargetFree(builder, llvmModule);
10542 // Get the corresponding device value in llvm
10543 mlir::Value deviceNum = freeMemOp.getDevice();
10544 llvm::Value *llvmDeviceNum = moduleTranslation.lookupValue(deviceNum);
10545 // Get the corresponding heapref value in llvm
10546 mlir::Value heapref = freeMemOp.getHeapref();
10547 llvm::Value *llvmHeapref = moduleTranslation.lookupValue(heapref);
10548 // Convert heapref int to ptr and call "omp_target_free"
10549 llvm::Value *intToPtr =
10550 builder.CreateIntToPtr(llvmHeapref, builder.getPtrTy(0));
10551 builder.CreateCall(ompTragetFreeFunc, {intToPtr, llvmDeviceNum});
10552 return success();
10553}
10554
10555static LogicalResult
10556convertFreeSharedMemOp(omp::FreeSharedMemOp freeMemOp,
10557 llvm::IRBuilderBase &builder,
10558 LLVM::ModuleTranslation &moduleTranslation) {
10559 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
10560 llvm::Value *size = getAllocationSize(builder, moduleTranslation, freeMemOp);
10561 ompBuilder->createOMPFreeShared(
10562 builder, moduleTranslation.lookupValue(freeMemOp.getHeapref()), size);
10563 return success();
10564}
10565
10566/// Converts an OpenMP groupprivate operation into LLVM IR.
10567static LogicalResult
10568convertOmpGroupprivate(Operation &opInst, llvm::IRBuilderBase &builder,
10569 LLVM::ModuleTranslation &moduleTranslation) {
10570 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
10571 auto groupprivateOp = cast<omp::GroupprivateOp>(opInst);
10572
10573 if (failed(checkImplementationStatus(opInst)))
10574 return failure();
10575
10576 bool isTargetDevice = ompBuilder->Config.isTargetDevice();
10577
10578 // Determine whether group-private storage should be allocated based on
10579 // device_type. When not specified, default to 'any' (allocate on both).
10580 bool shouldAllocate = true;
10581 switch (groupprivateOp.getDeviceType().value_or(
10582 mlir::omp::DeclareTargetDeviceType::any)) {
10583 case mlir::omp::DeclareTargetDeviceType::host:
10584 shouldAllocate = !isTargetDevice;
10585 break;
10586 case mlir::omp::DeclareTargetDeviceType::nohost:
10587 shouldAllocate = isTargetDevice;
10588 break;
10589 case mlir::omp::DeclareTargetDeviceType::any:
10590 shouldAllocate = true;
10591 break;
10592 }
10593
10594 // Look up the global variable directly by symbol name.
10596 &opInst, groupprivateOp.getSymNameAttr());
10597 if (!global)
10598 return opInst.emitError()
10599 << "expected symbol '" << groupprivateOp.getSymName()
10600 << "' to reference an LLVM global variable";
10601
10602 llvm::GlobalValue *globalValue = moduleTranslation.lookupGlobal(global);
10603 llvm::Type *varType = moduleTranslation.convertType(global.getType());
10604 std::string varName = globalValue->getName().str();
10605
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;
10615 else
10616 return opInst.emitError() << "groupprivate is not supported for target: "
10617 << targetTriple.str();
10618 llvm::GlobalVariable *sharedVar = new llvm::GlobalVariable(
10619 *llvmModule, varType, /*isConstant=*/false,
10620 llvm::GlobalValue::InternalLinkage, llvm::PoisonValue::get(varType),
10621 varName, /*InsertBefore=*/nullptr, llvm::GlobalValue::NotThreadLocal,
10622 sharedAddressSpace,
10623 /*isExternallyInitialized=*/false);
10624 resultPtr = sharedVar;
10625 } else {
10626 if (shouldAllocate && !isTargetDevice)
10627 opInst.emitWarning("groupprivate directive is currently ignored on the "
10628 "host, using original global");
10629 resultPtr = globalValue;
10630 }
10631
10632 moduleTranslation.mapValue(opInst.getResult(0), resultPtr);
10633 return success();
10634}
10635
10636/// Given an OpenMP MLIR operation, create the corresponding LLVM IR (including
10637/// OpenMP runtime calls).
10638LogicalResult OpenMPDialectLLVMIRTranslationInterface::convertOperation(
10639 Operation *op, llvm::IRBuilderBase &builder,
10640 LLVM::ModuleTranslation &moduleTranslation) const {
10641 llvm::OpenMPIRBuilder *ompBuilder = moduleTranslation.getOpenMPBuilder();
10642
10643 if (ompBuilder->Config.isTargetDevice() &&
10644 !isa<omp::TargetOp, omp::MapInfoOp, omp::TerminatorOp, omp::YieldOp>(
10645 op) &&
10646 isHostDeviceOp(op))
10647 return op->emitOpError() << "unsupported host op found in device";
10648
10649 // For each loop, introduce one stack frame to hold loop information. Ensure
10650 // this is only done for the outermost loop wrapper to prevent introducing
10651 // multiple stack frames for a single loop. Initially set to null, the loop
10652 // information structure is initialized during translation of the nested
10653 // omp.loop_nest operation, making it available to translation of all loop
10654 // wrappers after their body has been successfully translated.
10655 bool isOutermostLoopWrapper =
10656 isa_and_present<omp::LoopWrapperInterface>(op) &&
10657 !dyn_cast_if_present<omp::LoopWrapperInterface>(op->getParentOp());
10658
10659 // The TASKLOOP construct is implemented with an outer taskloop.context
10660 // operation which is not a loop wrapper, containing an inner taskloop
10661 // operation which is a loop wrapper. The stack frame should be pushed when
10662 // translating the outer taskloop.context and popped when translating the
10663 // inner taskloop which is a loop wrapper. We need access to the loop
10664 // information in the outer taskloop context so we need to create it and pop
10665 // it around the taskloop context not the inner loop wrapper.
10666 if (isa<omp::TaskloopContextOp>(op))
10667 isOutermostLoopWrapper = true;
10668 else if (isa<omp::TaskloopWrapperOp>(op))
10669 isOutermostLoopWrapper = false;
10670
10671 if (isOutermostLoopWrapper)
10672 moduleTranslation.stackPush<OpenMPLoopInfoStackFrame>();
10673
10674 auto result =
10675 llvm::TypeSwitch<Operation *, LogicalResult>(op)
10676 .Case([&](omp::BarrierOp op) -> LogicalResult {
10678 return failure();
10679
10680 llvm::OpenMPIRBuilder::InsertPointOrErrorTy afterIP =
10681 ompBuilder->createBarrier(builder, llvm::omp::OMPD_barrier);
10682 LogicalResult res = handleError(afterIP, *op);
10683 if (res.succeeded()) {
10684 // If the barrier generated a cancellation check, the insertion
10685 // point might now need to be changed to a new continuation block
10686 builder.restoreIP(*afterIP);
10687 }
10688 return res;
10689 })
10690 .Case([&](omp::TaskyieldOp op) {
10692 return failure();
10693
10694 ompBuilder->createTaskyield(builder);
10695 return success();
10696 })
10697 .Case([&](omp::FlushOp op) {
10699 return failure();
10700
10701 // No support in Openmp runtime function (__kmpc_flush) to accept
10702 // the argument list.
10703 // OpenMP standard states the following:
10704 // "An implementation may implement a flush with a list by ignoring
10705 // the list, and treating it the same as a flush without a list."
10706 //
10707 // The argument list is discarded so that, flush with a list is
10708 // treated same as a flush without a list.
10709 ompBuilder->createFlush(builder);
10710 return success();
10711 })
10712 .Case([&](omp::ErrorOp op) {
10714 return failure();
10715
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);
10725 return success();
10726 })
10727 .Case([&](omp::ParallelOp op) {
10728 return convertOmpParallel(op, builder, moduleTranslation);
10729 })
10730 .Case([&](omp::DispatchOp) {
10731 return convertOmpDispatch(*op, builder, moduleTranslation);
10732 })
10733 .Case([&](omp::MaskedOp) {
10734 return convertOmpMasked(*op, builder, moduleTranslation);
10735 })
10736 .Case([&](omp::MasterOp) {
10737 return convertOmpMaster(*op, builder, moduleTranslation);
10738 })
10739 .Case([&](omp::CriticalOp) {
10740 return convertOmpCritical(*op, builder, moduleTranslation);
10741 })
10742 .Case([&](omp::OrderedRegionOp) {
10743 return convertOmpOrderedRegion(*op, builder, moduleTranslation);
10744 })
10745 .Case([&](omp::OrderedOp) {
10746 return convertOmpOrdered(*op, builder, moduleTranslation);
10747 })
10748 .Case([&](omp::WsloopOp) {
10749 return convertOmpWsloop(*op, builder, moduleTranslation);
10750 })
10751 .Case([&](omp::SimdOp) {
10752 return convertOmpSimd(*op, builder, moduleTranslation);
10753 })
10754 .Case([&](omp::AtomicReadOp) {
10755 return convertOmpAtomicRead(*op, builder, moduleTranslation);
10756 })
10757 .Case([&](omp::AtomicWriteOp) {
10758 return convertOmpAtomicWrite(*op, builder, moduleTranslation);
10759 })
10760 .Case([&](omp::AtomicUpdateOp op) {
10761 return convertOmpAtomicUpdate(op, builder, moduleTranslation);
10762 })
10763 .Case([&](omp::AtomicCaptureOp op) {
10764 return convertOmpAtomicCapture(op, builder, moduleTranslation);
10765 })
10766 .Case([&](omp::AtomicCompareOp op) {
10767 return convertOmpAtomicCompare(op, builder, moduleTranslation);
10768 })
10769 .Case([&](omp::CancelOp op) {
10770 return convertOmpCancel(op, builder, moduleTranslation);
10771 })
10772 .Case([&](omp::CancellationPointOp op) {
10773 return convertOmpCancellationPoint(op, builder, moduleTranslation);
10774 })
10775 .Case([&](omp::SectionsOp) {
10776 return convertOmpSections(*op, builder, moduleTranslation);
10777 })
10778 .Case([&](omp::ScopeOp op) {
10779 return convertOmpScope(op, builder, moduleTranslation);
10780 })
10781 .Case([&](omp::SingleOp op) {
10782 return convertOmpSingle(op, builder, moduleTranslation);
10783 })
10784 .Case([&](omp::TeamsOp op) {
10785 return convertOmpTeams(op, builder, moduleTranslation);
10786 })
10787 .Case([&](omp::TaskOp op) {
10788 return convertOmpTaskOp(op, builder, moduleTranslation);
10789 })
10790 .Case([&](omp::TaskloopWrapperOp op) {
10791 return convertOmpTaskloopWrapperOp(op, builder, moduleTranslation);
10792 })
10793 .Case([&](omp::TaskloopContextOp op) {
10794 return convertOmpTaskloopContextOp(op, builder, moduleTranslation);
10795 })
10796 .Case([&](omp::TaskgroupOp op) {
10797 return convertOmpTaskgroupOp(op, builder, moduleTranslation);
10798 })
10799 .Case([&](omp::TaskwaitOp op) {
10800 return convertOmpTaskwaitOp(op, builder, moduleTranslation);
10801 })
10802 .Case([&](omp::InteropInitOp op) {
10803 return convertOmpInteropInitOp(op, builder, moduleTranslation);
10804 })
10805 .Case([&](omp::InteropDestroyOp op) {
10806 return convertOmpInteropDestroyOp(op, builder, moduleTranslation);
10807 })
10808 .Case([&](omp::InteropUseOp op) {
10809 return convertOmpInteropUseOp(op, builder, moduleTranslation);
10810 })
10811 .Case<omp::YieldOp, omp::TerminatorOp, omp::DeclareMapperOp,
10812 omp::DeclareMapperInfoOp, omp::DeclareReductionOp,
10813 omp::CriticalDeclareOp>([](auto op) {
10814 // `yield` and `terminator` can be just omitted. The block structure
10815 // was created in the region that handles their parent operation.
10816 // `declare_reduction` will be used by reductions and is not
10817 // converted directly, skip it.
10818 // `declare_mapper` and `declare_mapper.info` are handled whenever
10819 // they are referred to through a `map` clause.
10820 // `critical.declare` is only used to declare names of critical
10821 // sections which will be used by `critical` ops and hence can be
10822 // ignored for lowering. The OpenMP IRBuilder will create unique
10823 // name for critical section names.
10824 return success();
10825 })
10826 .Case([&](omp::ThreadprivateOp) {
10827 return convertOmpThreadprivate(*op, builder, moduleTranslation);
10828 })
10829 .Case<omp::TargetDataOp, omp::TargetEnterDataOp,
10830 omp::TargetExitDataOp, omp::TargetUpdateOp>([&](auto op) {
10831 return convertOmpTargetData(op, builder, moduleTranslation);
10832 })
10833 .Case([&](omp::TargetOp) {
10834 return convertOmpTarget(*op, builder, moduleTranslation);
10835 })
10836 .Case([&](omp::DistributeOp) {
10837 return convertOmpDistribute(*op, builder, moduleTranslation);
10838 })
10839 .Case([&](omp::LoopNestOp) {
10840 return convertOmpLoopNest(*op, builder, moduleTranslation);
10841 })
10842 .Case<omp::MapInfoOp, omp::MapBoundsOp, omp::PrivateClauseOp,
10843 omp::AffinityEntryOp, omp::IteratorOp>([&](auto op) {
10844 // No-op, should be handled by relevant owning operations e.g.
10845 // TargetOp, TargetEnterDataOp, TargetExitDataOp, TargetDataOp
10846 // etc. and then discarded
10847 return success();
10848 })
10849 .Case([&](omp::NewCliOp op) {
10850 // Meta-operation: Doesn't do anything by itself, but used to
10851 // identify a loop.
10852 return success();
10853 })
10854 .Case([&](omp::CanonicalLoopOp op) {
10855 return convertOmpCanonicalLoopOp(op, builder, moduleTranslation);
10856 })
10857 .Case([&](omp::UnrollHeuristicOp op) {
10858 // FIXME: Handling omp.unroll_heuristic as an executable requires
10859 // that the generator (e.g. omp.canonical_loop) has been seen first.
10860 // For construct that require all codegen to occur inside a callback
10861 // (e.g. OpenMPIRBilder::createParallel), all codegen of that
10862 // contained region including their transformations must occur at
10863 // the omp.canonical_loop.
10864 return applyUnrollHeuristic(op, builder, moduleTranslation);
10865 })
10866 .Case([&](omp::UnrollFullOp op) {
10867 return applyUnrollFull(op, builder, moduleTranslation);
10868 })
10869 .Case([&](omp::UnrollPartialOp op) {
10870 return applyUnrollPartial(op, builder, moduleTranslation);
10871 })
10872 .Case([&](omp::TileOp op) {
10873 return applyTile(op, builder, moduleTranslation);
10874 })
10875 .Case([&](omp::FuseOp op) {
10876 return applyFuse(op, builder, moduleTranslation);
10877 })
10878 .Case([&](omp::TargetAllocMemOp) {
10879 return convertTargetAllocMemOp(*op, builder, moduleTranslation);
10880 })
10881 .Case([&](omp::TargetFreeMemOp) {
10882 return convertTargetFreeMemOp(*op, builder, moduleTranslation);
10883 })
10884 .Case([&](omp::AllocateDirOp) {
10885 return convertAllocateDirOp(*op, builder, moduleTranslation, *this);
10886 })
10887 .Case([&](omp::AllocateFreeOp) {
10888 return convertAllocateFreeOp(*op, builder, moduleTranslation,
10889 *this);
10890 })
10891 .Case([&](omp::AllocSharedMemOp op) {
10892 return convertAllocSharedMemOp(op, builder, moduleTranslation);
10893 })
10894 .Case([&](omp::FreeSharedMemOp op) {
10895 return convertFreeSharedMemOp(op, builder, moduleTranslation);
10896 })
10897 .Case([&](omp::GroupprivateOp) {
10898 return convertOmpGroupprivate(*op, builder, moduleTranslation);
10899 })
10900 .Default([&](Operation *inst) {
10901 return inst->emitError()
10902 << "not yet implemented: " << inst->getName();
10903 });
10904
10905 if (isOutermostLoopWrapper)
10906 moduleTranslation.stackPop();
10907
10908 return result;
10909}
10910
10912 registry.insert<omp::OpenMPDialect>();
10913 registry.addExtension(+[](MLIRContext *ctx, omp::OpenMPDialect *dialect) {
10914 dialect->addInterfaces<OpenMPDialectLLVMIRTranslationInterface>();
10915 });
10916}
10917
10919 DialectRegistry registry;
10921 context.appendDialectRegistry(registry);
10922}
for(Operation *op :ops)
return success()
if(failed(verifyVectorMemoryOp(getOperation(), memrefType, getVectorType()))) return failure()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
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 &region, 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 &region, 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)
Definition TypeID.h:331
#define div(a, b)
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class represents an argument of a Block.
Definition Value.h:306
Block represents an ordered list of Operations.
Definition Block.h:34
BlockArgument getArgument(unsigned i)
Definition Block.h:154
unsigned getNumArguments()
Definition Block.h:153
OpListType & getOperations()
Definition Block.h:162
Operation & front()
Definition Block.h:178
Operation & back()
Definition Block.h:177
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
iterator_range< iterator > without_terminator()
Return an iterator range over the operation within this block excluding the terminator operation at t...
Definition Block.h:222
iterator begin()
Definition Block.h:168
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.
Definition Location.h:174
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 &region)
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.
Definition TypeToLLVM.h:39
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.
Definition Location.h:45
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
void appendDialectRegistry(const DialectRegistry &registry)
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.
Definition Attributes.h:179
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
StringAttr getIdentifier() const
Return the name of this operation as a StringAttr.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
Value getOperand(unsigned idx)
Definition Operation.h:375
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.
Definition Operation.h:432
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:726
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
unsigned getNumOperands()
Definition Operation.h:371
OperandRange operand_range
Definition Operation.h:396
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'.
Definition Operation.h:255
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
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),...
Definition Operation.h:849
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
result_range getResults()
Definition Operation.h:440
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
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.
Definition Operation.h:429
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
Block & front()
Definition Region.h:65
BlockArgListType getArguments()
Definition Region.h:94
bool empty()
Definition Region.h:60
unsigned getNumArguments()
Definition Region.h:136
iterator begin()
Definition Region.h:55
Operation * getParentOp()
Return the parent operation this region is attached to.
Definition Region.h:198
BlockListType & getBlocks()
Definition Region.h:45
bool hasOneBlock()
Return true if this region has exactly one block.
Definition Region.h:68
Concrete CRTP base class for StateStack frames.
Definition StateStack.h:47
@ Private
The symbol is private and may only be referenced by SymbolRefAttrs local to the operations within the...
Definition SymbolTable.h:92
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...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
user_range getUsers() const
Definition Value.h:218
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
A utility result that is used to signal how to proceed with an ongoing walk:
Definition WalkResult.h:29
static WalkResult skip()
Definition WalkResult.h:48
static WalkResult advance()
Definition WalkResult.h:47
bool wasInterrupted() const
Returns true if the walk was interrupted.
Definition WalkResult.h:51
static WalkResult interrupt()
Definition WalkResult.h:46
The OpAsmOpInterface, see OpAsmInterface.td for more details.
Definition CallGraph.h:227
void connectPHINodes(Region &region, 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.
Definition Utils.cpp:56
bool opInSharedDeviceContext(Operation &op)
Check whether the given operation is located in a context where an allocation to be used by multiple ...
Definition Utils.cpp:113
bool allocaUsesRequireSharedMem(Value alloc)
Check whether the value representing an allocation, assumed to have been defined in a shared device c...
Definition Utils.cpp:98
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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 &region)
Gets a list of blocks that is sorted according to dominance.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
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 &registry)
Register the OpenMP dialect and the translation from it to the LLVM IR in the given registry;.
llvm::SetVector< T, Vector, Set, N > SetVector
Definition LLVM.h:125
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...
Definition Utils.cpp:1380
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
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.
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.