1//===- OpenMPIRBuilder.cpp - Builder for LLVM-IR for OpenMP directives ----===//
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/// \file
9///
10/// This file implements the OpenMPIRBuilder class, which is used as a
11/// convenient way to create LLVM instructions for OpenMP directives.
12///
13//===----------------------------------------------------------------------===//
14
15#include "llvm/Frontend/OpenMP/OMPIRBuilder.h"
16#include "llvm/ADT/SmallBitVector.h"
17#include "llvm/ADT/SmallSet.h"
18#include "llvm/ADT/SmallVectorExtras.h"
19#include "llvm/ADT/StringExtras.h"
20#include "llvm/ADT/StringRef.h"
21#include "llvm/Analysis/AssumptionCache.h"
22#include "llvm/Analysis/CodeMetrics.h"
23#include "llvm/Analysis/LoopInfo.h"
24#include "llvm/Analysis/OptimizationRemarkEmitter.h"
25#include "llvm/Analysis/PostDominators.h"
26#include "llvm/Analysis/ScalarEvolution.h"
27#include "llvm/Analysis/TargetLibraryInfo.h"
28#include "llvm/Bitcode/BitcodeReader.h"
29#include "llvm/Frontend/Offloading/Utility.h"
30#include "llvm/Frontend/OpenMP/OMPGridValues.h"
31#include "llvm/IR/Attributes.h"
32#include "llvm/IR/BasicBlock.h"
33#include "llvm/IR/CFG.h"
34#include "llvm/IR/CallingConv.h"
35#include "llvm/IR/Constant.h"
36#include "llvm/IR/Constants.h"
37#include "llvm/IR/DIBuilder.h"
38#include "llvm/IR/DebugInfoMetadata.h"
39#include "llvm/IR/DerivedTypes.h"
40#include "llvm/IR/Function.h"
41#include "llvm/IR/GlobalVariable.h"
42#include "llvm/IR/IRBuilder.h"
43#include "llvm/IR/InstIterator.h"
44#include "llvm/IR/IntrinsicInst.h"
45#include "llvm/IR/LLVMContext.h"
46#include "llvm/IR/MDBuilder.h"
47#include "llvm/IR/Metadata.h"
48#include "llvm/IR/PassInstrumentation.h"
49#include "llvm/IR/PassManager.h"
50#include "llvm/IR/ReplaceConstant.h"
51#include "llvm/IR/Value.h"
52#include "llvm/MC/TargetRegistry.h"
53#include "llvm/Support/CommandLine.h"
54#include "llvm/Support/Error.h"
55#include "llvm/Support/ErrorHandling.h"
56#include "llvm/Support/NVVMAttributes.h"
57#include "llvm/Support/VirtualFileSystem.h"
58#include "llvm/Target/TargetMachine.h"
59#include "llvm/Target/TargetOptions.h"
60#include "llvm/Transforms/Utils/BasicBlockUtils.h"
61#include "llvm/Transforms/Utils/Cloning.h"
62#include "llvm/Transforms/Utils/CodeExtractor.h"
63#include "llvm/Transforms/Utils/LoopPeel.h"
64#include "llvm/Transforms/Utils/UnrollLoop.h"
65
66#include <cstdint>
67#include <optional>
68
69#define DEBUG_TYPE "openmp-ir-builder"
70
71using namespace llvm;
72using namespace omp;
73
74static cl::opt<bool>
75 OptimisticAttributes("openmp-ir-builder-optimistic-attributes", cl::Hidden,
76 cl::desc("Use optimistic attributes describing "
77 "'as-if' properties of runtime calls."),
78 cl::init(Val: false));
79
80static cl::opt<double> UnrollThresholdFactor(
81 "openmp-ir-builder-unroll-threshold-factor", cl::Hidden,
82 cl::desc("Factor for the unroll threshold to account for code "
83 "simplifications still taking place"),
84 cl::init(Val: 1.5));
85
86static cl::opt<bool> UseDefaultMaxThreads(
87 "openmp-ir-builder-use-default-max-threads", cl::Hidden,
88 cl::desc("Use a default max threads if none is provided."), cl::init(Val: true));
89
90#ifndef NDEBUG
91/// Return whether IP1 and IP2 are ambiguous, i.e. that inserting instructions
92/// at position IP1 may change the meaning of IP2 or vice-versa. This is because
93/// an InsertPoint stores the instruction before something is inserted. For
94/// instance, if both point to the same instruction, two IRBuilders alternating
95/// creating instruction will cause the instructions to be interleaved.
96static bool isConflictIP(IRBuilder<>::InsertPoint IP1,
97 IRBuilder<>::InsertPoint IP2) {
98 if (!IP1.isValid() || !IP2.isValid())
99 return false;
100 return IP1 == IP2;
101}
102
103static bool isValidWorkshareLoopScheduleType(OMPScheduleType SchedType) {
104 // Valid ordered/unordered and base algorithm combinations.
105 switch (SchedType & ~OMPScheduleType::MonotonicityMask) {
106 case OMPScheduleType::UnorderedStaticChunked:
107 case OMPScheduleType::UnorderedStatic:
108 case OMPScheduleType::UnorderedDynamicChunked:
109 case OMPScheduleType::UnorderedGuidedChunked:
110 case OMPScheduleType::UnorderedRuntime:
111 case OMPScheduleType::UnorderedAuto:
112 case OMPScheduleType::UnorderedTrapezoidal:
113 case OMPScheduleType::UnorderedGreedy:
114 case OMPScheduleType::UnorderedBalanced:
115 case OMPScheduleType::UnorderedGuidedIterativeChunked:
116 case OMPScheduleType::UnorderedGuidedAnalyticalChunked:
117 case OMPScheduleType::UnorderedSteal:
118 case OMPScheduleType::UnorderedStaticBalancedChunked:
119 case OMPScheduleType::UnorderedGuidedSimd:
120 case OMPScheduleType::UnorderedRuntimeSimd:
121 case OMPScheduleType::OrderedStaticChunked:
122 case OMPScheduleType::OrderedStatic:
123 case OMPScheduleType::OrderedDynamicChunked:
124 case OMPScheduleType::OrderedGuidedChunked:
125 case OMPScheduleType::OrderedRuntime:
126 case OMPScheduleType::OrderedAuto:
127 case OMPScheduleType::OrderdTrapezoidal:
128 case OMPScheduleType::NomergeUnorderedStaticChunked:
129 case OMPScheduleType::NomergeUnorderedStatic:
130 case OMPScheduleType::NomergeUnorderedDynamicChunked:
131 case OMPScheduleType::NomergeUnorderedGuidedChunked:
132 case OMPScheduleType::NomergeUnorderedRuntime:
133 case OMPScheduleType::NomergeUnorderedAuto:
134 case OMPScheduleType::NomergeUnorderedTrapezoidal:
135 case OMPScheduleType::NomergeUnorderedGreedy:
136 case OMPScheduleType::NomergeUnorderedBalanced:
137 case OMPScheduleType::NomergeUnorderedGuidedIterativeChunked:
138 case OMPScheduleType::NomergeUnorderedGuidedAnalyticalChunked:
139 case OMPScheduleType::NomergeUnorderedSteal:
140 case OMPScheduleType::NomergeOrderedStaticChunked:
141 case OMPScheduleType::NomergeOrderedStatic:
142 case OMPScheduleType::NomergeOrderedDynamicChunked:
143 case OMPScheduleType::NomergeOrderedGuidedChunked:
144 case OMPScheduleType::NomergeOrderedRuntime:
145 case OMPScheduleType::NomergeOrderedAuto:
146 case OMPScheduleType::NomergeOrderedTrapezoidal:
147 case OMPScheduleType::OrderedDistributeChunked:
148 case OMPScheduleType::OrderedDistribute:
149 break;
150 default:
151 return false;
152 }
153
154 // Must not set both monotonicity modifiers at the same time.
155 OMPScheduleType MonotonicityFlags =
156 SchedType & OMPScheduleType::MonotonicityMask;
157 if (MonotonicityFlags == OMPScheduleType::MonotonicityMask)
158 return false;
159
160 return true;
161}
162#endif
163
164/// This is a wrapper over IRBuilderBase::restoreIP that also restores a current
165/// debug location when the insert point is at the end of a block. It picks a
166/// location scoped to the current function: the block's last instruction
167/// location if the block is non-empty, otherwise a location synthesized from
168/// the function's subprogram (when the function has debug info).
169static void restoreIPandDebugLoc(llvm::IRBuilderBase &Builder,
170 llvm::IRBuilderBase::InsertPoint IP) {
171 Builder.restoreIP(IP);
172 // When IP points at a real instruction, restoreIP (SetInsertPoint) already
173 // set the debug location from that instruction, so leave it alone.
174 llvm::BasicBlock *BB = Builder.GetInsertBlock();
175 if (Builder.GetInsertPoint() != BB->end())
176 return;
177
178 // At the end of a block, pick a location guaranteed to belong to the current
179 // insertion function's subprogram. Prefer the block's own last instruction;
180 // otherwise synthesize a location from the function's subprogram.
181 if (!BB->empty())
182 Builder.SetCurrentDebugLocation(BB->back().getDebugLoc());
183 else if (llvm::DISubprogram *FSP =
184 BB->getParent() ? BB->getParent()->getSubprogram() : nullptr) {
185 unsigned Line = FSP->getScopeLine() ? FSP->getScopeLine() : FSP->getLine();
186 Builder.SetCurrentDebugLocation(
187 llvm::DILocation::get(Context&: FSP->getContext(), Line, /*Column=*/0, Scope: FSP));
188 }
189}
190
191static bool hasGridValue(const Triple &T) {
192 return T.isAMDGPU() || T.isNVPTX() || T.isSPIRV();
193}
194
195static const omp::GV &getGridValue(const Triple &T, Function *Kernel) {
196 if (T.isAMDGPU()) {
197 StringRef Features =
198 Kernel->getFnAttribute(Kind: "target-features").getValueAsString();
199 if (Features.count(Str: "+wavefrontsize64"))
200 return omp::getAMDGPUGridValues<64>();
201 return omp::getAMDGPUGridValues<32>();
202 }
203 if (T.isNVPTX())
204 return omp::NVPTXGridValues;
205 if (T.isSPIRV())
206 return omp::SPIRVGridValues;
207 llvm_unreachable("No grid value available for this architecture!");
208}
209
210/// Determine which scheduling algorithm to use, determined from schedule clause
211/// arguments.
212static OMPScheduleType
213getOpenMPBaseScheduleType(llvm::omp::ScheduleKind ClauseKind, bool HasChunks,
214 bool HasSimdModifier, bool HasDistScheduleChunks) {
215 // Currently, the default schedule it static.
216 switch (ClauseKind) {
217 case OMP_SCHEDULE_Default:
218 case OMP_SCHEDULE_Static:
219 return HasChunks ? OMPScheduleType::BaseStaticChunked
220 : OMPScheduleType::BaseStatic;
221 case OMP_SCHEDULE_Dynamic:
222 return OMPScheduleType::BaseDynamicChunked;
223 case OMP_SCHEDULE_Guided:
224 return HasSimdModifier ? OMPScheduleType::BaseGuidedSimd
225 : OMPScheduleType::BaseGuidedChunked;
226 case OMP_SCHEDULE_Auto:
227 return llvm::omp::OMPScheduleType::BaseAuto;
228 case OMP_SCHEDULE_Runtime:
229 return HasSimdModifier ? OMPScheduleType::BaseRuntimeSimd
230 : OMPScheduleType::BaseRuntime;
231 case OMP_SCHEDULE_Distribute:
232 return HasDistScheduleChunks ? OMPScheduleType::BaseDistributeChunked
233 : OMPScheduleType::BaseDistribute;
234 }
235 llvm_unreachable("unhandled schedule clause argument");
236}
237
238/// Adds ordering modifier flags to schedule type.
239static OMPScheduleType
240getOpenMPOrderingScheduleType(OMPScheduleType BaseScheduleType,
241 bool HasOrderedClause) {
242 assert((BaseScheduleType & OMPScheduleType::ModifierMask) ==
243 OMPScheduleType::None &&
244 "Must not have ordering nor monotonicity flags already set");
245
246 OMPScheduleType OrderingModifier = HasOrderedClause
247 ? OMPScheduleType::ModifierOrdered
248 : OMPScheduleType::ModifierUnordered;
249 OMPScheduleType OrderingScheduleType = BaseScheduleType | OrderingModifier;
250
251 // Unsupported combinations
252 if (OrderingScheduleType ==
253 (OMPScheduleType::BaseGuidedSimd | OMPScheduleType::ModifierOrdered))
254 return OMPScheduleType::OrderedGuidedChunked;
255 else if (OrderingScheduleType == (OMPScheduleType::BaseRuntimeSimd |
256 OMPScheduleType::ModifierOrdered))
257 return OMPScheduleType::OrderedRuntime;
258
259 return OrderingScheduleType;
260}
261
262/// Adds monotonicity modifier flags to schedule type.
263static OMPScheduleType
264getOpenMPMonotonicityScheduleType(OMPScheduleType ScheduleType,
265 bool HasSimdModifier, bool HasMonotonic,
266 bool HasNonmonotonic, bool HasOrderedClause) {
267 assert((ScheduleType & OMPScheduleType::MonotonicityMask) ==
268 OMPScheduleType::None &&
269 "Must not have monotonicity flags already set");
270 assert((!HasMonotonic || !HasNonmonotonic) &&
271 "Monotonic and Nonmonotonic are contradicting each other");
272
273 if (HasMonotonic) {
274 return ScheduleType | OMPScheduleType::ModifierMonotonic;
275 } else if (HasNonmonotonic) {
276 return ScheduleType | OMPScheduleType::ModifierNonmonotonic;
277 } else {
278 // OpenMP 5.1, 2.11.4 Worksharing-Loop Construct, Description.
279 // If the static schedule kind is specified or if the ordered clause is
280 // specified, and if the nonmonotonic modifier is not specified, the
281 // effect is as if the monotonic modifier is specified. Otherwise, unless
282 // the monotonic modifier is specified, the effect is as if the
283 // nonmonotonic modifier is specified.
284 OMPScheduleType BaseScheduleType =
285 ScheduleType & ~OMPScheduleType::ModifierMask;
286 if ((BaseScheduleType == OMPScheduleType::BaseStatic) ||
287 (BaseScheduleType == OMPScheduleType::BaseStaticChunked) ||
288 HasOrderedClause) {
289 // The monotonic is used by default in openmp runtime library, so no need
290 // to set it.
291 return ScheduleType;
292 } else {
293 return ScheduleType | OMPScheduleType::ModifierNonmonotonic;
294 }
295 }
296}
297
298/// Determine the schedule type using schedule and ordering clause arguments.
299static OMPScheduleType
300computeOpenMPScheduleType(ScheduleKind ClauseKind, bool HasChunks,
301 bool HasSimdModifier, bool HasMonotonicModifier,
302 bool HasNonmonotonicModifier, bool HasOrderedClause,
303 bool HasDistScheduleChunks) {
304 OMPScheduleType BaseSchedule = getOpenMPBaseScheduleType(
305 ClauseKind, HasChunks, HasSimdModifier, HasDistScheduleChunks);
306 OMPScheduleType OrderedSchedule =
307 getOpenMPOrderingScheduleType(BaseScheduleType: BaseSchedule, HasOrderedClause);
308 OMPScheduleType Result = getOpenMPMonotonicityScheduleType(
309 ScheduleType: OrderedSchedule, HasSimdModifier, HasMonotonic: HasMonotonicModifier,
310 HasNonmonotonic: HasNonmonotonicModifier, HasOrderedClause);
311
312 assert(isValidWorkshareLoopScheduleType(Result));
313 return Result;
314}
315
316/// Given a function, if it represents the entry point of a target kernel, this
317/// returns the execution mode flags associated with that kernel.
318static std::optional<omp::OMPTgtExecModeFlags>
319getTargetKernelExecMode(Function &Kernel) {
320 CallInst *TargetInitCall = nullptr;
321 for (Instruction &Inst : Kernel.getEntryBlock()) {
322 if (auto *Call = dyn_cast<CallInst>(Val: &Inst)) {
323 if (Call->getCalledFunction()->getName() == "__kmpc_target_init") {
324 TargetInitCall = Call;
325 break;
326 }
327 }
328 }
329
330 if (!TargetInitCall)
331 return std::nullopt;
332
333 // Get the kernel mode information from the global variable associated to the
334 // first argument to the call to __kmpc_target_init. Refer to
335 // createTargetInit() to see how this is initialized.
336 Value *InitOperand = TargetInitCall->getArgOperand(i: 0);
337 GlobalVariable *KernelEnv = nullptr;
338 if (auto *Cast = dyn_cast<ConstantExpr>(Val: InitOperand))
339 KernelEnv = cast<GlobalVariable>(Val: Cast->getOperand(i_nocapture: 0));
340 else
341 KernelEnv = cast<GlobalVariable>(Val: InitOperand);
342 auto *KernelEnvInit = cast<ConstantStruct>(Val: KernelEnv->getInitializer());
343 auto *ConfigEnv = cast<ConstantStruct>(Val: KernelEnvInit->getOperand(i_nocapture: 0));
344 auto *KernelMode = cast<ConstantInt>(Val: ConfigEnv->getOperand(i_nocapture: 2));
345 return static_cast<OMPTgtExecModeFlags>(KernelMode->getZExtValue());
346}
347
348static bool isGenericKernel(Function &Fn) {
349 std::optional<omp::OMPTgtExecModeFlags> ExecMode =
350 getTargetKernelExecMode(Kernel&: Fn);
351 return !ExecMode || (*ExecMode & OMP_TGT_EXEC_MODE_GENERIC);
352}
353
354/// Make \p Source branch to \p Target.
355///
356/// Handles two situations:
357/// * \p Source already has an unconditional branch.
358/// * \p Source is a degenerate block (no terminator because the BB is
359/// the current head of the IR construction).
360static void redirectTo(BasicBlock *Source, BasicBlock *Target, DebugLoc DL) {
361 if (Instruction *Term = Source->getTerminatorOrNull()) {
362 auto *Br = cast<UncondBrInst>(Val: Term);
363 BasicBlock *Succ = Br->getSuccessor();
364 Succ->removePredecessor(Pred: Source, /*KeepOneInputPHIs=*/true);
365 Br->setSuccessor(Target);
366 return;
367 }
368
369 auto *NewBr = UncondBrInst::Create(Target, InsertBefore: Source);
370 NewBr->setDebugLoc(DL);
371}
372
373void llvm::spliceBB(IRBuilderBase::InsertPoint IP, BasicBlock *New,
374 bool CreateBranch, DebugLoc DL) {
375 assert(New->getFirstInsertionPt() == New->begin() &&
376 "Target BB must not have PHI nodes");
377
378 // Move instructions to new block.
379 BasicBlock *Old = IP.getNodeParent();
380 // If the `Old` block is empty then there are no instructions to move. But in
381 // the new debug scheme, it could have trailing debug records which will be
382 // moved to `New` in `spliceDebugInfoEmptyBlock`. We dont want that for 2
383 // reasons:
384 // 1. If `New` is also empty, `BasicBlock::splice` crashes.
385 // 2. Even if `New` is not empty, the rationale to move those records to `New`
386 // (in `spliceDebugInfoEmptyBlock`) does not apply here. That function
387 // assumes that `Old` is optimized out and is going away. This is not the case
388 // here. The `Old` block is still being used e.g. a branch instruction is
389 // added to it later in this function.
390 // So we call `BasicBlock::splice` only when `Old` is not empty.
391 if (!Old->empty())
392 New->splice(ToIt: New->begin(), FromBB: Old, FromBeginIt: IP, FromEndIt: Old->end());
393
394 if (CreateBranch) {
395 auto *NewBr = UncondBrInst::Create(Target: New, InsertBefore: Old);
396 NewBr->setDebugLoc(DL);
397 }
398}
399
400void llvm::spliceBB(IRBuilder<> &Builder, BasicBlock *New, bool CreateBranch) {
401 DebugLoc DebugLoc = Builder.getCurrentDebugLocation();
402 BasicBlock *Old = Builder.GetInsertBlock();
403
404 spliceBB(IP: Builder.saveIP(), New, CreateBranch, DL: DebugLoc);
405 if (CreateBranch)
406 Builder.SetInsertPoint(Old->getTerminator());
407 else
408 Builder.SetInsertPoint(Old);
409
410 // SetInsertPoint also updates the Builder's debug location, but we want to
411 // keep the one the Builder was configured to use.
412 Builder.SetCurrentDebugLocation(DebugLoc);
413}
414
415BasicBlock *llvm::splitBB(IRBuilderBase::InsertPoint IP, bool CreateBranch,
416 DebugLoc DL, llvm::Twine Name) {
417 BasicBlock *Old = IP.getNodeParent();
418 BasicBlock *New = BasicBlock::Create(
419 Context&: Old->getContext(), Name: Name.isTriviallyEmpty() ? Old->getName() : Name,
420 Parent: Old->getParent(), InsertBefore: Old->getNextNode());
421 spliceBB(IP, New, CreateBranch, DL);
422 New->replaceSuccessorsPhiUsesWith(Old, New);
423 return New;
424}
425
426BasicBlock *llvm::splitBB(IRBuilderBase &Builder, bool CreateBranch,
427 llvm::Twine Name) {
428 DebugLoc DebugLoc = Builder.getCurrentDebugLocation();
429 BasicBlock *New = splitBB(IP: Builder.saveIP(), CreateBranch, DL: DebugLoc, Name);
430 if (CreateBranch)
431 Builder.SetInsertPoint(Builder.GetInsertBlock()->getTerminator());
432 else
433 Builder.SetInsertPoint(Builder.GetInsertBlock());
434 // SetInsertPoint also updates the Builder's debug location, but we want to
435 // keep the one the Builder was configured to use.
436 Builder.SetCurrentDebugLocation(DebugLoc);
437 return New;
438}
439
440BasicBlock *llvm::splitBB(IRBuilder<> &Builder, bool CreateBranch,
441 llvm::Twine Name) {
442 DebugLoc DebugLoc = Builder.getCurrentDebugLocation();
443 BasicBlock *New = splitBB(IP: Builder.saveIP(), CreateBranch, DL: DebugLoc, Name);
444 if (CreateBranch)
445 Builder.SetInsertPoint(Builder.GetInsertBlock()->getTerminator());
446 else
447 Builder.SetInsertPoint(Builder.GetInsertBlock());
448 // SetInsertPoint also updates the Builder's debug location, but we want to
449 // keep the one the Builder was configured to use.
450 Builder.SetCurrentDebugLocation(DebugLoc);
451 return New;
452}
453
454BasicBlock *llvm::splitBBWithSuffix(IRBuilderBase &Builder, bool CreateBranch,
455 llvm::Twine Suffix) {
456 BasicBlock *Old = Builder.GetInsertBlock();
457 return splitBB(Builder, CreateBranch, Name: Old->getName() + Suffix);
458}
459
460// This function creates a fake integer value and a fake use for the integer
461// value. It returns the fake value created. This is useful in modeling the
462// extra arguments to the outlined functions.
463Value *createFakeIntVal(IRBuilderBase &Builder,
464 OpenMPIRBuilder::InsertPointTy OuterAllocaIP,
465 llvm::SmallVectorImpl<Instruction *> &ToBeDeleted,
466 OpenMPIRBuilder::InsertPointTy InnerAllocaIP,
467 const Twine &Name = "", bool AsPtr = true,
468 bool Is64Bit = false) {
469 Builder.restoreIP(IP: OuterAllocaIP);
470 IntegerType *IntTy = Is64Bit ? Builder.getInt64Ty() : Builder.getInt32Ty();
471 Instruction *FakeVal;
472 AllocaInst *FakeValAddr =
473 Builder.CreateAlloca(Ty: IntTy, ArraySize: nullptr, Name: Name + ".addr");
474 ToBeDeleted.push_back(Elt: FakeValAddr);
475
476 if (AsPtr) {
477 FakeVal = FakeValAddr;
478 // The runtime passes these extra arguments to the outlined function as
479 // generic pointers, so cast away a non-zero alloca address space.
480 if (FakeValAddr->getAddressSpace() != 0) {
481 FakeVal = cast<Instruction>(Val: Builder.CreateAddrSpaceCast(
482 V: FakeValAddr, DestTy: Builder.getPtrTy(), Name: Name + ".ascast"));
483 ToBeDeleted.push_back(Elt: FakeVal);
484 }
485 } else {
486 FakeVal = Builder.CreateLoad(Ty: IntTy, Ptr: FakeValAddr, Name: Name + ".val");
487 ToBeDeleted.push_back(Elt: FakeVal);
488 }
489
490 // Generate a fake use of this value
491 Builder.restoreIP(IP: InnerAllocaIP);
492 Instruction *UseFakeVal;
493 if (AsPtr) {
494 UseFakeVal = Builder.CreateLoad(Ty: IntTy, Ptr: FakeVal, Name: Name + ".use");
495 } else {
496 UseFakeVal = cast<BinaryOperator>(Val: Builder.CreateAdd(
497 LHS: FakeVal, RHS: Is64Bit ? Builder.getInt64(C: 10) : Builder.getInt32(C: 10)));
498 }
499 ToBeDeleted.push_back(Elt: UseFakeVal);
500 return FakeVal;
501}
502
503//===----------------------------------------------------------------------===//
504// OpenMPIRBuilderConfig
505//===----------------------------------------------------------------------===//
506
507namespace {
508LLVM_ENABLE_BITMASK_ENUMS_IN_NAMESPACE();
509/// Values for bit flags for marking which requires clauses have been used.
510enum OpenMPOffloadingRequiresDirFlags {
511 /// flag undefined.
512 OMP_REQ_UNDEFINED = 0x000,
513 /// no requires directive present.
514 OMP_REQ_NONE = 0x001,
515 /// reverse_offload clause.
516 OMP_REQ_REVERSE_OFFLOAD = 0x002,
517 /// unified_address clause.
518 OMP_REQ_UNIFIED_ADDRESS = 0x004,
519 /// unified_shared_memory clause.
520 OMP_REQ_UNIFIED_SHARED_MEMORY = 0x008,
521 /// dynamic_allocators clause.
522 OMP_REQ_DYNAMIC_ALLOCATORS = 0x010,
523 LLVM_MARK_AS_BITMASK_ENUM(/*LargestValue=*/OMP_REQ_DYNAMIC_ALLOCATORS)
524};
525
526class OMPCodeExtractor : public CodeExtractor {
527public:
528 OMPCodeExtractor(OpenMPIRBuilder &OMPBuilder, ArrayRef<BasicBlock *> BBs,
529 DominatorTree *DT = nullptr, bool AggregateArgs = false,
530 BlockFrequencyInfo *BFI = nullptr,
531 BranchProbabilityInfo *BPI = nullptr,
532 AssumptionCache *AC = nullptr, bool AllowVarArgs = false,
533 bool AllowAlloca = false,
534 BasicBlock *AllocationBlock = nullptr,
535 ArrayRef<BasicBlock *> DeallocationBlocks = {},
536 std::string Suffix = "", bool ArgsInZeroAddressSpace = false)
537 : CodeExtractor(BBs, DT, AggregateArgs, BFI, BPI, AC, AllowVarArgs,
538 AllowAlloca, AllocationBlock, DeallocationBlocks, Suffix,
539 ArgsInZeroAddressSpace),
540 OMPBuilder(OMPBuilder) {}
541
542 virtual ~OMPCodeExtractor() = default;
543
544protected:
545 OpenMPIRBuilder &OMPBuilder;
546};
547
548class DeviceSharedMemCodeExtractor : public OMPCodeExtractor {
549public:
550 using OMPCodeExtractor::OMPCodeExtractor;
551 virtual ~DeviceSharedMemCodeExtractor() = default;
552
553protected:
554 virtual Instruction *
555 allocateVar(IRBuilder<>::InsertPoint AllocaIP, DebugLoc DL, Type *VarType,
556 const Twine &Name = Twine(""),
557 AddrSpaceCastInst **CastedAlloc = nullptr) override {
558 return OMPBuilder.createOMPAllocShared(Loc: {AllocaIP, DL}, VarType, Name);
559 }
560
561 virtual Instruction *deallocateVar(IRBuilder<>::InsertPoint DeallocIP,
562 DebugLoc DL, Value *Var,
563 Type *VarType) override {
564 return OMPBuilder.createOMPFreeShared(Loc: {DeallocIP, DL}, Addr: Var, VarType);
565 }
566};
567
568/// Helper storing information about regions to outline using device shared
569/// memory for intermediate allocations.
570struct DeviceSharedMemOutlineInfo : public OpenMPIRBuilder::OutlineInfo {
571 OpenMPIRBuilder &OMPBuilder;
572
573 DeviceSharedMemOutlineInfo(OpenMPIRBuilder &OMPBuilder)
574 : OMPBuilder(OMPBuilder) {}
575 virtual ~DeviceSharedMemOutlineInfo() = default;
576
577 virtual std::unique_ptr<CodeExtractor>
578 createCodeExtractor(ArrayRef<BasicBlock *> Blocks,
579 bool ArgsInZeroAddressSpace,
580 Twine Suffix = Twine("")) override;
581};
582
583} // anonymous namespace
584
585OpenMPIRBuilderConfig::OpenMPIRBuilderConfig()
586 : RequiresFlags(OMP_REQ_UNDEFINED) {}
587
588OpenMPIRBuilderConfig::OpenMPIRBuilderConfig(
589 bool IsTargetDevice, bool IsGPU, bool OpenMPOffloadMandatory,
590 bool HasRequiresReverseOffload, bool HasRequiresUnifiedAddress,
591 bool HasRequiresUnifiedSharedMemory, bool HasRequiresDynamicAllocators)
592 : IsTargetDevice(IsTargetDevice), IsGPU(IsGPU),
593 OpenMPOffloadMandatory(OpenMPOffloadMandatory),
594 RequiresFlags(OMP_REQ_UNDEFINED) {
595 if (HasRequiresReverseOffload)
596 RequiresFlags |= OMP_REQ_REVERSE_OFFLOAD;
597 if (HasRequiresUnifiedAddress)
598 RequiresFlags |= OMP_REQ_UNIFIED_ADDRESS;
599 if (HasRequiresUnifiedSharedMemory)
600 RequiresFlags |= OMP_REQ_UNIFIED_SHARED_MEMORY;
601 if (HasRequiresDynamicAllocators)
602 RequiresFlags |= OMP_REQ_DYNAMIC_ALLOCATORS;
603}
604
605bool OpenMPIRBuilderConfig::hasRequiresReverseOffload() const {
606 return RequiresFlags & OMP_REQ_REVERSE_OFFLOAD;
607}
608
609bool OpenMPIRBuilderConfig::hasRequiresUnifiedAddress() const {
610 return RequiresFlags & OMP_REQ_UNIFIED_ADDRESS;
611}
612
613bool OpenMPIRBuilderConfig::hasRequiresUnifiedSharedMemory() const {
614 return RequiresFlags & OMP_REQ_UNIFIED_SHARED_MEMORY;
615}
616
617bool OpenMPIRBuilderConfig::hasRequiresDynamicAllocators() const {
618 return RequiresFlags & OMP_REQ_DYNAMIC_ALLOCATORS;
619}
620
621int64_t OpenMPIRBuilderConfig::getRequiresFlags() const {
622 return hasRequiresFlags() ? RequiresFlags
623 : static_cast<int64_t>(OMP_REQ_NONE);
624}
625
626void OpenMPIRBuilderConfig::setHasRequiresReverseOffload(bool Value) {
627 if (Value)
628 RequiresFlags |= OMP_REQ_REVERSE_OFFLOAD;
629 else
630 RequiresFlags &= ~OMP_REQ_REVERSE_OFFLOAD;
631}
632
633void OpenMPIRBuilderConfig::setHasRequiresUnifiedAddress(bool Value) {
634 if (Value)
635 RequiresFlags |= OMP_REQ_UNIFIED_ADDRESS;
636 else
637 RequiresFlags &= ~OMP_REQ_UNIFIED_ADDRESS;
638}
639
640void OpenMPIRBuilderConfig::setHasRequiresUnifiedSharedMemory(bool Value) {
641 if (Value)
642 RequiresFlags |= OMP_REQ_UNIFIED_SHARED_MEMORY;
643 else
644 RequiresFlags &= ~OMP_REQ_UNIFIED_SHARED_MEMORY;
645}
646
647void OpenMPIRBuilderConfig::setHasRequiresDynamicAllocators(bool Value) {
648 if (Value)
649 RequiresFlags |= OMP_REQ_DYNAMIC_ALLOCATORS;
650 else
651 RequiresFlags &= ~OMP_REQ_DYNAMIC_ALLOCATORS;
652}
653
654//===----------------------------------------------------------------------===//
655// OpenMPIRBuilder
656//===----------------------------------------------------------------------===//
657
658void OpenMPIRBuilder::getKernelArgsVector(TargetKernelArgs &KernelArgs,
659 IRBuilderBase &Builder,
660 SmallVector<Value *> &ArgsVector) {
661 Value *Version = Builder.getInt32(OMP_KERNEL_ARG_VERSION);
662 Value *PointerNum = Builder.getInt32(C: KernelArgs.NumTargetItems);
663 auto Int32Ty = Type::getInt32Ty(C&: Builder.getContext());
664 constexpr size_t MaxDim = 3;
665 Value *ZeroArray = Constant::getNullValue(Ty: ArrayType::get(ElementType: Int32Ty, NumElements: MaxDim));
666
667 Value *HasNoWaitFlag = Builder.getInt64(C: KernelArgs.HasNoWait);
668
669 Value *DynCGroupMemFallbackFlag =
670 Builder.getInt64(C: static_cast<uint64_t>(KernelArgs.DynCGroupMemFallback));
671 DynCGroupMemFallbackFlag = Builder.CreateShl(LHS: DynCGroupMemFallbackFlag, RHS: 2);
672
673 Value *StrictBlocksFlag = Builder.getInt64(C: KernelArgs.StrictBlocks);
674 Value *StrictThreadsFlag = Builder.getInt64(C: KernelArgs.StrictThreads);
675
676 StrictBlocksFlag = Builder.CreateShl(LHS: StrictBlocksFlag, RHS: 6);
677 StrictThreadsFlag = Builder.CreateShl(LHS: StrictThreadsFlag, RHS: 7);
678
679 Value *Flags = Builder.CreateOr(LHS: HasNoWaitFlag, RHS: DynCGroupMemFallbackFlag);
680 Flags = Builder.CreateOr(LHS: Flags, RHS: StrictBlocksFlag);
681 Flags = Builder.CreateOr(LHS: Flags, RHS: StrictThreadsFlag);
682
683 assert(!KernelArgs.NumTeams.empty() && !KernelArgs.NumThreads.empty());
684
685 Value *NumTeams3D =
686 Builder.CreateInsertValue(Agg: ZeroArray, Val: KernelArgs.NumTeams[0], Idxs: {0});
687 Value *NumThreads3D =
688 Builder.CreateInsertValue(Agg: ZeroArray, Val: KernelArgs.NumThreads[0], Idxs: {0});
689 for (unsigned I :
690 seq<unsigned>(Begin: 1, End: std::min(a: KernelArgs.NumTeams.size(), b: MaxDim)))
691 NumTeams3D =
692 Builder.CreateInsertValue(Agg: NumTeams3D, Val: KernelArgs.NumTeams[I], Idxs: {I});
693 for (unsigned I :
694 seq<unsigned>(Begin: 1, End: std::min(a: KernelArgs.NumThreads.size(), b: MaxDim)))
695 NumThreads3D =
696 Builder.CreateInsertValue(Agg: NumThreads3D, Val: KernelArgs.NumThreads[I], Idxs: {I});
697
698 ArgsVector = {Version,
699 PointerNum,
700 KernelArgs.RTArgs.BasePointersArray,
701 KernelArgs.RTArgs.PointersArray,
702 KernelArgs.RTArgs.SizesArray,
703 KernelArgs.RTArgs.MapTypesArray,
704 KernelArgs.RTArgs.MapNamesArray,
705 KernelArgs.RTArgs.MappersArray,
706 KernelArgs.NumIterations,
707 Flags,
708 NumTeams3D,
709 NumThreads3D,
710 KernelArgs.DynCGroupMem};
711}
712
713void OpenMPIRBuilder::addAttributes(omp::RuntimeFunction FnID, Function &Fn) {
714 LLVMContext &Ctx = Fn.getContext();
715
716 // Get the function's current attributes.
717 auto Attrs = Fn.getAttributes();
718 auto FnAttrs = Attrs.getFnAttrs();
719 auto RetAttrs = Attrs.getRetAttrs();
720 SmallVector<AttributeSet, 4> ArgAttrs;
721 for (size_t ArgNo = 0; ArgNo < Fn.arg_size(); ++ArgNo)
722 ArgAttrs.emplace_back(Args: Attrs.getParamAttrs(ArgNo));
723
724 // Add AS to FnAS while taking special care with integer extensions.
725 auto addAttrSet = [&](AttributeSet &FnAS, const AttributeSet &AS,
726 bool Param = true) -> void {
727 bool HasSignExt = AS.hasAttribute(Kind: Attribute::SExt);
728 bool HasZeroExt = AS.hasAttribute(Kind: Attribute::ZExt);
729 if (HasSignExt || HasZeroExt) {
730 assert(AS.getNumAttributes() == 1 &&
731 "Currently not handling extension attr combined with others.");
732 if (Param) {
733 if (auto AK = TargetLibraryInfo::getExtAttrForI32Param(T, Signed: HasSignExt))
734 FnAS = FnAS.addAttribute(C&: Ctx, Kind: AK);
735 } else if (auto AK =
736 TargetLibraryInfo::getExtAttrForI32Return(T, Signed: HasSignExt))
737 FnAS = FnAS.addAttribute(C&: Ctx, Kind: AK);
738 } else {
739 FnAS = FnAS.addAttributes(C&: Ctx, AS);
740 }
741 };
742
743#define OMP_ATTRS_SET(VarName, AttrSet) AttributeSet VarName = AttrSet;
744#include "llvm/Frontend/OpenMP/OMPKinds.def"
745
746 // Add attributes to the function declaration.
747 switch (FnID) {
748#define OMP_RTL_ATTRS(Enum, FnAttrSet, RetAttrSet, ArgAttrSets) \
749 case Enum: \
750 FnAttrs = FnAttrs.addAttributes(Ctx, FnAttrSet); \
751 addAttrSet(RetAttrs, RetAttrSet, /*Param*/ false); \
752 for (size_t ArgNo = 0; ArgNo < ArgAttrSets.size(); ++ArgNo) \
753 addAttrSet(ArgAttrs[ArgNo], ArgAttrSets[ArgNo]); \
754 Fn.setAttributes(AttributeList::get(Ctx, FnAttrs, RetAttrs, ArgAttrs)); \
755 break;
756#include "llvm/Frontend/OpenMP/OMPKinds.def"
757 default:
758 // Attributes are optional.
759 break;
760 }
761}
762
763FunctionCallee
764OpenMPIRBuilder::getOrCreateRuntimeFunction(Module &M, RuntimeFunction FnID) {
765 FunctionType *FnTy = nullptr;
766 Function *Fn = nullptr;
767
768 // Try to find the declation in the module first.
769 switch (FnID) {
770#define OMP_RTL(Enum, Str, IsVarArg, ReturnType, ...) \
771 case Enum: \
772 FnTy = FunctionType::get(ReturnType, ArrayRef<Type *>{__VA_ARGS__}, \
773 IsVarArg); \
774 Fn = M.getFunction(Str); \
775 break;
776#include "llvm/Frontend/OpenMP/OMPKinds.def"
777 }
778
779 if (!Fn) {
780 // Create a new declaration if we need one.
781 switch (FnID) {
782#define OMP_RTL(Enum, Str, ...) \
783 case Enum: \
784 Fn = Function::Create(FnTy, GlobalValue::ExternalLinkage, Str, M); \
785 break;
786#include "llvm/Frontend/OpenMP/OMPKinds.def"
787 }
788 Fn->setCallingConv(Config.getRuntimeCC());
789 // Add information if the runtime function takes a callback function
790 if (FnID == OMPRTL___kmpc_fork_call || FnID == OMPRTL___kmpc_fork_teams) {
791 if (!Fn->hasMetadata(KindID: LLVMContext::MD_callback)) {
792 LLVMContext &Ctx = Fn->getContext();
793 MDBuilder MDB(Ctx);
794 // Annotate the callback behavior of the runtime function:
795 // - The callback callee is argument number 2 (microtask).
796 // - The first two arguments of the callback callee are unknown (-1).
797 // - All variadic arguments to the runtime function are passed to the
798 // callback callee.
799 Fn->addMetadata(
800 KindID: LLVMContext::MD_callback,
801 MD&: *MDNode::get(Context&: Ctx, MDs: {MDB.createCallbackEncoding(
802 CalleeArgNo: 2, Arguments: {-1, -1}, /* VarArgsArePassed */ true)}));
803 }
804 }
805
806 LLVM_DEBUG(dbgs() << "Created OpenMP runtime function " << Fn->getName()
807 << " with type " << *Fn->getFunctionType() << "\n");
808 addAttributes(FnID, Fn&: *Fn);
809
810 } else {
811 LLVM_DEBUG(dbgs() << "Found OpenMP runtime function " << Fn->getName()
812 << " with type " << *Fn->getFunctionType() << "\n");
813 }
814
815 assert(Fn && "Failed to create OpenMP runtime function");
816
817 return {FnTy, Fn};
818}
819
820Expected<BasicBlock *>
821OpenMPIRBuilder::FinalizationInfo::getFiniBB(IRBuilderBase &Builder) {
822 if (!FiniBB) {
823 Function *ParentFunc = Builder.GetInsertBlock()->getParent();
824 IRBuilderBase::InsertPointGuard Guard(Builder);
825 FiniBB = BasicBlock::Create(Context&: Builder.getContext(), Name: ".fini", Parent: ParentFunc);
826 Builder.SetInsertPoint(FiniBB);
827 // FiniCB adds the branch to the exit stub.
828 if (Error Err = FiniCB(Builder.saveIP()))
829 return Err;
830 }
831 return FiniBB;
832}
833
834Error OpenMPIRBuilder::FinalizationInfo::mergeFiniBB(IRBuilderBase &Builder,
835 BasicBlock *OtherFiniBB) {
836 // Simple case: FiniBB does not exist yet: re-use OtherFiniBB.
837 if (!FiniBB) {
838 FiniBB = OtherFiniBB;
839
840 Builder.SetInsertPoint(FiniBB->getFirstNonPHIIt());
841 if (Error Err = FiniCB(Builder.saveIP()))
842 return Err;
843
844 return Error::success();
845 }
846
847 // Move instructions from FiniBB to the start of OtherFiniBB.
848 auto EndIt = FiniBB->end();
849 if (FiniBB->size() >= 1)
850 if (auto Prev = std::prev(x: EndIt); Prev->isTerminator())
851 EndIt = Prev;
852 OtherFiniBB->splice(ToIt: OtherFiniBB->getFirstNonPHIIt(), FromBB: FiniBB, FromBeginIt: FiniBB->begin(),
853 FromEndIt: EndIt);
854
855 FiniBB->replaceAllUsesWith(V: OtherFiniBB);
856 FiniBB->eraseFromParent();
857 FiniBB = OtherFiniBB;
858 return Error::success();
859}
860
861Function *OpenMPIRBuilder::getOrCreateRuntimeFunctionPtr(RuntimeFunction FnID) {
862 FunctionCallee RTLFn = getOrCreateRuntimeFunction(M, FnID);
863 auto *Fn = dyn_cast<llvm::Function>(Val: RTLFn.getCallee());
864 assert(Fn && "Failed to create OpenMP runtime function pointer");
865 return Fn;
866}
867
868CallInst *OpenMPIRBuilder::createRuntimeFunctionCall(FunctionCallee Callee,
869 ArrayRef<Value *> Args,
870 StringRef Name) {
871 CallInst *Call = Builder.CreateCall(Callee, Args, Name);
872 Call->setCallingConv(Config.getRuntimeCC());
873 return Call;
874}
875
876void OpenMPIRBuilder::initialize() { initializeTypes(M); }
877
878static void raiseUserConstantDataAllocasToEntryBlock(IRBuilderBase &Builder,
879 Function *Function) {
880 BasicBlock &EntryBlock = Function->getEntryBlock();
881 BasicBlock::iterator MoveLocInst = EntryBlock.getFirstNonPHIIt();
882
883 // Loop over blocks looking for constant allocas, skipping the entry block
884 // as any allocas there are already in the desired location.
885 for (auto Block = std::next(x: Function->begin(), n: 1); Block != Function->end();
886 Block++) {
887 for (auto Inst = Block->getReverseIterator()->begin();
888 Inst != Block->getReverseIterator()->end();) {
889 if (auto *AllocaInst = dyn_cast_if_present<llvm::AllocaInst>(Val&: Inst)) {
890 Inst++;
891 if (!isa<ConstantData>(Val: AllocaInst->getArraySize()))
892 continue;
893 AllocaInst->moveBeforePreserving(MovePos: MoveLocInst);
894 } else {
895 Inst++;
896 }
897 }
898 }
899}
900
901static void hoistNonEntryAllocasToEntryBlock(llvm::BasicBlock &Block) {
902 llvm::SmallVector<llvm::Instruction *> AllocasToMove;
903
904 auto ShouldHoistAlloca = [](const llvm::AllocaInst &AllocaInst) {
905 // TODO: For now, we support simple static allocations, we might need to
906 // move non-static ones as well. However, this will need further analysis to
907 // move the lenght arguments as well.
908 return !AllocaInst.isArrayAllocation();
909 };
910
911 for (llvm::Instruction &Inst : Block)
912 if (auto *AllocaInst = llvm::dyn_cast<llvm::AllocaInst>(Val: &Inst))
913 if (ShouldHoistAlloca(*AllocaInst))
914 AllocasToMove.push_back(Elt: AllocaInst);
915
916 auto InsertPoint =
917 Block.getParent()->getEntryBlock().getTerminator()->getIterator();
918
919 for (llvm::Instruction *AllocaInst : AllocasToMove)
920 AllocaInst->moveBefore(InsertPos: InsertPoint);
921}
922
923static void hoistNonEntryAllocasToEntryBlock(llvm::Function *Func) {
924 PostDominatorTree PostDomTree(*Func);
925 for (llvm::BasicBlock &BB : *Func)
926 if (PostDomTree.properlyDominates(A: &BB, B: &Func->getEntryBlock()))
927 hoistNonEntryAllocasToEntryBlock(Block&: BB);
928}
929
930void OpenMPIRBuilder::finalize(Function *Fn) {
931 SmallPtrSet<BasicBlock *, 32> ParallelRegionBlockSet;
932 SmallVector<BasicBlock *, 32> Blocks;
933 SmallVector<std::unique_ptr<OutlineInfo>, 16> DeferredOutlines;
934 for (std::unique_ptr<OutlineInfo> &OI : OutlineInfos) {
935 // Skip functions that have not finalized yet; may happen with nested
936 // function generation.
937 if (Fn && OI->getFunction() != Fn) {
938 DeferredOutlines.push_back(Elt: std::move(OI));
939 continue;
940 }
941
942 ParallelRegionBlockSet.clear();
943 Blocks.clear();
944 OI->collectBlocks(BlockSet&: ParallelRegionBlockSet, BlockVector&: Blocks);
945
946 Function *OuterFn = OI->getFunction();
947 CodeExtractorAnalysisCache CEAC(*OuterFn);
948 // If we generate code for the target device, we need to allocate
949 // struct for aggregate params in the device default alloca address space.
950 // OpenMP runtime requires that the params of the extracted functions are
951 // passed as zero address space pointers. This flag ensures that
952 // CodeExtractor generates correct code for extracted functions
953 // which are used by OpenMP runtime.
954 bool ArgsInZeroAddressSpace = Config.isTargetDevice();
955 std::unique_ptr<CodeExtractor> Extractor =
956 OI->createCodeExtractor(Blocks, ArgsInZeroAddressSpace, Suffix: ".omp_par");
957
958 LLVM_DEBUG(dbgs() << "Before outlining: " << *OuterFn << "\n");
959 LLVM_DEBUG(dbgs() << "Entry " << OI->EntryBB->getName()
960 << " Exit: " << OI->ExitBB->getName() << "\n");
961 assert(Extractor->isEligible() &&
962 "Expected OpenMP outlining to be possible!");
963
964 for (auto *V : OI->ExcludeArgsFromAggregate)
965 Extractor->excludeArgFromAggregate(Arg: V);
966
967 Function *OutlinedFn =
968 Extractor->extractCodeRegion(CEAC, Inputs&: OI->Inputs, Outputs&: OI->Outputs);
969
970 // Forward target-cpu, target-features attributes to the outlined function.
971 auto TargetCpuAttr = OuterFn->getFnAttribute(Kind: "target-cpu");
972 if (TargetCpuAttr.isStringAttribute())
973 OutlinedFn->addFnAttr(Attr: TargetCpuAttr);
974
975 auto TargetFeaturesAttr = OuterFn->getFnAttribute(Kind: "target-features");
976 if (TargetFeaturesAttr.isStringAttribute())
977 OutlinedFn->addFnAttr(Attr: TargetFeaturesAttr);
978
979 LLVM_DEBUG(dbgs() << "After outlining: " << *OuterFn << "\n");
980 LLVM_DEBUG(dbgs() << " Outlined function: " << *OutlinedFn << "\n");
981 assert(OutlinedFn->getReturnType()->isVoidTy() &&
982 "OpenMP outlined functions should not return a value!");
983
984 // For compability with the clang CG we move the outlined function after the
985 // one with the parallel region.
986 OutlinedFn->removeFromParent();
987 M.getFunctionList().insertAfter(where: OuterFn->getIterator(), New: OutlinedFn);
988
989 // Remove the artificial entry introduced by the extractor right away, we
990 // made our own entry block after all.
991 {
992 BasicBlock &ArtificialEntry = OutlinedFn->getEntryBlock();
993 assert(ArtificialEntry.getUniqueSuccessor() == OI->EntryBB);
994 assert(OI->EntryBB->getUniquePredecessor() == &ArtificialEntry);
995 // Move instructions from the to-be-deleted ArtificialEntry to the entry
996 // basic block of the parallel region. CodeExtractor generates
997 // instructions to unwrap the aggregate argument and may sink
998 // allocas/bitcasts for values that are solely used in the outlined region
999 // and do not escape.
1000 assert(!ArtificialEntry.empty() &&
1001 "Expected instructions to add in the outlined region entry");
1002 for (BasicBlock::reverse_iterator It = ArtificialEntry.rbegin(),
1003 End = ArtificialEntry.rend();
1004 It != End;) {
1005 Instruction &I = *It;
1006 It++;
1007
1008 if (I.isTerminator()) {
1009 // Absorb any debug value that terminator may have
1010 if (Instruction *TI = OI->EntryBB->getTerminatorOrNull())
1011 TI->adoptDbgRecords(BB: &ArtificialEntry, It: I.getIterator(), InsertAtHead: false);
1012 continue;
1013 }
1014
1015 I.moveBeforePreserving(BB&: *OI->EntryBB,
1016 I: OI->EntryBB->getFirstInsertionPt());
1017 }
1018
1019 OI->EntryBB->moveBefore(MovePos: &ArtificialEntry);
1020 ArtificialEntry.eraseFromParent();
1021 }
1022 assert(&OutlinedFn->getEntryBlock() == OI->EntryBB);
1023 assert(OutlinedFn && OutlinedFn->hasNUses(1));
1024
1025 // Run a user callback, e.g. to add attributes.
1026 if (OI->PostOutlineCB)
1027 OI->PostOutlineCB(*OutlinedFn);
1028
1029 if (OI->FixUpNonEntryAllocas)
1030 hoistNonEntryAllocasToEntryBlock(Func: OutlinedFn);
1031 }
1032
1033 // Remove work items that have been completed.
1034 OutlineInfos = std::move(DeferredOutlines);
1035
1036 // The createTarget functions embeds user written code into
1037 // the target region which may inject allocas which need to
1038 // be moved to the entry block of our target or risk malformed
1039 // optimisations by later passes, this is only relevant for
1040 // the device pass which appears to be a little more delicate
1041 // when it comes to optimisations (however, we do not block on
1042 // that here, it's up to the inserter to the list to do so).
1043 // This notbaly has to occur after the OutlinedInfo candidates
1044 // have been extracted so we have an end product that will not
1045 // be implicitly adversely affected by any raises unless
1046 // intentionally appended to the list.
1047 // NOTE: This only does so for ConstantData, it could be extended
1048 // to ConstantExpr's with further effort, however, they should
1049 // largely be folded when they get here. Extending it to runtime
1050 // defined/read+writeable allocation sizes would be non-trivial
1051 // (need to factor in movement of any stores to variables the
1052 // allocation size depends on, as well as the usual loads,
1053 // otherwise it'll yield the wrong result after movement) and
1054 // likely be more suitable as an LLVM optimisation pass.
1055 for (Function *F : ConstantAllocaRaiseCandidates)
1056 raiseUserConstantDataAllocasToEntryBlock(Builder, Function: F);
1057
1058 EmitMetadataErrorReportFunctionTy &&ErrorReportFn =
1059 [](EmitMetadataErrorKind Kind,
1060 const TargetRegionEntryInfo &EntryInfo) -> void {
1061 errs() << "Error of kind: " << Kind
1062 << " when emitting offload entries and metadata during "
1063 "OMPIRBuilder finalization \n";
1064 };
1065
1066 if (!OffloadInfoManager.empty())
1067 createOffloadEntriesAndInfoMetadata(ErrorReportFunction&: ErrorReportFn);
1068
1069 // Rewrite uses of globals to their replacement declare target globals if
1070 // we are processing a device module.
1071 if (Config.isTargetDevice())
1072 applyDeclareTargetGlobalReplacements();
1073
1074 if (Config.EmitLLVMUsedMetaInfo.value_or(u: false)) {
1075 std::vector<WeakTrackingVH> LLVMCompilerUsed = {
1076 M.getGlobalVariable(Name: "__openmp_nvptx_data_transfer_temporary_storage")};
1077 emitUsed(Name: "llvm.compiler.used", List: LLVMCompilerUsed);
1078 }
1079
1080 IsFinalized = true;
1081}
1082
1083bool OpenMPIRBuilder::isFinalized() { return IsFinalized; }
1084
1085void OpenMPIRBuilder::registerDeclareTargetGlobalReplacement(
1086 GlobalValue *Original, GlobalValue *Replacement) {
1087 assert(Original && Replacement &&
1088 "Null values provided to registerDeclareTargetGlobalReplacement");
1089 DeclareTargetGlobalReplacements.push_back(Elt: {.Original: Original, .Replacement: Replacement});
1090}
1091
1092void OpenMPIRBuilder::applyDeclareTargetGlobalReplacements() {
1093 for (DeclareTargetGlobalReplacement &R : DeclareTargetGlobalReplacements) {
1094 GlobalValue *OldGV = R.Original;
1095 GlobalValue *NewGV = R.Replacement;
1096
1097 assert(OldGV && NewGV &&
1098 "A null value was inserted into DeclareTargetGlobalReplacements");
1099
1100 // The assert above should catch this case, but this is kept to attempt
1101 // to proceed without issue when asserts are off.
1102 if (!OldGV || !NewGV)
1103 continue;
1104
1105 // The replacement global is a reference pointer that holds the
1106 // address of the device-resident storage. Every use must load the
1107 // reference pointer first and use the loaded address.
1108 //
1109 // Constant expression users (e.g. a constant GEP embedded in another
1110 // global's initializer or in an instruction) cannot have a load inserted
1111 // in place, so first expand any constant-expression users that live inside
1112 // functions into instructions. Any remaining constant users are handled
1113 // via a direct constant rewrite below as we cannot materialize a load
1114 // there.
1115 //
1116 // NOTE: We extend the constant rewrite to module scope, as we replace all
1117 // usages.
1118 if (auto *OldConst = dyn_cast<Constant>(Val: OldGV))
1119 convertUsersOfConstantsToInstructions(Consts: OldConst,
1120 /*RestrictToFunc=*/nullptr,
1121 /*RemoveDeadConstants=*/false);
1122
1123 IRBuilderBase::InsertPointGuard Guard(Builder);
1124 SmallVector<User *, 16> Users(OldGV->users());
1125 for (User *U : Users) {
1126 auto *Insn = dyn_cast<Instruction>(Val: U);
1127 if (!Insn)
1128 continue;
1129
1130 // A PHI node cannot have a load inserted immediately before it, as PHIs
1131 // must remain grouped at the top of their basic block. So we need to
1132 // make sure any loads we emit are generated in the preceding edge, a
1133 // PHI may reference the global on more than one edge, so every matching
1134 // slot must be handled.
1135 if (auto *PHI = dyn_cast<PHINode>(Val: Insn)) {
1136 for (unsigned I = 0, E = PHI->getNumIncomingValues(); I < E; ++I) {
1137 if (PHI->getIncomingValue(i: I) != OldGV)
1138 continue;
1139
1140 BasicBlock *IncomingBB = PHI->getIncomingBlock(i: I);
1141 Builder.SetInsertPoint(IncomingBB->getTerminator());
1142 Builder.SetCurrentDebugLocation(PHI->getDebugLoc());
1143 LoadInst *EdgeLoad = Builder.CreateLoad(Ty: NewGV->getType(), Ptr: NewGV);
1144 PHI->setIncomingValue(i: I, V: EdgeLoad);
1145 }
1146 continue;
1147 }
1148
1149 Builder.SetInsertPoint(Insn);
1150 Builder.SetCurrentDebugLocation(Insn->getDebugLoc());
1151 LoadInst *Load = Builder.CreateLoad(Ty: NewGV->getType(), Ptr: NewGV);
1152
1153 // The replacement declare target global lives in the default address
1154 // space, whereas the original global may reside in a non-default
1155 // address space. In that case the initial lowering may have
1156 // emitted an addrspacecast that is no longer valid. Replace the
1157 // whole addrspacecast with the load and erase it rather than
1158 // feeding the load back into the (now pointless) cast.
1159 // NOTE: If we end up with replacement declare target globals in
1160 // non-zero AS's the below will need some minor extensions to have the
1161 // option to alter the address space cast to the new address space where
1162 // required rather than just replacing it.
1163 if (auto *ASC = dyn_cast<AddrSpaceCastInst>(Val: Insn)) {
1164 unsigned NewGVAS = NewGV->getType()->getPointerAddressSpace();
1165 assert(NewGVAS == 0 &&
1166 "Non-default address space declare target global");
1167 unsigned OldGVAS = OldGV->getType()->getPointerAddressSpace();
1168 unsigned DestAS = ASC->getType()->getPointerAddressSpace();
1169 if (DestAS == 0 && NewGVAS != OldGVAS) {
1170 ASC->replaceAllUsesWith(V: Load);
1171 ASC->eraseFromParent();
1172 continue;
1173 }
1174 }
1175
1176 Insn->replaceUsesOfWith(From: OldGV, To: Load);
1177 }
1178 }
1179
1180 DeclareTargetGlobalReplacements.clear();
1181}
1182
1183OpenMPIRBuilder::~OpenMPIRBuilder() {
1184 assert(OutlineInfos.empty() && "There must be no outstanding outlinings");
1185}
1186
1187GlobalValue *OpenMPIRBuilder::createGlobalFlag(unsigned Value, StringRef Name) {
1188 IntegerType *I32Ty = Type::getInt32Ty(C&: M.getContext());
1189 auto *GV =
1190 new GlobalVariable(M, I32Ty,
1191 /* isConstant = */ true, GlobalValue::WeakODRLinkage,
1192 ConstantInt::get(Ty: I32Ty, V: Value), Name);
1193 GV->setVisibility(GlobalValue::HiddenVisibility);
1194
1195 return GV;
1196}
1197
1198void OpenMPIRBuilder::emitUsed(StringRef Name, ArrayRef<WeakTrackingVH> List) {
1199 if (List.empty())
1200 return;
1201
1202 // Convert List to what ConstantArray needs.
1203 SmallVector<Constant *, 8> UsedArray;
1204 UsedArray.resize(N: List.size());
1205 for (unsigned I = 0, E = List.size(); I != E; ++I)
1206 UsedArray[I] = ConstantExpr::getPointerBitCastOrAddrSpaceCast(
1207 C: cast<Constant>(Val: &*List[I]), Ty: Builder.getPtrTy());
1208
1209 if (UsedArray.empty())
1210 return;
1211 ArrayType *ATy = ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: UsedArray.size());
1212
1213 auto *GV = new GlobalVariable(M, ATy, false, GlobalValue::AppendingLinkage,
1214 ConstantArray::get(T: ATy, V: UsedArray), Name);
1215
1216 GV->setSection("llvm.metadata");
1217}
1218
1219GlobalVariable *
1220OpenMPIRBuilder::emitKernelExecutionMode(StringRef KernelName,
1221 OMPTgtExecModeFlags Mode) {
1222 auto *Int8Ty = Builder.getInt8Ty();
1223 auto *GVMode = new GlobalVariable(
1224 M, Int8Ty, /*isConstant=*/true, GlobalValue::WeakAnyLinkage,
1225 ConstantInt::get(Ty: Int8Ty, V: Mode), Twine(KernelName, "_exec_mode"));
1226 GVMode->setVisibility(GlobalVariable::ProtectedVisibility);
1227 return GVMode;
1228}
1229
1230Constant *OpenMPIRBuilder::getOrCreateIdent(Constant *SrcLocStr,
1231 uint32_t SrcLocStrSize,
1232 IdentFlag LocFlags,
1233 unsigned Reserve2Flags) {
1234 // Enable "C-mode".
1235 LocFlags |= OMP_IDENT_FLAG_KMPC;
1236
1237 Constant *&Ident =
1238 IdentMap[{SrcLocStr, uint64_t(LocFlags) << 31 | Reserve2Flags}];
1239 if (!Ident) {
1240 Constant *I32Null = ConstantInt::getNullValue(Ty: Int32);
1241 Constant *IdentData[] = {I32Null,
1242 ConstantInt::get(Ty: Int32, V: uint32_t(LocFlags)),
1243 ConstantInt::get(Ty: Int32, V: Reserve2Flags),
1244 ConstantInt::get(Ty: Int32, V: SrcLocStrSize), SrcLocStr};
1245
1246 size_t SrcLocStrArgIdx = 4;
1247 if (OpenMPIRBuilder::Ident->getElementType(N: SrcLocStrArgIdx)
1248 ->getPointerAddressSpace() !=
1249 IdentData[SrcLocStrArgIdx]->getType()->getPointerAddressSpace())
1250 IdentData[SrcLocStrArgIdx] = ConstantExpr::getAddrSpaceCast(
1251 C: SrcLocStr, Ty: OpenMPIRBuilder::Ident->getElementType(N: SrcLocStrArgIdx));
1252 Constant *Initializer =
1253 ConstantStruct::get(T: OpenMPIRBuilder::Ident, V: IdentData);
1254
1255 // Look for existing encoding of the location + flags, not needed but
1256 // minimizes the difference to the existing solution while we transition.
1257 for (GlobalVariable &GV : M.globals())
1258 if (GV.getValueType() == OpenMPIRBuilder::Ident && GV.hasInitializer())
1259 if (GV.getInitializer() == Initializer)
1260 Ident = &GV;
1261
1262 if (!Ident) {
1263 auto *GV = new GlobalVariable(
1264 M, OpenMPIRBuilder::Ident,
1265 /* isConstant = */ true, GlobalValue::PrivateLinkage, Initializer, "",
1266 nullptr, GlobalValue::NotThreadLocal,
1267 M.getDataLayout().getDefaultGlobalsAddressSpace());
1268 GV->setUnnamedAddr(GlobalValue::UnnamedAddr::Global);
1269 GV->setAlignment(Align(8));
1270 Ident = GV;
1271 }
1272 }
1273
1274 return ConstantExpr::getPointerBitCastOrAddrSpaceCast(C: Ident, Ty: IdentPtr);
1275}
1276
1277Constant *OpenMPIRBuilder::getOrCreateSrcLocStr(StringRef LocStr,
1278 uint32_t &SrcLocStrSize) {
1279 SrcLocStrSize = LocStr.size();
1280 Constant *&SrcLocStr = SrcLocStrMap[LocStr];
1281 if (!SrcLocStr) {
1282 Constant *Initializer =
1283 ConstantDataArray::getString(Context&: M.getContext(), Initializer: LocStr);
1284
1285 // Look for existing encoding of the location, not needed but minimizes the
1286 // difference to the existing solution while we transition.
1287 for (GlobalVariable &GV : M.globals())
1288 if (GV.isConstant() && GV.hasInitializer() &&
1289 GV.getInitializer() == Initializer)
1290 return SrcLocStr = ConstantExpr::getPointerCast(C: &GV, Ty: Int8Ptr);
1291
1292 SrcLocStr = Builder.CreateGlobalString(
1293 Str: LocStr, /*Name=*/"", AddressSpace: M.getDataLayout().getDefaultGlobalsAddressSpace(),
1294 M: &M);
1295 }
1296 return SrcLocStr;
1297}
1298
1299Constant *OpenMPIRBuilder::getOrCreateSrcLocStr(StringRef FunctionName,
1300 StringRef FileName,
1301 unsigned Line, unsigned Column,
1302 uint32_t &SrcLocStrSize) {
1303 SmallString<128> Buffer;
1304 Buffer.push_back(Elt: ';');
1305 Buffer.append(RHS: FileName);
1306 Buffer.push_back(Elt: ';');
1307 Buffer.append(RHS: FunctionName);
1308 Buffer.push_back(Elt: ';');
1309 Buffer.append(RHS: std::to_string(val: Line));
1310 Buffer.push_back(Elt: ';');
1311 Buffer.append(RHS: std::to_string(val: Column));
1312 Buffer.push_back(Elt: ';');
1313 Buffer.push_back(Elt: ';');
1314 return getOrCreateSrcLocStr(LocStr: Buffer.str(), SrcLocStrSize);
1315}
1316
1317Constant *
1318OpenMPIRBuilder::getOrCreateDefaultSrcLocStr(uint32_t &SrcLocStrSize) {
1319 StringRef UnknownLoc = ";unknown;unknown;0;0;;";
1320 return getOrCreateSrcLocStr(LocStr: UnknownLoc, SrcLocStrSize);
1321}
1322
1323Constant *OpenMPIRBuilder::getOrCreateSrcLocStr(DebugLoc DL,
1324 uint32_t &SrcLocStrSize,
1325 const Function *F) {
1326 DILocation *DIL = DL.get();
1327 if (!DIL)
1328 return getOrCreateDefaultSrcLocStr(SrcLocStrSize);
1329 StringRef FileName =
1330 !DIL->getFilename().empty() ? DIL->getFilename() : M.getName();
1331 StringRef Function = DIL->getScope()->getSubprogram()->getName();
1332 if (Function.empty() && F)
1333 Function = F->getName();
1334 return getOrCreateSrcLocStr(FunctionName: Function, FileName, Line: DIL->getLine(),
1335 Column: DIL->getColumn(), SrcLocStrSize);
1336}
1337
1338Constant *OpenMPIRBuilder::getOrCreateSrcLocStr(const LocationDescription &Loc,
1339 uint32_t &SrcLocStrSize) {
1340 return getOrCreateSrcLocStr(DL: Loc.DL, SrcLocStrSize,
1341 F: Loc.IP.getNodeParent()->getParent());
1342}
1343
1344Value *OpenMPIRBuilder::getOrCreateThreadID(Value *Ident) {
1345 return createRuntimeFunctionCall(
1346 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_global_thread_num), Args: Ident,
1347 Name: "omp_global_thread_num");
1348}
1349
1350OpenMPIRBuilder::InsertPointTy OpenMPIRBuilder::createTargetInReduction(
1351 const LocationDescription &Loc, ArrayRef<Value *> OrigPtrs,
1352 ArrayRef<Type *> ResultPtrTys,
1353 function_ref<void(unsigned, Value *)> MapPrivateCB) {
1354 assert(OrigPtrs.size() == ResultPtrTys.size() &&
1355 "expected one result pointer type per in_reduction item");
1356 if (!updateToLocation(Loc))
1357 return Loc.IP;
1358 if (OrigPtrs.empty())
1359 return Builder.saveIP();
1360
1361 // Compute the executing thread's gtid once for the whole target body and
1362 // reuse it for every in_reduction lookup, so a target with several
1363 // in_reduction items does not emit a redundant __kmpc_global_thread_num per
1364 // item.
1365 uint32_t SrcLocStrSize;
1366 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
1367 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
1368 Value *Gtid = getOrCreateThreadID(Ident);
1369
1370 // The runtime entry point takes (and returns) a generic, default-address-
1371 // space `ptr`. A NULL descriptor makes the runtime walk the enclosing
1372 // taskgroups to find the matching task_reduction registration for the item.
1373 Type *PtrTy = PointerType::getUnqual(C&: M.getContext());
1374 Value *NullDesc = ConstantPointerNull::get(T: PtrTy);
1375 FunctionCallee GetThData =
1376 getOrCreateRuntimeFunction(M, FnID: OMPRTL___kmpc_task_reduction_get_th_data);
1377
1378 for (unsigned Idx = 0; Idx < OrigPtrs.size(); ++Idx) {
1379 // Normalize a non-default-address-space original pointer to the generic
1380 // address space before the call.
1381 Value *OrigPtr = OrigPtrs[Idx];
1382 if (auto *OrigPtrTy = dyn_cast<PointerType>(Val: OrigPtr->getType());
1383 OrigPtrTy && OrigPtrTy->getAddressSpace() != 0)
1384 OrigPtr = Builder.CreateAddrSpaceCast(V: OrigPtr, DestTy: PtrTy);
1385
1386 Value *Priv = Builder.CreateCall(Callee: GetThData, Args: {Gtid, NullDesc, OrigPtr},
1387 Name: "omp.inred.priv");
1388
1389 // Cast the returned private pointer back to the requested address space
1390 // when it differs.
1391 if (auto *ResPtrTy = dyn_cast<PointerType>(Val: ResultPtrTys[Idx]);
1392 ResPtrTy && ResPtrTy->getAddressSpace() != 0)
1393 Priv = Builder.CreateAddrSpaceCast(V: Priv, DestTy: ResultPtrTys[Idx]);
1394
1395 MapPrivateCB(Idx, Priv);
1396 }
1397 return Builder.saveIP();
1398}
1399
1400OpenMPIRBuilder::InsertPointOrErrorTy
1401OpenMPIRBuilder::createBarrier(const LocationDescription &Loc, Directive Kind,
1402 bool ForceSimpleCall, bool CheckCancelFlag) {
1403 if (!updateToLocation(Loc))
1404 return Loc.IP;
1405
1406 // Build call __kmpc_cancel_barrier(loc, thread_id) or
1407 // __kmpc_barrier(loc, thread_id);
1408
1409 IdentFlag BarrierLocFlags;
1410 switch (Kind) {
1411 case OMPD_for:
1412 BarrierLocFlags = OMP_IDENT_FLAG_BARRIER_IMPL_FOR;
1413 break;
1414 case OMPD_sections:
1415 BarrierLocFlags = OMP_IDENT_FLAG_BARRIER_IMPL_SECTIONS;
1416 break;
1417 case OMPD_single:
1418 BarrierLocFlags = OMP_IDENT_FLAG_BARRIER_IMPL_SINGLE;
1419 break;
1420 case OMPD_barrier:
1421 BarrierLocFlags = OMP_IDENT_FLAG_BARRIER_EXPL;
1422 break;
1423 default:
1424 BarrierLocFlags = OMP_IDENT_FLAG_BARRIER_IMPL;
1425 break;
1426 }
1427
1428 uint32_t SrcLocStrSize;
1429 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
1430 Value *Args[] = {
1431 getOrCreateIdent(SrcLocStr, SrcLocStrSize, LocFlags: BarrierLocFlags),
1432 getOrCreateThreadID(Ident: getOrCreateIdent(SrcLocStr, SrcLocStrSize))};
1433
1434 // If we are in a cancellable parallel region, barriers are cancellation
1435 // points.
1436 // TODO: Check why we would force simple calls or to ignore the cancel flag.
1437 bool UseCancelBarrier =
1438 !ForceSimpleCall && isLastFinalizationInfoCancellable(DK: OMPD_parallel);
1439
1440 Value *Result = createRuntimeFunctionCall(
1441 Callee: getOrCreateRuntimeFunctionPtr(FnID: UseCancelBarrier
1442 ? OMPRTL___kmpc_cancel_barrier
1443 : OMPRTL___kmpc_barrier),
1444 Args);
1445
1446 if (UseCancelBarrier && CheckCancelFlag)
1447 if (Error Err = emitCancelationCheckImpl(CancelFlag: Result, CanceledDirective: OMPD_parallel))
1448 return Err;
1449
1450 return Builder.saveIP();
1451}
1452
1453OpenMPIRBuilder::InsertPointOrErrorTy
1454OpenMPIRBuilder::createCancel(const LocationDescription &Loc,
1455 Value *IfCondition,
1456 omp::Directive CanceledDirective) {
1457 if (!updateToLocation(Loc))
1458 return Loc.IP;
1459
1460 // LLVM utilities like blocks with terminators.
1461 auto *UI = Builder.CreateUnreachable();
1462
1463 Instruction *ThenTI = UI, *ElseTI = nullptr;
1464 if (IfCondition) {
1465 SplitBlockAndInsertIfThenElse(Cond: IfCondition, SplitBefore: UI, ThenTerm: &ThenTI, ElseTerm: &ElseTI);
1466
1467 // Even if the if condition evaluates to false, this should count as a
1468 // cancellation point
1469 Builder.SetInsertPoint(ElseTI);
1470 auto ElseIP = Builder.saveIP();
1471
1472 InsertPointOrErrorTy IPOrErr = createCancellationPoint(
1473 Loc: LocationDescription{ElseIP, Loc.DL}, CanceledDirective);
1474 if (!IPOrErr)
1475 return IPOrErr;
1476 }
1477
1478 Builder.SetInsertPoint(ThenTI);
1479
1480 Value *CancelKind = nullptr;
1481 switch (CanceledDirective) {
1482#define OMP_CANCEL_KIND(Enum, Str, DirectiveEnum, Value) \
1483 case DirectiveEnum: \
1484 CancelKind = Builder.getInt32(Value); \
1485 break;
1486#include "llvm/Frontend/OpenMP/OMPKinds.def"
1487 default:
1488 llvm_unreachable("Unknown cancel kind!");
1489 }
1490
1491 uint32_t SrcLocStrSize;
1492 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
1493 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
1494 Value *Args[] = {Ident, getOrCreateThreadID(Ident), CancelKind};
1495 Value *Result = createRuntimeFunctionCall(
1496 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_cancel), Args);
1497
1498 // The actual cancel logic is shared with others, e.g., cancel_barriers.
1499 if (Error Err = emitCancelationCheckImpl(CancelFlag: Result, CanceledDirective))
1500 return Err;
1501
1502 // Update the insertion point and remove the terminator we introduced.
1503 Builder.SetInsertPoint(UI->getParent());
1504 UI->eraseFromParent();
1505
1506 return Builder.saveIP();
1507}
1508
1509OpenMPIRBuilder::InsertPointOrErrorTy
1510OpenMPIRBuilder::createCancellationPoint(const LocationDescription &Loc,
1511 omp::Directive CanceledDirective) {
1512 if (!updateToLocation(Loc))
1513 return Loc.IP;
1514
1515 // LLVM utilities like blocks with terminators.
1516 auto *UI = Builder.CreateUnreachable();
1517 Builder.SetInsertPoint(UI);
1518
1519 Value *CancelKind = nullptr;
1520 switch (CanceledDirective) {
1521#define OMP_CANCEL_KIND(Enum, Str, DirectiveEnum, Value) \
1522 case DirectiveEnum: \
1523 CancelKind = Builder.getInt32(Value); \
1524 break;
1525#include "llvm/Frontend/OpenMP/OMPKinds.def"
1526 default:
1527 llvm_unreachable("Unknown cancel kind!");
1528 }
1529
1530 uint32_t SrcLocStrSize;
1531 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
1532 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
1533 Value *Args[] = {Ident, getOrCreateThreadID(Ident), CancelKind};
1534 Value *Result = createRuntimeFunctionCall(
1535 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_cancellationpoint), Args);
1536
1537 // The actual cancel logic is shared with others, e.g., cancel_barriers.
1538 if (Error Err = emitCancelationCheckImpl(CancelFlag: Result, CanceledDirective))
1539 return Err;
1540
1541 // Update the insertion point and remove the terminator we introduced.
1542 Builder.SetInsertPoint(UI->getParent());
1543 UI->eraseFromParent();
1544
1545 return Builder.saveIP();
1546}
1547
1548OpenMPIRBuilder::InsertPointTy OpenMPIRBuilder::emitTargetKernel(
1549 const LocationDescription &Loc, InsertPointTy AllocaIP, Value *&Return,
1550 Value *Ident, Value *DeviceID, Value *NumTeams, Value *NumThreads,
1551 Value *HostPtr, ArrayRef<Value *> KernelArgs) {
1552 if (!updateToLocation(Loc))
1553 return Loc.IP;
1554
1555 Builder.restoreIP(IP: AllocaIP);
1556 auto *KernelArgsPtr =
1557 Builder.CreateAlloca(Ty: OpenMPIRBuilder::KernelArgs, ArraySize: nullptr, Name: "kernel_args");
1558 updateToLocation(Loc);
1559
1560 for (unsigned I = 0, Size = KernelArgs.size(); I != Size; ++I) {
1561 llvm::Value *Arg =
1562 Builder.CreateStructGEP(Ty: OpenMPIRBuilder::KernelArgs, Ptr: KernelArgsPtr, Idx: I);
1563 Builder.CreateAlignedStore(
1564 Val: KernelArgs[I], Ptr: Arg,
1565 Align: M.getDataLayout().getPrefTypeAlign(Ty: KernelArgs[I]->getType()));
1566 }
1567
1568 SmallVector<Value *> OffloadingArgs{Ident, DeviceID, NumTeams,
1569 NumThreads, HostPtr, KernelArgsPtr};
1570
1571 Return = createRuntimeFunctionCall(
1572 Callee: getOrCreateRuntimeFunction(M, FnID: OMPRTL___tgt_target_kernel),
1573 Args: OffloadingArgs);
1574
1575 return Builder.saveIP();
1576}
1577
1578OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::emitKernelLaunch(
1579 const LocationDescription &Loc, Value *OutlinedFnID,
1580 EmitFallbackCallbackTy EmitTargetCallFallbackCB, TargetKernelArgs &Args,
1581 Value *DeviceID, Value *RTLoc, InsertPointTy AllocaIP) {
1582
1583 if (!updateToLocation(Loc))
1584 return Loc.IP;
1585
1586 // On top of the arrays that were filled up, the target offloading call
1587 // takes as arguments the device id as well as the host pointer. The host
1588 // pointer is used by the runtime library to identify the current target
1589 // region, so it only has to be unique and not necessarily point to
1590 // anything. It could be the pointer to the outlined function that
1591 // implements the target region, but we aren't using that so that the
1592 // compiler doesn't need to keep that, and could therefore inline the host
1593 // function if proven worthwhile during optimization.
1594
1595 // From this point on, we need to have an ID of the target region defined.
1596 assert(OutlinedFnID && "Invalid outlined function ID!");
1597 (void)OutlinedFnID;
1598
1599 // Return value of the runtime offloading call.
1600 Value *Return = nullptr;
1601
1602 // Arguments for the target kernel.
1603 SmallVector<Value *> ArgsVector;
1604 getKernelArgsVector(KernelArgs&: Args, Builder, ArgsVector);
1605
1606 // The target region is an outlined function launched by the runtime
1607 // via calls to __tgt_target_kernel().
1608 //
1609 // Note that on the host and CPU targets, the runtime implementation of
1610 // these calls simply call the outlined function without forking threads.
1611 // The outlined functions themselves have runtime calls to
1612 // __kmpc_fork_teams() and __kmpc_fork() for this purpose, codegen'd by
1613 // the compiler in emitTeamsCall() and emitParallelCall().
1614 //
1615 // In contrast, on the NVPTX target, the implementation of
1616 // __tgt_target_teams() launches a GPU kernel with the requested number
1617 // of teams and threads so no additional calls to the runtime are required.
1618 // Check the error code and execute the host version if required.
1619 Builder.restoreIP(IP: emitTargetKernel(
1620 Loc: Builder, AllocaIP, Return, Ident: RTLoc, DeviceID, NumTeams: Args.NumTeams.front(),
1621 NumThreads: Args.NumThreads.front(), HostPtr: OutlinedFnID, KernelArgs: ArgsVector));
1622
1623 BasicBlock *OffloadFailedBlock =
1624 BasicBlock::Create(Context&: Builder.getContext(), Name: "omp_offload.failed");
1625 BasicBlock *OffloadContBlock =
1626 BasicBlock::Create(Context&: Builder.getContext(), Name: "omp_offload.cont");
1627 Value *Failed = Builder.CreateIsNotNull(Arg: Return);
1628 Builder.CreateCondBr(Cond: Failed, True: OffloadFailedBlock, False: OffloadContBlock);
1629
1630 auto CurFn = Builder.GetInsertBlock()->getParent();
1631 emitBlock(BB: OffloadFailedBlock, CurFn);
1632 InsertPointOrErrorTy AfterIP = EmitTargetCallFallbackCB(Builder.saveIP());
1633 if (!AfterIP)
1634 return AfterIP.takeError();
1635 Builder.restoreIP(IP: *AfterIP);
1636 emitBranch(Target: OffloadContBlock);
1637 emitBlock(BB: OffloadContBlock, CurFn, /*IsFinished=*/true);
1638 return Builder.saveIP();
1639}
1640
1641Error OpenMPIRBuilder::emitCancelationCheckImpl(
1642 Value *CancelFlag, omp::Directive CanceledDirective) {
1643 assert(isLastFinalizationInfoCancellable(CanceledDirective) &&
1644 "Unexpected cancellation!");
1645
1646 // For a cancel barrier we create two new blocks.
1647 BasicBlock *BB = Builder.GetInsertBlock();
1648 BasicBlock *NonCancellationBlock;
1649 if (Builder.GetInsertPoint() == BB->end()) {
1650 // TODO: This branch will not be needed once we moved to the
1651 // OpenMPIRBuilder codegen completely.
1652 NonCancellationBlock = BasicBlock::Create(
1653 Context&: BB->getContext(), Name: BB->getName() + ".cont", Parent: BB->getParent());
1654 } else {
1655 NonCancellationBlock = SplitBlock(Old: BB, SplitPt: &*Builder.GetInsertPoint());
1656 BB->getTerminator()->eraseFromParent();
1657 Builder.SetInsertPoint(BB);
1658 }
1659 BasicBlock *CancellationBlock = BasicBlock::Create(
1660 Context&: BB->getContext(), Name: BB->getName() + ".cncl", Parent: BB->getParent());
1661
1662 // Jump to them based on the return value.
1663 Value *Cmp = Builder.CreateIsNull(Arg: CancelFlag);
1664 Builder.CreateCondBr(Cond: Cmp, True: NonCancellationBlock, False: CancellationBlock,
1665 /* TODO weight */ BranchWeights: nullptr, Unpredictable: nullptr);
1666
1667 // From the cancellation block we finalize all variables and go to the
1668 // post finalization block that is known to the FiniCB callback.
1669 auto &FI = FinalizationStack.back();
1670 Expected<BasicBlock *> FiniBBOrErr = FI.getFiniBB(Builder);
1671 if (!FiniBBOrErr)
1672 return FiniBBOrErr.takeError();
1673 Builder.SetInsertPoint(CancellationBlock);
1674 Builder.CreateBr(Dest: *FiniBBOrErr);
1675
1676 // The continuation block is where code generation continues.
1677 Builder.SetInsertPoint(NonCancellationBlock->begin());
1678 return Error::success();
1679}
1680
1681/// Create wrapper function used to gather the outlined function's argument
1682/// structure from a shared buffer and to forward them to it when running in
1683/// Generic mode.
1684///
1685/// The outlined function is expected to receive 2 integer arguments followed by
1686/// an optional pointer argument to an argument structure holding the rest.
1687static Function *createTargetParallelWrapper(OpenMPIRBuilder *OMPIRBuilder,
1688 Function &OutlinedFn) {
1689 size_t NumArgs = OutlinedFn.arg_size();
1690 assert((NumArgs == 2 || NumArgs == 3) &&
1691 "expected a 2-3 argument parallel outlined function");
1692 bool UseArgStruct = NumArgs == 3;
1693
1694 IRBuilder<> &Builder = OMPIRBuilder->Builder;
1695 IRBuilder<>::InsertPointGuard IPG(Builder);
1696 auto *FnTy = FunctionType::get(Result: Builder.getVoidTy(),
1697 Params: {Builder.getInt16Ty(), Builder.getInt32Ty()},
1698 /*isVarArg=*/false);
1699 auto *WrapperFn =
1700 Function::Create(Ty: FnTy, Linkage: GlobalValue::InternalLinkage,
1701 N: OutlinedFn.getName() + ".wrapper", M&: OMPIRBuilder->M);
1702
1703 WrapperFn->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
1704 WrapperFn->addParamAttr(ArgNo: 0, Kind: Attribute::ZExt);
1705 WrapperFn->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
1706
1707 BasicBlock *EntryBB =
1708 BasicBlock::Create(Context&: OMPIRBuilder->M.getContext(), Name: "entry", Parent: WrapperFn);
1709 Builder.SetInsertPoint(EntryBB);
1710
1711 // Allocation.
1712 Value *AddrAlloca = Builder.CreateAlloca(Ty: Builder.getInt32Ty(),
1713 /*ArraySize=*/nullptr, Name: "addr");
1714 AddrAlloca = Builder.CreatePointerBitCastOrAddrSpaceCast(
1715 V: AddrAlloca, DestTy: Builder.getPtrTy(/*AddrSpace=*/0),
1716 Name: AddrAlloca->getName() + ".ascast");
1717
1718 Value *ZeroAlloca = Builder.CreateAlloca(Ty: Builder.getInt32Ty(),
1719 /*ArraySize=*/nullptr, Name: "zero");
1720 ZeroAlloca = Builder.CreatePointerBitCastOrAddrSpaceCast(
1721 V: ZeroAlloca, DestTy: Builder.getPtrTy(/*AddrSpace=*/0),
1722 Name: ZeroAlloca->getName() + ".ascast");
1723
1724 Value *ArgsAlloca = nullptr;
1725 if (UseArgStruct) {
1726 ArgsAlloca = Builder.CreateAlloca(Ty: Builder.getPtrTy(),
1727 /*ArraySize=*/nullptr, Name: "global_args");
1728 ArgsAlloca = Builder.CreatePointerBitCastOrAddrSpaceCast(
1729 V: ArgsAlloca, DestTy: Builder.getPtrTy(/*AddrSpace=*/0),
1730 Name: ArgsAlloca->getName() + ".ascast");
1731 }
1732
1733 // Initialization.
1734 Builder.CreateStore(Val: WrapperFn->getArg(i: 1), Ptr: AddrAlloca);
1735 Builder.CreateStore(Val: Builder.getInt32(C: 0), Ptr: ZeroAlloca);
1736 if (UseArgStruct) {
1737 Builder.CreateCall(
1738 Callee: OMPIRBuilder->getOrCreateRuntimeFunctionPtr(
1739 FnID: llvm::omp::RuntimeFunction::OMPRTL___kmpc_get_shared_variables),
1740 Args: {ArgsAlloca});
1741 }
1742
1743 SmallVector<Value *, 3> Args{AddrAlloca, ZeroAlloca};
1744
1745 // Load structArg from global_args.
1746 if (UseArgStruct) {
1747 Value *StructArg = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ArgsAlloca);
1748 StructArg = Builder.CreateInBoundsGEP(Ty: Builder.getPtrTy(), Ptr: StructArg,
1749 IdxList: {Builder.getInt64(C: 0)});
1750 StructArg = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: StructArg, Name: "structArg");
1751 Args.push_back(Elt: StructArg);
1752 }
1753
1754 // Call the outlined function holding the parallel body.
1755 Builder.CreateCall(Callee: &OutlinedFn, Args);
1756 Builder.CreateRetVoid();
1757
1758 return WrapperFn;
1759}
1760
1761// Callback used to create OpenMP runtime calls to support
1762// omp parallel clause for the device.
1763// We need to use this callback to replace call to the OutlinedFn in OuterFn
1764// by the call to the OpenMP DeviceRTL runtime function (kmpc_parallel_60)
1765static void targetParallelCallback(
1766 OpenMPIRBuilder *OMPIRBuilder, Function &OutlinedFn, Function *OuterFn,
1767 BasicBlock *OuterAllocaBB, Value *Ident, Value *IfCondition,
1768 Value *NumThreads, Instruction *PrivTID, AllocaInst *PrivTIDAddr,
1769 Value *ThreadID, const SmallVector<Instruction *, 4> &ToBeDeleted) {
1770 assert(OutlinedFn.arg_size() >= 2 &&
1771 "Expected at least tid and bounded tid as arguments");
1772 unsigned NumCapturedVars = OutlinedFn.arg_size() - /* tid & bounded tid */ 2;
1773
1774 // Add some known attributes.
1775 IRBuilder<> &Builder = OMPIRBuilder->Builder;
1776 OutlinedFn.addParamAttr(ArgNo: 0, Kind: Attribute::NoAlias);
1777 OutlinedFn.addParamAttr(ArgNo: 1, Kind: Attribute::NoAlias);
1778 OutlinedFn.addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
1779 OutlinedFn.addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
1780 OutlinedFn.addFnAttr(Kind: Attribute::NoUnwind);
1781
1782 CallInst *CI = cast<CallInst>(Val: OutlinedFn.user_back());
1783 assert(CI && "Expected call instruction to outlined function");
1784 CI->getParent()->setName("omp_parallel");
1785
1786 Builder.SetInsertPoint(CI);
1787 Type *PtrTy = OMPIRBuilder->VoidPtr;
1788
1789 // Add alloca for kernel args
1790 OpenMPIRBuilder ::InsertPointTy CurrentIP = Builder.saveIP();
1791 Builder.SetInsertPoint(OuterAllocaBB->getFirstInsertionPt());
1792 AllocaInst *ArgsAlloca =
1793 Builder.CreateAlloca(Ty: ArrayType::get(ElementType: PtrTy, NumElements: NumCapturedVars));
1794 Value *Args = ArgsAlloca;
1795 // Add address space cast if array for storing arguments is not allocated
1796 // in address space 0
1797 if (ArgsAlloca->getAddressSpace())
1798 Args = Builder.CreatePointerCast(V: ArgsAlloca, DestTy: PtrTy);
1799 Builder.restoreIP(IP: CurrentIP);
1800
1801 // Store captured vars which are used by kmpc_parallel_60
1802 for (unsigned Idx = 0; Idx < NumCapturedVars; Idx++) {
1803 Value *V = *(CI->arg_begin() + 2 + Idx);
1804 Value *StoreAddress = Builder.CreateConstInBoundsGEP2_64(
1805 Ty: ArrayType::get(ElementType: PtrTy, NumElements: NumCapturedVars), Ptr: Args, Idx0: 0, Idx1: Idx);
1806 Builder.CreateStore(Val: V, Ptr: StoreAddress);
1807 }
1808
1809 Value *Cond =
1810 IfCondition ? Builder.CreateSExtOrTrunc(V: IfCondition, DestTy: OMPIRBuilder->Int32)
1811 : Builder.getInt32(C: 1);
1812 Value *NumThreadsArg =
1813 NumThreads ? Builder.CreateZExtOrTrunc(V: NumThreads, DestTy: OMPIRBuilder->Int32)
1814 : Builder.getInt32(C: -1);
1815
1816 // If this is not a Generic kernel, we can skip generating the wrapper.
1817 Value *WrapperFn;
1818 if (isGenericKernel(Fn&: *OuterFn))
1819 WrapperFn = createTargetParallelWrapper(OMPIRBuilder, OutlinedFn);
1820 else
1821 WrapperFn = Constant::getNullValue(Ty: PtrTy);
1822
1823 // Build kmpc_parallel_60 call
1824 Value *Parallel60CallArgs[] = {
1825 /* identifier*/ Ident,
1826 /* global thread num*/ ThreadID,
1827 /* if expression */ Cond,
1828 /* number of threads */ NumThreadsArg,
1829 /* Proc bind */ Builder.getInt32(C: -1),
1830 /* outlined function */ &OutlinedFn,
1831 /* wrapper function */ WrapperFn,
1832 /* arguments of the outlined funciton*/ Args,
1833 /* number of arguments */ Builder.getInt64(C: NumCapturedVars),
1834 /* strict for number of threads */ Builder.getInt32(C: 0)};
1835
1836 FunctionCallee RTLFn =
1837 OMPIRBuilder->getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_parallel_60);
1838
1839 OMPIRBuilder->createRuntimeFunctionCall(Callee: RTLFn, Args: Parallel60CallArgs);
1840
1841 LLVM_DEBUG(dbgs() << "With kmpc_parallel_60 placed: "
1842 << *Builder.GetInsertBlock()->getParent() << "\n");
1843
1844 // Initialize the local TID stack location with the argument value.
1845 Builder.SetInsertPoint(PrivTID);
1846 Function::arg_iterator OutlinedAI = OutlinedFn.arg_begin();
1847 Builder.CreateStore(Val: Builder.CreateLoad(Ty: OMPIRBuilder->Int32, Ptr: OutlinedAI),
1848 Ptr: PrivTIDAddr);
1849
1850 // Remove redundant call to the outlined function.
1851 CI->eraseFromParent();
1852
1853 for (Instruction *I : ToBeDeleted) {
1854 I->eraseFromParent();
1855 }
1856}
1857
1858// Callback used to create OpenMP runtime calls to support
1859// omp parallel clause for the host.
1860// We need to use this callback to replace call to the OutlinedFn in OuterFn
1861// by the call to the OpenMP host runtime function ( __kmpc_fork_call[_if])
1862static void
1863hostParallelCallback(OpenMPIRBuilder *OMPIRBuilder, Function &OutlinedFn,
1864 Function *OuterFn, Value *Ident, Value *IfCondition,
1865 Instruction *PrivTID, AllocaInst *PrivTIDAddr,
1866 const SmallVector<Instruction *, 4> &ToBeDeleted) {
1867 IRBuilder<> &Builder = OMPIRBuilder->Builder;
1868 FunctionCallee RTLFn;
1869 if (IfCondition) {
1870 RTLFn =
1871 OMPIRBuilder->getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_fork_call_if);
1872 } else {
1873 RTLFn =
1874 OMPIRBuilder->getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_fork_call);
1875 }
1876 if (auto *F = dyn_cast<Function>(Val: RTLFn.getCallee())) {
1877 if (!F->hasMetadata(KindID: LLVMContext::MD_callback)) {
1878 LLVMContext &Ctx = F->getContext();
1879 MDBuilder MDB(Ctx);
1880 // Annotate the callback behavior of the __kmpc_fork_call:
1881 // - The callback callee is argument number 2 (microtask).
1882 // - The first two arguments of the callback callee are unknown (-1).
1883 // - All variadic arguments to the __kmpc_fork_call are passed to the
1884 // callback callee.
1885 F->addMetadata(KindID: LLVMContext::MD_callback,
1886 MD&: *MDNode::get(Context&: Ctx, MDs: {MDB.createCallbackEncoding(
1887 CalleeArgNo: 2, Arguments: {-1, -1},
1888 /* VarArgsArePassed */ true)}));
1889 }
1890 }
1891 // Add some known attributes.
1892 OutlinedFn.addParamAttr(ArgNo: 0, Kind: Attribute::NoAlias);
1893 OutlinedFn.addParamAttr(ArgNo: 1, Kind: Attribute::NoAlias);
1894 OutlinedFn.addFnAttr(Kind: Attribute::NoUnwind);
1895
1896 assert(OutlinedFn.arg_size() >= 2 &&
1897 "Expected at least tid and bounded tid as arguments");
1898 unsigned NumCapturedVars = OutlinedFn.arg_size() - /* tid & bounded tid */ 2;
1899
1900 CallInst *CI = cast<CallInst>(Val: OutlinedFn.user_back());
1901 CI->getParent()->setName("omp_parallel");
1902 Builder.SetInsertPoint(CI);
1903
1904 // Build call __kmpc_fork_call[_if](Ident, n, microtask, var1, .., varn);
1905 Value *ForkCallArgs[] = {Ident, Builder.getInt32(C: NumCapturedVars),
1906 &OutlinedFn};
1907
1908 SmallVector<Value *, 16> RealArgs;
1909 RealArgs.append(in_start: std::begin(arr&: ForkCallArgs), in_end: std::end(arr&: ForkCallArgs));
1910 if (IfCondition) {
1911 Value *Cond = Builder.CreateSExtOrTrunc(V: IfCondition, DestTy: OMPIRBuilder->Int32);
1912 RealArgs.push_back(Elt: Cond);
1913 }
1914 RealArgs.append(in_start: CI->arg_begin() + /* tid & bound tid */ 2, in_end: CI->arg_end());
1915
1916 // __kmpc_fork_call_if always expects a void ptr as the last argument
1917 // If there are no arguments, pass a null pointer.
1918 auto PtrTy = OMPIRBuilder->VoidPtr;
1919 if (IfCondition && NumCapturedVars == 0) {
1920 Value *NullPtrValue = Constant::getNullValue(Ty: PtrTy);
1921 RealArgs.push_back(Elt: NullPtrValue);
1922 }
1923
1924 OMPIRBuilder->createRuntimeFunctionCall(Callee: RTLFn, Args: RealArgs);
1925
1926 LLVM_DEBUG(dbgs() << "With fork_call placed: "
1927 << *Builder.GetInsertBlock()->getParent() << "\n");
1928
1929 // Initialize the local TID stack location with the argument value.
1930 Builder.SetInsertPoint(PrivTID);
1931 Function::arg_iterator OutlinedAI = OutlinedFn.arg_begin();
1932 Builder.CreateStore(Val: Builder.CreateLoad(Ty: OMPIRBuilder->Int32, Ptr: OutlinedAI),
1933 Ptr: PrivTIDAddr);
1934
1935 // Remove redundant call to the outlined function.
1936 CI->eraseFromParent();
1937
1938 for (Instruction *I : ToBeDeleted) {
1939 I->eraseFromParent();
1940 }
1941}
1942
1943OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createParallel(
1944 const LocationDescription &Loc, InsertPointTy OuterAllocIP,
1945 ArrayRef<BasicBlock *> OuterDeallocBlocks, BodyGenCallbackTy BodyGenCB,
1946 PrivatizeCallbackTy PrivCB, FinalizeCallbackTy FiniCB, Value *IfCondition,
1947 Value *NumThreads, omp::ProcBindKind ProcBind, bool IsCancellable) {
1948 assert(!isConflictIP(Loc.IP, OuterAllocIP) && "IPs must not be ambiguous");
1949
1950 if (!updateToLocation(Loc))
1951 return Loc.IP;
1952
1953 uint32_t SrcLocStrSize;
1954 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
1955 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
1956 const bool NeedThreadID = NumThreads || Config.isTargetDevice() ||
1957 (ProcBind != OMP_PROC_BIND_default);
1958 Value *ThreadID = NeedThreadID ? getOrCreateThreadID(Ident) : nullptr;
1959 // If we generate code for the target device, we need to allocate
1960 // struct for aggregate params in the device default alloca address space.
1961 // OpenMP runtime requires that the params of the extracted functions are
1962 // passed as zero address space pointers. This flag ensures that extracted
1963 // function arguments are declared in zero address space
1964 bool ArgsInZeroAddressSpace = Config.isTargetDevice();
1965
1966 // Build call __kmpc_push_num_threads(&Ident, global_tid, num_threads)
1967 // only if we compile for host side.
1968 if (NumThreads && !Config.isTargetDevice()) {
1969 Value *Args[] = {
1970 Ident, ThreadID,
1971 Builder.CreateIntCast(V: NumThreads, DestTy: Int32, /*isSigned*/ false)};
1972 createRuntimeFunctionCall(
1973 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_push_num_threads), Args);
1974 }
1975
1976 if (ProcBind != OMP_PROC_BIND_default) {
1977 // Build call __kmpc_push_proc_bind(&Ident, global_tid, proc_bind)
1978 Value *Args[] = {
1979 Ident, ThreadID,
1980 ConstantInt::get(Ty: Int32, V: unsigned(ProcBind), /*isSigned=*/IsSigned: true)};
1981 createRuntimeFunctionCall(
1982 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_push_proc_bind), Args);
1983 }
1984
1985 BasicBlock *InsertBB = Builder.GetInsertBlock();
1986 Function *OuterFn = InsertBB->getParent();
1987
1988 // Save the outer alloca block because the insertion iterator may get
1989 // invalidated and we still need this later.
1990 BasicBlock *OuterAllocaBlock = OuterAllocIP.getNodeParent();
1991
1992 // Vector to remember instructions we used only during the modeling but which
1993 // we want to delete at the end.
1994 SmallVector<Instruction *, 4> ToBeDeleted;
1995
1996 // Change the location to the outer alloca insertion point to create and
1997 // initialize the allocas we pass into the parallel region.
1998 InsertPointTy NewOuter(OuterAllocaBlock->begin());
1999 Builder.restoreIP(IP: NewOuter);
2000 AllocaInst *TIDAddrAlloca = Builder.CreateAlloca(Ty: Int32, ArraySize: nullptr, Name: "tid.addr");
2001 AllocaInst *ZeroAddrAlloca =
2002 Builder.CreateAlloca(Ty: Int32, ArraySize: nullptr, Name: "zero.addr");
2003 Instruction *TIDAddr = TIDAddrAlloca;
2004 Instruction *ZeroAddr = ZeroAddrAlloca;
2005 if (ArgsInZeroAddressSpace && M.getDataLayout().getAllocaAddrSpace() != 0) {
2006 // Add additional casts to enforce pointers in zero address space
2007 TIDAddr = new AddrSpaceCastInst(
2008 TIDAddrAlloca, PointerType ::get(C&: M.getContext(), AddressSpace: 0), "tid.addr.ascast");
2009 TIDAddr->insertAfter(InsertPos: TIDAddrAlloca->getIterator());
2010 ToBeDeleted.push_back(Elt: TIDAddr);
2011 ZeroAddr = new AddrSpaceCastInst(ZeroAddrAlloca,
2012 PointerType ::get(C&: M.getContext(), AddressSpace: 0),
2013 "zero.addr.ascast");
2014 ZeroAddr->insertAfter(InsertPos: ZeroAddrAlloca->getIterator());
2015 ToBeDeleted.push_back(Elt: ZeroAddr);
2016 }
2017
2018 // We only need TIDAddr and ZeroAddr for modeling purposes to get the
2019 // associated arguments in the outlined function, so we delete them later.
2020 ToBeDeleted.push_back(Elt: TIDAddrAlloca);
2021 ToBeDeleted.push_back(Elt: ZeroAddrAlloca);
2022
2023 // Create an artificial insertion point that will also ensure the blocks we
2024 // are about to split are not degenerated.
2025 auto *UI = new UnreachableInst(Builder.getContext(), InsertBB);
2026
2027 BasicBlock *EntryBB = UI->getParent();
2028 BasicBlock *PRegEntryBB = EntryBB->splitBasicBlock(I: UI, BBName: "omp.par.entry");
2029 BasicBlock *PRegBodyBB = PRegEntryBB->splitBasicBlock(I: UI, BBName: "omp.par.region");
2030 BasicBlock *PRegPreFiniBB =
2031 PRegBodyBB->splitBasicBlock(I: UI, BBName: "omp.par.pre_finalize");
2032 BasicBlock *PRegExitBB = PRegPreFiniBB->splitBasicBlock(I: UI, BBName: "omp.par.exit");
2033
2034 auto FiniCBWrapper = [&](InsertPointTy IP) {
2035 // Hide "open-ended" blocks from the given FiniCB by setting the right jump
2036 // target to the region exit block.
2037 if (IP == IP.getNodeParent()->end()) {
2038 IRBuilder<>::InsertPointGuard IPG(Builder);
2039 Builder.restoreIP(IP);
2040 Instruction *I = Builder.CreateBr(Dest: PRegExitBB);
2041 IP = I->getIterator();
2042 }
2043 assert(IP.getNodeParent()->getTerminator()->getNumSuccessors() == 1 &&
2044 IP.getNodeParent()->getTerminator()->getSuccessor(0) == PRegExitBB &&
2045 "Unexpected insertion point for finalization call!");
2046 return FiniCB(IP);
2047 };
2048
2049 FinalizationStack.push_back(Elt: {FiniCBWrapper, OMPD_parallel, IsCancellable});
2050
2051 // Generate the privatization allocas in the block that will become the entry
2052 // of the outlined function.
2053 Builder.SetInsertPoint(PRegEntryBB->getTerminator());
2054 InsertPointTy InnerAllocaIP = Builder.saveIP();
2055
2056 AllocaInst *PrivTIDAddr =
2057 Builder.CreateAlloca(Ty: Int32, ArraySize: nullptr, Name: "tid.addr.local");
2058 Instruction *PrivTID = Builder.CreateLoad(Ty: Int32, Ptr: PrivTIDAddr, Name: "tid");
2059
2060 // Add some fake uses for OpenMP provided arguments.
2061 ToBeDeleted.push_back(Elt: Builder.CreateLoad(Ty: Int32, Ptr: TIDAddr, Name: "tid.addr.use"));
2062 Instruction *ZeroAddrUse =
2063 Builder.CreateLoad(Ty: Int32, Ptr: ZeroAddr, Name: "zero.addr.use");
2064 ToBeDeleted.push_back(Elt: ZeroAddrUse);
2065
2066 // EntryBB
2067 // |
2068 // V
2069 // PRegionEntryBB <- Privatization allocas are placed here.
2070 // |
2071 // V
2072 // PRegionBodyBB <- BodeGen is invoked here.
2073 // |
2074 // V
2075 // PRegPreFiniBB <- The block we will start finalization from.
2076 // |
2077 // V
2078 // PRegionExitBB <- A common exit to simplify block collection.
2079 //
2080
2081 LLVM_DEBUG(dbgs() << "Before body codegen: " << *OuterFn << "\n");
2082
2083 // Let the caller create the body.
2084 assert(BodyGenCB && "Expected body generation callback!");
2085 InsertPointTy CodeGenIP(PRegBodyBB->begin());
2086 if (Error Err = BodyGenCB(InnerAllocaIP, CodeGenIP, PRegExitBB))
2087 return Err;
2088
2089 LLVM_DEBUG(dbgs() << "After body codegen: " << *OuterFn << "\n");
2090
2091 // If OuterFn is a Generic kernel, we need to use device shared memory to
2092 // allocate argument structures. Otherwise, we use stack allocations as usual.
2093 bool UsesDeviceSharedMemory =
2094 Config.isTargetDevice() && isGenericKernel(Fn&: *OuterFn);
2095 std::unique_ptr<OutlineInfo> OI =
2096 UsesDeviceSharedMemory
2097 ? std::make_unique<DeviceSharedMemOutlineInfo>(args&: *this)
2098 : std::make_unique<OutlineInfo>();
2099
2100 if (Config.isTargetDevice()) {
2101 // Generate OpenMP target specific runtime call
2102 OI->PostOutlineCB = [=, ToBeDeletedVec =
2103 std::move(ToBeDeleted)](Function &OutlinedFn) {
2104 targetParallelCallback(OMPIRBuilder: this, OutlinedFn, OuterFn, OuterAllocaBB: OuterAllocaBlock, Ident,
2105 IfCondition, NumThreads, PrivTID, PrivTIDAddr,
2106 ThreadID, ToBeDeleted: ToBeDeletedVec);
2107 };
2108 } else {
2109 // Generate OpenMP host runtime call
2110 OI->PostOutlineCB = [=, ToBeDeletedVec =
2111 std::move(ToBeDeleted)](Function &OutlinedFn) {
2112 hostParallelCallback(OMPIRBuilder: this, OutlinedFn, OuterFn, Ident, IfCondition,
2113 PrivTID, PrivTIDAddr, ToBeDeleted: ToBeDeletedVec);
2114 };
2115 }
2116
2117 OI->FixUpNonEntryAllocas = true;
2118 OI->OuterAllocBB = OuterAllocaBlock;
2119 OI->EntryBB = PRegEntryBB;
2120 OI->ExitBB = PRegExitBB;
2121 OI->OuterDeallocBBs.reserve(N: OuterDeallocBlocks.size());
2122 copy(Range&: OuterDeallocBlocks, Out: OI->OuterDeallocBBs.end());
2123
2124 SmallPtrSet<BasicBlock *, 32> ParallelRegionBlockSet;
2125 SmallVector<BasicBlock *, 32> Blocks;
2126 OI->collectBlocks(BlockSet&: ParallelRegionBlockSet, BlockVector&: Blocks);
2127
2128 CodeExtractorAnalysisCache CEAC(*OuterFn);
2129 CodeExtractor Extractor(Blocks, /* DominatorTree */ nullptr,
2130 /* AggregateArgs */ false,
2131 /* BlockFrequencyInfo */ nullptr,
2132 /* BranchProbabilityInfo */ nullptr,
2133 /* AssumptionCache */ nullptr,
2134 /* AllowVarArgs */ true,
2135 /* AllowAlloca */ true,
2136 /* AllocationBlock */ OuterAllocaBlock,
2137 /* DeallocationBlocks */ {},
2138 /* Suffix */ ".omp_par", ArgsInZeroAddressSpace);
2139
2140 // Find inputs to, outputs from the code region.
2141 BasicBlock *CommonExit = nullptr;
2142 SetVector<Value *> Inputs, Outputs, SinkingCands, HoistingCands;
2143 Extractor.findAllocas(CEAC, SinkCands&: SinkingCands, HoistCands&: HoistingCands, ExitBlock&: CommonExit);
2144
2145 Extractor.findInputsOutputs(Inputs, Outputs, Allocas: SinkingCands,
2146 /*CollectGlobalInputs=*/true);
2147
2148 Inputs.remove_if(P: [&](Value *I) {
2149 if (auto *GV = dyn_cast_if_present<GlobalVariable>(Val: I))
2150 return GV->getValueType() == OpenMPIRBuilder::Ident;
2151
2152 return false;
2153 });
2154
2155 LLVM_DEBUG(dbgs() << "Before privatization: " << *OuterFn << "\n");
2156
2157 FunctionCallee TIDRTLFn =
2158 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_global_thread_num);
2159
2160 auto PrivHelper = [&](Value &V) -> Error {
2161 if (&V == TIDAddr || &V == ZeroAddr) {
2162 OI->ExcludeArgsFromAggregate.push_back(Elt: &V);
2163 return Error::success();
2164 }
2165
2166 SetVector<Use *> Uses;
2167 for (Use &U : V.uses())
2168 if (auto *UserI = dyn_cast<Instruction>(Val: U.getUser()))
2169 if (ParallelRegionBlockSet.count(Ptr: UserI->getParent()))
2170 Uses.insert(X: &U);
2171
2172 // __kmpc_fork_call expects extra arguments as pointers. If the input
2173 // already has a pointer type, everything is fine. Otherwise, store the
2174 // value onto stack and load it back inside the to-be-outlined region. This
2175 // will ensure only the pointer will be passed to the function.
2176 // FIXME: if there are more than 15 trailing arguments, they must be
2177 // additionally packed in a struct.
2178 Value *Inner = &V;
2179 if (!V.getType()->isPointerTy()) {
2180 IRBuilder<>::InsertPointGuard Guard(Builder);
2181 LLVM_DEBUG(llvm::dbgs() << "Forwarding input as pointer: " << V << "\n");
2182
2183 Builder.restoreIP(IP: OuterAllocIP);
2184 Value *Ptr;
2185 if (UsesDeviceSharedMemory) {
2186 // Use device shared memory instead, if needed.
2187 Ptr = createOMPAllocShared(Loc: Builder, VarType: V.getType(),
2188 Name: V.getName() + ".reloaded");
2189 for (BasicBlock *DeallocBlock : OuterDeallocBlocks) {
2190 assert(DeallocBlock->getParent() ==
2191 OuterAllocIP.getNodeParent()->getParent() &&
2192 "Dealloc block must be in the allocation's function to reuse "
2193 "its debug location");
2194 createOMPFreeShared(Loc: {DeallocBlock->getFirstInsertionPt(),
2195 Builder.getCurrentDebugLocation()},
2196 Addr: Ptr, VarType: V.getType());
2197 }
2198 } else {
2199 Ptr = Builder.CreateAlloca(Ty: V.getType(), ArraySize: nullptr,
2200 Name: V.getName() + ".reloaded");
2201 }
2202
2203 // Store to stack at end of the block that currently branches to the entry
2204 // block of the to-be-outlined region.
2205 Builder.SetInsertPoint(InsertBB->getTerminator()->getIterator());
2206 Builder.CreateStore(Val: &V, Ptr);
2207
2208 // Load back next to allocations in the to-be-outlined region.
2209 Builder.restoreIP(IP: InnerAllocaIP);
2210 Inner = Builder.CreateLoad(Ty: V.getType(), Ptr);
2211 }
2212
2213 Value *ReplacementValue = nullptr;
2214 CallInst *CI = dyn_cast<CallInst>(Val: &V);
2215 if (CI && CI->getCalledFunction() == TIDRTLFn.getCallee()) {
2216 ReplacementValue = PrivTID;
2217 } else {
2218 InsertPointOrErrorTy AfterIP =
2219 PrivCB(InnerAllocaIP, Builder.saveIP(), V, *Inner, ReplacementValue);
2220 if (!AfterIP)
2221 return AfterIP.takeError();
2222 Builder.restoreIP(IP: *AfterIP);
2223 InnerAllocaIP =
2224 InnerAllocaIP.getNodeParent()->getTerminator()->getIterator();
2225
2226 assert(ReplacementValue &&
2227 "Expected copy/create callback to set replacement value!");
2228 if (ReplacementValue == &V)
2229 return Error::success();
2230 }
2231
2232 for (Use *UPtr : Uses)
2233 UPtr->set(ReplacementValue);
2234
2235 return Error::success();
2236 };
2237
2238 // Reset the inner alloca insertion as it will be used for loading the values
2239 // wrapped into pointers before passing them into the to-be-outlined region.
2240 // Configure it to insert immediately after the fake use of zero address so
2241 // that they are available in the generated body and so that the
2242 // OpenMP-related values (thread ID and zero address pointers) remain leading
2243 // in the argument list.
2244 InnerAllocaIP = ZeroAddrUse->getNextNode()->getIterator();
2245
2246 // Reset the outer alloca insertion point to the entry of the relevant block
2247 // in case it was invalidated.
2248 OuterAllocIP = OuterAllocaBlock->getFirstInsertionPt();
2249
2250 for (Value *Input : Inputs) {
2251 LLVM_DEBUG(dbgs() << "Captured input: " << *Input << "\n");
2252 if (Error Err = PrivHelper(*Input))
2253 return Err;
2254 }
2255 LLVM_DEBUG({
2256 for (Value *Output : Outputs)
2257 LLVM_DEBUG(dbgs() << "Captured output: " << *Output << "\n");
2258 });
2259 assert(Outputs.empty() &&
2260 "OpenMP outlining should not produce live-out values!");
2261
2262 LLVM_DEBUG(dbgs() << "After privatization: " << *OuterFn << "\n");
2263 LLVM_DEBUG({
2264 for (auto *BB : Blocks)
2265 dbgs() << " PBR: " << BB->getName() << "\n";
2266 });
2267
2268 // Adjust the finalization stack, verify the adjustment, and call the
2269 // finalize function a last time to finalize values between the pre-fini
2270 // block and the exit block if we left the parallel "the normal way".
2271 auto FiniInfo = FinalizationStack.pop_back_val();
2272 (void)FiniInfo;
2273 assert(FiniInfo.DK == OMPD_parallel &&
2274 "Unexpected finalization stack state!");
2275
2276 Instruction *PRegPreFiniTI = PRegPreFiniBB->getTerminator();
2277
2278 InsertPointTy PreFiniIP(PRegPreFiniTI->getIterator());
2279 Expected<BasicBlock *> FiniBBOrErr = FiniInfo.getFiniBB(Builder);
2280 if (!FiniBBOrErr)
2281 return FiniBBOrErr.takeError();
2282 {
2283 IRBuilderBase::InsertPointGuard Guard(Builder);
2284 Builder.restoreIP(IP: PreFiniIP);
2285 Builder.CreateBr(Dest: *FiniBBOrErr);
2286 // There's currently a branch to omp.par.exit. Delete it. We will get there
2287 // via the fini block
2288 if (Instruction *Term = Builder.GetInsertBlock()->getTerminator())
2289 Term->eraseFromParent();
2290 }
2291
2292 // Register the outlined info.
2293 addOutlineInfo(OI: std::move(OI));
2294
2295 InsertPointTy AfterIP(UI->getParent()->end());
2296 UI->eraseFromParent();
2297
2298 return AfterIP;
2299}
2300
2301void OpenMPIRBuilder::emitFlush(const LocationDescription &Loc) {
2302 // Build call void __kmpc_flush(ident_t *loc)
2303 uint32_t SrcLocStrSize;
2304 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
2305 Value *Args[] = {getOrCreateIdent(SrcLocStr, SrcLocStrSize)};
2306
2307 createRuntimeFunctionCall(Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_flush),
2308 Args);
2309}
2310
2311void OpenMPIRBuilder::createFlush(const LocationDescription &Loc) {
2312 if (!updateToLocation(Loc))
2313 return;
2314 emitFlush(Loc);
2315}
2316
2317void OpenMPIRBuilder::createError(const LocationDescription &Loc, bool IsFatal,
2318 Value *Message) {
2319 if (!updateToLocation(Loc))
2320 return;
2321
2322 // Build call void __kmpc_error(ident_t *loc, int severity,
2323 // const char *message)
2324 uint32_t SrcLocStrSize;
2325 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
2326 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
2327 // Severity: 1 = warning, 2 = fatal.
2328 Value *Severity = ConstantInt::get(Ty: Int32, V: IsFatal ? 2 : 1);
2329 Value *MessageArg = Message ? Message : ConstantPointerNull::get(T: Int8Ptr);
2330 Value *Args[] = {Ident, Severity, MessageArg};
2331
2332 createRuntimeFunctionCall(Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_error),
2333 Args);
2334}
2335
2336void OpenMPIRBuilder::emitTaskyieldImpl(const LocationDescription &Loc) {
2337 // Build call __kmpc_omp_taskyield(loc, thread_id, 0);
2338 uint32_t SrcLocStrSize;
2339 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
2340 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
2341 Constant *I32Null = ConstantInt::getNullValue(Ty: Int32);
2342 Value *Args[] = {Ident, getOrCreateThreadID(Ident), I32Null};
2343
2344 createRuntimeFunctionCall(
2345 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_taskyield), Args);
2346}
2347
2348void OpenMPIRBuilder::createTaskyield(const LocationDescription &Loc) {
2349 if (!updateToLocation(Loc))
2350 return;
2351 emitTaskyieldImpl(Loc);
2352}
2353
2354void OpenMPIRBuilder::emitTaskDependency(IRBuilderBase &Builder, Value *Entry,
2355 const DependData &Dep) {
2356 // Store the pointer to the variable
2357 Value *Addr = Builder.CreateStructGEP(
2358 Ty: DependInfo, Ptr: Entry,
2359 Idx: static_cast<unsigned int>(RTLDependInfoFields::BaseAddr));
2360 Value *DepValPtr = Builder.CreatePtrToInt(V: Dep.DepVal, DestTy: SizeTy);
2361 Builder.CreateStore(Val: DepValPtr, Ptr: Addr);
2362 // Store the size of the variable
2363 Value *Size = Builder.CreateStructGEP(
2364 Ty: DependInfo, Ptr: Entry, Idx: static_cast<unsigned int>(RTLDependInfoFields::Len));
2365 Builder.CreateStore(
2366 Val: ConstantInt::get(Ty: SizeTy,
2367 V: M.getDataLayout().getTypeStoreSize(Ty: Dep.DepValueType)),
2368 Ptr: Size);
2369 // Store the dependency kind
2370 Value *Flags = Builder.CreateStructGEP(
2371 Ty: DependInfo, Ptr: Entry, Idx: static_cast<unsigned int>(RTLDependInfoFields::Flags));
2372 Builder.CreateStore(Val: ConstantInt::get(Ty: Builder.getInt8Ty(),
2373 V: static_cast<unsigned int>(Dep.DepKind)),
2374 Ptr: Flags);
2375}
2376
2377// Processes the dependencies in Dependencies and does the following
2378// - Allocates space on the stack of an array of DependInfo objects
2379// - Populates each DependInfo object with relevant information of
2380// the corresponding dependence.
2381// - All code is inserted in the entry block of the current function.
2382static Value *emitTaskDependencies(
2383 OpenMPIRBuilder &OMPBuilder,
2384 const SmallVectorImpl<OpenMPIRBuilder::DependData> &Dependencies) {
2385 // Early return if we have no dependencies to process
2386 if (Dependencies.empty())
2387 return nullptr;
2388
2389 // Given a vector of DependData objects, in this function we create an
2390 // array on the stack that holds kmp_depend_info objects corresponding
2391 // to each dependency. This is then passed to the OpenMP runtime.
2392 // For example, if there are 'n' dependencies then the following psedo
2393 // code is generated. Assume the first dependence is on a variable 'a'
2394 //
2395 // \code{c}
2396 // DepArray = alloc(n x sizeof(kmp_depend_info);
2397 // idx = 0;
2398 // DepArray[idx].base_addr = ptrtoint(&a);
2399 // DepArray[idx].len = 8;
2400 // DepArray[idx].flags = Dep.DepKind; /*(See OMPContants.h for DepKind)*/
2401 // ++idx;
2402 // DepArray[idx].base_addr = ...;
2403 // \endcode
2404
2405 IRBuilderBase &Builder = OMPBuilder.Builder;
2406 Type *DependInfo = OMPBuilder.DependInfo;
2407
2408 Value *DepArray = nullptr;
2409 Type *DepArrayTy = ArrayType::get(ElementType: DependInfo, NumElements: Dependencies.size());
2410 {
2411 // Use a InsertPointGuard to restore the location back along with the
2412 // insertion point.
2413 IRBuilderBase::InsertPointGuard IPGuard(Builder);
2414 Builder.SetInsertPoint(
2415 Builder.GetInsertBlock()->getParent()->getEntryBlock().getTerminator());
2416 DepArray = Builder.CreateAlloca(Ty: DepArrayTy, ArraySize: nullptr, Name: ".dep.arr.addr");
2417 }
2418
2419 for (const auto &[DepIdx, Dep] : enumerate(First: Dependencies)) {
2420 Value *Base =
2421 Builder.CreateConstInBoundsGEP2_64(Ty: DepArrayTy, Ptr: DepArray, Idx0: 0, Idx1: DepIdx);
2422 OMPBuilder.emitTaskDependency(Builder, Entry: Base, Dep);
2423 }
2424 return DepArray;
2425}
2426
2427void OpenMPIRBuilder::emitTaskwaitImpl(const LocationDescription &Loc) {
2428 // Build call kmp_int32 __kmpc_omp_taskwait(ident_t *loc, kmp_int32
2429 // global_tid);
2430 uint32_t SrcLocStrSize;
2431 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
2432 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
2433 Value *Args[] = {Ident, getOrCreateThreadID(Ident)};
2434
2435 // Ignore return result until untied tasks are supported.
2436 createRuntimeFunctionCall(
2437 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_taskwait), Args);
2438}
2439
2440void OpenMPIRBuilder::createTaskwait(const LocationDescription &Loc,
2441 DependenciesInfo Dependencies,
2442 bool IsNowait) {
2443 if (!updateToLocation(Loc))
2444 return;
2445
2446 Value *DepArray = nullptr;
2447 Type *DepArrayTy = nullptr;
2448 Value *NumDeps = nullptr;
2449 if (Dependencies.DepArray) {
2450 DepArray = Dependencies.DepArray;
2451 NumDeps = Dependencies.NumDeps;
2452 } else if (!Dependencies.Deps.empty()) {
2453 DepArrayTy = ArrayType::get(ElementType: DependInfo, NumElements: Dependencies.Deps.size());
2454 NumDeps = Builder.getInt32(C: Dependencies.Deps.size());
2455 {
2456 IRBuilderBase::InsertPointGuard IPGuard(Builder);
2457 BasicBlock &entryBB =
2458 Builder.GetInsertBlock()->getParent()->getEntryBlock();
2459 Builder.SetInsertPoint(entryBB.getFirstInsertionPt());
2460 DepArray = Builder.CreateAlloca(Ty: DepArrayTy, ArraySize: nullptr, Name: ".dep.arr.addr");
2461 }
2462
2463 for (const auto &[DepIdx, Dep] : enumerate(First&: Dependencies.Deps)) {
2464 Value *Base =
2465 Builder.CreateConstInBoundsGEP2_64(Ty: DepArrayTy, Ptr: DepArray, Idx0: 0, Idx1: DepIdx);
2466 this->emitTaskDependency(Builder, Entry: Base, Dep);
2467 }
2468 }
2469
2470 if (DepArray) {
2471 uint32_t SrcLocStrSize;
2472 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
2473 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
2474 Value *Args[] = {
2475 Ident,
2476 getOrCreateThreadID(Ident),
2477 NumDeps,
2478 DepArray,
2479 ConstantInt::get(Ty: Builder.getInt32Ty(), V: 0),
2480 ConstantPointerNull::get(T: PointerType::getUnqual(C&: M.getContext())),
2481 ConstantInt::get(Ty: Builder.getInt32Ty(), V: IsNowait)};
2482 createRuntimeFunctionCall(
2483 Callee: getOrCreateRuntimeFunctionPtr(
2484 FnID: omp::RuntimeFunction::OMPRTL___kmpc_omp_taskwait_deps_51),
2485 Args);
2486 } else {
2487 emitTaskwaitImpl(Loc);
2488 }
2489}
2490
2491/// Create the task duplication function passed to kmpc_taskloop.
2492Expected<Value *> OpenMPIRBuilder::createTaskDuplicationFunction(
2493 Type *PrivatesTy, int32_t PrivatesIndex, TaskDupCallbackTy DupCB) {
2494 unsigned ProgramAddressSpace = M.getDataLayout().getProgramAddressSpace();
2495 if (!DupCB)
2496 return Constant::getNullValue(
2497 Ty: PointerType::get(C&: Builder.getContext(), AddressSpace: ProgramAddressSpace));
2498
2499 // From OpenMP Runtime p_task_dup_t:
2500 // Routine optionally generated by the compiler for setting the lastprivate
2501 // flag and calling needed constructors for private/firstprivate objects (used
2502 // to form taskloop tasks from pattern task) Parameters: dest task, src task,
2503 // lastprivate flag.
2504 // typedef void (*p_task_dup_t)(kmp_task_t *, kmp_task_t *, kmp_int32);
2505
2506 auto *VoidPtrTy = PointerType::get(C&: Builder.getContext(), AddressSpace: ProgramAddressSpace);
2507
2508 FunctionType *DupFuncTy = FunctionType::get(
2509 Result: Builder.getVoidTy(), Params: {VoidPtrTy, VoidPtrTy, Builder.getInt32Ty()},
2510 /*isVarArg=*/false);
2511
2512 Function *DupFunction = Function::Create(Ty: DupFuncTy, Linkage: Function::InternalLinkage,
2513 N: "omp_taskloop_dup", M);
2514 Value *DestTaskArg = DupFunction->getArg(i: 0);
2515 Value *SrcTaskArg = DupFunction->getArg(i: 1);
2516 Value *LastprivateFlagArg = DupFunction->getArg(i: 2);
2517 DestTaskArg->setName("dest_task");
2518 SrcTaskArg->setName("src_task");
2519 LastprivateFlagArg->setName("lastprivate_flag");
2520
2521 IRBuilderBase::InsertPointGuard Guard(Builder);
2522 Builder.SetInsertPoint(
2523 BasicBlock::Create(Context&: Builder.getContext(), Name: "entry", Parent: DupFunction));
2524
2525 auto GetTaskContextPtrFromArg = [&](Value *Arg) -> Value * {
2526 Type *TaskWithPrivatesTy =
2527 StructType::get(Context&: Builder.getContext(), Elements: {Task, PrivatesTy});
2528 Value *TaskPrivates = Builder.CreateGEP(
2529 Ty: TaskWithPrivatesTy, Ptr: Arg, IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: 1)});
2530 Value *ContextPtr = Builder.CreateGEP(
2531 Ty: PrivatesTy, Ptr: TaskPrivates,
2532 IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: PrivatesIndex)});
2533 return ContextPtr;
2534 };
2535
2536 Value *DestTaskContextPtr = GetTaskContextPtrFromArg(DestTaskArg);
2537 Value *SrcTaskContextPtr = GetTaskContextPtrFromArg(SrcTaskArg);
2538
2539 DestTaskContextPtr->setName("destPtr");
2540 SrcTaskContextPtr->setName("srcPtr");
2541
2542 InsertPointTy AllocaIP(DupFunction->getEntryBlock().begin());
2543 InsertPointTy CodeGenIP = Builder.saveIP();
2544 Expected<IRBuilderBase::InsertPoint> AfterIPOrError =
2545 DupCB(AllocaIP, CodeGenIP, DestTaskContextPtr, SrcTaskContextPtr);
2546 if (!AfterIPOrError)
2547 return AfterIPOrError.takeError();
2548 Builder.restoreIP(IP: *AfterIPOrError);
2549
2550 Builder.CreateRetVoid();
2551
2552 return DupFunction;
2553}
2554
2555OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createTaskloop(
2556 const LocationDescription &Loc, InsertPointTy AllocaIP,
2557 ArrayRef<BasicBlock *> DeallocBlocks, BodyGenCallbackTy BodyGenCB,
2558 llvm::function_ref<llvm::Expected<llvm::CanonicalLoopInfo *>()> LoopInfo,
2559 Value *LBVal, Value *UBVal, Value *StepVal, bool Untied, Value *IfCond,
2560 Value *GrainSize, bool NoGroup, int Sched, Value *Final, bool Mergeable,
2561 Value *Priority, uint64_t NumOfCollapseLoops, TaskDupCallbackTy DupCB,
2562 Value *TaskContextStructPtrVal, bool FreeAgent) {
2563
2564 if (!updateToLocation(Loc))
2565 return InsertPointTy();
2566
2567 uint32_t SrcLocStrSize;
2568 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
2569 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
2570
2571 BasicBlock *TaskloopExitBB =
2572 splitBB(Builder, /*CreateBranch=*/true, Name: "taskloop.exit");
2573 BasicBlock *TaskloopBodyBB =
2574 splitBB(Builder, /*CreateBranch=*/true, Name: "taskloop.body");
2575 BasicBlock *TaskloopAllocaBB =
2576 splitBB(Builder, /*CreateBranch=*/true, Name: "taskloop.alloca");
2577
2578 InsertPointTy TaskloopAllocaIP = TaskloopAllocaBB->begin();
2579 InsertPointTy TaskloopBodyIP = TaskloopBodyBB->begin();
2580
2581 if (Error Err = BodyGenCB(TaskloopAllocaIP, TaskloopBodyIP, TaskloopExitBB))
2582 return Err;
2583
2584 llvm::Expected<llvm::CanonicalLoopInfo *> result = LoopInfo();
2585 if (!result) {
2586 return result.takeError();
2587 }
2588
2589 llvm::CanonicalLoopInfo *CLI = result.get();
2590 auto OI = std::make_unique<OutlineInfo>();
2591 OI->EntryBB = TaskloopAllocaBB;
2592 OI->OuterAllocBB = AllocaIP.getNodeParent();
2593 OI->ExitBB = TaskloopExitBB;
2594 OI->OuterDeallocBBs.reserve(N: DeallocBlocks.size());
2595 copy(Range&: DeallocBlocks, Out: OI->OuterDeallocBBs.end());
2596
2597 // Add the thread ID argument.
2598 SmallVector<Instruction *> ToBeDeleted;
2599 // dummy instruction to be used as a fake argument
2600 OI->ExcludeArgsFromAggregate.push_back(Elt: createFakeIntVal(
2601 Builder, OuterAllocaIP: AllocaIP, ToBeDeleted, InnerAllocaIP: TaskloopAllocaIP, Name: "global.tid", AsPtr: false));
2602 Value *FakeLB = createFakeIntVal(Builder, OuterAllocaIP: AllocaIP, ToBeDeleted,
2603 InnerAllocaIP: TaskloopAllocaIP, Name: "lb", AsPtr: false, Is64Bit: true);
2604 Value *FakeUB = createFakeIntVal(Builder, OuterAllocaIP: AllocaIP, ToBeDeleted,
2605 InnerAllocaIP: TaskloopAllocaIP, Name: "ub", AsPtr: false, Is64Bit: true);
2606 Value *FakeStep = createFakeIntVal(Builder, OuterAllocaIP: AllocaIP, ToBeDeleted,
2607 InnerAllocaIP: TaskloopAllocaIP, Name: "step", AsPtr: false, Is64Bit: true);
2608 // For Taskloop, we want to force the bounds being the first 3 inputs in the
2609 // aggregate struct
2610 OI->Inputs.insert(X: FakeLB);
2611 OI->Inputs.insert(X: FakeUB);
2612 OI->Inputs.insert(X: FakeStep);
2613 if (TaskContextStructPtrVal)
2614 OI->Inputs.insert(X: TaskContextStructPtrVal);
2615 assert(((TaskContextStructPtrVal && DupCB) ||
2616 (!TaskContextStructPtrVal && !DupCB)) &&
2617 "Task context struct ptr and duplication callback must be both set "
2618 "or both null");
2619
2620 // It isn't safe to run the duplication bodygen callback inside the post
2621 // outlining callback so this has to be run now before we know the real task
2622 // shareds structure type.
2623 unsigned ProgramAddressSpace = M.getDataLayout().getProgramAddressSpace();
2624 Type *PointerTy = PointerType::get(C&: Builder.getContext(), AddressSpace: ProgramAddressSpace);
2625 Type *FakeSharedsTy = StructType::get(
2626 Context&: Builder.getContext(),
2627 Elements: {FakeLB->getType(), FakeUB->getType(), FakeStep->getType(), PointerTy});
2628 Expected<Value *> TaskDupFnOrErr = createTaskDuplicationFunction(
2629 PrivatesTy: FakeSharedsTy,
2630 /*PrivatesIndex: the pointer after the three indices above*/ PrivatesIndex: 3, DupCB);
2631 if (!TaskDupFnOrErr) {
2632 return TaskDupFnOrErr.takeError();
2633 }
2634 Value *TaskDupFn = *TaskDupFnOrErr;
2635
2636 OI->PostOutlineCB = [this, Ident, LBVal, UBVal, StepVal, Untied,
2637 TaskloopAllocaBB, CLI, TaskDupFn, ToBeDeleted, IfCond,
2638 GrainSize, NoGroup, Sched, FakeLB, FakeUB, FakeStep,
2639 FakeSharedsTy, Final, Mergeable, Priority,
2640 NumOfCollapseLoops,
2641 FreeAgent](Function &OutlinedFn) mutable {
2642 // Replace the Stale CI by appropriate RTL function call.
2643 assert(OutlinedFn.hasOneUse() &&
2644 "there must be a single user for the outlined function");
2645 CallInst *StaleCI = cast<CallInst>(Val: OutlinedFn.user_back());
2646
2647 /* Create the casting for the Bounds Values that can be used when outlining
2648 * to replace the uses of the fakes with real values */
2649 BasicBlock *CodeReplBB = StaleCI->getParent();
2650 Builder.SetInsertPoint(CodeReplBB->getFirstInsertionPt());
2651 Value *CastedLBVal =
2652 Builder.CreateIntCast(V: LBVal, DestTy: Builder.getInt64Ty(), isSigned: true, Name: "lb64");
2653 Value *CastedUBVal =
2654 Builder.CreateIntCast(V: UBVal, DestTy: Builder.getInt64Ty(), isSigned: true, Name: "ub64");
2655 Value *CastedStepVal =
2656 Builder.CreateIntCast(V: StepVal, DestTy: Builder.getInt64Ty(), isSigned: true, Name: "step64");
2657
2658 Builder.SetInsertPoint(StaleCI);
2659
2660 // Gather the arguments for emitting the runtime call for
2661 // @__kmpc_omp_task_alloc
2662 Function *TaskAllocFn =
2663 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_alloc);
2664
2665 Value *ThreadID = getOrCreateThreadID(Ident);
2666
2667 if (!NoGroup) {
2668 // Emit runtime call for @__kmpc_taskgroup
2669 Function *TaskgroupFn =
2670 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_taskgroup);
2671 Builder.CreateCall(Callee: TaskgroupFn, Args: {Ident, ThreadID});
2672 }
2673
2674 // `flags` Argument Configuration
2675 // Task is tied if (Flags & 1) == 1.
2676 // Task is untied if (Flags & 1) == 0.
2677 // Task is final if (Flags & 2) == 2.
2678 // Task is not final if (Flags & 2) == 0.
2679 // Task is mergeable if (Flags & 4) == 4.
2680 // Task is not mergeable if (Flags & 4) == 0.
2681 // Task is priority if (Flags & 32) == 32.
2682 // Task is not priority if (Flags & 32) == 0.
2683 // Task is free-agent eligible if (Flags & 128) == 128.
2684 // Task is not free-agent eligible if (Flags & 128) == 0.
2685 Value *Flags = Builder.getInt32(C: Untied ? 0 : 1);
2686 if (Final)
2687 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 2), RHS: Flags);
2688 if (Mergeable)
2689 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 4), RHS: Flags);
2690 if (Priority)
2691 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 32), RHS: Flags);
2692 if (FreeAgent)
2693 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 128), RHS: Flags);
2694
2695 Value *TaskSize = Builder.getInt64(
2696 C: divideCeil(Numerator: M.getDataLayout().getTypeSizeInBits(Ty: Task), Denominator: 8));
2697
2698 AllocaInst *ArgStructAlloca =
2699 dyn_cast<AllocaInst>(Val: StaleCI->getArgOperand(i: 1));
2700 assert(ArgStructAlloca &&
2701 "Unable to find the alloca instruction corresponding to arguments "
2702 "for extracted function");
2703 std::optional<TypeSize> ArgAllocSize =
2704 ArgStructAlloca->getAllocationSize(DL: M.getDataLayout());
2705 assert(ArgAllocSize &&
2706 "Unable to determine size of arguments for extracted function");
2707 Value *SharedsSize = Builder.getInt64(C: ArgAllocSize->getFixedValue());
2708
2709 // Emit the @__kmpc_omp_task_alloc runtime call
2710 // The runtime call returns a pointer to an area where the task captured
2711 // variables must be copied before the task is run (TaskData)
2712 CallInst *TaskData = Builder.CreateCall(
2713 Callee: TaskAllocFn, Args: {/*loc_ref=*/Ident, /*gtid=*/ThreadID, /*flags=*/Flags,
2714 /*sizeof_task=*/TaskSize, /*sizeof_shared=*/SharedsSize,
2715 /*task_func=*/&OutlinedFn});
2716
2717 Value *Shareds = StaleCI->getArgOperand(i: 1);
2718 Align Alignment = TaskData->getPointerAlignment(DL: M.getDataLayout());
2719 Value *TaskShareds = Builder.CreateLoad(Ty: VoidPtr, Ptr: TaskData);
2720 Builder.CreateMemCpy(Dst: TaskShareds, DstAlign: Alignment, Src: Shareds, SrcAlign: Alignment,
2721 Size: SharedsSize);
2722 // Get the pointer to loop lb, ub, step from task ptr
2723 // and set up the lowerbound,upperbound and step values
2724 llvm::Value *Lb = Builder.CreateGEP(
2725 Ty: FakeSharedsTy, Ptr: TaskShareds, IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: 0)});
2726
2727 llvm::Value *Ub = Builder.CreateGEP(
2728 Ty: FakeSharedsTy, Ptr: TaskShareds, IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: 1)});
2729
2730 llvm::Value *Step = Builder.CreateGEP(
2731 Ty: FakeSharedsTy, Ptr: TaskShareds, IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: 2)});
2732 llvm::Value *Loadstep = Builder.CreateLoad(Ty: Builder.getInt64Ty(), Ptr: Step);
2733
2734 // set up the arguments for emitting kmpc_taskloop runtime call
2735 // setting values for ifval, nogroup, sched, grainsize, task_dup
2736 Value *IfCondVal =
2737 IfCond ? Builder.CreateIntCast(V: IfCond, DestTy: Builder.getInt32Ty(), isSigned: true)
2738 : Builder.getInt32(C: 1);
2739 // As __kmpc_taskgroup is called manually in OMPIRBuilder, NoGroupVal should
2740 // always be 1 when calling __kmpc_taskloop to ensure it is not called again
2741 Value *NoGroupVal = Builder.getInt32(C: 1);
2742 Value *SchedVal = Builder.getInt32(C: Sched);
2743 Value *GrainSizeVal =
2744 GrainSize ? Builder.CreateIntCast(V: GrainSize, DestTy: Builder.getInt64Ty(), isSigned: true)
2745 : Builder.getInt64(C: 0);
2746 Value *TaskDup = TaskDupFn;
2747
2748 Value *Args[] = {Ident, ThreadID, TaskData, IfCondVal, Lb, Ub,
2749 Loadstep, NoGroupVal, SchedVal, GrainSizeVal, TaskDup};
2750
2751 // taskloop runtime call
2752 Function *TaskloopFn =
2753 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_taskloop);
2754 Builder.CreateCall(Callee: TaskloopFn, Args);
2755
2756 // Emit the @__kmpc_end_taskgroup runtime call to end the taskgroup if
2757 // nogroup is not defined
2758 if (!NoGroup) {
2759 Function *EndTaskgroupFn =
2760 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_end_taskgroup);
2761 Builder.CreateCall(Callee: EndTaskgroupFn, Args: {Ident, ThreadID});
2762 }
2763
2764 StaleCI->eraseFromParent();
2765
2766 Builder.SetInsertPoint(TaskloopAllocaBB->begin());
2767
2768 LoadInst *SharedsOutlined =
2769 Builder.CreateLoad(Ty: VoidPtr, Ptr: OutlinedFn.getArg(i: 1));
2770 OutlinedFn.getArg(i: 1)->replaceUsesWithIf(
2771 New: SharedsOutlined,
2772 ShouldReplace: [SharedsOutlined](Use &U) { return U.getUser() != SharedsOutlined; });
2773
2774 Value *IV = CLI->getIndVar();
2775 Type *IVTy = IV->getType();
2776 Constant *One = ConstantInt::get(Ty: Builder.getInt64Ty(), V: 1);
2777
2778 // When outlining, CodeExtractor will create GEP's to the LowerBound and
2779 // UpperBound. These GEP's can be reused for loading the tasks respective
2780 // bounds.
2781 Value *TaskLB = nullptr;
2782 Value *TaskUB = nullptr;
2783 Value *TaskStep = nullptr;
2784 Value *LoadTaskLB = nullptr;
2785 Value *LoadTaskUB = nullptr;
2786 Value *LoadTaskStep = nullptr;
2787 for (Instruction &I : *TaskloopAllocaBB) {
2788 if (I.getOpcode() == Instruction::GetElementPtr) {
2789 GetElementPtrInst &Gep = cast<GetElementPtrInst>(Val&: I);
2790 if (ConstantInt *CI = dyn_cast<ConstantInt>(Val: Gep.getOperand(i_nocapture: 2))) {
2791 switch (CI->getZExtValue()) {
2792 case 0:
2793 TaskLB = &I;
2794 break;
2795 case 1:
2796 TaskUB = &I;
2797 break;
2798 case 2:
2799 TaskStep = &I;
2800 break;
2801 }
2802 }
2803 } else if (I.getOpcode() == Instruction::Load) {
2804 LoadInst &Load = cast<LoadInst>(Val&: I);
2805 if (Load.getPointerOperand() == TaskLB) {
2806 assert(TaskLB != nullptr && "Expected value for TaskLB");
2807 LoadTaskLB = &I;
2808 } else if (Load.getPointerOperand() == TaskUB) {
2809 assert(TaskUB != nullptr && "Expected value for TaskUB");
2810 LoadTaskUB = &I;
2811 } else if (Load.getPointerOperand() == TaskStep) {
2812 assert(TaskStep != nullptr && "Expected value for TaskStep");
2813 LoadTaskStep = &I;
2814 }
2815 }
2816 }
2817
2818 Builder.SetInsertPoint(CLI->getPreheader()->getTerminator());
2819
2820 assert(LoadTaskLB != nullptr && "Expected value for LoadTaskLB");
2821 assert(LoadTaskUB != nullptr && "Expected value for LoadTaskUB");
2822 assert(LoadTaskStep != nullptr && "Expected value for LoadTaskStep");
2823 Value *TripCountMinusOne = Builder.CreateSDiv(
2824 LHS: Builder.CreateSub(LHS: LoadTaskUB, RHS: LoadTaskLB), RHS: LoadTaskStep);
2825 Value *TripCount = Builder.CreateAdd(LHS: TripCountMinusOne, RHS: One, Name: "trip_cnt");
2826 Value *CastedTripCount = Builder.CreateIntCast(V: TripCount, DestTy: IVTy, isSigned: true);
2827 Value *CastedTaskLB = Builder.CreateIntCast(V: LoadTaskLB, DestTy: IVTy, isSigned: true);
2828 // set the trip count in the CLI
2829 CLI->setTripCount(CastedTripCount);
2830
2831 Builder.SetInsertPoint(CLI->getBody()->getFirstInsertionPt());
2832
2833 if (NumOfCollapseLoops > 1) {
2834 llvm::SmallVector<User *> UsersToReplace;
2835 // When using the collapse clause, the bounds of the loop have to be
2836 // adjusted to properly represent the iterator of the outer loop.
2837 Value *IVPlusTaskLB = Builder.CreateAdd(
2838 LHS: CLI->getIndVar(),
2839 RHS: Builder.CreateSub(LHS: CastedTaskLB, RHS: ConstantInt::get(Ty: IVTy, V: 1)));
2840 // To ensure every Use is correctly captured, we first want to record
2841 // which users to replace the value in, and then replace the value.
2842 for (auto IVUse = CLI->getIndVar()->uses().begin();
2843 IVUse != CLI->getIndVar()->uses().end(); IVUse++) {
2844 User *IVUser = IVUse->getUser();
2845 if (auto *Op = dyn_cast<BinaryOperator>(Val: IVUser)) {
2846 if (Op->getOpcode() == Instruction::URem ||
2847 Op->getOpcode() == Instruction::UDiv) {
2848 UsersToReplace.push_back(Elt: IVUser);
2849 }
2850 }
2851 }
2852 for (User *User : UsersToReplace) {
2853 User->replaceUsesOfWith(From: CLI->getIndVar(), To: IVPlusTaskLB);
2854 }
2855 } else {
2856 // The canonical loop is generated with a fixed lower bound. We need to
2857 // update the index calculation code to use the task's lower bound. The
2858 // generated code looks like this:
2859 // %omp_loop.iv = phi ...
2860 // ...
2861 // %tmp = mul [type] %omp_loop.iv, step
2862 // %user_index = add [type] tmp, lb
2863 // OpenMPIRBuilder constructs canonical loops to have exactly three uses
2864 // of the normalised induction variable:
2865 // 1. This one: converting the normalised IV to the user IV
2866 // 2. The increment (add)
2867 // 3. The comparison against the trip count (icmp)
2868 // (1) is the only use that is a mul followed by an add so this cannot
2869 // match other IR.
2870 assert(CLI->getIndVar()->getNumUses() == 3 &&
2871 "Canonical loop should have exactly three uses of the ind var");
2872 for (User *IVUser : CLI->getIndVar()->users()) {
2873 if (auto *Mul = dyn_cast<BinaryOperator>(Val: IVUser)) {
2874 if (Mul->getOpcode() == Instruction::Mul) {
2875 for (User *MulUser : Mul->users()) {
2876 if (auto *Add = dyn_cast<BinaryOperator>(Val: MulUser)) {
2877 if (Add->getOpcode() == Instruction::Add) {
2878 Add->setOperand(i_nocapture: 1, Val_nocapture: CastedTaskLB);
2879 }
2880 }
2881 }
2882 }
2883 }
2884 }
2885 }
2886
2887 FakeLB->replaceAllUsesWith(V: CastedLBVal);
2888 FakeUB->replaceAllUsesWith(V: CastedUBVal);
2889 FakeStep->replaceAllUsesWith(V: CastedStepVal);
2890 for (Instruction *I : llvm::reverse(C&: ToBeDeleted)) {
2891 I->eraseFromParent();
2892 }
2893 };
2894
2895 addOutlineInfo(OI: std::move(OI));
2896 Builder.SetInsertPoint(TaskloopExitBB->begin());
2897 return Builder.saveIP();
2898}
2899
2900llvm::StructType *OpenMPIRBuilder::getKmpTaskAffinityInfoTy() {
2901 llvm::Type *IntPtrTy = llvm::Type::getIntNTy(
2902 C&: M.getContext(), N: M.getDataLayout().getPointerSizeInBits());
2903 return llvm::StructType::get(elt1: IntPtrTy, elts: IntPtrTy,
2904 elts: llvm::Type::getInt32Ty(C&: M.getContext()));
2905}
2906
2907OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createTask(
2908 const LocationDescription &Loc, InsertPointTy AllocaIP,
2909 ArrayRef<BasicBlock *> DeallocBlocks, BodyGenCallbackTy BodyGenCB,
2910 bool Tied, Value *Final, Value *IfCondition,
2911 const DependenciesInfo &Dependencies, const AffinityData &Affinities,
2912 bool Mergeable, Value *EventHandle, Value *Priority, bool FreeAgent) {
2913
2914 if (!updateToLocation(Loc))
2915 return InsertPointTy();
2916
2917 uint32_t SrcLocStrSize;
2918 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
2919 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
2920 // The current basic block is split into four basic blocks. After outlining,
2921 // they will be mapped as follows:
2922 // ```
2923 // def current_fn() {
2924 // current_basic_block:
2925 // br label %task.exit
2926 // task.exit:
2927 // ; instructions after task
2928 // }
2929 // def outlined_fn() {
2930 // task.alloca:
2931 // br label %task.body
2932 // task.body:
2933 // ret void
2934 // }
2935 // ```
2936 BasicBlock *TaskExitBB = splitBB(Builder, /*CreateBranch=*/true, Name: "task.exit");
2937 BasicBlock *TaskBodyBB = splitBB(Builder, /*CreateBranch=*/true, Name: "task.body");
2938 BasicBlock *TaskAllocaBB =
2939 splitBB(Builder, /*CreateBranch=*/true, Name: "task.alloca");
2940
2941 InsertPointTy TaskAllocaIP = TaskAllocaBB->begin();
2942 InsertPointTy TaskBodyIP = TaskBodyBB->begin();
2943 if (Error Err = BodyGenCB(TaskAllocaIP, TaskBodyIP, TaskExitBB))
2944 return Err;
2945
2946 auto OI = std::make_unique<OutlineInfo>();
2947 OI->EntryBB = TaskAllocaBB;
2948 OI->OuterAllocBB = AllocaIP.getNodeParent();
2949 OI->ExitBB = TaskExitBB;
2950 OI->OuterDeallocBBs.reserve(N: DeallocBlocks.size());
2951 copy(Range&: DeallocBlocks, Out: OI->OuterDeallocBBs.end());
2952
2953 // Add the thread ID argument.
2954 SmallVector<Instruction *, 4> ToBeDeleted;
2955 OI->ExcludeArgsFromAggregate.push_back(Elt: createFakeIntVal(
2956 Builder, OuterAllocaIP: AllocaIP, ToBeDeleted, InnerAllocaIP: TaskAllocaIP, Name: "global.tid", AsPtr: false));
2957
2958 OI->PostOutlineCB = [this, Ident, Tied, Final, IfCondition, Dependencies,
2959 Affinities, Mergeable, Priority, EventHandle, FreeAgent,
2960 TaskAllocaBB,
2961 ToBeDeleted](Function &OutlinedFn) mutable {
2962 // Replace the Stale CI by appropriate RTL function call.
2963 assert(OutlinedFn.hasOneUse() &&
2964 "there must be a single user for the outlined function");
2965 CallInst *StaleCI = cast<CallInst>(Val: OutlinedFn.user_back());
2966
2967 // HasShareds is true if any variables are captured in the outlined region,
2968 // false otherwise.
2969 bool HasShareds = StaleCI->arg_size() > 1;
2970 Builder.SetInsertPoint(StaleCI);
2971
2972 // Gather the arguments for emitting the runtime call for
2973 // @__kmpc_omp_task_alloc
2974 Function *TaskAllocFn =
2975 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_alloc);
2976
2977 // Arguments - `loc_ref` (Ident) and `gtid` (ThreadID)
2978 // call.
2979 Value *ThreadID = getOrCreateThreadID(Ident);
2980
2981 // Argument - `flags`
2982 // Task is tied iff (Flags & 1) == 1.
2983 // Task is untied iff (Flags & 1) == 0.
2984 // Task is final iff (Flags & 2) == 2.
2985 // Task is not final iff (Flags & 2) == 0.
2986 // Task is mergeable or merged-if0 iff (Flags & 4) == 4.
2987 // Task is neither mergeable nor merged-if0 iff (Flags & 4) == 0.
2988 // Task is detachable iff (Flags & 64) == 64.
2989 // Task is not detachable iff (Flags & 64) == 0.
2990 // Task is priority iff (Flags & 32) == 32.
2991 // Task is not priority iff (Flags & 32) == 0.
2992 // Task is free-agent eligible iff (Flags & 128) == 128.
2993 // Task is not free-agent eligible iff (Flags & 128) == 0.
2994 // TODO: Handle the other flags.
2995 Value *Flags = Builder.getInt32(C: Tied);
2996 auto *ConstIfCondition = dyn_cast_or_null<ConstantInt>(Val: IfCondition);
2997 bool UseMergedIf0Path = ConstIfCondition && ConstIfCondition->isZero();
2998 if (Final) {
2999 Value *FinalFlag =
3000 Builder.CreateSelect(C: Final, True: Builder.getInt32(C: 2), False: Builder.getInt32(C: 0));
3001 Flags = Builder.CreateOr(LHS: FinalFlag, RHS: Flags);
3002 }
3003
3004 if (Mergeable || UseMergedIf0Path)
3005 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 4), RHS: Flags);
3006 if (EventHandle)
3007 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 64), RHS: Flags);
3008 if (Priority)
3009 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 32), RHS: Flags);
3010 if (FreeAgent)
3011 Flags = Builder.CreateOr(LHS: Builder.getInt32(C: 128), RHS: Flags);
3012
3013 // Argument - `sizeof_kmp_task_t` (TaskSize)
3014 // Tasksize refers to the size in bytes of kmp_task_t data structure
3015 // including private vars accessed in task.
3016 // TODO: add kmp_task_t_with_privates (privates)
3017 Value *TaskSize = Builder.getInt64(
3018 C: divideCeil(Numerator: M.getDataLayout().getTypeSizeInBits(Ty: Task), Denominator: 8));
3019
3020 // Argument - `sizeof_shareds` (SharedsSize)
3021 // SharedsSize refers to the shareds array size in the kmp_task_t data
3022 // structure.
3023 Value *SharedsSize = Builder.getInt64(C: 0);
3024 if (HasShareds) {
3025 AllocaInst *ArgStructAlloca =
3026 dyn_cast<AllocaInst>(Val: StaleCI->getArgOperand(i: 1));
3027 assert(ArgStructAlloca &&
3028 "Unable to find the alloca instruction corresponding to arguments "
3029 "for extracted function");
3030 std::optional<TypeSize> ArgAllocSize =
3031 ArgStructAlloca->getAllocationSize(DL: M.getDataLayout());
3032 assert(ArgAllocSize &&
3033 "Unable to determine size of arguments for extracted function");
3034 SharedsSize = Builder.getInt64(C: ArgAllocSize->getFixedValue());
3035 }
3036 // Emit the @__kmpc_omp_task_alloc runtime call
3037 // The runtime call returns a pointer to an area where the task captured
3038 // variables must be copied before the task is run (TaskData)
3039 CallInst *TaskData = createRuntimeFunctionCall(
3040 Callee: TaskAllocFn, Args: {/*loc_ref=*/Ident, /*gtid=*/ThreadID, /*flags=*/Flags,
3041 /*sizeof_task=*/TaskSize, /*sizeof_shared=*/SharedsSize,
3042 /*task_func=*/&OutlinedFn});
3043
3044 if (Affinities.Count && Affinities.Info) {
3045 Function *RegAffFn = getOrCreateRuntimeFunctionPtr(
3046 FnID: OMPRTL___kmpc_omp_reg_task_with_affinity);
3047
3048 createRuntimeFunctionCall(Callee: RegAffFn, Args: {Ident, ThreadID, TaskData,
3049 Affinities.Count, Affinities.Info});
3050 }
3051
3052 // Emit detach clause initialization.
3053 // evt = (typeof(evt))__kmpc_task_allow_completion_event(loc, tid,
3054 // task_descriptor);
3055 if (EventHandle) {
3056 Function *TaskDetachFn = getOrCreateRuntimeFunctionPtr(
3057 FnID: OMPRTL___kmpc_task_allow_completion_event);
3058 llvm::Value *EventVal =
3059 createRuntimeFunctionCall(Callee: TaskDetachFn, Args: {Ident, ThreadID, TaskData});
3060 llvm::Value *EventHandleAddr =
3061 Builder.CreatePointerBitCastOrAddrSpaceCast(V: EventHandle,
3062 DestTy: Builder.getPtrTy(AddrSpace: 0));
3063 EventVal = Builder.CreatePtrToInt(V: EventVal, DestTy: Builder.getInt64Ty());
3064 Builder.CreateStore(Val: EventVal, Ptr: EventHandleAddr);
3065 }
3066 // Copy the arguments for outlined function
3067 if (HasShareds) {
3068 Value *Shareds = StaleCI->getArgOperand(i: 1);
3069 Align Alignment = TaskData->getPointerAlignment(DL: M.getDataLayout());
3070 Value *TaskShareds = Builder.CreateLoad(Ty: VoidPtr, Ptr: TaskData);
3071 Builder.CreateMemCpy(Dst: TaskShareds, DstAlign: Alignment, Src: Shareds, SrcAlign: Alignment,
3072 Size: SharedsSize);
3073 }
3074
3075 if (Priority) {
3076 //
3077 // The return type of "__kmpc_omp_task_alloc" is "kmp_task_t *",
3078 // we populate the priority information into the "kmp_task_t" here
3079 //
3080 // The struct "kmp_task_t" definition is available in kmp.h
3081 // kmp_task_t = { shareds, routine, part_id, data1, data2 }
3082 // data2 is used for priority
3083 //
3084 Type *Int32Ty = Builder.getInt32Ty();
3085 Constant *Zero = ConstantInt::get(Ty: Int32Ty, V: 0);
3086 // kmp_task_t* => { ptr }
3087 Type *TaskPtr = StructType::get(elt1: VoidPtr);
3088 Value *TaskGEP =
3089 Builder.CreateInBoundsGEP(Ty: TaskPtr, Ptr: TaskData, IdxList: {Zero, Zero});
3090 // kmp_task_t => { ptr, ptr, i32, ptr, ptr }
3091 Type *TaskStructType = StructType::get(
3092 elt1: VoidPtr, elts: VoidPtr, elts: Builder.getInt32Ty(), elts: VoidPtr, elts: VoidPtr);
3093 Value *PriorityData = Builder.CreateInBoundsGEP(
3094 Ty: TaskStructType, Ptr: TaskGEP, IdxList: {Zero, ConstantInt::get(Ty: Int32Ty, V: 4)});
3095 // kmp_cmplrdata_t => { ptr, ptr }
3096 Type *CmplrStructType = StructType::get(elt1: VoidPtr, elts: VoidPtr);
3097 Value *CmplrData = Builder.CreateInBoundsGEP(Ty: CmplrStructType,
3098 Ptr: PriorityData, IdxList: {Zero, Zero});
3099 Builder.CreateStore(Val: Priority, Ptr: CmplrData);
3100 }
3101
3102 Value *DepArray = nullptr;
3103 Value *NumDeps = nullptr;
3104 if (Dependencies.DepArray) {
3105 DepArray = Dependencies.DepArray;
3106 NumDeps = Dependencies.NumDeps;
3107 } else if (!Dependencies.Deps.empty()) {
3108 DepArray = emitTaskDependencies(OMPBuilder&: *this, Dependencies: Dependencies.Deps);
3109 NumDeps = Builder.getInt32(C: Dependencies.Deps.size());
3110 }
3111
3112 // In the presence of the `if` clause, the following IR is generated:
3113 // ...
3114 // %data = call @__kmpc_omp_task_alloc(...)
3115 // br i1 %if_condition, label %then, label %else
3116 // then:
3117 // call @__kmpc_omp_task(...)
3118 // br label %exit
3119 // else:
3120 // ;; Wait for resolution of dependencies, if any, before
3121 // ;; beginning the task
3122 // call @__kmpc_omp_wait_deps(...)
3123 // call @__kmpc_omp_task_begin_if0(...)
3124 // call @outlined_fn(...)
3125 // call @__kmpc_omp_task_complete_if0(...)
3126 // br label %exit
3127 // exit:
3128 // ...
3129 if (IfCondition && !UseMergedIf0Path) {
3130 // `SplitBlockAndInsertIfThenElse` requires the block to have a
3131 // terminator.
3132 splitBB(Builder, /*CreateBranch=*/true, Name: "if.end");
3133 Instruction *IfTerminator =
3134 Builder.GetInsertPoint()->getParent()->getTerminator();
3135 Instruction *ThenTI = IfTerminator, *ElseTI = nullptr;
3136 Builder.SetInsertPoint(IfTerminator);
3137 SplitBlockAndInsertIfThenElse(Cond: IfCondition, SplitBefore: IfTerminator, ThenTerm: &ThenTI,
3138 ElseTerm: &ElseTI);
3139 Builder.SetInsertPoint(ElseTI);
3140
3141 if (DepArray) {
3142 Function *TaskWaitFn =
3143 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_wait_deps);
3144 createRuntimeFunctionCall(
3145 Callee: TaskWaitFn,
3146 Args: {Ident, ThreadID, NumDeps, DepArray,
3147 ConstantInt::get(Ty: Builder.getInt32Ty(), V: 0),
3148 ConstantPointerNull::get(T: PointerType::getUnqual(C&: M.getContext()))});
3149 }
3150 Function *TaskBeginFn =
3151 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_begin_if0);
3152 Function *TaskCompleteFn =
3153 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_complete_if0);
3154 createRuntimeFunctionCall(Callee: TaskBeginFn, Args: {Ident, ThreadID, TaskData});
3155 CallInst *CI = nullptr;
3156 if (HasShareds)
3157 CI = createRuntimeFunctionCall(Callee: &OutlinedFn, Args: {ThreadID, TaskData});
3158 else
3159 CI = createRuntimeFunctionCall(Callee: &OutlinedFn, Args: {ThreadID});
3160 CI->setDebugLoc(StaleCI->getDebugLoc());
3161 createRuntimeFunctionCall(Callee: TaskCompleteFn, Args: {Ident, ThreadID, TaskData});
3162 Builder.SetInsertPoint(ThenTI);
3163 }
3164
3165 if (DepArray) {
3166 Function *TaskFn =
3167 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_with_deps);
3168 createRuntimeFunctionCall(
3169 Callee: TaskFn,
3170 Args: {Ident, ThreadID, TaskData, NumDeps, DepArray,
3171 ConstantInt::get(Ty: Builder.getInt32Ty(), V: 0),
3172 ConstantPointerNull::get(T: PointerType::getUnqual(C&: M.getContext()))});
3173
3174 } else {
3175 // Emit the @__kmpc_omp_task runtime call to spawn the task
3176 Function *TaskFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task);
3177 createRuntimeFunctionCall(Callee: TaskFn, Args: {Ident, ThreadID, TaskData});
3178 }
3179
3180 StaleCI->eraseFromParent();
3181
3182 Builder.SetInsertPoint(TaskAllocaBB->begin());
3183 if (HasShareds) {
3184 LoadInst *Shareds = Builder.CreateLoad(Ty: VoidPtr, Ptr: OutlinedFn.getArg(i: 1));
3185 OutlinedFn.getArg(i: 1)->replaceUsesWithIf(
3186 New: Shareds, ShouldReplace: [Shareds](Use &U) { return U.getUser() != Shareds; });
3187 }
3188
3189 // The insert point may refer to one of the instructions about to be
3190 // deleted. It is not needed anymore so clear it instead of leaving it
3191 // dangling.
3192 Builder.ClearInsertionPoint();
3193 for (Instruction *I : llvm::reverse(C&: ToBeDeleted))
3194 I->eraseFromParent();
3195 };
3196
3197 addOutlineInfo(OI: std::move(OI));
3198 Builder.SetInsertPoint(TaskExitBB->begin());
3199
3200 return Builder.saveIP();
3201}
3202
3203OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createTaskgroup(
3204 const LocationDescription &Loc, InsertPointTy AllocaIP,
3205 ArrayRef<BasicBlock *> DeallocBlocks, BodyGenCallbackTy BodyGenCB) {
3206 if (!updateToLocation(Loc))
3207 return InsertPointTy();
3208
3209 uint32_t SrcLocStrSize;
3210 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
3211 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
3212 Value *ThreadID = getOrCreateThreadID(Ident);
3213
3214 // Emit the @__kmpc_taskgroup runtime call to start the taskgroup
3215 Function *TaskgroupFn =
3216 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_taskgroup);
3217 createRuntimeFunctionCall(Callee: TaskgroupFn, Args: {Ident, ThreadID});
3218
3219 BasicBlock *TaskgroupExitBB = splitBB(Builder, CreateBranch: true, Name: "taskgroup.exit");
3220 if (Error Err = BodyGenCB(AllocaIP, Builder.saveIP(), DeallocBlocks))
3221 return Err;
3222
3223 Builder.SetInsertPoint(TaskgroupExitBB);
3224 // Emit the @__kmpc_end_taskgroup runtime call to end the taskgroup
3225 Function *EndTaskgroupFn =
3226 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_end_taskgroup);
3227 createRuntimeFunctionCall(Callee: EndTaskgroupFn, Args: {Ident, ThreadID});
3228
3229 return Builder.saveIP();
3230}
3231
3232OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createSections(
3233 const LocationDescription &Loc, InsertPointTy AllocaIP,
3234 ArrayRef<StorableBodyGenCallbackTy> SectionCBs, PrivatizeCallbackTy PrivCB,
3235 FinalizeCallbackTy FiniCB, bool IsCancellable, bool IsNowait) {
3236 assert(!isConflictIP(AllocaIP, Loc.IP) && "Dedicated IP allocas required");
3237
3238 if (!updateToLocation(Loc))
3239 return Loc.IP;
3240
3241 FinalizationStack.push_back(Elt: {FiniCB, OMPD_sections, IsCancellable});
3242
3243 // Each section is emitted as a switch case
3244 // Each finalization callback is handled from clang.EmitOMPSectionDirective()
3245 // -> OMP.createSection() which generates the IR for each section
3246 // Iterate through all sections and emit a switch construct:
3247 // switch (IV) {
3248 // case 0:
3249 // <SectionStmt[0]>;
3250 // break;
3251 // ...
3252 // case <NumSection> - 1:
3253 // <SectionStmt[<NumSection> - 1]>;
3254 // break;
3255 // }
3256 // ...
3257 // section_loop.after:
3258 // <FiniCB>;
3259 auto LoopBodyGenCB = [&](InsertPointTy CodeGenIP, Value *IndVar) -> Error {
3260 Builder.restoreIP(IP: CodeGenIP);
3261 BasicBlock *Continue =
3262 splitBBWithSuffix(Builder, /*CreateBranch=*/false, Suffix: ".sections.after");
3263 Function *CurFn = Continue->getParent();
3264 SwitchInst *SwitchStmt = Builder.CreateSwitch(V: IndVar, Dest: Continue);
3265
3266 unsigned CaseNumber = 0;
3267 for (auto SectionCB : SectionCBs) {
3268 BasicBlock *CaseBB = BasicBlock::Create(
3269 Context&: M.getContext(), Name: "omp_section_loop.body.case", Parent: CurFn, InsertBefore: Continue);
3270 SwitchStmt->addCase(OnVal: Builder.getInt32(C: CaseNumber), Dest: CaseBB);
3271 Builder.SetInsertPoint(CaseBB);
3272 UncondBrInst *CaseEndBr = Builder.CreateBr(Dest: Continue);
3273 if (Error Err = SectionCB(InsertPointTy(), CaseEndBr->getIterator(), {}))
3274 return Err;
3275 CaseNumber++;
3276 }
3277 // remove the existing terminator from body BB since there can be no
3278 // terminators after switch/case
3279 return Error::success();
3280 };
3281 // Loop body ends here
3282 // LowerBound, UpperBound, and STride for createCanonicalLoop
3283 Type *I32Ty = Type::getInt32Ty(C&: M.getContext());
3284 Value *LB = ConstantInt::get(Ty: I32Ty, V: 0);
3285 Value *UB = ConstantInt::get(Ty: I32Ty, V: SectionCBs.size());
3286 Value *ST = ConstantInt::get(Ty: I32Ty, V: 1);
3287 Expected<CanonicalLoopInfo *> LoopInfo = createCanonicalLoop(
3288 Loc, BodyGenCB: LoopBodyGenCB, Start: LB, Stop: UB, Step: ST, IsSigned: true, InclusiveStop: false, ComputeIP: AllocaIP, Name: "section_loop");
3289 if (!LoopInfo)
3290 return LoopInfo.takeError();
3291
3292 InsertPointOrErrorTy WsloopIP =
3293 applyStaticWorkshareLoop(DL: Loc.DL, CLI: *LoopInfo, AllocaIP,
3294 LoopType: WorksharingLoopType::ForStaticLoop, NeedsBarrier: !IsNowait);
3295 if (!WsloopIP)
3296 return WsloopIP.takeError();
3297 InsertPointTy AfterIP = *WsloopIP;
3298
3299 BasicBlock *LoopFini = AfterIP.getNodeParent()->getSinglePredecessor();
3300 assert(LoopFini && "Bad structure of static workshare loop finalization");
3301
3302 // Apply the finalization callback in LoopAfterBB
3303 auto FiniInfo = FinalizationStack.pop_back_val();
3304 assert(FiniInfo.DK == OMPD_sections &&
3305 "Unexpected finalization stack state!");
3306 if (Error Err = FiniInfo.mergeFiniBB(Builder, OtherFiniBB: LoopFini))
3307 return Err;
3308
3309 return AfterIP;
3310}
3311
3312OpenMPIRBuilder::InsertPointOrErrorTy
3313OpenMPIRBuilder::createSection(const LocationDescription &Loc,
3314 BodyGenCallbackTy BodyGenCB,
3315 FinalizeCallbackTy FiniCB) {
3316 if (!updateToLocation(Loc))
3317 return Loc.IP;
3318
3319 auto FiniCBWrapper = [&](InsertPointTy IP) {
3320 if (IP != IP.getNodeParent()->end())
3321 return FiniCB(IP);
3322 // This must be done otherwise any nested constructs using FinalizeOMPRegion
3323 // will fail because that function requires the Finalization Basic Block to
3324 // have a terminator, which is already removed by EmitOMPRegionBody.
3325 // IP is currently at cancelation block.
3326 // We need to backtrack to the condition block to fetch
3327 // the exit block and create a branch from cancelation
3328 // to exit block.
3329 IRBuilder<>::InsertPointGuard IPG(Builder);
3330 Builder.restoreIP(IP);
3331 auto *CaseBB = Loc.IP.getNodeParent();
3332 auto *CondBB = CaseBB->getSinglePredecessor()->getSinglePredecessor();
3333 auto *ExitBB = CondBB->getTerminator()->getSuccessor(Idx: 1);
3334 Instruction *I = Builder.CreateBr(Dest: ExitBB);
3335 IP = I->getIterator();
3336 return FiniCB(IP);
3337 };
3338
3339 Directive OMPD = Directive::OMPD_sections;
3340 // Since we are using Finalization Callback here, HasFinalize
3341 // and IsCancellable have to be true
3342 return EmitOMPInlinedRegion(OMPD, EntryCall: nullptr, ExitCall: nullptr, BodyGenCB, FiniCB: FiniCBWrapper,
3343 /*Conditional*/ false, /*hasFinalize*/ HasFinalize: true,
3344 /*IsCancellable*/ true);
3345}
3346
3347static OpenMPIRBuilder::InsertPointTy getInsertPointAfterInstr(Instruction *I) {
3348 BasicBlock::iterator IT(I);
3349 IT++;
3350 return IT;
3351}
3352
3353Value *OpenMPIRBuilder::getGPUThreadID() {
3354 return createRuntimeFunctionCall(
3355 Callee: getOrCreateRuntimeFunction(M,
3356 FnID: OMPRTL___kmpc_get_hardware_thread_id_in_block),
3357 Args: {});
3358}
3359
3360Value *OpenMPIRBuilder::getGPUWarpSize() {
3361 return createRuntimeFunctionCall(
3362 Callee: getOrCreateRuntimeFunction(M, FnID: OMPRTL___kmpc_get_warp_size), Args: {});
3363}
3364
3365Value *OpenMPIRBuilder::getNVPTXWarpID() {
3366 unsigned LaneIDBits = Log2_32(Value: Config.getGridValue().GV_Warp_Size);
3367 return Builder.CreateAShr(LHS: getGPUThreadID(), RHS: LaneIDBits, Name: "nvptx_warp_id");
3368}
3369
3370Value *OpenMPIRBuilder::getNVPTXLaneID() {
3371 unsigned LaneIDBits = Log2_32(Value: Config.getGridValue().GV_Warp_Size);
3372 assert(LaneIDBits < 32 && "Invalid LaneIDBits size in NVPTX device.");
3373 unsigned LaneIDMask = ~0u >> (32u - LaneIDBits);
3374 return Builder.CreateAnd(LHS: getGPUThreadID(), RHS: Builder.getInt32(C: LaneIDMask),
3375 Name: "nvptx_lane_id");
3376}
3377
3378Value *OpenMPIRBuilder::castValueToType(InsertPointTy AllocaIP, Value *From,
3379 Type *ToType) {
3380 Type *FromType = From->getType();
3381 uint64_t FromSize = M.getDataLayout().getTypeStoreSize(Ty: FromType);
3382 uint64_t ToSize = M.getDataLayout().getTypeStoreSize(Ty: ToType);
3383 assert(FromSize > 0 && "From size must be greater than zero");
3384 assert(ToSize > 0 && "To size must be greater than zero");
3385 if (FromType == ToType)
3386 return From;
3387 if (FromSize == ToSize)
3388 return Builder.CreateBitCast(V: From, DestTy: ToType);
3389 if (ToType->isIntegerTy() && FromType->isIntegerTy())
3390 return Builder.CreateIntCast(V: From, DestTy: ToType, /*isSigned*/ true);
3391 InsertPointTy SaveIP = Builder.saveIP();
3392 Builder.restoreIP(IP: AllocaIP);
3393 Value *CastItem = Builder.CreateAlloca(Ty: ToType);
3394 Builder.restoreIP(IP: SaveIP);
3395
3396 Value *ValCastItem = Builder.CreatePointerBitCastOrAddrSpaceCast(
3397 V: CastItem, DestTy: Builder.getPtrTy(AddrSpace: 0));
3398 Builder.CreateStore(Val: From, Ptr: ValCastItem);
3399 return Builder.CreateLoad(Ty: ToType, Ptr: CastItem);
3400}
3401
3402Value *OpenMPIRBuilder::createRuntimeShuffleFunction(InsertPointTy AllocaIP,
3403 Value *Element,
3404 Type *ElementType,
3405 Value *Offset) {
3406 uint64_t Size = M.getDataLayout().getTypeStoreSize(Ty: ElementType);
3407 assert(Size <= 8 && "Unsupported bitwidth in shuffle instruction");
3408
3409 // Cast all types to 32- or 64-bit values before calling shuffle routines.
3410 Type *CastTy = Builder.getIntNTy(N: Size <= 4 ? 32 : 64);
3411 Value *ElemCast = castValueToType(AllocaIP, From: Element, ToType: CastTy);
3412 Value *WarpSize =
3413 Builder.CreateIntCast(V: getGPUWarpSize(), DestTy: Builder.getInt16Ty(), isSigned: true);
3414 Function *ShuffleFunc = getOrCreateRuntimeFunctionPtr(
3415 FnID: Size <= 4 ? RuntimeFunction::OMPRTL___kmpc_shuffle_int32
3416 : RuntimeFunction::OMPRTL___kmpc_shuffle_int64);
3417 Value *WarpSizeCast =
3418 Builder.CreateIntCast(V: WarpSize, DestTy: Builder.getInt16Ty(), /*isSigned=*/true);
3419 Value *ShuffleCall =
3420 createRuntimeFunctionCall(Callee: ShuffleFunc, Args: {ElemCast, Offset, WarpSizeCast});
3421 // The shuffle runtime functions return a 32- or 64-bit value. Cast it back
3422 // down to the requested element type, otherwise storing the result would
3423 // write past the end of an element narrower than the shuffle width.
3424 return castValueToType(AllocaIP, From: ShuffleCall, ToType: ElementType);
3425}
3426
3427void OpenMPIRBuilder::shuffleAndStore(InsertPointTy AllocaIP, Value *SrcAddr,
3428 Value *DstAddr, Type *ElemType,
3429 Value *Offset, Type *ReductionArrayTy,
3430 bool IsByRefElem) {
3431 uint64_t Size = M.getDataLayout().getTypeStoreSize(Ty: ElemType);
3432 // Create the loop over the big sized data.
3433 // ptr = (void*)Elem;
3434 // ptrEnd = (void*) Elem + 1;
3435 // Step = 8;
3436 // while (ptr + Step < ptrEnd)
3437 // shuffle((int64_t)*ptr);
3438 // Step = 4;
3439 // while (ptr + Step < ptrEnd)
3440 // shuffle((int32_t)*ptr);
3441 // ...
3442 Type *IndexTy = Builder.getIndexTy(
3443 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
3444 Value *ElemPtr = DstAddr;
3445 Value *Ptr = SrcAddr;
3446 for (unsigned IntSize = 8; IntSize >= 1; IntSize /= 2) {
3447 if (Size < IntSize)
3448 continue;
3449 Type *IntType = Builder.getIntNTy(N: IntSize * 8);
3450 Ptr = Builder.CreatePointerBitCastOrAddrSpaceCast(
3451 V: Ptr, DestTy: Builder.getPtrTy(AddrSpace: 0), Name: Ptr->getName() + ".ascast");
3452 Value *SrcAddrGEP =
3453 Builder.CreateGEP(Ty: ElemType, Ptr: SrcAddr, IdxList: {ConstantInt::get(Ty: IndexTy, V: 1)});
3454 ElemPtr = Builder.CreatePointerBitCastOrAddrSpaceCast(
3455 V: ElemPtr, DestTy: Builder.getPtrTy(AddrSpace: 0), Name: ElemPtr->getName() + ".ascast");
3456
3457 Function *CurFunc = Builder.GetInsertBlock()->getParent();
3458 if ((Size / IntSize) > 1) {
3459 Value *PtrEnd = Builder.CreatePointerBitCastOrAddrSpaceCast(
3460 V: SrcAddrGEP, DestTy: Builder.getPtrTy());
3461 BasicBlock *PreCondBB =
3462 BasicBlock::Create(Context&: M.getContext(), Name: ".shuffle.pre_cond");
3463 BasicBlock *ThenBB = BasicBlock::Create(Context&: M.getContext(), Name: ".shuffle.then");
3464 BasicBlock *ExitBB = BasicBlock::Create(Context&: M.getContext(), Name: ".shuffle.exit");
3465 BasicBlock *CurrentBB = Builder.GetInsertBlock();
3466 emitBlock(BB: PreCondBB, CurFn: CurFunc);
3467 PHINode *PhiSrc =
3468 Builder.CreatePHI(Ty: Ptr->getType(), /*NumReservedValues=*/2);
3469 PhiSrc->addIncoming(V: Ptr, BB: CurrentBB);
3470 PHINode *PhiDest =
3471 Builder.CreatePHI(Ty: ElemPtr->getType(), /*NumReservedValues=*/2);
3472 PhiDest->addIncoming(V: ElemPtr, BB: CurrentBB);
3473 Ptr = PhiSrc;
3474 ElemPtr = PhiDest;
3475 Value *PtrDiff = Builder.CreatePtrDiff(
3476 ElemTy: Builder.getInt8Ty(), LHS: PtrEnd,
3477 RHS: Builder.CreatePointerBitCastOrAddrSpaceCast(V: Ptr, DestTy: Builder.getPtrTy()));
3478 Builder.CreateCondBr(
3479 Cond: Builder.CreateICmpSGT(LHS: PtrDiff, RHS: Builder.getInt64(C: IntSize - 1)), True: ThenBB,
3480 False: ExitBB);
3481 emitBlock(BB: ThenBB, CurFn: CurFunc);
3482 Value *Res = createRuntimeShuffleFunction(
3483 AllocaIP,
3484 Element: Builder.CreateAlignedLoad(
3485 Ty: IntType, Ptr, Align: M.getDataLayout().getPrefTypeAlign(Ty: ElemType)),
3486 ElementType: IntType, Offset);
3487 Builder.CreateAlignedStore(Val: Res, Ptr: ElemPtr,
3488 Align: M.getDataLayout().getPrefTypeAlign(Ty: ElemType));
3489 Value *LocalPtr =
3490 Builder.CreateGEP(Ty: IntType, Ptr, IdxList: {ConstantInt::get(Ty: IndexTy, V: 1)});
3491 Value *LocalElemPtr =
3492 Builder.CreateGEP(Ty: IntType, Ptr: ElemPtr, IdxList: {ConstantInt::get(Ty: IndexTy, V: 1)});
3493 PhiSrc->addIncoming(V: LocalPtr, BB: ThenBB);
3494 PhiDest->addIncoming(V: LocalElemPtr, BB: ThenBB);
3495 emitBranch(Target: PreCondBB);
3496 emitBlock(BB: ExitBB, CurFn: CurFunc);
3497 } else {
3498 // The shuffled value comes back as the chunk's integer type, so the
3499 // store covers exactly this chunk regardless of what ElemType is.
3500 Value *Res = createRuntimeShuffleFunction(
3501 AllocaIP, Element: Builder.CreateLoad(Ty: IntType, Ptr), ElementType: IntType, Offset);
3502 Builder.CreateStore(Val: Res, Ptr: ElemPtr);
3503 Ptr = Builder.CreateGEP(Ty: IntType, Ptr, IdxList: {ConstantInt::get(Ty: IndexTy, V: 1)});
3504 ElemPtr =
3505 Builder.CreateGEP(Ty: IntType, Ptr: ElemPtr, IdxList: {ConstantInt::get(Ty: IndexTy, V: 1)});
3506 }
3507 Size = Size % IntSize;
3508 }
3509}
3510
3511Error OpenMPIRBuilder::emitReductionListCopy(
3512 InsertPointTy AllocaIP, CopyAction Action, Type *ReductionArrayTy,
3513 ArrayRef<ReductionInfo> ReductionInfos, Value *SrcBase, Value *DestBase,
3514 ArrayRef<bool> IsByRef, CopyOptionsTy CopyOptions) {
3515 Type *IndexTy = Builder.getIndexTy(
3516 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
3517 Value *RemoteLaneOffset = CopyOptions.RemoteLaneOffset;
3518
3519 // Iterates, element-by-element, through the source Reduce list and
3520 // make a copy.
3521 for (auto En : enumerate(First&: ReductionInfos)) {
3522 const ReductionInfo &RI = En.value();
3523 Value *SrcElementAddr = nullptr;
3524 AllocaInst *DestAlloca = nullptr;
3525 Value *DestElementAddr = nullptr;
3526 Value *DestElementPtrAddr = nullptr;
3527 // Should we shuffle in an element from a remote lane?
3528 bool ShuffleInElement = false;
3529 // Set to true to update the pointer in the dest Reduce list to a
3530 // newly created element.
3531 bool UpdateDestListPtr = false;
3532
3533 // Step 1.1: Get the address for the src element in the Reduce list.
3534 Value *SrcElementPtrAddr = Builder.CreateInBoundsGEP(
3535 Ty: ReductionArrayTy, Ptr: SrcBase,
3536 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
3537 SrcElementAddr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: SrcElementPtrAddr);
3538
3539 // Step 1.2: Create a temporary to store the element in the destination
3540 // Reduce list.
3541 DestElementPtrAddr = Builder.CreateInBoundsGEP(
3542 Ty: ReductionArrayTy, Ptr: DestBase,
3543 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
3544 bool IsByRefElem = (!IsByRef.empty() && IsByRef[En.index()]);
3545 switch (Action) {
3546 case CopyAction::RemoteLaneToThread: {
3547 InsertPointTy CurIP = Builder.saveIP();
3548 Builder.restoreIP(IP: AllocaIP);
3549
3550 Type *DestAllocaType =
3551 IsByRefElem ? RI.ByRefAllocatedType : RI.ElementType;
3552 DestAlloca = Builder.CreateAlloca(Ty: DestAllocaType, ArraySize: nullptr,
3553 Name: ".omp.reduction.element");
3554 DestAlloca->setAlignment(
3555 M.getDataLayout().getPrefTypeAlign(Ty: DestAllocaType));
3556 DestElementAddr = DestAlloca;
3557 DestElementAddr =
3558 Builder.CreateAddrSpaceCast(V: DestElementAddr, DestTy: Builder.getPtrTy(),
3559 Name: DestElementAddr->getName() + ".ascast");
3560 Builder.restoreIP(IP: CurIP);
3561 ShuffleInElement = true;
3562 UpdateDestListPtr = true;
3563 break;
3564 }
3565 case CopyAction::ThreadCopy: {
3566 DestElementAddr =
3567 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: DestElementPtrAddr);
3568 break;
3569 }
3570 }
3571
3572 // Now that all active lanes have read the element in the
3573 // Reduce list, shuffle over the value from the remote lane.
3574 if (ShuffleInElement) {
3575 Type *ShuffleType = RI.ElementType;
3576 Value *ShuffleSrcAddr = SrcElementAddr;
3577 Value *ShuffleDestAddr = DestElementAddr;
3578 AllocaInst *LocalStorage = nullptr;
3579
3580 if (IsByRefElem) {
3581 assert(RI.ByRefElementType && "Expected by-ref element type to be set");
3582 assert(RI.ByRefAllocatedType &&
3583 "Expected by-ref allocated type to be set");
3584 // For by-ref reductions, we need to copy from the remote lane the
3585 // actual value of the partial reduction computed by that remote lane;
3586 // rather than, for example, a pointer to that data or, even worse, a
3587 // pointer to the descriptor of the by-ref reduction element.
3588 ShuffleType = RI.ByRefElementType;
3589
3590 if (RI.DataPtrPtrGen) {
3591 // Descriptor-based by-ref: extract data pointer from descriptor.
3592 InsertPointOrErrorTy GenResult = RI.DataPtrPtrGen(
3593 Builder.saveIP(), ShuffleSrcAddr, ShuffleSrcAddr);
3594
3595 if (!GenResult)
3596 return GenResult.takeError();
3597
3598 ShuffleSrcAddr =
3599 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ShuffleSrcAddr);
3600
3601 {
3602 InsertPointTy OldIP = Builder.saveIP();
3603 Builder.restoreIP(IP: AllocaIP);
3604
3605 LocalStorage = Builder.CreateAlloca(Ty: ShuffleType);
3606 Builder.restoreIP(IP: OldIP);
3607 ShuffleDestAddr = LocalStorage;
3608 }
3609 } else {
3610 // Non-descriptor by-ref: the pointer already references data
3611 // directly. Shuffle into the destination alloca.
3612 ShuffleDestAddr = DestElementAddr;
3613 }
3614 }
3615
3616 shuffleAndStore(AllocaIP, SrcAddr: ShuffleSrcAddr, DstAddr: ShuffleDestAddr, ElemType: ShuffleType,
3617 Offset: RemoteLaneOffset, ReductionArrayTy, IsByRefElem);
3618
3619 if (IsByRefElem && RI.DataPtrPtrGen) {
3620 // Copy descriptor from source and update base_ptr to shuffled data
3621 Value *DestDescriptorAddr = Builder.CreatePointerBitCastOrAddrSpaceCast(
3622 V: DestAlloca, DestTy: Builder.getPtrTy(), Name: ".ascast");
3623
3624 InsertPointOrErrorTy GenResult = generateReductionDescriptor(
3625 DescriptorAddr: DestDescriptorAddr, DataPtr: LocalStorage, SrcDescriptorAddr: SrcElementAddr,
3626 DescriptorType: RI.ByRefAllocatedType, DataPtrPtrGen: RI.DataPtrPtrGen);
3627
3628 if (!GenResult)
3629 return GenResult.takeError();
3630 }
3631 } else {
3632 switch (RI.EvaluationKind) {
3633 case EvalKind::Scalar: {
3634 Value *Elem = Builder.CreateLoad(Ty: RI.ElementType, Ptr: SrcElementAddr);
3635 // Store the source element value to the dest element address.
3636 Builder.CreateStore(Val: Elem, Ptr: DestElementAddr);
3637 break;
3638 }
3639 case EvalKind::Complex: {
3640 Value *SrcRealPtr = Builder.CreateConstInBoundsGEP2_32(
3641 Ty: RI.ElementType, Ptr: SrcElementAddr, Idx0: 0, Idx1: 0, Name: ".realp");
3642 Value *SrcReal = Builder.CreateLoad(
3643 Ty: RI.ElementType->getStructElementType(N: 0), Ptr: SrcRealPtr, Name: ".real");
3644 Value *SrcImgPtr = Builder.CreateConstInBoundsGEP2_32(
3645 Ty: RI.ElementType, Ptr: SrcElementAddr, Idx0: 0, Idx1: 1, Name: ".imagp");
3646 Value *SrcImg = Builder.CreateLoad(
3647 Ty: RI.ElementType->getStructElementType(N: 1), Ptr: SrcImgPtr, Name: ".imag");
3648
3649 Value *DestRealPtr = Builder.CreateConstInBoundsGEP2_32(
3650 Ty: RI.ElementType, Ptr: DestElementAddr, Idx0: 0, Idx1: 0, Name: ".realp");
3651 Value *DestImgPtr = Builder.CreateConstInBoundsGEP2_32(
3652 Ty: RI.ElementType, Ptr: DestElementAddr, Idx0: 0, Idx1: 1, Name: ".imagp");
3653 Builder.CreateStore(Val: SrcReal, Ptr: DestRealPtr);
3654 Builder.CreateStore(Val: SrcImg, Ptr: DestImgPtr);
3655 break;
3656 }
3657 case EvalKind::Aggregate: {
3658 Value *SizeVal = Builder.getInt64(
3659 C: M.getDataLayout().getTypeStoreSize(Ty: RI.ElementType));
3660 Builder.CreateMemCpy(
3661 Dst: DestElementAddr, DstAlign: M.getDataLayout().getPrefTypeAlign(Ty: RI.ElementType),
3662 Src: SrcElementAddr, SrcAlign: M.getDataLayout().getPrefTypeAlign(Ty: RI.ElementType),
3663 Size: SizeVal, isVolatile: false);
3664 break;
3665 }
3666 };
3667 }
3668
3669 // Step 3.1: Modify reference in dest Reduce list as needed.
3670 // Modifying the reference in Reduce list to point to the newly
3671 // created element. The element is live in the current function
3672 // scope and that of functions it invokes (i.e., reduce_function).
3673 // RemoteReduceData[i] = (void*)&RemoteElem
3674 if (UpdateDestListPtr) {
3675 Value *CastDestAddr = Builder.CreatePointerBitCastOrAddrSpaceCast(
3676 V: DestElementAddr, DestTy: Builder.getPtrTy(),
3677 Name: DestElementAddr->getName() + ".ascast");
3678 Builder.CreateStore(Val: CastDestAddr, Ptr: DestElementPtrAddr);
3679 }
3680 }
3681
3682 return Error::success();
3683}
3684
3685Expected<Function *> OpenMPIRBuilder::emitInterWarpCopyFunction(
3686 const LocationDescription &Loc, ArrayRef<ReductionInfo> ReductionInfos,
3687 AttributeList FuncAttrs, ArrayRef<bool> IsByRef) {
3688 IRBuilder<>::InsertPointGuard IPG(Builder);
3689 LLVMContext &Ctx = M.getContext();
3690 FunctionType *FuncTy = FunctionType::get(
3691 Result: Builder.getVoidTy(), Params: {Builder.getPtrTy(), Builder.getInt32Ty()},
3692 /* IsVarArg */ isVarArg: false);
3693 Function *WcFunc =
3694 Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage,
3695 N: "_omp_reduction_inter_warp_copy_func", M: &M);
3696 WcFunc->setCallingConv(Config.getRuntimeCC());
3697 WcFunc->setAttributes(FuncAttrs);
3698 WcFunc->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
3699 WcFunc->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
3700 BasicBlock *EntryBB = BasicBlock::Create(Context&: M.getContext(), Name: "entry", Parent: WcFunc);
3701 Builder.SetInsertPoint(EntryBB);
3702 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
3703
3704 // ReduceList: thread local Reduce list.
3705 // At the stage of the computation when this function is called, partially
3706 // aggregated values reside in the first lane of every active warp.
3707 Argument *ReduceListArg = WcFunc->getArg(i: 0);
3708 // NumWarps: number of warps active in the parallel region. This could
3709 // be smaller than 32 (max warps in a CTA) for partial block reduction.
3710 Argument *NumWarpsArg = WcFunc->getArg(i: 1);
3711
3712 // This array is used as a medium to transfer, one reduce element at a time,
3713 // the data from the first lane of every warp to lanes in the first warp
3714 // in order to perform the final step of a reduction in a parallel region
3715 // (reduction across warps). The array is placed in NVPTX __shared__ memory
3716 // for reduced latency, as well as to have a distinct copy for concurrently
3717 // executing target regions. The array is declared with common linkage so
3718 // as to be shared across compilation units.
3719 StringRef TransferMediumName =
3720 "__openmp_nvptx_data_transfer_temporary_storage";
3721 GlobalVariable *TransferMedium = M.getGlobalVariable(Name: TransferMediumName);
3722 unsigned WarpSize = Config.getGridValue().GV_Warp_Size;
3723 ArrayType *ArrayTy = ArrayType::get(ElementType: Builder.getInt32Ty(), NumElements: WarpSize);
3724 if (!TransferMedium) {
3725 TransferMedium = new GlobalVariable(
3726 M, ArrayTy, /*isConstant=*/false, GlobalVariable::WeakAnyLinkage,
3727 UndefValue::get(T: ArrayTy), TransferMediumName,
3728 /*InsertBefore=*/nullptr, GlobalVariable::NotThreadLocal,
3729 /*AddressSpace=*/3);
3730 }
3731
3732 // Get the CUDA thread id of the current OpenMP thread on the GPU.
3733 Value *GPUThreadID = getGPUThreadID();
3734 // nvptx_lane_id = nvptx_id % warpsize
3735 Value *LaneID = getNVPTXLaneID();
3736 // nvptx_warp_id = nvptx_id / warpsize
3737 Value *WarpID = getNVPTXWarpID();
3738
3739 InsertPointTy AllocaIP = Builder.GetInsertBlock()->getFirstInsertionPt();
3740 Type *Arg0Type = ReduceListArg->getType();
3741 Type *Arg1Type = NumWarpsArg->getType();
3742 Builder.restoreIP(IP: AllocaIP);
3743 AllocaInst *ReduceListAlloca = Builder.CreateAlloca(
3744 Ty: Arg0Type, ArraySize: nullptr, Name: ReduceListArg->getName() + ".addr");
3745 AllocaInst *NumWarpsAlloca =
3746 Builder.CreateAlloca(Ty: Arg1Type, ArraySize: nullptr, Name: NumWarpsArg->getName() + ".addr");
3747 Value *ReduceListAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
3748 V: ReduceListAlloca, DestTy: Arg0Type, Name: ReduceListAlloca->getName() + ".ascast");
3749 Value *NumWarpsAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
3750 V: NumWarpsAlloca, DestTy: Builder.getPtrTy(AddrSpace: 0),
3751 Name: NumWarpsAlloca->getName() + ".ascast");
3752 Builder.CreateStore(Val: ReduceListArg, Ptr: ReduceListAddrCast);
3753 Builder.CreateStore(Val: NumWarpsArg, Ptr: NumWarpsAddrCast);
3754 AllocaIP = getInsertPointAfterInstr(I: NumWarpsAlloca);
3755 InsertPointTy CodeGenIP =
3756 getInsertPointAfterInstr(I: &Builder.GetInsertBlock()->back());
3757 Builder.restoreIP(IP: CodeGenIP);
3758
3759 Value *ReduceList =
3760 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ReduceListAddrCast);
3761
3762 for (auto En : enumerate(First&: ReductionInfos)) {
3763 //
3764 // Warp master copies reduce element to transfer medium in __shared__
3765 // memory.
3766 //
3767 const ReductionInfo &RI = En.value();
3768 bool IsByRefElem = !IsByRef.empty() && IsByRef[En.index()];
3769 unsigned RealTySize = M.getDataLayout().getTypeAllocSize(
3770 Ty: IsByRefElem ? RI.ByRefElementType : RI.ElementType);
3771 for (unsigned TySize = 4; TySize > 0 && RealTySize > 0; TySize /= 2) {
3772 Type *CType = Builder.getIntNTy(N: TySize * 8);
3773
3774 unsigned NumIters = RealTySize / TySize;
3775 if (NumIters == 0)
3776 continue;
3777 Value *Cnt = nullptr;
3778 Value *CntAddr = nullptr;
3779 BasicBlock *PrecondBB = nullptr;
3780 BasicBlock *ExitBB = nullptr;
3781 if (NumIters > 1) {
3782 CodeGenIP = Builder.saveIP();
3783 Builder.restoreIP(IP: AllocaIP);
3784 CntAddr =
3785 Builder.CreateAlloca(Ty: Builder.getInt32Ty(), ArraySize: nullptr, Name: ".cnt.addr");
3786
3787 CntAddr = Builder.CreateAddrSpaceCast(V: CntAddr, DestTy: Builder.getPtrTy(),
3788 Name: CntAddr->getName() + ".ascast");
3789 Builder.restoreIP(IP: CodeGenIP);
3790 Builder.CreateStore(Val: Constant::getNullValue(Ty: Builder.getInt32Ty()),
3791 Ptr: CntAddr,
3792 /*Volatile=*/isVolatile: false);
3793 PrecondBB = BasicBlock::Create(Context&: Ctx, Name: "precond");
3794 ExitBB = BasicBlock::Create(Context&: Ctx, Name: "exit");
3795 BasicBlock *BodyBB = BasicBlock::Create(Context&: Ctx, Name: "body");
3796 emitBlock(BB: PrecondBB, CurFn: Builder.GetInsertBlock()->getParent());
3797 Cnt = Builder.CreateLoad(Ty: Builder.getInt32Ty(), Ptr: CntAddr,
3798 /*Volatile=*/isVolatile: false);
3799 Value *Cmp = Builder.CreateICmpULT(
3800 LHS: Cnt, RHS: ConstantInt::get(Ty: Builder.getInt32Ty(), V: NumIters));
3801 Builder.CreateCondBr(Cond: Cmp, True: BodyBB, False: ExitBB);
3802 emitBlock(BB: BodyBB, CurFn: Builder.GetInsertBlock()->getParent());
3803 }
3804
3805 // kmpc_barrier.
3806 InsertPointOrErrorTy BarrierIP1 =
3807 createBarrier(Loc: LocationDescription(Builder.saveIP(), DebugLoc()),
3808 Kind: omp::Directive::OMPD_unknown,
3809 /* ForceSimpleCall */ false,
3810 /* CheckCancelFlag */ true);
3811 if (!BarrierIP1)
3812 return BarrierIP1.takeError();
3813 BasicBlock *ThenBB = BasicBlock::Create(Context&: Ctx, Name: "then");
3814 BasicBlock *ElseBB = BasicBlock::Create(Context&: Ctx, Name: "else");
3815 BasicBlock *MergeBB = BasicBlock::Create(Context&: Ctx, Name: "ifcont");
3816
3817 // if (lane_id == 0)
3818 Value *IsWarpMaster = Builder.CreateIsNull(Arg: LaneID, Name: "warp_master");
3819 Builder.CreateCondBr(Cond: IsWarpMaster, True: ThenBB, False: ElseBB);
3820 emitBlock(BB: ThenBB, CurFn: Builder.GetInsertBlock()->getParent());
3821
3822 // Reduce element = LocalReduceList[i]
3823 auto *RedListArrayTy =
3824 ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: ReductionInfos.size());
3825 Type *IndexTy = Builder.getIndexTy(
3826 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
3827 Value *ElemPtrPtr =
3828 Builder.CreateInBoundsGEP(Ty: RedListArrayTy, Ptr: ReduceList,
3829 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0),
3830 ConstantInt::get(Ty: IndexTy, V: En.index())});
3831 // elemptr = ((CopyType*)(elemptrptr)) + I
3832 Value *ElemPtr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ElemPtrPtr);
3833
3834 if (IsByRefElem && RI.DataPtrPtrGen) {
3835 InsertPointOrErrorTy GenRes =
3836 RI.DataPtrPtrGen(Builder.saveIP(), ElemPtr, ElemPtr);
3837
3838 if (!GenRes)
3839 return GenRes.takeError();
3840
3841 ElemPtr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ElemPtr);
3842 }
3843
3844 if (NumIters > 1)
3845 ElemPtr = Builder.CreateGEP(Ty: Builder.getInt32Ty(), Ptr: ElemPtr, IdxList: Cnt);
3846
3847 // Get pointer to location in transfer medium.
3848 // MediumPtr = &medium[warp_id]
3849 Value *MediumPtr = Builder.CreateInBoundsGEP(
3850 Ty: ArrayTy, Ptr: TransferMedium, IdxList: {Builder.getInt64(C: 0), WarpID});
3851 // elem = *elemptr
3852 //*MediumPtr = elem
3853 Value *Elem = Builder.CreateLoad(Ty: CType, Ptr: ElemPtr);
3854 // Store the source element value to the dest element address.
3855 Builder.CreateStore(Val: Elem, Ptr: MediumPtr,
3856 /*IsVolatile*/ isVolatile: true);
3857 Builder.CreateBr(Dest: MergeBB);
3858
3859 // else
3860 emitBlock(BB: ElseBB, CurFn: Builder.GetInsertBlock()->getParent());
3861 Builder.CreateBr(Dest: MergeBB);
3862
3863 // endif
3864 emitBlock(BB: MergeBB, CurFn: Builder.GetInsertBlock()->getParent());
3865 InsertPointOrErrorTy BarrierIP2 =
3866 createBarrier(Loc: LocationDescription(Builder.saveIP(), DebugLoc()),
3867 Kind: omp::Directive::OMPD_unknown,
3868 /* ForceSimpleCall */ false,
3869 /* CheckCancelFlag */ true);
3870 if (!BarrierIP2)
3871 return BarrierIP2.takeError();
3872
3873 // Warp 0 copies reduce element from transfer medium
3874 BasicBlock *W0ThenBB = BasicBlock::Create(Context&: Ctx, Name: "then");
3875 BasicBlock *W0ElseBB = BasicBlock::Create(Context&: Ctx, Name: "else");
3876 BasicBlock *W0MergeBB = BasicBlock::Create(Context&: Ctx, Name: "ifcont");
3877
3878 Value *NumWarpsVal =
3879 Builder.CreateLoad(Ty: Builder.getInt32Ty(), Ptr: NumWarpsAddrCast);
3880 // Up to 32 threads in warp 0 are active.
3881 Value *IsActiveThread =
3882 Builder.CreateICmpULT(LHS: GPUThreadID, RHS: NumWarpsVal, Name: "is_active_thread");
3883 Builder.CreateCondBr(Cond: IsActiveThread, True: W0ThenBB, False: W0ElseBB);
3884
3885 emitBlock(BB: W0ThenBB, CurFn: Builder.GetInsertBlock()->getParent());
3886
3887 // SecMediumPtr = &medium[tid]
3888 // SrcMediumVal = *SrcMediumPtr
3889 Value *SrcMediumPtrVal = Builder.CreateInBoundsGEP(
3890 Ty: ArrayTy, Ptr: TransferMedium, IdxList: {Builder.getInt64(C: 0), GPUThreadID});
3891 // TargetElemPtr = (CopyType*)(SrcDataAddr[i]) + I
3892 Value *TargetElemPtrPtr =
3893 Builder.CreateInBoundsGEP(Ty: RedListArrayTy, Ptr: ReduceList,
3894 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0),
3895 ConstantInt::get(Ty: IndexTy, V: En.index())});
3896 Value *TargetElemPtrVal =
3897 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: TargetElemPtrPtr);
3898 Value *TargetElemPtr = TargetElemPtrVal;
3899
3900 if (IsByRefElem && RI.DataPtrPtrGen) {
3901 InsertPointOrErrorTy GenRes =
3902 RI.DataPtrPtrGen(Builder.saveIP(), TargetElemPtr, TargetElemPtr);
3903
3904 if (!GenRes)
3905 return GenRes.takeError();
3906
3907 TargetElemPtr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: TargetElemPtr);
3908 }
3909
3910 if (NumIters > 1)
3911 TargetElemPtr =
3912 Builder.CreateGEP(Ty: Builder.getInt32Ty(), Ptr: TargetElemPtr, IdxList: Cnt);
3913
3914 // *TargetElemPtr = SrcMediumVal;
3915 Value *SrcMediumValue =
3916 Builder.CreateLoad(Ty: CType, Ptr: SrcMediumPtrVal, /*IsVolatile*/ isVolatile: true);
3917 Builder.CreateStore(Val: SrcMediumValue, Ptr: TargetElemPtr);
3918 Builder.CreateBr(Dest: W0MergeBB);
3919
3920 emitBlock(BB: W0ElseBB, CurFn: Builder.GetInsertBlock()->getParent());
3921 Builder.CreateBr(Dest: W0MergeBB);
3922
3923 emitBlock(BB: W0MergeBB, CurFn: Builder.GetInsertBlock()->getParent());
3924
3925 if (NumIters > 1) {
3926 Cnt = Builder.CreateNSWAdd(
3927 LHS: Cnt, RHS: ConstantInt::get(Ty: Builder.getInt32Ty(), /*V=*/1));
3928 Builder.CreateStore(Val: Cnt, Ptr: CntAddr, /*Volatile=*/isVolatile: false);
3929
3930 auto *CurFn = Builder.GetInsertBlock()->getParent();
3931 emitBranch(Target: PrecondBB);
3932 emitBlock(BB: ExitBB, CurFn);
3933 }
3934 RealTySize %= TySize;
3935 }
3936 }
3937
3938 Builder.CreateRetVoid();
3939
3940 return WcFunc;
3941}
3942
3943Expected<Function *> OpenMPIRBuilder::emitShuffleAndReduceFunction(
3944 ArrayRef<ReductionInfo> ReductionInfos, Function *ReduceFn,
3945 AttributeList FuncAttrs, ArrayRef<bool> IsByRef) {
3946 LLVMContext &Ctx = M.getContext();
3947 IRBuilder<>::InsertPointGuard IPG(Builder);
3948 FunctionType *FuncTy =
3949 FunctionType::get(Result: Builder.getVoidTy(),
3950 Params: {Builder.getPtrTy(), Builder.getInt16Ty(),
3951 Builder.getInt16Ty(), Builder.getInt16Ty()},
3952 /* IsVarArg */ isVarArg: false);
3953 Function *SarFunc =
3954 Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage,
3955 N: "_omp_reduction_shuffle_and_reduce_func", M: &M);
3956 SarFunc->setCallingConv(Config.getRuntimeCC());
3957 SarFunc->setAttributes(FuncAttrs);
3958 SarFunc->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
3959 SarFunc->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
3960 SarFunc->addParamAttr(ArgNo: 2, Kind: Attribute::NoUndef);
3961 SarFunc->addParamAttr(ArgNo: 3, Kind: Attribute::NoUndef);
3962 SarFunc->addParamAttr(ArgNo: 1, Kind: Attribute::SExt);
3963 SarFunc->addParamAttr(ArgNo: 2, Kind: Attribute::SExt);
3964 SarFunc->addParamAttr(ArgNo: 3, Kind: Attribute::SExt);
3965 BasicBlock *EntryBB = BasicBlock::Create(Context&: M.getContext(), Name: "entry", Parent: SarFunc);
3966 Builder.SetInsertPoint(EntryBB);
3967 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
3968
3969 // Thread local Reduce list used to host the values of data to be reduced.
3970 Argument *ReduceListArg = SarFunc->getArg(i: 0);
3971 // Current lane id; could be logical.
3972 Argument *LaneIDArg = SarFunc->getArg(i: 1);
3973 // Offset of the remote source lane relative to the current lane.
3974 Argument *RemoteLaneOffsetArg = SarFunc->getArg(i: 2);
3975 // Algorithm version. This is expected to be known at compile time.
3976 Argument *AlgoVerArg = SarFunc->getArg(i: 3);
3977
3978 Type *ReduceListArgType = ReduceListArg->getType();
3979 Type *LaneIDArgType = LaneIDArg->getType();
3980 Type *LaneIDArgPtrType = Builder.getPtrTy(AddrSpace: 0);
3981 Value *ReduceListAlloca = Builder.CreateAlloca(
3982 Ty: ReduceListArgType, ArraySize: nullptr, Name: ReduceListArg->getName() + ".addr");
3983 Value *LaneIdAlloca = Builder.CreateAlloca(Ty: LaneIDArgType, ArraySize: nullptr,
3984 Name: LaneIDArg->getName() + ".addr");
3985 Value *RemoteLaneOffsetAlloca = Builder.CreateAlloca(
3986 Ty: LaneIDArgType, ArraySize: nullptr, Name: RemoteLaneOffsetArg->getName() + ".addr");
3987 Value *AlgoVerAlloca = Builder.CreateAlloca(Ty: LaneIDArgType, ArraySize: nullptr,
3988 Name: AlgoVerArg->getName() + ".addr");
3989 ArrayType *RedListArrayTy =
3990 ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: ReductionInfos.size());
3991
3992 // Create a local thread-private variable to host the Reduce list
3993 // from a remote lane.
3994 Instruction *RemoteReductionListAlloca = Builder.CreateAlloca(
3995 Ty: RedListArrayTy, ArraySize: nullptr, Name: ".omp.reduction.remote_reduce_list");
3996
3997 Value *ReduceListAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
3998 V: ReduceListAlloca, DestTy: ReduceListArgType,
3999 Name: ReduceListAlloca->getName() + ".ascast");
4000 Value *LaneIdAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4001 V: LaneIdAlloca, DestTy: LaneIDArgPtrType, Name: LaneIdAlloca->getName() + ".ascast");
4002 Value *RemoteLaneOffsetAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4003 V: RemoteLaneOffsetAlloca, DestTy: LaneIDArgPtrType,
4004 Name: RemoteLaneOffsetAlloca->getName() + ".ascast");
4005 Value *AlgoVerAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4006 V: AlgoVerAlloca, DestTy: LaneIDArgPtrType, Name: AlgoVerAlloca->getName() + ".ascast");
4007 Value *RemoteListAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4008 V: RemoteReductionListAlloca, DestTy: Builder.getPtrTy(),
4009 Name: RemoteReductionListAlloca->getName() + ".ascast");
4010
4011 Builder.CreateStore(Val: ReduceListArg, Ptr: ReduceListAddrCast);
4012 Builder.CreateStore(Val: LaneIDArg, Ptr: LaneIdAddrCast);
4013 Builder.CreateStore(Val: RemoteLaneOffsetArg, Ptr: RemoteLaneOffsetAddrCast);
4014 Builder.CreateStore(Val: AlgoVerArg, Ptr: AlgoVerAddrCast);
4015
4016 Value *ReduceList = Builder.CreateLoad(Ty: ReduceListArgType, Ptr: ReduceListAddrCast);
4017 Value *LaneId = Builder.CreateLoad(Ty: LaneIDArgType, Ptr: LaneIdAddrCast);
4018 Value *RemoteLaneOffset =
4019 Builder.CreateLoad(Ty: LaneIDArgType, Ptr: RemoteLaneOffsetAddrCast);
4020 Value *AlgoVer = Builder.CreateLoad(Ty: LaneIDArgType, Ptr: AlgoVerAddrCast);
4021
4022 InsertPointTy AllocaIP = getInsertPointAfterInstr(I: RemoteReductionListAlloca);
4023
4024 // This loop iterates through the list of reduce elements and copies,
4025 // element by element, from a remote lane in the warp to RemoteReduceList,
4026 // hosted on the thread's stack.
4027 Error EmitRedLsCpRes = emitReductionListCopy(
4028 AllocaIP, Action: CopyAction::RemoteLaneToThread, ReductionArrayTy: RedListArrayTy, ReductionInfos,
4029 SrcBase: ReduceList, DestBase: RemoteListAddrCast, IsByRef,
4030 CopyOptions: {.RemoteLaneOffset: RemoteLaneOffset, .ScratchpadIndex: nullptr, .ScratchpadWidth: nullptr});
4031
4032 if (EmitRedLsCpRes)
4033 return EmitRedLsCpRes;
4034
4035 // The actions to be performed on the Remote Reduce list is dependent
4036 // on the algorithm version.
4037 //
4038 // if (AlgoVer==0) || (AlgoVer==1 && (LaneId < Offset)) || (AlgoVer==2 &&
4039 // LaneId % 2 == 0 && Offset > 0):
4040 // do the reduction value aggregation
4041 //
4042 // The thread local variable Reduce list is mutated in place to host the
4043 // reduced data, which is the aggregated value produced from local and
4044 // remote lanes.
4045 //
4046 // Note that AlgoVer is expected to be a constant integer known at compile
4047 // time.
4048 // When AlgoVer==0, the first conjunction evaluates to true, making
4049 // the entire predicate true during compile time.
4050 // When AlgoVer==1, the second conjunction has only the second part to be
4051 // evaluated during runtime. Other conjunctions evaluates to false
4052 // during compile time.
4053 // When AlgoVer==2, the third conjunction has only the second part to be
4054 // evaluated during runtime. Other conjunctions evaluates to false
4055 // during compile time.
4056 Value *CondAlgo0 = Builder.CreateIsNull(Arg: AlgoVer);
4057 Value *Algo1 = Builder.CreateICmpEQ(LHS: AlgoVer, RHS: Builder.getInt16(C: 1));
4058 Value *LaneComp = Builder.CreateICmpULT(LHS: LaneId, RHS: RemoteLaneOffset);
4059 Value *CondAlgo1 = Builder.CreateAnd(LHS: Algo1, RHS: LaneComp);
4060 Value *Algo2 = Builder.CreateICmpEQ(LHS: AlgoVer, RHS: Builder.getInt16(C: 2));
4061 Value *LaneIdAnd1 = Builder.CreateAnd(LHS: LaneId, RHS: Builder.getInt16(C: 1));
4062 Value *LaneIdComp = Builder.CreateIsNull(Arg: LaneIdAnd1);
4063 Value *Algo2AndLaneIdComp = Builder.CreateAnd(LHS: Algo2, RHS: LaneIdComp);
4064 Value *RemoteOffsetComp =
4065 Builder.CreateICmpSGT(LHS: RemoteLaneOffset, RHS: Builder.getInt16(C: 0));
4066 Value *CondAlgo2 = Builder.CreateAnd(LHS: Algo2AndLaneIdComp, RHS: RemoteOffsetComp);
4067 Value *CA0OrCA1 = Builder.CreateOr(LHS: CondAlgo0, RHS: CondAlgo1);
4068 Value *CondReduce = Builder.CreateOr(LHS: CA0OrCA1, RHS: CondAlgo2);
4069
4070 BasicBlock *ThenBB = BasicBlock::Create(Context&: Ctx, Name: "then");
4071 BasicBlock *ElseBB = BasicBlock::Create(Context&: Ctx, Name: "else");
4072 BasicBlock *MergeBB = BasicBlock::Create(Context&: Ctx, Name: "ifcont");
4073
4074 Builder.CreateCondBr(Cond: CondReduce, True: ThenBB, False: ElseBB);
4075 emitBlock(BB: ThenBB, CurFn: Builder.GetInsertBlock()->getParent());
4076 Value *LocalReduceListPtr = Builder.CreatePointerBitCastOrAddrSpaceCast(
4077 V: ReduceList, DestTy: Builder.getPtrTy());
4078 Value *RemoteReduceListPtr = Builder.CreatePointerBitCastOrAddrSpaceCast(
4079 V: RemoteListAddrCast, DestTy: Builder.getPtrTy());
4080 createRuntimeFunctionCall(Callee: ReduceFn, Args: {LocalReduceListPtr, RemoteReduceListPtr})
4081 ->addFnAttr(Kind: Attribute::NoUnwind);
4082 Builder.CreateBr(Dest: MergeBB);
4083
4084 emitBlock(BB: ElseBB, CurFn: Builder.GetInsertBlock()->getParent());
4085 Builder.CreateBr(Dest: MergeBB);
4086
4087 emitBlock(BB: MergeBB, CurFn: Builder.GetInsertBlock()->getParent());
4088
4089 // if (AlgoVer==1 && (LaneId >= Offset)) copy Remote Reduce list to local
4090 // Reduce list.
4091 Algo1 = Builder.CreateICmpEQ(LHS: AlgoVer, RHS: Builder.getInt16(C: 1));
4092 Value *LaneIdGtOffset = Builder.CreateICmpUGE(LHS: LaneId, RHS: RemoteLaneOffset);
4093 Value *CondCopy = Builder.CreateAnd(LHS: Algo1, RHS: LaneIdGtOffset);
4094
4095 BasicBlock *CpyThenBB = BasicBlock::Create(Context&: Ctx, Name: "then");
4096 BasicBlock *CpyElseBB = BasicBlock::Create(Context&: Ctx, Name: "else");
4097 BasicBlock *CpyMergeBB = BasicBlock::Create(Context&: Ctx, Name: "ifcont");
4098 Builder.CreateCondBr(Cond: CondCopy, True: CpyThenBB, False: CpyElseBB);
4099
4100 emitBlock(BB: CpyThenBB, CurFn: Builder.GetInsertBlock()->getParent());
4101
4102 EmitRedLsCpRes = emitReductionListCopy(
4103 AllocaIP, Action: CopyAction::ThreadCopy, ReductionArrayTy: RedListArrayTy, ReductionInfos,
4104 SrcBase: RemoteListAddrCast, DestBase: ReduceList, IsByRef);
4105
4106 if (EmitRedLsCpRes)
4107 return EmitRedLsCpRes;
4108
4109 Builder.CreateBr(Dest: CpyMergeBB);
4110
4111 emitBlock(BB: CpyElseBB, CurFn: Builder.GetInsertBlock()->getParent());
4112 Builder.CreateBr(Dest: CpyMergeBB);
4113
4114 emitBlock(BB: CpyMergeBB, CurFn: Builder.GetInsertBlock()->getParent());
4115
4116 Builder.CreateRetVoid();
4117
4118 return SarFunc;
4119}
4120
4121OpenMPIRBuilder::InsertPointOrErrorTy
4122OpenMPIRBuilder::generateReductionDescriptor(
4123 Value *DescriptorAddr, Value *DataPtr, Value *SrcDescriptorAddr,
4124 Type *DescriptorType,
4125 function_ref<InsertPointOrErrorTy(InsertPointTy, Value *, Value *&)>
4126 DataPtrPtrGen) {
4127
4128 // Copy the source descriptor to preserve all metadata (rank, extents,
4129 // strides, etc.)
4130 Value *DescriptorSize =
4131 Builder.getInt64(C: M.getDataLayout().getTypeStoreSize(Ty: DescriptorType));
4132 Builder.CreateMemCpy(
4133 Dst: DescriptorAddr, DstAlign: M.getDataLayout().getPrefTypeAlign(Ty: DescriptorType),
4134 Src: SrcDescriptorAddr, SrcAlign: M.getDataLayout().getPrefTypeAlign(Ty: DescriptorType),
4135 Size: DescriptorSize);
4136
4137 // Update the base pointer field to point to the local shuffled data
4138 Value *DataPtrField;
4139 InsertPointOrErrorTy GenResult =
4140 DataPtrPtrGen(Builder.saveIP(), DescriptorAddr, DataPtrField);
4141
4142 if (!GenResult)
4143 return GenResult.takeError();
4144
4145 Builder.CreateStore(Val: Builder.CreatePointerBitCastOrAddrSpaceCast(
4146 V: DataPtr, DestTy: Builder.getPtrTy(), Name: ".ascast"),
4147 Ptr: DataPtrField);
4148
4149 return Builder.saveIP();
4150}
4151
4152Expected<Value *> OpenMPIRBuilder::createReductionDescriptorCopy(
4153 InsertPointTy AllocaIP, const ReductionInfo &RI, Value *DataPtr,
4154 Value *SrcDescriptorAddr, Type *DescriptorPtrTy, const Twine &Name) {
4155 InsertPointTy OldIP = Builder.saveIP();
4156 Builder.restoreIP(IP: AllocaIP);
4157
4158 AllocaInst *DescriptorAlloca =
4159 Builder.CreateAlloca(Ty: RI.ByRefAllocatedType, ArraySize: nullptr, Name);
4160 DescriptorAlloca->setAlignment(
4161 M.getDataLayout().getPrefTypeAlign(Ty: RI.ByRefAllocatedType));
4162 Value *DescriptorAddr = Builder.CreatePointerBitCastOrAddrSpaceCast(
4163 V: DescriptorAlloca, DestTy: DescriptorPtrTy,
4164 Name: DescriptorAlloca->getName() + ".ascast");
4165
4166 Builder.restoreIP(IP: OldIP);
4167
4168 InsertPointOrErrorTy GenResult =
4169 generateReductionDescriptor(DescriptorAddr, DataPtr, SrcDescriptorAddr,
4170 DescriptorType: RI.ByRefAllocatedType, DataPtrPtrGen: RI.DataPtrPtrGen);
4171 if (!GenResult)
4172 return GenResult.takeError();
4173
4174 return DescriptorAddr;
4175}
4176
4177Expected<Function *> OpenMPIRBuilder::emitListToGlobalCopyFunction(
4178 ArrayRef<ReductionInfo> ReductionInfos, Type *ReductionsBufferTy,
4179 AttributeList FuncAttrs, ArrayRef<bool> IsByRef) {
4180 IRBuilder<>::InsertPointGuard IPG(Builder);
4181 LLVMContext &Ctx = M.getContext();
4182 FunctionType *FuncTy = FunctionType::get(
4183 Result: Builder.getVoidTy(),
4184 Params: {Builder.getPtrTy(), Builder.getInt32Ty(), Builder.getPtrTy()},
4185 /* IsVarArg */ isVarArg: false);
4186 Function *LtGCFunc =
4187 Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage,
4188 N: "_omp_reduction_list_to_global_copy_func", M: &M);
4189 LtGCFunc->setAttributes(FuncAttrs);
4190 LtGCFunc->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
4191 LtGCFunc->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
4192 LtGCFunc->addParamAttr(ArgNo: 2, Kind: Attribute::NoUndef);
4193
4194 BasicBlock *EntryBlock = BasicBlock::Create(Context&: Ctx, Name: "entry", Parent: LtGCFunc);
4195 Builder.SetInsertPoint(EntryBlock);
4196 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
4197
4198 // Buffer: global reduction buffer.
4199 Argument *BufferArg = LtGCFunc->getArg(i: 0);
4200 // Idx: index of the buffer.
4201 Argument *IdxArg = LtGCFunc->getArg(i: 1);
4202 // ReduceList: thread local Reduce list.
4203 Argument *ReduceListArg = LtGCFunc->getArg(i: 2);
4204
4205 Value *BufferArgAlloca = Builder.CreateAlloca(Ty: Builder.getPtrTy(), ArraySize: nullptr,
4206 Name: BufferArg->getName() + ".addr");
4207 Value *IdxArgAlloca = Builder.CreateAlloca(Ty: Builder.getInt32Ty(), ArraySize: nullptr,
4208 Name: IdxArg->getName() + ".addr");
4209 Value *ReduceListArgAlloca = Builder.CreateAlloca(
4210 Ty: Builder.getPtrTy(), ArraySize: nullptr, Name: ReduceListArg->getName() + ".addr");
4211 Value *BufferArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4212 V: BufferArgAlloca, DestTy: Builder.getPtrTy(),
4213 Name: BufferArgAlloca->getName() + ".ascast");
4214 Value *IdxArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4215 V: IdxArgAlloca, DestTy: Builder.getPtrTy(), Name: IdxArgAlloca->getName() + ".ascast");
4216 Value *ReduceListArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4217 V: ReduceListArgAlloca, DestTy: Builder.getPtrTy(),
4218 Name: ReduceListArgAlloca->getName() + ".ascast");
4219
4220 Builder.CreateStore(Val: BufferArg, Ptr: BufferArgAddrCast);
4221 Builder.CreateStore(Val: IdxArg, Ptr: IdxArgAddrCast);
4222 Builder.CreateStore(Val: ReduceListArg, Ptr: ReduceListArgAddrCast);
4223
4224 Value *LocalReduceList =
4225 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ReduceListArgAddrCast);
4226 Value *BufferArgVal =
4227 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BufferArgAddrCast);
4228 Value *Idxs[] = {Builder.CreateLoad(Ty: Builder.getInt32Ty(), Ptr: IdxArgAddrCast)};
4229 Type *IndexTy = Builder.getIndexTy(
4230 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
4231 for (auto En : enumerate(First&: ReductionInfos)) {
4232 const ReductionInfo &RI = En.value();
4233 auto *RedListArrayTy =
4234 ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: ReductionInfos.size());
4235 // Reduce element = LocalReduceList[i]
4236 Value *ElemPtrPtr = Builder.CreateInBoundsGEP(
4237 Ty: RedListArrayTy, Ptr: LocalReduceList,
4238 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
4239 // elemptr = ((CopyType*)(elemptrptr)) + I
4240 Value *ElemPtr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ElemPtrPtr);
4241
4242 // Global = Buffer.VD[Idx];
4243 Value *BufferVD =
4244 Builder.CreateInBoundsGEP(Ty: ReductionsBufferTy, Ptr: BufferArgVal, IdxList: Idxs);
4245 Value *GlobVal = Builder.CreateConstInBoundsGEP2_32(
4246 Ty: ReductionsBufferTy, Ptr: BufferVD, Idx0: 0, Idx1: En.index());
4247
4248 switch (RI.EvaluationKind) {
4249 case EvalKind::Scalar: {
4250 Value *TargetElement;
4251
4252 if (IsByRef.empty() || !IsByRef[En.index()]) {
4253 TargetElement = Builder.CreateLoad(Ty: RI.ElementType, Ptr: ElemPtr);
4254 } else {
4255 if (RI.DataPtrPtrGen) {
4256 InsertPointOrErrorTy GenResult =
4257 RI.DataPtrPtrGen(Builder.saveIP(), ElemPtr, ElemPtr);
4258
4259 if (!GenResult)
4260 return GenResult.takeError();
4261
4262 ElemPtr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ElemPtr);
4263 }
4264 TargetElement = Builder.CreateLoad(Ty: RI.ByRefElementType, Ptr: ElemPtr);
4265 }
4266
4267 Builder.CreateStore(Val: TargetElement, Ptr: GlobVal);
4268 break;
4269 }
4270 case EvalKind::Complex: {
4271 Value *SrcRealPtr = Builder.CreateConstInBoundsGEP2_32(
4272 Ty: RI.ElementType, Ptr: ElemPtr, Idx0: 0, Idx1: 0, Name: ".realp");
4273 Value *SrcReal = Builder.CreateLoad(
4274 Ty: RI.ElementType->getStructElementType(N: 0), Ptr: SrcRealPtr, Name: ".real");
4275 Value *SrcImgPtr = Builder.CreateConstInBoundsGEP2_32(
4276 Ty: RI.ElementType, Ptr: ElemPtr, Idx0: 0, Idx1: 1, Name: ".imagp");
4277 Value *SrcImg = Builder.CreateLoad(
4278 Ty: RI.ElementType->getStructElementType(N: 1), Ptr: SrcImgPtr, Name: ".imag");
4279
4280 Value *DestRealPtr = Builder.CreateConstInBoundsGEP2_32(
4281 Ty: RI.ElementType, Ptr: GlobVal, Idx0: 0, Idx1: 0, Name: ".realp");
4282 Value *DestImgPtr = Builder.CreateConstInBoundsGEP2_32(
4283 Ty: RI.ElementType, Ptr: GlobVal, Idx0: 0, Idx1: 1, Name: ".imagp");
4284 Builder.CreateStore(Val: SrcReal, Ptr: DestRealPtr);
4285 Builder.CreateStore(Val: SrcImg, Ptr: DestImgPtr);
4286 break;
4287 }
4288 case EvalKind::Aggregate: {
4289 Value *SizeVal =
4290 Builder.getInt64(C: M.getDataLayout().getTypeStoreSize(Ty: RI.ElementType));
4291 Builder.CreateMemCpy(
4292 Dst: GlobVal, DstAlign: M.getDataLayout().getPrefTypeAlign(Ty: RI.ElementType), Src: ElemPtr,
4293 SrcAlign: M.getDataLayout().getPrefTypeAlign(Ty: RI.ElementType), Size: SizeVal, isVolatile: false);
4294 break;
4295 }
4296 }
4297 }
4298
4299 Builder.CreateRetVoid();
4300 return LtGCFunc;
4301}
4302
4303Expected<Function *> OpenMPIRBuilder::emitListToGlobalReduceFunction(
4304 ArrayRef<ReductionInfo> ReductionInfos, Function *ReduceFn,
4305 Type *ReductionsBufferTy, AttributeList FuncAttrs, ArrayRef<bool> IsByRef) {
4306 IRBuilder<>::InsertPointGuard IPG(Builder);
4307 LLVMContext &Ctx = M.getContext();
4308 FunctionType *FuncTy = FunctionType::get(
4309 Result: Builder.getVoidTy(),
4310 Params: {Builder.getPtrTy(), Builder.getInt32Ty(), Builder.getPtrTy()},
4311 /* IsVarArg */ isVarArg: false);
4312 Function *LtGRFunc =
4313 Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage,
4314 N: "_omp_reduction_list_to_global_reduce_func", M: &M);
4315 LtGRFunc->setAttributes(FuncAttrs);
4316 LtGRFunc->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
4317 LtGRFunc->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
4318 LtGRFunc->addParamAttr(ArgNo: 2, Kind: Attribute::NoUndef);
4319
4320 BasicBlock *EntryBlock = BasicBlock::Create(Context&: Ctx, Name: "entry", Parent: LtGRFunc);
4321 Builder.SetInsertPoint(EntryBlock);
4322 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
4323
4324 // Buffer: global reduction buffer.
4325 Argument *BufferArg = LtGRFunc->getArg(i: 0);
4326 // Idx: index of the buffer.
4327 Argument *IdxArg = LtGRFunc->getArg(i: 1);
4328 // ReduceList: thread local Reduce list.
4329 Argument *ReduceListArg = LtGRFunc->getArg(i: 2);
4330
4331 Value *BufferArgAlloca = Builder.CreateAlloca(Ty: Builder.getPtrTy(), ArraySize: nullptr,
4332 Name: BufferArg->getName() + ".addr");
4333 Value *IdxArgAlloca = Builder.CreateAlloca(Ty: Builder.getInt32Ty(), ArraySize: nullptr,
4334 Name: IdxArg->getName() + ".addr");
4335 Value *ReduceListArgAlloca = Builder.CreateAlloca(
4336 Ty: Builder.getPtrTy(), ArraySize: nullptr, Name: ReduceListArg->getName() + ".addr");
4337 auto *RedListArrayTy =
4338 ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: ReductionInfos.size());
4339
4340 // 1. Build a list of reduction variables.
4341 // void *RedList[<n>] = {<ReductionVars>[0], ..., <ReductionVars>[<n>-1]};
4342 Value *LocalReduceList =
4343 Builder.CreateAlloca(Ty: RedListArrayTy, ArraySize: nullptr, Name: ".omp.reduction.red_list");
4344
4345 InsertPointTy AllocaIP(EntryBlock->begin());
4346
4347 Value *BufferArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4348 V: BufferArgAlloca, DestTy: Builder.getPtrTy(),
4349 Name: BufferArgAlloca->getName() + ".ascast");
4350 Value *IdxArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4351 V: IdxArgAlloca, DestTy: Builder.getPtrTy(), Name: IdxArgAlloca->getName() + ".ascast");
4352 Value *ReduceListArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4353 V: ReduceListArgAlloca, DestTy: Builder.getPtrTy(),
4354 Name: ReduceListArgAlloca->getName() + ".ascast");
4355 Value *LocalReduceListAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4356 V: LocalReduceList, DestTy: Builder.getPtrTy(),
4357 Name: LocalReduceList->getName() + ".ascast");
4358
4359 Builder.CreateStore(Val: BufferArg, Ptr: BufferArgAddrCast);
4360 Builder.CreateStore(Val: IdxArg, Ptr: IdxArgAddrCast);
4361 Builder.CreateStore(Val: ReduceListArg, Ptr: ReduceListArgAddrCast);
4362
4363 Value *BufferVal = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BufferArgAddrCast);
4364 Value *Idxs[] = {Builder.CreateLoad(Ty: Builder.getInt32Ty(), Ptr: IdxArgAddrCast)};
4365 Type *IndexTy = Builder.getIndexTy(
4366 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
4367 for (auto En : enumerate(First&: ReductionInfos)) {
4368 const ReductionInfo &RI = En.value();
4369
4370 Value *TargetElementPtrPtr = Builder.CreateInBoundsGEP(
4371 Ty: RedListArrayTy, Ptr: LocalReduceListAddrCast,
4372 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
4373 Value *BufferVD =
4374 Builder.CreateInBoundsGEP(Ty: ReductionsBufferTy, Ptr: BufferVal, IdxList: Idxs);
4375 // Global = Buffer.VD[Idx];
4376 Value *GlobValPtr = Builder.CreateConstInBoundsGEP2_32(
4377 Ty: ReductionsBufferTy, Ptr: BufferVD, Idx0: 0, Idx1: En.index());
4378
4379 if (!IsByRef.empty() && IsByRef[En.index()] && RI.DataPtrPtrGen) {
4380 // Get source descriptor from the reduce list argument
4381 Value *ReduceList =
4382 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ReduceListArgAddrCast);
4383 Value *SrcElementPtrPtr =
4384 Builder.CreateInBoundsGEP(Ty: RedListArrayTy, Ptr: ReduceList,
4385 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0),
4386 ConstantInt::get(Ty: IndexTy, V: En.index())});
4387 Value *SrcDescriptorAddr =
4388 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: SrcElementPtrPtr);
4389
4390 // Copy descriptor from source and update base_ptr to global buffer data
4391 Expected<Value *> ByRefAlloc = createReductionDescriptorCopy(
4392 AllocaIP, RI, DataPtr: GlobValPtr, SrcDescriptorAddr, DescriptorPtrTy: Builder.getPtrTy());
4393 if (!ByRefAlloc)
4394 return ByRefAlloc.takeError();
4395
4396 Builder.CreateStore(Val: *ByRefAlloc, Ptr: TargetElementPtrPtr);
4397 } else {
4398 Builder.CreateStore(Val: GlobValPtr, Ptr: TargetElementPtrPtr);
4399 }
4400 }
4401
4402 // Call reduce_function(GlobalReduceList, ReduceList)
4403 Value *ReduceList =
4404 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ReduceListArgAddrCast);
4405 createRuntimeFunctionCall(Callee: ReduceFn, Args: {LocalReduceListAddrCast, ReduceList})
4406 ->addFnAttr(Kind: Attribute::NoUnwind);
4407 Builder.CreateRetVoid();
4408 return LtGRFunc;
4409}
4410
4411Expected<Function *> OpenMPIRBuilder::emitGlobalToListCopyFunction(
4412 ArrayRef<ReductionInfo> ReductionInfos, Type *ReductionsBufferTy,
4413 AttributeList FuncAttrs, ArrayRef<bool> IsByRef) {
4414 IRBuilder<>::InsertPointGuard IPG(Builder);
4415 LLVMContext &Ctx = M.getContext();
4416 FunctionType *FuncTy = FunctionType::get(
4417 Result: Builder.getVoidTy(),
4418 Params: {Builder.getPtrTy(), Builder.getInt32Ty(), Builder.getPtrTy()},
4419 /* IsVarArg */ isVarArg: false);
4420 Function *GtLCFunc =
4421 Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage,
4422 N: "_omp_reduction_global_to_list_copy_func", M: &M);
4423 GtLCFunc->setAttributes(FuncAttrs);
4424 GtLCFunc->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
4425 GtLCFunc->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
4426 GtLCFunc->addParamAttr(ArgNo: 2, Kind: Attribute::NoUndef);
4427
4428 BasicBlock *EntryBlock = BasicBlock::Create(Context&: Ctx, Name: "entry", Parent: GtLCFunc);
4429 Builder.SetInsertPoint(EntryBlock);
4430 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
4431
4432 // Buffer: global reduction buffer.
4433 Argument *BufferArg = GtLCFunc->getArg(i: 0);
4434 // Idx: index of the buffer.
4435 Argument *IdxArg = GtLCFunc->getArg(i: 1);
4436 // ReduceList: thread local Reduce list.
4437 Argument *ReduceListArg = GtLCFunc->getArg(i: 2);
4438
4439 Value *BufferArgAlloca = Builder.CreateAlloca(Ty: Builder.getPtrTy(), ArraySize: nullptr,
4440 Name: BufferArg->getName() + ".addr");
4441 Value *IdxArgAlloca = Builder.CreateAlloca(Ty: Builder.getInt32Ty(), ArraySize: nullptr,
4442 Name: IdxArg->getName() + ".addr");
4443 Value *ReduceListArgAlloca = Builder.CreateAlloca(
4444 Ty: Builder.getPtrTy(), ArraySize: nullptr, Name: ReduceListArg->getName() + ".addr");
4445 Value *BufferArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4446 V: BufferArgAlloca, DestTy: Builder.getPtrTy(),
4447 Name: BufferArgAlloca->getName() + ".ascast");
4448 Value *IdxArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4449 V: IdxArgAlloca, DestTy: Builder.getPtrTy(), Name: IdxArgAlloca->getName() + ".ascast");
4450 Value *ReduceListArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4451 V: ReduceListArgAlloca, DestTy: Builder.getPtrTy(),
4452 Name: ReduceListArgAlloca->getName() + ".ascast");
4453 Builder.CreateStore(Val: BufferArg, Ptr: BufferArgAddrCast);
4454 Builder.CreateStore(Val: IdxArg, Ptr: IdxArgAddrCast);
4455 Builder.CreateStore(Val: ReduceListArg, Ptr: ReduceListArgAddrCast);
4456
4457 Value *LocalReduceList =
4458 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ReduceListArgAddrCast);
4459 Value *BufferVal = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BufferArgAddrCast);
4460 Value *Idxs[] = {Builder.CreateLoad(Ty: Builder.getInt32Ty(), Ptr: IdxArgAddrCast)};
4461 Type *IndexTy = Builder.getIndexTy(
4462 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
4463 for (auto En : enumerate(First&: ReductionInfos)) {
4464 const OpenMPIRBuilder::ReductionInfo &RI = En.value();
4465 auto *RedListArrayTy =
4466 ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: ReductionInfos.size());
4467 // Reduce element = LocalReduceList[i]
4468 Value *ElemPtrPtr = Builder.CreateInBoundsGEP(
4469 Ty: RedListArrayTy, Ptr: LocalReduceList,
4470 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
4471 // elemptr = ((CopyType*)(elemptrptr)) + I
4472 Value *ElemPtr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ElemPtrPtr);
4473 // Global = Buffer.VD[Idx];
4474 Value *BufferVD =
4475 Builder.CreateInBoundsGEP(Ty: ReductionsBufferTy, Ptr: BufferVal, IdxList: Idxs);
4476 Value *GlobValPtr = Builder.CreateConstInBoundsGEP2_32(
4477 Ty: ReductionsBufferTy, Ptr: BufferVD, Idx0: 0, Idx1: En.index());
4478
4479 switch (RI.EvaluationKind) {
4480 case EvalKind::Scalar: {
4481 Type *ElemType = RI.ElementType;
4482
4483 if (!IsByRef.empty() && IsByRef[En.index()]) {
4484 ElemType = RI.ByRefElementType;
4485 if (RI.DataPtrPtrGen) {
4486 InsertPointOrErrorTy GenResult =
4487 RI.DataPtrPtrGen(Builder.saveIP(), ElemPtr, ElemPtr);
4488
4489 if (!GenResult)
4490 return GenResult.takeError();
4491
4492 ElemPtr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ElemPtr);
4493 }
4494 }
4495
4496 Value *TargetElement = Builder.CreateLoad(Ty: ElemType, Ptr: GlobValPtr);
4497 Builder.CreateStore(Val: TargetElement, Ptr: ElemPtr);
4498 break;
4499 }
4500 case EvalKind::Complex: {
4501 Value *SrcRealPtr = Builder.CreateConstInBoundsGEP2_32(
4502 Ty: RI.ElementType, Ptr: GlobValPtr, Idx0: 0, Idx1: 0, Name: ".realp");
4503 Value *SrcReal = Builder.CreateLoad(
4504 Ty: RI.ElementType->getStructElementType(N: 0), Ptr: SrcRealPtr, Name: ".real");
4505 Value *SrcImgPtr = Builder.CreateConstInBoundsGEP2_32(
4506 Ty: RI.ElementType, Ptr: GlobValPtr, Idx0: 0, Idx1: 1, Name: ".imagp");
4507 Value *SrcImg = Builder.CreateLoad(
4508 Ty: RI.ElementType->getStructElementType(N: 1), Ptr: SrcImgPtr, Name: ".imag");
4509
4510 Value *DestRealPtr = Builder.CreateConstInBoundsGEP2_32(
4511 Ty: RI.ElementType, Ptr: ElemPtr, Idx0: 0, Idx1: 0, Name: ".realp");
4512 Value *DestImgPtr = Builder.CreateConstInBoundsGEP2_32(
4513 Ty: RI.ElementType, Ptr: ElemPtr, Idx0: 0, Idx1: 1, Name: ".imagp");
4514 Builder.CreateStore(Val: SrcReal, Ptr: DestRealPtr);
4515 Builder.CreateStore(Val: SrcImg, Ptr: DestImgPtr);
4516 break;
4517 }
4518 case EvalKind::Aggregate: {
4519 Value *SizeVal =
4520 Builder.getInt64(C: M.getDataLayout().getTypeStoreSize(Ty: RI.ElementType));
4521 Builder.CreateMemCpy(
4522 Dst: ElemPtr, DstAlign: M.getDataLayout().getPrefTypeAlign(Ty: RI.ElementType),
4523 Src: GlobValPtr, SrcAlign: M.getDataLayout().getPrefTypeAlign(Ty: RI.ElementType),
4524 Size: SizeVal, isVolatile: false);
4525 break;
4526 }
4527 }
4528 }
4529
4530 Builder.CreateRetVoid();
4531 return GtLCFunc;
4532}
4533
4534Expected<Function *> OpenMPIRBuilder::emitGlobalToListReduceFunction(
4535 ArrayRef<ReductionInfo> ReductionInfos, Function *ReduceFn,
4536 Type *ReductionsBufferTy, AttributeList FuncAttrs, ArrayRef<bool> IsByRef) {
4537 IRBuilder<>::InsertPointGuard IPG(Builder);
4538 LLVMContext &Ctx = M.getContext();
4539 auto *FuncTy = FunctionType::get(
4540 Result: Builder.getVoidTy(),
4541 Params: {Builder.getPtrTy(), Builder.getInt32Ty(), Builder.getPtrTy()},
4542 /* IsVarArg */ isVarArg: false);
4543 Function *GtLRFunc =
4544 Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage,
4545 N: "_omp_reduction_global_to_list_reduce_func", M: &M);
4546 GtLRFunc->setAttributes(FuncAttrs);
4547 GtLRFunc->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
4548 GtLRFunc->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
4549 GtLRFunc->addParamAttr(ArgNo: 2, Kind: Attribute::NoUndef);
4550
4551 BasicBlock *EntryBlock = BasicBlock::Create(Context&: Ctx, Name: "entry", Parent: GtLRFunc);
4552 Builder.SetInsertPoint(EntryBlock);
4553 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
4554
4555 // Buffer: global reduction buffer.
4556 Argument *BufferArg = GtLRFunc->getArg(i: 0);
4557 // Idx: index of the buffer.
4558 Argument *IdxArg = GtLRFunc->getArg(i: 1);
4559 // ReduceList: thread local Reduce list.
4560 Argument *ReduceListArg = GtLRFunc->getArg(i: 2);
4561
4562 Value *BufferArgAlloca = Builder.CreateAlloca(Ty: Builder.getPtrTy(), ArraySize: nullptr,
4563 Name: BufferArg->getName() + ".addr");
4564 Value *IdxArgAlloca = Builder.CreateAlloca(Ty: Builder.getInt32Ty(), ArraySize: nullptr,
4565 Name: IdxArg->getName() + ".addr");
4566 Value *ReduceListArgAlloca = Builder.CreateAlloca(
4567 Ty: Builder.getPtrTy(), ArraySize: nullptr, Name: ReduceListArg->getName() + ".addr");
4568 ArrayType *RedListArrayTy =
4569 ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: ReductionInfos.size());
4570
4571 // 1. Build a list of reduction variables.
4572 // void *RedList[<n>] = {<ReductionVars>[0], ..., <ReductionVars>[<n>-1]};
4573 Value *LocalReduceList =
4574 Builder.CreateAlloca(Ty: RedListArrayTy, ArraySize: nullptr, Name: ".omp.reduction.red_list");
4575
4576 InsertPointTy AllocaIP(EntryBlock->begin());
4577
4578 Value *BufferArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4579 V: BufferArgAlloca, DestTy: Builder.getPtrTy(),
4580 Name: BufferArgAlloca->getName() + ".ascast");
4581 Value *IdxArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4582 V: IdxArgAlloca, DestTy: Builder.getPtrTy(), Name: IdxArgAlloca->getName() + ".ascast");
4583 Value *ReduceListArgAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4584 V: ReduceListArgAlloca, DestTy: Builder.getPtrTy(),
4585 Name: ReduceListArgAlloca->getName() + ".ascast");
4586 Value *ReductionList = Builder.CreatePointerBitCastOrAddrSpaceCast(
4587 V: LocalReduceList, DestTy: Builder.getPtrTy(),
4588 Name: LocalReduceList->getName() + ".ascast");
4589
4590 Builder.CreateStore(Val: BufferArg, Ptr: BufferArgAddrCast);
4591 Builder.CreateStore(Val: IdxArg, Ptr: IdxArgAddrCast);
4592 Builder.CreateStore(Val: ReduceListArg, Ptr: ReduceListArgAddrCast);
4593
4594 Value *BufferVal = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BufferArgAddrCast);
4595 Value *Idxs[] = {Builder.CreateLoad(Ty: Builder.getInt32Ty(), Ptr: IdxArgAddrCast)};
4596 Type *IndexTy = Builder.getIndexTy(
4597 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
4598 for (auto En : enumerate(First&: ReductionInfos)) {
4599 const ReductionInfo &RI = En.value();
4600
4601 Value *TargetElementPtrPtr = Builder.CreateInBoundsGEP(
4602 Ty: RedListArrayTy, Ptr: ReductionList,
4603 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
4604 // Global = Buffer.VD[Idx];
4605 Value *BufferVD =
4606 Builder.CreateInBoundsGEP(Ty: ReductionsBufferTy, Ptr: BufferVal, IdxList: Idxs);
4607 Value *GlobValPtr = Builder.CreateConstInBoundsGEP2_32(
4608 Ty: ReductionsBufferTy, Ptr: BufferVD, Idx0: 0, Idx1: En.index());
4609
4610 if (!IsByRef.empty() && IsByRef[En.index()] && RI.DataPtrPtrGen) {
4611 // Get source descriptor from the reduce list
4612 Value *ReduceListVal =
4613 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ReduceListArgAddrCast);
4614 Value *SrcElementPtrPtr =
4615 Builder.CreateInBoundsGEP(Ty: RedListArrayTy, Ptr: ReduceListVal,
4616 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0),
4617 ConstantInt::get(Ty: IndexTy, V: En.index())});
4618 Value *SrcDescriptorAddr =
4619 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: SrcElementPtrPtr);
4620
4621 // Copy descriptor from source and update base_ptr to global buffer data
4622 Expected<Value *> ByRefAlloc = createReductionDescriptorCopy(
4623 AllocaIP, RI, DataPtr: GlobValPtr, SrcDescriptorAddr, DescriptorPtrTy: Builder.getPtrTy());
4624 if (!ByRefAlloc)
4625 return ByRefAlloc.takeError();
4626
4627 Builder.CreateStore(Val: *ByRefAlloc, Ptr: TargetElementPtrPtr);
4628 } else {
4629 Builder.CreateStore(Val: GlobValPtr, Ptr: TargetElementPtrPtr);
4630 }
4631 }
4632
4633 // Call reduce_function(ReduceList, GlobalReduceList)
4634 Value *ReduceList =
4635 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: ReduceListArgAddrCast);
4636 createRuntimeFunctionCall(Callee: ReduceFn, Args: {ReduceList, ReductionList})
4637 ->addFnAttr(Kind: Attribute::NoUnwind);
4638 Builder.CreateRetVoid();
4639 return GtLRFunc;
4640}
4641
4642std::string OpenMPIRBuilder::getReductionFuncName(StringRef Name) const {
4643 std::string Suffix =
4644 createPlatformSpecificName(Parts: {"omp", "reduction", "reduction_func"});
4645 return (Name + Suffix).str();
4646}
4647
4648Expected<Function *> OpenMPIRBuilder::createReductionFunction(
4649 StringRef ReducerName, ArrayRef<ReductionInfo> ReductionInfos,
4650 ArrayRef<bool> IsByRef, ReductionGenCBKind ReductionGenCBKind,
4651 AttributeList FuncAttrs) {
4652 IRBuilder<>::InsertPointGuard IPG(Builder);
4653 auto *FuncTy = FunctionType::get(Result: Builder.getVoidTy(),
4654 Params: {Builder.getPtrTy(), Builder.getPtrTy()},
4655 /* IsVarArg */ isVarArg: false);
4656 std::string Name = getReductionFuncName(Name: ReducerName);
4657 Function *ReductionFunc =
4658 Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage, N: Name, M: &M);
4659 ReductionFunc->setCallingConv(Config.getRuntimeCC());
4660 ReductionFunc->setAttributes(FuncAttrs);
4661 ReductionFunc->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
4662 ReductionFunc->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
4663 BasicBlock *EntryBB =
4664 BasicBlock::Create(Context&: M.getContext(), Name: "entry", Parent: ReductionFunc);
4665 Builder.SetInsertPoint(EntryBB);
4666 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
4667
4668 // Need to alloca memory here and deal with the pointers before getting
4669 // LHS/RHS pointers out
4670 Value *LHSArrayPtr = nullptr;
4671 Value *RHSArrayPtr = nullptr;
4672 Argument *Arg0 = ReductionFunc->getArg(i: 0);
4673 Argument *Arg1 = ReductionFunc->getArg(i: 1);
4674 Type *Arg0Type = Arg0->getType();
4675 Type *Arg1Type = Arg1->getType();
4676
4677 Value *LHSAlloca =
4678 Builder.CreateAlloca(Ty: Arg0Type, ArraySize: nullptr, Name: Arg0->getName() + ".addr");
4679 Value *RHSAlloca =
4680 Builder.CreateAlloca(Ty: Arg1Type, ArraySize: nullptr, Name: Arg1->getName() + ".addr");
4681 Value *LHSAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4682 V: LHSAlloca, DestTy: Arg0Type, Name: LHSAlloca->getName() + ".ascast");
4683 Value *RHSAddrCast = Builder.CreatePointerBitCastOrAddrSpaceCast(
4684 V: RHSAlloca, DestTy: Arg1Type, Name: RHSAlloca->getName() + ".ascast");
4685 Builder.CreateStore(Val: Arg0, Ptr: LHSAddrCast);
4686 Builder.CreateStore(Val: Arg1, Ptr: RHSAddrCast);
4687 LHSArrayPtr = Builder.CreateLoad(Ty: Arg0Type, Ptr: LHSAddrCast);
4688 RHSArrayPtr = Builder.CreateLoad(Ty: Arg1Type, Ptr: RHSAddrCast);
4689
4690 Type *RedArrayTy = ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: ReductionInfos.size());
4691 Type *IndexTy = Builder.getIndexTy(
4692 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
4693 SmallVector<Value *> LHSPtrs, RHSPtrs;
4694 for (auto En : enumerate(First&: ReductionInfos)) {
4695 const ReductionInfo &RI = En.value();
4696 Value *RHSI8PtrPtr = Builder.CreateInBoundsGEP(
4697 Ty: RedArrayTy, Ptr: RHSArrayPtr,
4698 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
4699 Value *RHSI8Ptr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: RHSI8PtrPtr);
4700 Value *RHSPtr = Builder.CreatePointerBitCastOrAddrSpaceCast(
4701 V: RHSI8Ptr, DestTy: RI.PrivateVariable->getType(),
4702 Name: RHSI8Ptr->getName() + ".ascast");
4703
4704 Value *LHSI8PtrPtr = Builder.CreateInBoundsGEP(
4705 Ty: RedArrayTy, Ptr: LHSArrayPtr,
4706 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
4707 Value *LHSI8Ptr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: LHSI8PtrPtr);
4708 Value *LHSPtr = Builder.CreatePointerBitCastOrAddrSpaceCast(
4709 V: LHSI8Ptr, DestTy: RI.Variable->getType(), Name: LHSI8Ptr->getName() + ".ascast");
4710
4711 if (ReductionGenCBKind == ReductionGenCBKind::Clang) {
4712 LHSPtrs.emplace_back(Args&: LHSPtr);
4713 RHSPtrs.emplace_back(Args&: RHSPtr);
4714 } else {
4715 Value *LHS = LHSPtr;
4716 Value *RHS = RHSPtr;
4717
4718 if (!IsByRef.empty() && !IsByRef[En.index()]) {
4719 LHS = Builder.CreateLoad(Ty: RI.ElementType, Ptr: LHSPtr);
4720 RHS = Builder.CreateLoad(Ty: RI.ElementType, Ptr: RHSPtr);
4721 }
4722
4723 Value *Reduced;
4724 InsertPointOrErrorTy AfterIP =
4725 RI.ReductionGen(Builder.saveIP(), LHS, RHS, Reduced);
4726 if (!AfterIP)
4727 return AfterIP.takeError();
4728 if (!Builder.GetInsertBlock())
4729 return ReductionFunc;
4730
4731 Builder.restoreIP(IP: *AfterIP);
4732
4733 if (!IsByRef.empty() && !IsByRef[En.index()])
4734 Builder.CreateStore(Val: Reduced, Ptr: LHSPtr);
4735 }
4736 }
4737
4738 if (ReductionGenCBKind == ReductionGenCBKind::Clang)
4739 for (auto En : enumerate(First&: ReductionInfos)) {
4740 unsigned Index = En.index();
4741 const ReductionInfo &RI = En.value();
4742 Value *LHSFixupPtr, *RHSFixupPtr;
4743 Builder.restoreIP(IP: RI.ReductionGenClang(
4744 Builder.saveIP(), Index, &LHSFixupPtr, &RHSFixupPtr, ReductionFunc));
4745
4746 // Fix the CallBack code genereated to use the correct Values for the LHS
4747 // and RHS
4748 LHSFixupPtr->replaceUsesWithIf(
4749 New: LHSPtrs[Index], ShouldReplace: [ReductionFunc](const Use &U) {
4750 return cast<Instruction>(Val: U.getUser())->getParent()->getParent() ==
4751 ReductionFunc;
4752 });
4753 RHSFixupPtr->replaceUsesWithIf(
4754 New: RHSPtrs[Index], ShouldReplace: [ReductionFunc](const Use &U) {
4755 return cast<Instruction>(Val: U.getUser())->getParent()->getParent() ==
4756 ReductionFunc;
4757 });
4758 }
4759
4760 Builder.CreateRetVoid();
4761 // Compiling with `-O0`, `alloca`s emitted in non-entry blocks are not hoisted
4762 // to the entry block (this is dones for higher opt levels by later passes in
4763 // the pipeline). This has caused issues because non-entry `alloca`s force the
4764 // function to use dynamic stack allocations and we might run out of scratch
4765 // memory.
4766 hoistNonEntryAllocasToEntryBlock(Func: ReductionFunc);
4767
4768 return ReductionFunc;
4769}
4770
4771static void
4772checkReductionInfos(ArrayRef<OpenMPIRBuilder::ReductionInfo> ReductionInfos,
4773 bool IsGPU) {
4774 for (const OpenMPIRBuilder::ReductionInfo &RI : ReductionInfos) {
4775 (void)RI;
4776 assert(RI.Variable && "expected non-null variable");
4777 assert(RI.PrivateVariable && "expected non-null private variable");
4778 assert((RI.ReductionGen || RI.ReductionGenClang) &&
4779 "expected non-null reduction generator callback");
4780 if (!IsGPU) {
4781 assert(
4782 RI.Variable->getType() == RI.PrivateVariable->getType() &&
4783 "expected variables and their private equivalents to have the same "
4784 "type");
4785 }
4786 assert(RI.Variable->getType()->isPointerTy() &&
4787 "expected variables to be pointers");
4788 }
4789}
4790
4791// The atomic cross-team reduction fast path applies when every reduction in the
4792// set can be represented by an atomicrmw. Clang only populates it for scalar
4793// reductions with a supported atomic operator.
4794static bool isAtomicableReductionSet(
4795 ArrayRef<OpenMPIRBuilder::ReductionInfo> ReductionInfos) {
4796 return all_of(Range&: ReductionInfos, P: [](const OpenMPIRBuilder::ReductionInfo &RI) {
4797 return static_cast<bool>(RI.AtomicReductionGen);
4798 });
4799}
4800
4801OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createReductionsGPU(
4802 const LocationDescription &Loc, InsertPointTy AllocaIP,
4803 InsertPointTy CodeGenIP, ArrayRef<ReductionInfo> ReductionInfos,
4804 ArrayRef<bool> IsByRef, bool IsNoWait, bool IsTeamsReduction, bool IsSPMD,
4805 ReductionGenCBKind ReductionGenCBKind, std::optional<omp::GV> GridValue,
4806 Value *SrcLocInfo) {
4807 if (!updateToLocation(Loc))
4808 return InsertPointTy();
4809 Builder.restoreIP(IP: CodeGenIP);
4810 Builder.SetCurrentDebugLocation(Loc.DL);
4811 checkReductionInfos(ReductionInfos, /*IsGPU*/ true);
4812 LLVMContext &Ctx = M.getContext();
4813
4814 // Source location for the ident struct
4815 if (!SrcLocInfo) {
4816 uint32_t SrcLocStrSize;
4817 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
4818 SrcLocInfo = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
4819 }
4820
4821 if (ReductionInfos.size() == 0)
4822 return Builder.saveIP();
4823
4824 BasicBlock *ContinuationBlock = nullptr;
4825 if (ReductionGenCBKind != ReductionGenCBKind::Clang) {
4826 // Copied code from createReductions
4827 BasicBlock *InsertBlock = Loc.IP.getNodeParent();
4828 ContinuationBlock = InsertBlock->splitBasicBlock(I: Loc.IP, BBName: "reduce.finalize");
4829 InsertBlock->getTerminator()->eraseFromParent();
4830 Builder.SetInsertPoint(InsertBlock->end());
4831 }
4832
4833 Function *CurFunc = Builder.GetInsertBlock()->getParent();
4834 AttributeList FuncAttrs;
4835 AttrBuilder AttrBldr(Ctx);
4836 for (auto Attr : CurFunc->getAttributes().getFnAttrs())
4837 AttrBldr.addAttribute(A: Attr);
4838 AttrBldr.removeAttribute(Val: Attribute::OptimizeNone);
4839 FuncAttrs = FuncAttrs.addFnAttributes(C&: Ctx, B: AttrBldr);
4840
4841 Expected<Function *> ReductionResult = createReductionFunction(
4842 ReducerName: Builder.GetInsertBlock()->getParent()->getName(), ReductionInfos, IsByRef,
4843 ReductionGenCBKind, FuncAttrs);
4844 if (!ReductionResult)
4845 return ReductionResult.takeError();
4846 Function *ReductionFunc = *ReductionResult;
4847
4848 // Set the grid value in the config needed for lowering later on
4849 if (GridValue.has_value())
4850 Config.setGridValue(GridValue.value());
4851 else
4852 Config.setGridValue(getGridValue(T, Kernel: ReductionFunc));
4853
4854 // Build res = __kmpc_reduce{_nowait}(<gtid>, <n>, sizeof(RedList),
4855 // RedList, shuffle_reduce_func, interwarp_copy_func);
4856 // or
4857 // Build res = __kmpc_reduce_teams_nowait_simple(<loc>, <gtid>, <lck>);
4858 Value *Res;
4859
4860 // 1. Build a list of reduction variables.
4861 // void *RedList[<n>] = {<ReductionVars>[0], ..., <ReductionVars>[<n>-1]};
4862 auto Size = ReductionInfos.size();
4863 Type *PtrTy = PointerType::get(C&: Ctx, AddressSpace: Config.getDefaultTargetAS());
4864 Type *FuncPtrTy =
4865 Builder.getPtrTy(AddrSpace: M.getDataLayout().getProgramAddressSpace());
4866 Type *RedArrayTy = ArrayType::get(ElementType: PtrTy, NumElements: Size);
4867 Value *ReductionList;
4868 {
4869 IRBuilder<>::InsertPointGuard IPG(Builder);
4870 Builder.restoreIP(IP: AllocaIP);
4871 Value *ReductionListAlloca =
4872 Builder.CreateAlloca(Ty: RedArrayTy, ArraySize: nullptr, Name: ".omp.reduction.red_list");
4873 ReductionList = Builder.CreatePointerBitCastOrAddrSpaceCast(
4874 V: ReductionListAlloca, DestTy: PtrTy, Name: ReductionListAlloca->getName() + ".ascast");
4875 }
4876 Type *IndexTy = Builder.getIndexTy(
4877 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
4878 for (auto En : enumerate(First&: ReductionInfos)) {
4879 const ReductionInfo &RI = En.value();
4880 Value *ElemPtr = Builder.CreateInBoundsGEP(
4881 Ty: RedArrayTy, Ptr: ReductionList,
4882 IdxList: {ConstantInt::get(Ty: IndexTy, V: 0), ConstantInt::get(Ty: IndexTy, V: En.index())});
4883
4884 Value *PrivateVar = RI.PrivateVariable;
4885 bool IsByRefElem = !IsByRef.empty() && IsByRef[En.index()];
4886 if (IsByRefElem)
4887 PrivateVar = Builder.CreateLoad(Ty: RI.ElementType, Ptr: PrivateVar);
4888
4889 Value *CastElem =
4890 Builder.CreatePointerBitCastOrAddrSpaceCast(V: PrivateVar, DestTy: PtrTy);
4891 Builder.CreateStore(Val: CastElem, Ptr: ElemPtr);
4892 }
4893 Expected<Function *> SarFunc = emitShuffleAndReduceFunction(
4894 ReductionInfos, ReduceFn: ReductionFunc, FuncAttrs, IsByRef);
4895
4896 if (!SarFunc)
4897 return SarFunc.takeError();
4898
4899 Expected<Function *> CopyResult =
4900 emitInterWarpCopyFunction(Loc, ReductionInfos, FuncAttrs, IsByRef);
4901 if (!CopyResult)
4902 return CopyResult.takeError();
4903 Function *WcFunc = *CopyResult;
4904
4905 Value *RL = Builder.CreatePointerBitCastOrAddrSpaceCast(V: ReductionList, DestTy: PtrTy);
4906
4907 // NOTE: ReductionDataSize is passed as the reduce_data_size argument to
4908 // __kmpc_nvptx_parallel_reduce_nowait_v2, but the runtime implementations do
4909 // not currently use it. It is computed here conservatively as max(element
4910 // sizes) * N rather than the exact sum, which over-calculates the size for
4911 // mixed reduction types but is harmless given the argument is unused.
4912 // TODO: Consider dropping this computation if the runtime API is ever revised
4913 // to remove the unused parameter.
4914 unsigned MaxDataSize = 0;
4915 SmallVector<Type *> ReductionTypeArgs;
4916 for (auto En : enumerate(First&: ReductionInfos)) {
4917 // Use ByRefElementType for by-ref reductions so that MaxDataSize matches
4918 // the actual data size stored in the global reduction buffer, consistent
4919 // with the ReductionsBufferTy struct used for GEP offsets below.
4920 Type *RedTypeArg = (!IsByRef.empty() && IsByRef[En.index()])
4921 ? En.value().ByRefElementType
4922 : En.value().ElementType;
4923 auto Size = M.getDataLayout().getTypeStoreSize(Ty: RedTypeArg);
4924 if (Size > MaxDataSize)
4925 MaxDataSize = Size;
4926 ReductionTypeArgs.emplace_back(Args&: RedTypeArg);
4927 }
4928 Value *ReductionDataSize =
4929 Builder.getInt64(C: MaxDataSize * ReductionInfos.size());
4930
4931 // Helper function to copy thread-local data back to the original reduction
4932 // list.
4933 Function *CopyScratchToListFunc = nullptr;
4934 // Thread-local storage for the reduction variables.
4935 Value *ScratchForCopyBack = nullptr;
4936 // RL pointer to which the final value from the per-thread scratch should be
4937 // copied back. (Basically RL, appropriately casted if necessary.)
4938 Value *RLForCopyBack = RL;
4939
4940 bool IsAtomicReduction =
4941 IsTeamsReduction && isAtomicableReductionSet(ReductionInfos);
4942
4943 if (!IsTeamsReduction) {
4944 Value *SarFuncCast =
4945 Builder.CreatePointerBitCastOrAddrSpaceCast(V: *SarFunc, DestTy: FuncPtrTy);
4946 Value *WcFuncCast =
4947 Builder.CreatePointerBitCastOrAddrSpaceCast(V: WcFunc, DestTy: FuncPtrTy);
4948 Value *Args[] = {SrcLocInfo, ReductionDataSize, RL, SarFuncCast,
4949 WcFuncCast};
4950 Function *Pv2Ptr = getOrCreateRuntimeFunctionPtr(
4951 FnID: RuntimeFunction::OMPRTL___kmpc_nvptx_parallel_reduce_nowait_v2);
4952 Res = createRuntimeFunctionCall(Callee: Pv2Ptr, Args);
4953 } else if (IsAtomicReduction) {
4954 // Atomic cross-team reduction fast path: determine the team's main thread
4955 // that is later to fold its value atomically into the mapped variable.
4956 Function *IsMainThreadFn = getOrCreateRuntimeFunctionPtr(
4957 FnID: RuntimeFunction::OMPRTL___kmpc_is_team_main_thread);
4958 Res = createRuntimeFunctionCall(Callee: IsMainThreadFn, Args: {});
4959 } else {
4960 StructType *ReductionsBufferTy = StructType::create(
4961 Context&: Ctx, Elements: ReductionTypeArgs, Name: "struct._globalized_locals_ty");
4962
4963 Expected<Function *> LtGCFunc = emitListToGlobalCopyFunction(
4964 ReductionInfos, ReductionsBufferTy, FuncAttrs, IsByRef);
4965 if (!LtGCFunc)
4966 return LtGCFunc.takeError();
4967
4968 Expected<Function *> GtLCFunc = emitGlobalToListCopyFunction(
4969 ReductionInfos, ReductionsBufferTy, FuncAttrs, IsByRef);
4970 if (!GtLCFunc)
4971 return GtLCFunc.takeError();
4972
4973 Expected<Function *> GtLRFunc = emitGlobalToListReduceFunction(
4974 ReductionInfos, ReduceFn: ReductionFunc, ReductionsBufferTy, FuncAttrs, IsByRef);
4975 if (!GtLRFunc)
4976 return GtLRFunc.takeError();
4977
4978 // The runtime's cross-team final aggregate uses the storage pointed at by
4979 // its reduce-list argument as per-thread scratch. When the surrounding
4980 // kernel is already in SPMD execution mode, clang emitted each reduction
4981 // private as a per-thread `alloca addrspace(5)`, so the original red_list
4982 // (RL) is already per-thread and nothing else is needed.
4983 //
4984 // When the kernel is in Non-SPMD execution mode at codegen time, clang's
4985 // Generic-mode globalization put the reduction private into team-shared
4986 // LDS. OpenMPOpt may later upgrade the kernel to Generic-SPMD, at which
4987 // point all threads of the last team would race on the shared LDS slot.
4988 // Emit a per-thread scratch buffer and a per-thread RL, copy the team-local
4989 // value in, and hand the per-thread RL to the runtime instead. The writer
4990 // thread copies the final value from that per-thread scratch back to RL
4991 // before running the existing combine path below.
4992
4993 // Thread-local RL (might need localization below before being passed to the
4994 // runtime).
4995 Value *RuntimeRL = RL;
4996
4997 if (!IsSPMD) {
4998 Value *PerThreadScratch;
4999
5000 {
5001 IRBuilder<>::InsertPointGuard IPG(Builder);
5002 Builder.restoreIP(IP: AllocaIP);
5003 // Allocate thread-local buffer for the reduction variables.
5004 Value *PerThreadScratchAlloca =
5005 Builder.CreateAlloca(Ty: ReductionsBufferTy, /*ArraySize=*/nullptr,
5006 Name: ".omp.reduction.scratch");
5007 PerThreadScratch = Builder.CreatePointerBitCastOrAddrSpaceCast(
5008 V: PerThreadScratchAlloca, DestTy: PtrTy,
5009 Name: PerThreadScratchAlloca->getName() + ".ascast");
5010 // Allocate thread-local buffer for the pointers to the reduction
5011 // variables.
5012 Value *PerThreadRedListAlloca =
5013 Builder.CreateAlloca(Ty: RedArrayTy, /*ArraySize=*/nullptr,
5014 Name: ".omp.reduction.per_thread_red_list");
5015 RuntimeRL = Builder.CreatePointerBitCastOrAddrSpaceCast(
5016 V: PerThreadRedListAlloca, DestTy: PtrTy,
5017 Name: PerThreadRedListAlloca->getName() + ".ascast");
5018 }
5019
5020 // Iterate over the reduction variables and copy the team-local value to
5021 // the thread-local buffer.
5022 for (auto En : enumerate(First&: ReductionInfos)) {
5023 const ReductionInfo &RI = En.value();
5024 bool IsByRefElem = !IsByRef.empty() && IsByRef[En.index()];
5025
5026 Value *FieldPtr = Builder.CreateConstInBoundsGEP2_32(
5027 Ty: ReductionsBufferTy, Ptr: PerThreadScratch, Idx0: 0, Idx1: En.index());
5028 Value *Slot = Builder.CreateConstInBoundsGEP2_32(Ty: RedArrayTy, Ptr: RuntimeRL,
5029 Idx0: 0, Idx1: En.index());
5030
5031 Value *RuntimeListEntry = FieldPtr;
5032 if (IsByRefElem && RI.DataPtrPtrGen) {
5033 Value *SrcDescriptor =
5034 Builder.CreateLoad(Ty: RI.ElementType, Ptr: RI.PrivateVariable);
5035 Expected<Value *> Descriptor = createReductionDescriptorCopy(
5036 AllocaIP, RI, DataPtr: FieldPtr, SrcDescriptorAddr: SrcDescriptor, DescriptorPtrTy: PtrTy);
5037 if (!Descriptor)
5038 return Descriptor.takeError();
5039 RuntimeListEntry = *Descriptor;
5040 }
5041 Builder.CreateStore(Val: RuntimeListEntry, Ptr: Slot);
5042 }
5043 // The copy helpers were emitted with default-AS (AS 0) pointer params
5044 // (see emitListToGlobalCopyFunction / emitGlobalToListCopyFunction),
5045 // but PerThreadScratch and RL live in the target's default AS, which
5046 // is non-zero on e.g. SPIRV. (See Config.getDefaultTargetAS().)
5047 Type *CopyArg0Ty = (*LtGCFunc)->getFunctionType()->getParamType(i: 0);
5048 Type *CopyArg2Ty = (*LtGCFunc)->getFunctionType()->getParamType(i: 2);
5049 ScratchForCopyBack = Builder.CreatePointerBitCastOrAddrSpaceCast(
5050 V: PerThreadScratch, DestTy: CopyArg0Ty);
5051 RLForCopyBack =
5052 Builder.CreatePointerBitCastOrAddrSpaceCast(V: RL, DestTy: CopyArg2Ty);
5053 // Use index 0 because there is no array of target values to index into,
5054 // there is only one thread-local memory slot.
5055 // restoreIP above left a stale/empty debug location; this inlinable call
5056 // to a debug-info-bearing helper needs one or the verifier rejects the
5057 // module ("!dbg attachment points at wrong subprogram") after inlining.
5058 Builder.SetCurrentDebugLocation(Loc.DL);
5059 Builder.CreateCall(
5060 Callee: *LtGCFunc, Args: {ScratchForCopyBack, Builder.getInt32(C: 0), RLForCopyBack});
5061 CopyScratchToListFunc = *GtLCFunc;
5062 }
5063
5064 Value *Args3[] = {SrcLocInfo, RuntimeRL, *SarFunc, WcFunc,
5065 *LtGCFunc, *GtLCFunc, *GtLRFunc};
5066
5067 Function *TeamsReduceFn = getOrCreateRuntimeFunctionPtr(
5068 FnID: RuntimeFunction::OMPRTL___kmpc_gpu_xteam_reduce_nowait);
5069 Res = createRuntimeFunctionCall(Callee: TeamsReduceFn, Args: Args3);
5070 }
5071
5072 // 5. Build if (res == 1)
5073 BasicBlock *ExitBB = BasicBlock::Create(Context&: Ctx, Name: ".omp.reduction.done");
5074 BasicBlock *ThenBB = BasicBlock::Create(Context&: Ctx, Name: ".omp.reduction.then");
5075 Value *Cond = Builder.CreateICmpEQ(LHS: Res, RHS: Builder.getInt32(C: 1));
5076 Builder.CreateCondBr(Cond, True: ThenBB, False: ExitBB);
5077
5078 // 6. Build then branch: where we have reduced values in the master
5079 // thread in each team.
5080 // __kmpc_end_reduce{_nowait}(<gtid>);
5081 // break;
5082 emitBlock(BB: ThenBB, CurFn: CurFunc);
5083
5084 // Copy the writer thread's per-thread scratch result back into the original
5085 // red-list storage before the existing combine path reads RI.PrivateVariable.
5086 // Set a debug location: this inlinable call to a debug-info-bearing helper
5087 // needs one or the verifier rejects the module after inlining.
5088 if (ScratchForCopyBack) {
5089 Builder.SetCurrentDebugLocation(Loc.DL);
5090 Builder.CreateCall(
5091 Callee: CopyScratchToListFunc,
5092 Args: {ScratchForCopyBack, Builder.getInt32(C: 0), RLForCopyBack});
5093 }
5094
5095 // Add emission of __kmpc_end_reduce{_nowait}(<gtid>);
5096 for (auto En : enumerate(First&: ReductionInfos)) {
5097 const ReductionInfo &RI = En.value();
5098
5099 // Atomic cross-team fast path: each team's main thread folds its
5100 // team-reduced value directly into the mapped reduction variable with a
5101 // single atomicrmw.
5102 if (IsAtomicReduction) {
5103 InsertPointOrErrorTy AfterIP = RI.AtomicReductionGen(
5104 Builder.saveIP(), RI.ElementType, RI.Variable, RI.PrivateVariable);
5105 if (!AfterIP)
5106 return AfterIP.takeError();
5107 Builder.restoreIP(IP: *AfterIP);
5108 continue;
5109 }
5110
5111 Type *ValueType = RI.ElementType;
5112 Value *RedValue = RI.Variable;
5113
5114 Value *RHS =
5115 Builder.CreatePointerBitCastOrAddrSpaceCast(V: RI.PrivateVariable, DestTy: PtrTy);
5116
5117 if (ReductionGenCBKind == ReductionGenCBKind::Clang) {
5118 Value *LHSPtr, *RHSPtr;
5119 Builder.restoreIP(IP: RI.ReductionGenClang(Builder.saveIP(), En.index(),
5120 &LHSPtr, &RHSPtr, CurFunc));
5121
5122 // Fix the CallBack code genereated to use the correct Values for the LHS
5123 // and RHS. Cast to match types before replacing (necessary to handle
5124 // different address spaces).
5125 if (LHSPtr->getType() != RedValue->getType())
5126 RedValue = Builder.CreatePointerBitCastOrAddrSpaceCast(
5127 V: RedValue, DestTy: LHSPtr->getType());
5128 if (RHSPtr->getType() != RHS->getType())
5129 RHS =
5130 Builder.CreatePointerBitCastOrAddrSpaceCast(V: RHS, DestTy: RHSPtr->getType());
5131
5132 LHSPtr->replaceUsesWithIf(New: RedValue, ShouldReplace: [ReductionFunc](const Use &U) {
5133 return cast<Instruction>(Val: U.getUser())->getParent()->getParent() ==
5134 ReductionFunc;
5135 });
5136 RHSPtr->replaceUsesWithIf(New: RHS, ShouldReplace: [ReductionFunc](const Use &U) {
5137 return cast<Instruction>(Val: U.getUser())->getParent()->getParent() ==
5138 ReductionFunc;
5139 });
5140 } else {
5141 if (IsByRef.empty() || !IsByRef[En.index()]) {
5142 RedValue = Builder.CreateLoad(Ty: ValueType, Ptr: RI.Variable,
5143 Name: "red.value." + Twine(En.index()));
5144 }
5145 Value *PrivateRedValue = Builder.CreateLoad(
5146 Ty: ValueType, Ptr: RHS, Name: "red.private.value" + Twine(En.index()));
5147 Value *Reduced;
5148 InsertPointOrErrorTy AfterIP =
5149 RI.ReductionGen(Builder.saveIP(), RedValue, PrivateRedValue, Reduced);
5150 if (!AfterIP)
5151 return AfterIP.takeError();
5152 Builder.restoreIP(IP: *AfterIP);
5153
5154 if (!IsByRef.empty() && !IsByRef[En.index()])
5155 Builder.CreateStore(Val: Reduced, Ptr: RI.Variable);
5156 }
5157 }
5158 emitBlock(BB: ExitBB, CurFn: CurFunc);
5159 if (ContinuationBlock) {
5160 Builder.CreateBr(Dest: ContinuationBlock);
5161 Builder.SetInsertPoint(ContinuationBlock);
5162 }
5163 Config.setEmitLLVMUsed();
5164
5165 return Builder.saveIP();
5166}
5167
5168static Function *getFreshReductionFunc(Module &M) {
5169 Type *VoidTy = Type::getVoidTy(C&: M.getContext());
5170 Type *Int8PtrTy = PointerType::getUnqual(C&: M.getContext());
5171 auto *FuncTy =
5172 FunctionType::get(Result: VoidTy, Params: {Int8PtrTy, Int8PtrTy}, /* IsVarArg */ isVarArg: false);
5173 return Function::Create(Ty: FuncTy, Linkage: GlobalVariable::InternalLinkage,
5174 N: ".omp.reduction.func", M: &M);
5175}
5176
5177static Error populateReductionFunction(
5178 Function *ReductionFunc,
5179 ArrayRef<OpenMPIRBuilder::ReductionInfo> ReductionInfos,
5180 IRBuilder<> &Builder, ArrayRef<bool> IsByRef, bool IsGPU) {
5181 IRBuilder<>::InsertPointGuard IPG(Builder);
5182 Module *Module = ReductionFunc->getParent();
5183 BasicBlock *ReductionFuncBlock =
5184 BasicBlock::Create(Context&: Module->getContext(), Name: "", Parent: ReductionFunc);
5185 Builder.SetInsertPoint(ReductionFuncBlock);
5186 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
5187 Value *LHSArrayPtr = nullptr;
5188 Value *RHSArrayPtr = nullptr;
5189 if (IsGPU) {
5190 // Need to alloca memory here and deal with the pointers before getting
5191 // LHS/RHS pointers out
5192 //
5193 Argument *Arg0 = ReductionFunc->getArg(i: 0);
5194 Argument *Arg1 = ReductionFunc->getArg(i: 1);
5195 Type *Arg0Type = Arg0->getType();
5196 Type *Arg1Type = Arg1->getType();
5197
5198 Value *LHSAlloca =
5199 Builder.CreateAlloca(Ty: Arg0Type, ArraySize: nullptr, Name: Arg0->getName() + ".addr");
5200 Value *RHSAlloca =
5201 Builder.CreateAlloca(Ty: Arg1Type, ArraySize: nullptr, Name: Arg1->getName() + ".addr");
5202 Value *LHSAddrCast =
5203 Builder.CreatePointerBitCastOrAddrSpaceCast(V: LHSAlloca, DestTy: Arg0Type);
5204 Value *RHSAddrCast =
5205 Builder.CreatePointerBitCastOrAddrSpaceCast(V: RHSAlloca, DestTy: Arg1Type);
5206 Builder.CreateStore(Val: Arg0, Ptr: LHSAddrCast);
5207 Builder.CreateStore(Val: Arg1, Ptr: RHSAddrCast);
5208 LHSArrayPtr = Builder.CreateLoad(Ty: Arg0Type, Ptr: LHSAddrCast);
5209 RHSArrayPtr = Builder.CreateLoad(Ty: Arg1Type, Ptr: RHSAddrCast);
5210 } else {
5211 LHSArrayPtr = ReductionFunc->getArg(i: 0);
5212 RHSArrayPtr = ReductionFunc->getArg(i: 1);
5213 }
5214
5215 unsigned NumReductions = ReductionInfos.size();
5216 Type *RedArrayTy = ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: NumReductions);
5217
5218 for (auto En : enumerate(First&: ReductionInfos)) {
5219 const OpenMPIRBuilder::ReductionInfo &RI = En.value();
5220 Value *LHSI8PtrPtr = Builder.CreateConstInBoundsGEP2_64(
5221 Ty: RedArrayTy, Ptr: LHSArrayPtr, Idx0: 0, Idx1: En.index());
5222 Value *LHSI8Ptr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: LHSI8PtrPtr);
5223 Value *LHSPtr = Builder.CreatePointerBitCastOrAddrSpaceCast(
5224 V: LHSI8Ptr, DestTy: RI.Variable->getType());
5225 Value *LHS = Builder.CreateLoad(Ty: RI.ElementType, Ptr: LHSPtr);
5226 Value *RHSI8PtrPtr = Builder.CreateConstInBoundsGEP2_64(
5227 Ty: RedArrayTy, Ptr: RHSArrayPtr, Idx0: 0, Idx1: En.index());
5228 Value *RHSI8Ptr = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: RHSI8PtrPtr);
5229 Value *RHSPtr = Builder.CreatePointerBitCastOrAddrSpaceCast(
5230 V: RHSI8Ptr, DestTy: RI.PrivateVariable->getType());
5231 Value *RHS = Builder.CreateLoad(Ty: RI.ElementType, Ptr: RHSPtr);
5232 Value *Reduced;
5233 OpenMPIRBuilder::InsertPointOrErrorTy AfterIP =
5234 RI.ReductionGen(Builder.saveIP(), LHS, RHS, Reduced);
5235 if (!AfterIP)
5236 return AfterIP.takeError();
5237
5238 Builder.restoreIP(IP: *AfterIP);
5239 // TODO: Consider flagging an error.
5240 if (!Builder.GetInsertBlock())
5241 return Error::success();
5242
5243 // store is inside of the reduction region when using by-ref
5244 if (!IsByRef[En.index()])
5245 Builder.CreateStore(Val: Reduced, Ptr: LHSPtr);
5246 }
5247 Builder.CreateRetVoid();
5248 return Error::success();
5249}
5250
5251OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createReductions(
5252 const LocationDescription &Loc, InsertPointTy AllocaIP,
5253 ArrayRef<ReductionInfo> ReductionInfos, ArrayRef<bool> IsByRef,
5254 bool IsNoWait, bool IsTeamsReduction) {
5255 assert(ReductionInfos.size() == IsByRef.size());
5256 if (Config.isGPU())
5257 return createReductionsGPU(Loc, AllocaIP, CodeGenIP: Builder.saveIP(), ReductionInfos,
5258 IsByRef, IsNoWait, IsTeamsReduction);
5259
5260 checkReductionInfos(ReductionInfos, /*IsGPU*/ false);
5261
5262 if (!updateToLocation(Loc))
5263 return InsertPointTy();
5264
5265 if (ReductionInfos.size() == 0)
5266 return Builder.saveIP();
5267
5268 BasicBlock *InsertBlock = Loc.IP.getNodeParent();
5269 BasicBlock *ContinuationBlock =
5270 InsertBlock->splitBasicBlock(I: Loc.IP, BBName: "reduce.finalize");
5271 InsertBlock->getTerminator()->eraseFromParent();
5272
5273 // Create and populate array of type-erased pointers to private reduction
5274 // values.
5275 unsigned NumReductions = ReductionInfos.size();
5276 Type *RedArrayTy = ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: NumReductions);
5277 Builder.SetInsertPoint(AllocaIP.getNodeParent()->getTerminator());
5278 Value *RedArray = Builder.CreateAlloca(Ty: RedArrayTy, ArraySize: nullptr, Name: "red.array");
5279
5280 Builder.SetInsertPoint(InsertBlock->end());
5281 // Emitting the alloca moved the insertion point into the alloca block and
5282 // can clear the debug loc. Restore back to Loc.DL.
5283 Builder.SetCurrentDebugLocation(Loc.DL);
5284
5285 for (auto En : enumerate(First&: ReductionInfos)) {
5286 unsigned Index = En.index();
5287 const ReductionInfo &RI = En.value();
5288 Value *RedArrayElemPtr = Builder.CreateConstInBoundsGEP2_64(
5289 Ty: RedArrayTy, Ptr: RedArray, Idx0: 0, Idx1: Index, Name: "red.array.elem." + Twine(Index));
5290 Builder.CreateStore(Val: RI.PrivateVariable, Ptr: RedArrayElemPtr);
5291 }
5292
5293 // Emit a call to the runtime function that orchestrates the reduction.
5294 // Declare the reduction function in the process.
5295 Type *IndexTy = Builder.getIndexTy(
5296 DL: M.getDataLayout(), AddrSpace: M.getDataLayout().getDefaultGlobalsAddressSpace());
5297 Function *Func = Builder.GetInsertBlock()->getParent();
5298 Module *Module = Func->getParent();
5299 uint32_t SrcLocStrSize;
5300 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
5301 bool CanGenerateAtomic = all_of(Range&: ReductionInfos, P: [](const ReductionInfo &RI) {
5302 return RI.AtomicReductionGen;
5303 });
5304 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize,
5305 LocFlags: CanGenerateAtomic
5306 ? IdentFlag::OMP_IDENT_FLAG_ATOMIC_REDUCE
5307 : IdentFlag(0));
5308 Value *ThreadId = getOrCreateThreadID(Ident);
5309 Constant *NumVariables = Builder.getInt32(C: NumReductions);
5310 const DataLayout &DL = Module->getDataLayout();
5311 unsigned RedArrayByteSize = DL.getTypeStoreSize(Ty: RedArrayTy);
5312 Constant *RedArraySize = ConstantInt::get(Ty: IndexTy, V: RedArrayByteSize);
5313 Function *ReductionFunc = getFreshReductionFunc(M&: *Module);
5314 Value *Lock = getOMPCriticalRegionLock(CriticalName: ".reduction");
5315 Function *ReduceFunc = getOrCreateRuntimeFunctionPtr(
5316 FnID: IsNoWait ? RuntimeFunction::OMPRTL___kmpc_reduce_nowait
5317 : RuntimeFunction::OMPRTL___kmpc_reduce);
5318 CallInst *ReduceCall =
5319 createRuntimeFunctionCall(Callee: ReduceFunc,
5320 Args: {Ident, ThreadId, NumVariables, RedArraySize,
5321 RedArray, ReductionFunc, Lock},
5322 Name: "reduce");
5323
5324 // Create final reduction entry blocks for the atomic and non-atomic case.
5325 // Emit IR that dispatches control flow to one of the blocks based on the
5326 // reduction supporting the atomic mode.
5327 BasicBlock *NonAtomicRedBlock =
5328 BasicBlock::Create(Context&: Module->getContext(), Name: "reduce.switch.nonatomic", Parent: Func);
5329 BasicBlock *AtomicRedBlock =
5330 BasicBlock::Create(Context&: Module->getContext(), Name: "reduce.switch.atomic", Parent: Func);
5331 SwitchInst *Switch =
5332 Builder.CreateSwitch(V: ReduceCall, Dest: ContinuationBlock, /* NumCases */ 2);
5333 Switch->addCase(OnVal: Builder.getInt32(C: 1), Dest: NonAtomicRedBlock);
5334 Switch->addCase(OnVal: Builder.getInt32(C: 2), Dest: AtomicRedBlock);
5335
5336 // Populate the non-atomic reduction using the elementwise reduction function.
5337 // This loads the elements from the global and private variables and reduces
5338 // them before storing back the result to the global variable.
5339 Builder.SetInsertPoint(NonAtomicRedBlock);
5340 for (auto En : enumerate(First&: ReductionInfos)) {
5341 const ReductionInfo &RI = En.value();
5342 Type *ValueType = RI.ElementType;
5343 // We have one less load for by-ref case because that load is now inside of
5344 // the reduction region
5345 Value *RedValue = RI.Variable;
5346 if (!IsByRef[En.index()]) {
5347 RedValue = Builder.CreateLoad(Ty: ValueType, Ptr: RI.Variable,
5348 Name: "red.value." + Twine(En.index()));
5349 }
5350 Value *PrivateRedValue =
5351 Builder.CreateLoad(Ty: ValueType, Ptr: RI.PrivateVariable,
5352 Name: "red.private.value." + Twine(En.index()));
5353 Value *Reduced;
5354 InsertPointOrErrorTy AfterIP =
5355 RI.ReductionGen(Builder.saveIP(), RedValue, PrivateRedValue, Reduced);
5356 if (!AfterIP)
5357 return AfterIP.takeError();
5358 Builder.restoreIP(IP: *AfterIP);
5359
5360 if (!Builder.GetInsertBlock())
5361 return InsertPointTy();
5362 // for by-ref case, the load is inside of the reduction region
5363 if (!IsByRef[En.index()])
5364 Builder.CreateStore(Val: Reduced, Ptr: RI.Variable);
5365 }
5366 Function *EndReduceFunc = getOrCreateRuntimeFunctionPtr(
5367 FnID: IsNoWait ? RuntimeFunction::OMPRTL___kmpc_end_reduce_nowait
5368 : RuntimeFunction::OMPRTL___kmpc_end_reduce);
5369 createRuntimeFunctionCall(Callee: EndReduceFunc, Args: {Ident, ThreadId, Lock});
5370 Builder.CreateBr(Dest: ContinuationBlock);
5371
5372 // Populate the atomic reduction using the atomic elementwise reduction
5373 // function. There are no loads/stores here because they will be happening
5374 // inside the atomic elementwise reduction.
5375 Builder.SetInsertPoint(AtomicRedBlock);
5376 if (CanGenerateAtomic && llvm::none_of(Range&: IsByRef, P: [](bool P) { return P; })) {
5377 for (const ReductionInfo &RI : ReductionInfos) {
5378 InsertPointOrErrorTy AfterIP = RI.AtomicReductionGen(
5379 Builder.saveIP(), RI.ElementType, RI.Variable, RI.PrivateVariable);
5380 if (!AfterIP)
5381 return AfterIP.takeError();
5382 Builder.restoreIP(IP: *AfterIP);
5383 if (!Builder.GetInsertBlock())
5384 return InsertPointTy();
5385 }
5386 Builder.CreateBr(Dest: ContinuationBlock);
5387 } else {
5388 Builder.CreateUnreachable();
5389 }
5390
5391 // Populate the outlined reduction function using the elementwise reduction
5392 // function. Partial values are extracted from the type-erased array of
5393 // pointers to private variables.
5394 Error Err = populateReductionFunction(ReductionFunc, ReductionInfos, Builder,
5395 IsByRef, /*isGPU=*/IsGPU: false);
5396 if (Err)
5397 return Err;
5398
5399 if (!Builder.GetInsertBlock())
5400 return InsertPointTy();
5401
5402 Builder.SetInsertPoint(ContinuationBlock);
5403 return Builder.saveIP();
5404}
5405
5406OpenMPIRBuilder::InsertPointOrErrorTy
5407OpenMPIRBuilder::createMaster(const LocationDescription &Loc,
5408 BodyGenCallbackTy BodyGenCB,
5409 FinalizeCallbackTy FiniCB) {
5410 if (!updateToLocation(Loc))
5411 return Loc.IP;
5412
5413 Directive OMPD = Directive::OMPD_master;
5414 uint32_t SrcLocStrSize;
5415 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
5416 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
5417 Value *ThreadId = getOrCreateThreadID(Ident);
5418 Value *Args[] = {Ident, ThreadId};
5419
5420 Function *EntryRTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_master);
5421 Instruction *EntryCall = createRuntimeFunctionCall(Callee: EntryRTLFn, Args);
5422
5423 Function *ExitRTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_end_master);
5424 Instruction *ExitCall = createRuntimeFunctionCall(Callee: ExitRTLFn, Args);
5425
5426 return EmitOMPInlinedRegion(OMPD, EntryCall, ExitCall, BodyGenCB, FiniCB,
5427 /*Conditional*/ true, /*hasFinalize*/ HasFinalize: true);
5428}
5429
5430OpenMPIRBuilder::InsertPointOrErrorTy
5431OpenMPIRBuilder::createMasked(const LocationDescription &Loc,
5432 BodyGenCallbackTy BodyGenCB,
5433 FinalizeCallbackTy FiniCB, Value *Filter) {
5434 IRBuilder<>::InsertPointGuard IPG(Builder);
5435 if (!updateToLocation(Loc))
5436 return Loc.IP;
5437
5438 Directive OMPD = Directive::OMPD_masked;
5439 uint32_t SrcLocStrSize;
5440 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
5441 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
5442 Value *ThreadId = getOrCreateThreadID(Ident);
5443 Value *Args[] = {Ident, ThreadId, Filter};
5444 Value *ArgsEnd[] = {Ident, ThreadId};
5445
5446 Function *EntryRTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_masked);
5447 Instruction *EntryCall = createRuntimeFunctionCall(Callee: EntryRTLFn, Args);
5448
5449 Function *ExitRTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_end_masked);
5450 Instruction *ExitCall = createRuntimeFunctionCall(Callee: ExitRTLFn, Args: ArgsEnd);
5451
5452 return EmitOMPInlinedRegion(OMPD, EntryCall, ExitCall, BodyGenCB, FiniCB,
5453 /*Conditional*/ true, /*hasFinalize*/ HasFinalize: true);
5454}
5455
5456static llvm::CallInst *emitNoUnwindRuntimeCall(IRBuilder<> &Builder,
5457 llvm::FunctionCallee Callee,
5458 ArrayRef<llvm::Value *> Args,
5459 const llvm::Twine &Name) {
5460 llvm::CallInst *Call = Builder.CreateCall(
5461 Callee, Args, OpBundles: SmallVector<llvm::OperandBundleDef, 1>(), Name);
5462 Call->setDoesNotThrow();
5463 return Call;
5464}
5465
5466// Expects input basic block is dominated by BeforeScanBB.
5467// Once Scan directive is encountered, the code after scan directive should be
5468// dominated by AfterScanBB. Scan directive splits the code sequence to
5469// scan and input phase. Based on whether inclusive or exclusive
5470// clause is used in the scan directive and whether input loop or scan loop
5471// is lowered, it adds jumps to input and scan phase. First Scan loop is the
5472// input loop and second is the scan loop. The code generated handles only
5473// inclusive scans now.
5474OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createScan(
5475 const LocationDescription &Loc, InsertPointTy AllocaIP,
5476 ArrayRef<llvm::Value *> ScanVars, ArrayRef<llvm::Type *> ScanVarsType,
5477 bool IsInclusive, ScanInfo *ScanRedInfo) {
5478 if (ScanRedInfo->OMPFirstScanLoop) {
5479 llvm::Error Err = emitScanBasedDirectiveDeclsIR(AllocaIP, ScanVars,
5480 ScanVarsType, ScanRedInfo);
5481 if (Err)
5482 return Err;
5483 }
5484 if (!updateToLocation(Loc))
5485 return Loc.IP;
5486
5487 llvm::Value *IV = ScanRedInfo->IV;
5488
5489 if (ScanRedInfo->OMPFirstScanLoop) {
5490 // Emit buffer[i] = red; at the end of the input phase.
5491 for (size_t i = 0; i < ScanVars.size(); i++) {
5492 Value *BuffPtr = (*(ScanRedInfo->ScanBuffPtrs))[ScanVars[i]];
5493 Value *Buff = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BuffPtr);
5494 Type *DestTy = ScanVarsType[i];
5495 Value *Val = Builder.CreateInBoundsGEP(Ty: DestTy, Ptr: Buff, IdxList: IV, Name: "arrayOffset");
5496 Value *Src = Builder.CreateLoad(Ty: DestTy, Ptr: ScanVars[i]);
5497
5498 Builder.CreateStore(Val: Src, Ptr: Val);
5499 }
5500 }
5501 Builder.CreateBr(Dest: ScanRedInfo->OMPScanLoopExit);
5502 emitBlock(BB: ScanRedInfo->OMPScanDispatch,
5503 CurFn: Builder.GetInsertBlock()->getParent());
5504
5505 if (!ScanRedInfo->OMPFirstScanLoop) {
5506 IV = ScanRedInfo->IV;
5507 // Emit red = buffer[i]; at the entrance to the scan phase.
5508 // TODO: if exclusive scan, the red = buffer[i-1] needs to be updated.
5509 for (size_t i = 0; i < ScanVars.size(); i++) {
5510 Value *BuffPtr = (*(ScanRedInfo->ScanBuffPtrs))[ScanVars[i]];
5511 Value *Buff = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BuffPtr);
5512 Type *DestTy = ScanVarsType[i];
5513 Value *SrcPtr =
5514 Builder.CreateInBoundsGEP(Ty: DestTy, Ptr: Buff, IdxList: IV, Name: "arrayOffset");
5515 Value *Src = Builder.CreateLoad(Ty: DestTy, Ptr: SrcPtr);
5516 Builder.CreateStore(Val: Src, Ptr: ScanVars[i]);
5517 }
5518 }
5519
5520 // TODO: Update it to CreateBr and remove dead blocks
5521 llvm::Value *CmpI = Builder.getInt1(V: true);
5522 if (ScanRedInfo->OMPFirstScanLoop == IsInclusive) {
5523 Builder.CreateCondBr(Cond: CmpI, True: ScanRedInfo->OMPBeforeScanBlock,
5524 False: ScanRedInfo->OMPAfterScanBlock);
5525 } else {
5526 Builder.CreateCondBr(Cond: CmpI, True: ScanRedInfo->OMPAfterScanBlock,
5527 False: ScanRedInfo->OMPBeforeScanBlock);
5528 }
5529 emitBlock(BB: ScanRedInfo->OMPAfterScanBlock,
5530 CurFn: Builder.GetInsertBlock()->getParent());
5531 Builder.SetInsertPoint(ScanRedInfo->OMPAfterScanBlock);
5532 return Builder.saveIP();
5533}
5534
5535Error OpenMPIRBuilder::emitScanBasedDirectiveDeclsIR(
5536 InsertPointTy AllocaIP, ArrayRef<Value *> ScanVars,
5537 ArrayRef<Type *> ScanVarsType, ScanInfo *ScanRedInfo) {
5538
5539 Builder.restoreIP(IP: AllocaIP);
5540 // Create the shared pointer at alloca IP.
5541 for (size_t i = 0; i < ScanVars.size(); i++) {
5542 llvm::Value *BuffPtr =
5543 Builder.CreateAlloca(Ty: Builder.getPtrTy(), ArraySize: nullptr, Name: "vla");
5544 (*(ScanRedInfo->ScanBuffPtrs))[ScanVars[i]] = BuffPtr;
5545 }
5546
5547 // Allocate temporary buffer by master thread
5548 auto BodyGenCB = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
5549 ArrayRef<BasicBlock *> DeallocBlocks) -> Error {
5550 Builder.restoreIP(IP: CodeGenIP);
5551 Value *AllocSpan =
5552 Builder.CreateAdd(LHS: ScanRedInfo->Span, RHS: Builder.getInt32(C: 1));
5553 for (size_t i = 0; i < ScanVars.size(); i++) {
5554 Type *IntPtrTy = Builder.getInt32Ty();
5555 Value *Allocsize = Builder.CreateTypeSize(
5556 Ty: IntPtrTy, Size: M.getDataLayout().getTypeAllocSize(Ty: ScanVarsType[i]));
5557 Value *Buff =
5558 Builder.CreateMalloc(IntPtrTy, AllocSize: Allocsize, ArraySize: AllocSpan, MallocF: nullptr, Name: "arr");
5559 Builder.CreateStore(Val: Buff, Ptr: (*(ScanRedInfo->ScanBuffPtrs))[ScanVars[i]]);
5560 }
5561 return Error::success();
5562 };
5563 // TODO: Perform finalization actions for variables. This has to be
5564 // called for variables which have destructors/finalizers.
5565 auto FiniCB = [&](InsertPointTy CodeGenIP) { return llvm::Error::success(); };
5566
5567 Builder.SetInsertPoint(ScanRedInfo->OMPScanInit->getTerminator());
5568 llvm::Value *FilterVal = Builder.getInt32(C: 0);
5569 llvm::OpenMPIRBuilder::InsertPointOrErrorTy AfterIP =
5570 createMasked(Loc: Builder, BodyGenCB, FiniCB, Filter: FilterVal);
5571
5572 if (!AfterIP)
5573 return AfterIP.takeError();
5574 Builder.restoreIP(IP: *AfterIP);
5575 BasicBlock *InputBB = Builder.GetInsertBlock();
5576 if (InputBB->hasTerminator())
5577 Builder.SetInsertPoint(InputBB->getTerminator());
5578 AfterIP = createBarrier(Loc: Builder, Kind: llvm::omp::OMPD_barrier);
5579 if (!AfterIP)
5580 return AfterIP.takeError();
5581 Builder.restoreIP(IP: *AfterIP);
5582
5583 return Error::success();
5584}
5585
5586Error OpenMPIRBuilder::emitScanBasedDirectiveFinalsIR(
5587 ArrayRef<ReductionInfo> ReductionInfos, ScanInfo *ScanRedInfo) {
5588 auto BodyGenCB = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
5589 ArrayRef<BasicBlock *> DeallocBlocks) -> Error {
5590 Builder.restoreIP(IP: CodeGenIP);
5591 for (ReductionInfo RedInfo : ReductionInfos) {
5592 Value *PrivateVar = RedInfo.PrivateVariable;
5593 Value *OrigVar = RedInfo.Variable;
5594 Value *BuffPtr = (*(ScanRedInfo->ScanBuffPtrs))[PrivateVar];
5595 Value *Buff = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BuffPtr);
5596
5597 Type *SrcTy = RedInfo.ElementType;
5598 Value *Val = Builder.CreateInBoundsGEP(Ty: SrcTy, Ptr: Buff, IdxList: ScanRedInfo->Span,
5599 Name: "arrayOffset");
5600 Value *Src = Builder.CreateLoad(Ty: SrcTy, Ptr: Val);
5601
5602 Builder.CreateStore(Val: Src, Ptr: OrigVar);
5603 Builder.CreateFree(Source: Buff);
5604 }
5605 return Error::success();
5606 };
5607 // TODO: Perform finalization actions for variables. This has to be
5608 // called for variables which have destructors/finalizers.
5609 auto FiniCB = [&](InsertPointTy CodeGenIP) { return llvm::Error::success(); };
5610
5611 if (Instruction *TI = ScanRedInfo->OMPScanFinish->getTerminatorOrNull())
5612 Builder.SetInsertPoint(TI);
5613 else
5614 Builder.SetInsertPoint(ScanRedInfo->OMPScanFinish);
5615
5616 llvm::Value *FilterVal = Builder.getInt32(C: 0);
5617 llvm::OpenMPIRBuilder::InsertPointOrErrorTy AfterIP =
5618 createMasked(Loc: Builder, BodyGenCB, FiniCB, Filter: FilterVal);
5619
5620 if (!AfterIP)
5621 return AfterIP.takeError();
5622 Builder.restoreIP(IP: *AfterIP);
5623 BasicBlock *InputBB = Builder.GetInsertBlock();
5624 if (InputBB->hasTerminator())
5625 Builder.SetInsertPoint(InputBB->getTerminator());
5626 AfterIP = createBarrier(Loc: Builder, Kind: llvm::omp::OMPD_barrier);
5627 if (!AfterIP)
5628 return AfterIP.takeError();
5629 Builder.restoreIP(IP: *AfterIP);
5630 return Error::success();
5631}
5632
5633OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::emitScanReduction(
5634 const LocationDescription &Loc,
5635 ArrayRef<llvm::OpenMPIRBuilder::ReductionInfo> ReductionInfos,
5636 ScanInfo *ScanRedInfo) {
5637
5638 if (!updateToLocation(Loc))
5639 return Loc.IP;
5640 auto BodyGenCB = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
5641 ArrayRef<BasicBlock *> DeallocBlocks) -> Error {
5642 Builder.restoreIP(IP: CodeGenIP);
5643 Function *CurFn = Builder.GetInsertBlock()->getParent();
5644 // for (int k = 0; k <= ceil(log2(n)); ++k)
5645 llvm::BasicBlock *LoopBB =
5646 BasicBlock::Create(Context&: CurFn->getContext(), Name: "omp.outer.log.scan.body");
5647 llvm::BasicBlock *ExitBB =
5648 splitBB(Builder, CreateBranch: false, Name: "omp.outer.log.scan.exit");
5649 llvm::Function *F = llvm::Intrinsic::getOrInsertDeclaration(
5650 M: Builder.getModule(), id: (llvm::Intrinsic::ID)llvm::Intrinsic::log2,
5651 OverloadTys: Builder.getDoubleTy());
5652 llvm::BasicBlock *InputBB = Builder.GetInsertBlock();
5653 llvm::Value *Arg =
5654 Builder.CreateUIToFP(V: ScanRedInfo->Span, DestTy: Builder.getDoubleTy());
5655 llvm::Value *LogVal = emitNoUnwindRuntimeCall(Builder, Callee: F, Args: Arg, Name: "");
5656 F = llvm::Intrinsic::getOrInsertDeclaration(
5657 M: Builder.getModule(), id: (llvm::Intrinsic::ID)llvm::Intrinsic::ceil,
5658 OverloadTys: Builder.getDoubleTy());
5659 LogVal = emitNoUnwindRuntimeCall(Builder, Callee: F, Args: LogVal, Name: "");
5660 LogVal = Builder.CreateFPToUI(V: LogVal, DestTy: Builder.getInt32Ty());
5661 llvm::Value *NMin1 = Builder.CreateNUWSub(
5662 LHS: ScanRedInfo->Span,
5663 RHS: llvm::ConstantInt::get(Ty: ScanRedInfo->Span->getType(), V: 1));
5664 Builder.SetInsertPoint(InputBB);
5665 Builder.CreateBr(Dest: LoopBB);
5666 emitBlock(BB: LoopBB, CurFn);
5667 Builder.SetInsertPoint(LoopBB);
5668
5669 PHINode *Counter = Builder.CreatePHI(Ty: Builder.getInt32Ty(), NumReservedValues: 2);
5670 // size pow2k = 1;
5671 PHINode *Pow2K = Builder.CreatePHI(Ty: Builder.getInt32Ty(), NumReservedValues: 2);
5672 Counter->addIncoming(V: llvm::ConstantInt::get(Ty: Builder.getInt32Ty(), V: 0),
5673 BB: InputBB);
5674 Pow2K->addIncoming(V: llvm::ConstantInt::get(Ty: Builder.getInt32Ty(), V: 1),
5675 BB: InputBB);
5676 // for (size i = n - 1; i >= 2 ^ k; --i)
5677 // tmp[i] op= tmp[i-pow2k];
5678 llvm::BasicBlock *InnerLoopBB =
5679 BasicBlock::Create(Context&: CurFn->getContext(), Name: "omp.inner.log.scan.body");
5680 llvm::BasicBlock *InnerExitBB =
5681 BasicBlock::Create(Context&: CurFn->getContext(), Name: "omp.inner.log.scan.exit");
5682 llvm::Value *CmpI = Builder.CreateICmpUGE(LHS: NMin1, RHS: Pow2K);
5683 Builder.CreateCondBr(Cond: CmpI, True: InnerLoopBB, False: InnerExitBB);
5684 emitBlock(BB: InnerLoopBB, CurFn);
5685 Builder.SetInsertPoint(InnerLoopBB);
5686 PHINode *IVal = Builder.CreatePHI(Ty: Builder.getInt32Ty(), NumReservedValues: 2);
5687 IVal->addIncoming(V: NMin1, BB: LoopBB);
5688 for (ReductionInfo RedInfo : ReductionInfos) {
5689 Value *ReductionVal = RedInfo.PrivateVariable;
5690 Value *BuffPtr = (*(ScanRedInfo->ScanBuffPtrs))[ReductionVal];
5691 Value *Buff = Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: BuffPtr);
5692 Type *DestTy = RedInfo.ElementType;
5693 Value *IV = Builder.CreateAdd(LHS: IVal, RHS: Builder.getInt32(C: 1));
5694 Value *LHSPtr =
5695 Builder.CreateInBoundsGEP(Ty: DestTy, Ptr: Buff, IdxList: IV, Name: "arrayOffset");
5696 Value *OffsetIval = Builder.CreateNUWSub(LHS: IV, RHS: Pow2K);
5697 Value *RHSPtr =
5698 Builder.CreateInBoundsGEP(Ty: DestTy, Ptr: Buff, IdxList: OffsetIval, Name: "arrayOffset");
5699 Value *LHS = Builder.CreateLoad(Ty: DestTy, Ptr: LHSPtr);
5700 Value *RHS = Builder.CreateLoad(Ty: DestTy, Ptr: RHSPtr);
5701 llvm::Value *Result;
5702 InsertPointOrErrorTy AfterIP =
5703 RedInfo.ReductionGen(Builder.saveIP(), LHS, RHS, Result);
5704 if (!AfterIP)
5705 return AfterIP.takeError();
5706 Builder.CreateStore(Val: Result, Ptr: LHSPtr);
5707 }
5708 llvm::Value *NextIVal = Builder.CreateNUWSub(
5709 LHS: IVal, RHS: llvm::ConstantInt::get(Ty: Builder.getInt32Ty(), V: 1));
5710 IVal->addIncoming(V: NextIVal, BB: Builder.GetInsertBlock());
5711 CmpI = Builder.CreateICmpUGE(LHS: NextIVal, RHS: Pow2K);
5712 Builder.CreateCondBr(Cond: CmpI, True: InnerLoopBB, False: InnerExitBB);
5713 emitBlock(BB: InnerExitBB, CurFn);
5714 llvm::Value *Next = Builder.CreateNUWAdd(
5715 LHS: Counter, RHS: llvm::ConstantInt::get(Ty: Counter->getType(), V: 1));
5716 Counter->addIncoming(V: Next, BB: Builder.GetInsertBlock());
5717 // pow2k <<= 1;
5718 llvm::Value *NextPow2K = Builder.CreateShl(LHS: Pow2K, RHS: 1, Name: "", /*HasNUW=*/true);
5719 Pow2K->addIncoming(V: NextPow2K, BB: Builder.GetInsertBlock());
5720 llvm::Value *Cmp = Builder.CreateICmpNE(LHS: Next, RHS: LogVal);
5721 Builder.CreateCondBr(Cond: Cmp, True: LoopBB, False: ExitBB);
5722 Builder.SetInsertPoint(ExitBB->getFirstInsertionPt());
5723 return Error::success();
5724 };
5725
5726 // TODO: Perform finalization actions for variables. This has to be
5727 // called for variables which have destructors/finalizers.
5728 auto FiniCB = [&](InsertPointTy CodeGenIP) { return llvm::Error::success(); };
5729
5730 llvm::Value *FilterVal = Builder.getInt32(C: 0);
5731 llvm::OpenMPIRBuilder::InsertPointOrErrorTy AfterIP =
5732 createMasked(Loc: Builder, BodyGenCB, FiniCB, Filter: FilterVal);
5733
5734 if (!AfterIP)
5735 return AfterIP.takeError();
5736 Builder.restoreIP(IP: *AfterIP);
5737 AfterIP = createBarrier(Loc: Builder, Kind: llvm::omp::OMPD_barrier);
5738
5739 if (!AfterIP)
5740 return AfterIP.takeError();
5741 Builder.restoreIP(IP: *AfterIP);
5742 Error Err = emitScanBasedDirectiveFinalsIR(ReductionInfos, ScanRedInfo);
5743 if (Err)
5744 return Err;
5745
5746 return AfterIP;
5747}
5748
5749Error OpenMPIRBuilder::emitScanBasedDirectiveIR(
5750 llvm::function_ref<Error()> InputLoopGen,
5751 llvm::function_ref<Error(LocationDescription Loc)> ScanLoopGen,
5752 ScanInfo *ScanRedInfo) {
5753
5754 {
5755 // Emit loop with input phase:
5756 // for (i: 0..<num_iters>) {
5757 // <input phase>;
5758 // buffer[i] = red;
5759 // }
5760 ScanRedInfo->OMPFirstScanLoop = true;
5761 Error Err = InputLoopGen();
5762 if (Err)
5763 return Err;
5764 }
5765 {
5766 // Emit loop with scan phase:
5767 // for (i: 0..<num_iters>) {
5768 // red = buffer[i];
5769 // <scan phase>;
5770 // }
5771 ScanRedInfo->OMPFirstScanLoop = false;
5772 Error Err = ScanLoopGen(Builder);
5773 if (Err)
5774 return Err;
5775 }
5776 return Error::success();
5777}
5778
5779void OpenMPIRBuilder::createScanBBs(ScanInfo *ScanRedInfo) {
5780 Function *Fun = Builder.GetInsertBlock()->getParent();
5781 ScanRedInfo->OMPScanDispatch =
5782 BasicBlock::Create(Context&: Fun->getContext(), Name: "omp.inscan.dispatch");
5783 ScanRedInfo->OMPAfterScanBlock =
5784 BasicBlock::Create(Context&: Fun->getContext(), Name: "omp.after.scan.bb");
5785 ScanRedInfo->OMPBeforeScanBlock =
5786 BasicBlock::Create(Context&: Fun->getContext(), Name: "omp.before.scan.bb");
5787 ScanRedInfo->OMPScanLoopExit =
5788 BasicBlock::Create(Context&: Fun->getContext(), Name: "omp.scan.loop.exit");
5789}
5790CanonicalLoopInfo *OpenMPIRBuilder::createLoopSkeleton(
5791 DebugLoc DL, Value *TripCount, Function *F, BasicBlock *PreInsertBefore,
5792 BasicBlock *PostInsertBefore, const Twine &Name, bool IsCollapsed) {
5793 Module *M = F->getParent();
5794 LLVMContext &Ctx = M->getContext();
5795 Type *IndVarTy = TripCount->getType();
5796
5797 // Create the basic block structure.
5798 BasicBlock *Preheader =
5799 BasicBlock::Create(Context&: Ctx, Name: "omp_" + Name + ".preheader", Parent: F, InsertBefore: PreInsertBefore);
5800 BasicBlock *Header =
5801 BasicBlock::Create(Context&: Ctx, Name: "omp_" + Name + ".header", Parent: F, InsertBefore: PreInsertBefore);
5802 BasicBlock *Cond =
5803 BasicBlock::Create(Context&: Ctx, Name: "omp_" + Name + ".cond", Parent: F, InsertBefore: PreInsertBefore);
5804 BasicBlock *Body =
5805 BasicBlock::Create(Context&: Ctx, Name: "omp_" + Name + ".body", Parent: F, InsertBefore: PreInsertBefore);
5806 BasicBlock *Latch =
5807 BasicBlock::Create(Context&: Ctx, Name: "omp_" + Name + ".inc", Parent: F, InsertBefore: PostInsertBefore);
5808 BasicBlock *Exit =
5809 BasicBlock::Create(Context&: Ctx, Name: "omp_" + Name + ".exit", Parent: F, InsertBefore: PostInsertBefore);
5810 BasicBlock *After =
5811 BasicBlock::Create(Context&: Ctx, Name: "omp_" + Name + ".after", Parent: F, InsertBefore: PostInsertBefore);
5812
5813 // Use specified DebugLoc for new instructions.
5814 Builder.SetCurrentDebugLocation(DL);
5815
5816 Builder.SetInsertPoint(Preheader);
5817 Builder.CreateBr(Dest: Header);
5818
5819 Builder.SetInsertPoint(Header);
5820 PHINode *IndVarPHI = Builder.CreatePHI(Ty: IndVarTy, NumReservedValues: 2, Name: "omp_" + Name + ".iv");
5821 IndVarPHI->addIncoming(V: ConstantInt::get(Ty: IndVarTy, V: 0), BB: Preheader);
5822 Builder.CreateBr(Dest: Cond);
5823
5824 Builder.SetInsertPoint(Cond);
5825 Value *Cmp =
5826 Builder.CreateICmpULT(LHS: IndVarPHI, RHS: TripCount, Name: "omp_" + Name + ".cmp");
5827 Builder.CreateCondBr(Cond: Cmp, True: Body, False: Exit);
5828
5829 Builder.SetInsertPoint(Body);
5830 Builder.CreateBr(Dest: Latch);
5831
5832 Builder.SetInsertPoint(Latch);
5833 // Decide whether the induction variable increment can carry nsw.
5834 //
5835 // Single loops: nsw is always kept (matching Clang). Any Fortran program
5836 // whose trip count overflows i32 is non-conforming per F2018 11.1.7.4.1, so
5837 // for valid programs 0 <= count <= INT_MAX always holds.
5838 //
5839 // Collapsed loops: the trip count is a product that can overflow i32 even for
5840 // a conforming program, so nsw is kept only when the product is a constant
5841 // that provably fits, dropped otherwise.
5842 bool HasNSW = Config.hasNoSignedWrap();
5843 if (HasNSW) {
5844 if (auto *CI = dyn_cast<ConstantInt>(Val: TripCount)) {
5845 unsigned BitWidth = CI->getType()->getIntegerBitWidth();
5846 APInt SignedMax = APInt::getSignedMaxValue(numBits: BitWidth);
5847 if (CI->getValue().ugt(RHS: SignedMax))
5848 HasNSW = false;
5849 } else if (IsCollapsed) {
5850 HasNSW = false;
5851 }
5852 }
5853 Value *Next =
5854 Builder.CreateAdd(LHS: IndVarPHI, RHS: ConstantInt::get(Ty: IndVarTy, V: 1),
5855 Name: "omp_" + Name + ".next", /*HasNUW=*/true, HasNSW);
5856 Builder.CreateBr(Dest: Header);
5857 IndVarPHI->addIncoming(V: Next, BB: Latch);
5858
5859 Builder.SetInsertPoint(Exit);
5860 Builder.CreateBr(Dest: After);
5861
5862 // Remember and return the canonical control flow.
5863 LoopInfos.emplace_front();
5864 CanonicalLoopInfo *CL = &LoopInfos.front();
5865
5866 CL->Header = Header;
5867 CL->Cond = Cond;
5868 CL->Latch = Latch;
5869 CL->Exit = Exit;
5870
5871#ifndef NDEBUG
5872 CL->assertOK();
5873#endif
5874 return CL;
5875}
5876
5877Expected<CanonicalLoopInfo *>
5878OpenMPIRBuilder::createCanonicalLoop(const LocationDescription &Loc,
5879 LoopBodyGenCallbackTy BodyGenCB,
5880 Value *TripCount, const Twine &Name) {
5881 BasicBlock *BB = Loc.IP.getNodeParent();
5882 BasicBlock *NextBB = BB->getNextNode();
5883
5884 CanonicalLoopInfo *CL = createLoopSkeleton(DL: Loc.DL, TripCount, F: BB->getParent(),
5885 PreInsertBefore: NextBB, PostInsertBefore: NextBB, Name);
5886 BasicBlock *After = CL->getAfter();
5887
5888 // If location is not set, don't connect the loop.
5889 if (updateToLocation(Loc)) {
5890 // Split the loop at the insertion point: Branch to the preheader and move
5891 // every following instruction to after the loop (the After BB). Also, the
5892 // new successor is the loop's after block.
5893 spliceBB(Builder, New: After, /*CreateBranch=*/false);
5894 Builder.CreateBr(Dest: CL->getPreheader());
5895 }
5896
5897 // Emit the body content. We do it after connecting the loop to the CFG to
5898 // avoid that the callback encounters degenerate BBs.
5899 if (Error Err = BodyGenCB(CL->getBodyIP(), CL->getIndVar()))
5900 return Err;
5901
5902#ifndef NDEBUG
5903 CL->assertOK();
5904#endif
5905 return CL;
5906}
5907
5908Expected<ScanInfo *> OpenMPIRBuilder::scanInfoInitialize() {
5909 ScanInfos.emplace_front();
5910 ScanInfo *Result = &ScanInfos.front();
5911 return Result;
5912}
5913
5914Expected<SmallVector<llvm::CanonicalLoopInfo *>>
5915OpenMPIRBuilder::createCanonicalScanLoops(
5916 const LocationDescription &Loc, LoopBodyGenCallbackTy BodyGenCB,
5917 Value *Start, Value *Stop, Value *Step, bool IsSigned, bool InclusiveStop,
5918 InsertPointTy ComputeIP, const Twine &Name, ScanInfo *ScanRedInfo) {
5919 LocationDescription ComputeLoc =
5920 ComputeIP.isValid() ? LocationDescription(ComputeIP, Loc.DL) : Loc;
5921 updateToLocation(Loc: ComputeLoc);
5922
5923 SmallVector<CanonicalLoopInfo *> Result;
5924
5925 Value *TripCount = calculateCanonicalLoopTripCount(
5926 Loc: ComputeLoc, Start, Stop, Step, IsSigned, InclusiveStop, Name);
5927 ScanRedInfo->Span = TripCount;
5928 ScanRedInfo->OMPScanInit = splitBB(Builder, CreateBranch: true, Name: "scan.init");
5929 Builder.SetInsertPoint(ScanRedInfo->OMPScanInit);
5930
5931 auto BodyGen = [=](InsertPointTy CodeGenIP, Value *IV) {
5932 Builder.restoreIP(IP: CodeGenIP);
5933 ScanRedInfo->IV = IV;
5934 createScanBBs(ScanRedInfo);
5935 BasicBlock *InputBlock = Builder.GetInsertBlock();
5936 Instruction *Terminator = InputBlock->getTerminator();
5937 assert(Terminator->getNumSuccessors() == 1);
5938 BasicBlock *ContinueBlock = Terminator->getSuccessor(Idx: 0);
5939 Terminator->setSuccessor(Idx: 0, BB: ScanRedInfo->OMPScanDispatch);
5940 emitBlock(BB: ScanRedInfo->OMPBeforeScanBlock,
5941 CurFn: Builder.GetInsertBlock()->getParent());
5942 Builder.CreateBr(Dest: ScanRedInfo->OMPScanLoopExit);
5943 emitBlock(BB: ScanRedInfo->OMPScanLoopExit,
5944 CurFn: Builder.GetInsertBlock()->getParent());
5945 Builder.CreateBr(Dest: ContinueBlock);
5946 Builder.SetInsertPoint(
5947 ScanRedInfo->OMPBeforeScanBlock->getFirstInsertionPt());
5948 return BodyGenCB(Builder.saveIP(), IV);
5949 };
5950
5951 const auto &&InputLoopGen = [&]() -> Error {
5952 Expected<CanonicalLoopInfo *> LoopInfo =
5953 createCanonicalLoop(Loc: Builder, BodyGenCB: BodyGen, Start, Stop, Step, IsSigned,
5954 InclusiveStop, ComputeIP, Name, InScan: true, ScanRedInfo);
5955 if (!LoopInfo)
5956 return LoopInfo.takeError();
5957 Result.push_back(Elt: *LoopInfo);
5958 Builder.restoreIP(IP: (*LoopInfo)->getAfterIP());
5959 return Error::success();
5960 };
5961 const auto &&ScanLoopGen = [&](LocationDescription Loc) -> Error {
5962 Expected<CanonicalLoopInfo *> LoopInfo =
5963 createCanonicalLoop(Loc, BodyGenCB: BodyGen, Start, Stop, Step, IsSigned,
5964 InclusiveStop, ComputeIP, Name, InScan: true, ScanRedInfo);
5965 if (!LoopInfo)
5966 return LoopInfo.takeError();
5967 Result.push_back(Elt: *LoopInfo);
5968 Builder.restoreIP(IP: (*LoopInfo)->getAfterIP());
5969 ScanRedInfo->OMPScanFinish = Builder.GetInsertBlock();
5970 return Error::success();
5971 };
5972 Error Err = emitScanBasedDirectiveIR(InputLoopGen, ScanLoopGen, ScanRedInfo);
5973 if (Err)
5974 return Err;
5975 return Result;
5976}
5977
5978Value *OpenMPIRBuilder::calculateCanonicalLoopTripCount(
5979 const LocationDescription &Loc, Value *Start, Value *Stop, Value *Step,
5980 bool IsSigned, bool InclusiveStop, const Twine &Name) {
5981
5982 // Consider the following difficulties (assuming 8-bit signed integers):
5983 // * Adding \p Step to the loop counter which passes \p Stop may overflow:
5984 // DO I = 1, 100, 50
5985 /// * A \p Step of INT_MIN cannot not be normalized to a positive direction:
5986 // DO I = 100, 0, -128
5987
5988 // Start, Stop and Step must be of the same integer type.
5989 auto *IndVarTy = cast<IntegerType>(Val: Start->getType());
5990 assert(IndVarTy == Stop->getType() && "Stop type mismatch");
5991 assert(IndVarTy == Step->getType() && "Step type mismatch");
5992
5993 updateToLocation(Loc);
5994
5995 ConstantInt *Zero = ConstantInt::get(Ty: IndVarTy, V: 0);
5996 ConstantInt *One = ConstantInt::get(Ty: IndVarTy, V: 1);
5997
5998 // Like Step, but always positive.
5999 Value *Incr = Step;
6000
6001 // Distance between Start and Stop; always positive.
6002 Value *Span;
6003
6004 // Condition whether there are no iterations are executed at all, e.g. because
6005 // UB < LB.
6006 Value *ZeroCmp;
6007
6008 if (IsSigned) {
6009 // Ensure that increment is positive. If not, negate and invert LB and UB.
6010 Value *IsNeg = Builder.CreateICmpSLT(LHS: Step, RHS: Zero);
6011 Incr = Builder.CreateSelect(C: IsNeg, True: Builder.CreateNeg(V: Step), False: Step);
6012 Value *LB = Builder.CreateSelect(C: IsNeg, True: Stop, False: Start);
6013 Value *UB = Builder.CreateSelect(C: IsNeg, True: Start, False: Stop);
6014 Span = Builder.CreateSub(LHS: UB, RHS: LB, Name: "", HasNUW: false, HasNSW: true);
6015 ZeroCmp = Builder.CreateICmp(
6016 P: InclusiveStop ? CmpInst::ICMP_SLT : CmpInst::ICMP_SLE, LHS: UB, RHS: LB);
6017 } else {
6018 Span = Builder.CreateSub(LHS: Stop, RHS: Start, Name: "", HasNUW: true);
6019 ZeroCmp = Builder.CreateICmp(
6020 P: InclusiveStop ? CmpInst::ICMP_ULT : CmpInst::ICMP_ULE, LHS: Stop, RHS: Start);
6021 }
6022
6023 Value *CountIfLooping;
6024 if (InclusiveStop) {
6025 CountIfLooping = Builder.CreateAdd(LHS: Builder.CreateUDiv(LHS: Span, RHS: Incr), RHS: One);
6026 } else {
6027 // Avoid incrementing past stop since it could overflow.
6028 Value *CountIfTwo = Builder.CreateAdd(
6029 LHS: Builder.CreateUDiv(LHS: Builder.CreateSub(LHS: Span, RHS: One), RHS: Incr), RHS: One);
6030 Value *OneCmp = Builder.CreateICmp(P: CmpInst::ICMP_ULE, LHS: Span, RHS: Incr);
6031 CountIfLooping = Builder.CreateSelect(C: OneCmp, True: One, False: CountIfTwo);
6032 }
6033
6034 return Builder.CreateSelect(C: ZeroCmp, True: Zero, False: CountIfLooping,
6035 Name: "omp_" + Name + ".tripcount");
6036}
6037
6038Expected<CanonicalLoopInfo *> OpenMPIRBuilder::createCanonicalLoop(
6039 const LocationDescription &Loc, LoopBodyGenCallbackTy BodyGenCB,
6040 Value *Start, Value *Stop, Value *Step, bool IsSigned, bool InclusiveStop,
6041 InsertPointTy ComputeIP, const Twine &Name, bool InScan,
6042 ScanInfo *ScanRedInfo) {
6043 LocationDescription ComputeLoc =
6044 ComputeIP.isValid() ? LocationDescription(ComputeIP, Loc.DL) : Loc;
6045
6046 Value *TripCount = calculateCanonicalLoopTripCount(
6047 Loc: ComputeLoc, Start, Stop, Step, IsSigned, InclusiveStop, Name);
6048
6049 auto BodyGen = [=](InsertPointTy CodeGenIP, Value *IV) {
6050 Builder.restoreIP(IP: CodeGenIP);
6051 Value *Span = Builder.CreateMul(LHS: IV, RHS: Step, Name: "", /*HasNUW=*/false,
6052 /*HasNSW=*/Config.hasNoSignedWrap());
6053 Value *IndVar = Builder.CreateAdd(LHS: Span, RHS: Start, Name: "", /*HasNUW=*/false,
6054 /*HasNSW=*/Config.hasNoSignedWrap());
6055 if (InScan)
6056 ScanRedInfo->IV = IndVar;
6057 return BodyGenCB(Builder.saveIP(), IndVar);
6058 };
6059 LocationDescription LoopLoc =
6060 ComputeIP.isValid()
6061 ? Loc
6062 : LocationDescription(Builder.saveIP(),
6063 Builder.getCurrentDebugLocation());
6064 return createCanonicalLoop(Loc: LoopLoc, BodyGenCB: BodyGen, TripCount, Name);
6065}
6066
6067// Returns an LLVM function to call for initializing loop bounds using OpenMP
6068// static scheduling for composite `distribute parallel for` depending on
6069// `type`. Only i32 and i64 are supported by the runtime. Always interpret
6070// integers as unsigned similarly to CanonicalLoopInfo.
6071static FunctionCallee
6072getKmpcDistForStaticInitForType(Type *Ty, Module &M,
6073 OpenMPIRBuilder &OMPBuilder) {
6074 unsigned Bitwidth = Ty->getIntegerBitWidth();
6075 if (Bitwidth == 32)
6076 return OMPBuilder.getOrCreateRuntimeFunction(
6077 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dist_for_static_init_4u);
6078 if (Bitwidth == 64)
6079 return OMPBuilder.getOrCreateRuntimeFunction(
6080 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dist_for_static_init_8u);
6081 llvm_unreachable("unknown OpenMP loop iterator bitwidth");
6082}
6083
6084// Returns an LLVM function to call for initializing loop bounds using OpenMP
6085// static scheduling depending on `type`. Only i32 and i64 are supported by the
6086// runtime. Always interpret integers as unsigned similarly to
6087// CanonicalLoopInfo.
6088static FunctionCallee getKmpcForStaticInitForType(Type *Ty, Module &M,
6089 OpenMPIRBuilder &OMPBuilder) {
6090 unsigned Bitwidth = Ty->getIntegerBitWidth();
6091 if (Bitwidth == 32)
6092 return OMPBuilder.getOrCreateRuntimeFunction(
6093 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_for_static_init_4u);
6094 if (Bitwidth == 64)
6095 return OMPBuilder.getOrCreateRuntimeFunction(
6096 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_for_static_init_8u);
6097 llvm_unreachable("unknown OpenMP loop iterator bitwidth");
6098}
6099
6100OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::applyStaticWorkshareLoop(
6101 DebugLoc DL, CanonicalLoopInfo *CLI, InsertPointTy AllocaIP,
6102 WorksharingLoopType LoopType, bool NeedsBarrier, bool HasDistSchedule,
6103 OMPScheduleType DistScheduleSchedType) {
6104 assert(CLI->isValid() && "Requires a valid canonical loop");
6105 assert(!isConflictIP(AllocaIP, CLI->getPreheaderIP()) &&
6106 "Require dedicated allocate IP");
6107
6108 // Set up the source location value for OpenMP runtime.
6109 Builder.restoreIP(IP: CLI->getPreheaderIP());
6110 Builder.SetCurrentDebugLocation(DL);
6111
6112 uint32_t SrcLocStrSize;
6113 Constant *SrcLocStr = getOrCreateSrcLocStr(DL, SrcLocStrSize);
6114 IdentFlag Flag = IdentFlag(0);
6115 switch (LoopType) {
6116 case WorksharingLoopType::ForStaticLoop:
6117 Flag = OMP_IDENT_FLAG_WORK_LOOP;
6118 break;
6119 case WorksharingLoopType::DistributeStaticLoop:
6120 Flag = OMP_IDENT_FLAG_WORK_DISTRIBUTE;
6121 break;
6122 case WorksharingLoopType::DistributeForStaticLoop:
6123 Flag = OMP_IDENT_FLAG_WORK_DISTRIBUTE | OMP_IDENT_FLAG_WORK_LOOP;
6124 break;
6125 }
6126 Value *SrcLoc = getOrCreateIdent(SrcLocStr, SrcLocStrSize, LocFlags: Flag);
6127
6128 // Declare useful OpenMP runtime functions.
6129 Value *IV = CLI->getIndVar();
6130 Type *IVTy = IV->getType();
6131 FunctionCallee StaticInit =
6132 LoopType == WorksharingLoopType::DistributeForStaticLoop
6133 ? getKmpcDistForStaticInitForType(Ty: IVTy, M, OMPBuilder&: *this)
6134 : getKmpcForStaticInitForType(Ty: IVTy, M, OMPBuilder&: *this);
6135 FunctionCallee StaticFini =
6136 getOrCreateRuntimeFunction(M, FnID: omp::OMPRTL___kmpc_for_static_fini);
6137
6138 // Allocate space for computed loop bounds as expected by the "init" function.
6139 Builder.SetInsertPoint(
6140 AllocaIP.getNodeParent()->getFirstNonPHIOrDbgOrAlloca());
6141
6142 Type *I32Type = Type::getInt32Ty(C&: M.getContext());
6143 Value *PLastIter = Builder.CreateAlloca(Ty: I32Type, ArraySize: nullptr, Name: "p.lastiter");
6144 Value *PLowerBound = Builder.CreateAlloca(Ty: IVTy, ArraySize: nullptr, Name: "p.lowerbound");
6145 Value *PUpperBound = Builder.CreateAlloca(Ty: IVTy, ArraySize: nullptr, Name: "p.upperbound");
6146 Value *PStride = Builder.CreateAlloca(Ty: IVTy, ArraySize: nullptr, Name: "p.stride");
6147 CLI->setLastIter(PLastIter);
6148
6149 // At the end of the preheader, prepare for calling the "init" function by
6150 // storing the current loop bounds into the allocated space. A canonical loop
6151 // always iterates from 0 to trip-count with step 1. Note that "init" expects
6152 // and produces an inclusive upper bound.
6153 Builder.SetInsertPoint(CLI->getPreheader()->getTerminator());
6154 Constant *Zero = ConstantInt::get(Ty: IVTy, V: 0);
6155 Constant *One = ConstantInt::get(Ty: IVTy, V: 1);
6156 Builder.CreateStore(Val: Zero, Ptr: PLowerBound);
6157 Value *UpperBound = Builder.CreateSub(LHS: CLI->getTripCount(), RHS: One);
6158 Builder.CreateStore(Val: UpperBound, Ptr: PUpperBound);
6159 Builder.CreateStore(Val: One, Ptr: PStride);
6160
6161 Value *ThreadNum =
6162 getOrCreateThreadID(Ident: getOrCreateIdent(SrcLocStr, SrcLocStrSize));
6163
6164 OMPScheduleType SchedType =
6165 (LoopType == WorksharingLoopType::DistributeStaticLoop)
6166 ? OMPScheduleType::OrderedDistribute
6167 : OMPScheduleType::UnorderedStatic;
6168 Constant *SchedulingType =
6169 ConstantInt::get(Ty: I32Type, V: static_cast<int>(SchedType));
6170
6171 // Call the "init" function and update the trip count of the loop with the
6172 // value it produced.
6173 auto BuildInitCall = [LoopType, SrcLoc, ThreadNum, PLastIter, PLowerBound,
6174 PUpperBound, IVTy, PStride, One, Zero, StaticInit,
6175 this](Value *SchedulingType, auto &Builder) {
6176 SmallVector<Value *, 10> Args({SrcLoc, ThreadNum, SchedulingType, PLastIter,
6177 PLowerBound, PUpperBound});
6178 if (LoopType == WorksharingLoopType::DistributeForStaticLoop) {
6179 Value *PDistUpperBound =
6180 Builder.CreateAlloca(IVTy, nullptr, "p.distupperbound");
6181 Args.push_back(Elt: PDistUpperBound);
6182 }
6183 Args.append(IL: {PStride, One, Zero});
6184 createRuntimeFunctionCall(Callee: StaticInit, Args);
6185 };
6186 BuildInitCall(SchedulingType, Builder);
6187 if (HasDistSchedule &&
6188 LoopType != WorksharingLoopType::DistributeStaticLoop) {
6189 Constant *DistScheduleSchedType = ConstantInt::get(
6190 Ty: I32Type, V: static_cast<int>(omp::OMPScheduleType::OrderedDistribute));
6191 // We want to emit a second init function call for the dist_schedule clause
6192 // to the Distribute construct. This should only be done however if a
6193 // Workshare Loop is nested within a Distribute Construct
6194 BuildInitCall(DistScheduleSchedType, Builder);
6195 }
6196 Value *LowerBound = Builder.CreateLoad(Ty: IVTy, Ptr: PLowerBound);
6197 Value *InclusiveUpperBound = Builder.CreateLoad(Ty: IVTy, Ptr: PUpperBound);
6198 Value *TripCountMinusOne = Builder.CreateSub(LHS: InclusiveUpperBound, RHS: LowerBound);
6199 Value *TripCount = Builder.CreateAdd(LHS: TripCountMinusOne, RHS: One);
6200 CLI->setTripCount(TripCount);
6201
6202 // Update all uses of the induction variable except the one in the condition
6203 // block that compares it with the actual upper bound, and the increment in
6204 // the latch block.
6205
6206 CLI->mapIndVar(Updater: [&](Instruction *OldIV) -> Value * {
6207 Builder.SetInsertPoint(CLI->getBody()->getFirstInsertionPt());
6208 Builder.SetCurrentDebugLocation(DL);
6209 return Builder.CreateAdd(LHS: OldIV, RHS: LowerBound, Name: "", /*HasNUW=*/false,
6210 /*HasNSW=*/Config.hasNoSignedWrap());
6211 });
6212
6213 // In the "exit" block, call the "fini" function.
6214 Builder.SetInsertPoint(CLI->getExit()->getTerminator()->getIterator());
6215 createRuntimeFunctionCall(Callee: StaticFini, Args: {SrcLoc, ThreadNum});
6216
6217 // Add the barrier if requested.
6218 if (NeedsBarrier) {
6219 InsertPointOrErrorTy BarrierIP =
6220 createBarrier(Loc: LocationDescription(Builder.saveIP(), DL),
6221 Kind: omp::Directive::OMPD_for, /* ForceSimpleCall */ false,
6222 /* CheckCancelFlag */ false);
6223 if (!BarrierIP)
6224 return BarrierIP.takeError();
6225 }
6226
6227 InsertPointTy AfterIP = CLI->getAfterIP();
6228 CLI->invalidate();
6229
6230 return AfterIP;
6231}
6232
6233static void addAccessGroupMetadata(BasicBlock *Block, MDNode *AccessGroup,
6234 LoopInfo &LI);
6235static void addLoopMetadata(CanonicalLoopInfo *Loop,
6236 ArrayRef<Metadata *> Properties);
6237
6238static void applyParallelAccessesMetadata(CanonicalLoopInfo *CLI,
6239 LLVMContext &Ctx, Loop *Loop,
6240 LoopInfo &LoopInfo,
6241 SmallVector<Metadata *> &LoopMDList) {
6242 SmallSet<BasicBlock *, 8> Reachable;
6243
6244 // Get the basic blocks from the loop in which memref instructions
6245 // can be found.
6246 // TODO: Generalize getting all blocks inside a CanonicalizeLoopInfo,
6247 // preferably without running any passes.
6248 for (BasicBlock *Block : Loop->getBlocks()) {
6249 if (Block == CLI->getCond() || Block == CLI->getHeader())
6250 continue;
6251 Reachable.insert(Ptr: Block);
6252 }
6253
6254 // Add access group metadata to memory-access instructions.
6255 MDNode *AccessGroup = MDNode::getDistinct(Context&: Ctx, MDs: {});
6256 for (BasicBlock *BB : Reachable)
6257 addAccessGroupMetadata(Block: BB, AccessGroup, LI&: LoopInfo);
6258 // TODO: If the loop has existing parallel access metadata, have
6259 // to combine two lists.
6260 LoopMDList.push_back(Elt: MDNode::get(
6261 Context&: Ctx, MDs: {MDString::get(Context&: Ctx, Str: "llvm.loop.parallel_accesses"), AccessGroup}));
6262}
6263
6264OpenMPIRBuilder::InsertPointOrErrorTy
6265OpenMPIRBuilder::applyStaticChunkedWorkshareLoop(
6266 DebugLoc DL, CanonicalLoopInfo *CLI, InsertPointTy AllocaIP,
6267 bool NeedsBarrier, Value *ChunkSize, OMPScheduleType SchedType,
6268 Value *DistScheduleChunkSize, OMPScheduleType DistScheduleSchedType) {
6269 assert(CLI->isValid() && "Requires a valid canonical loop");
6270 assert((ChunkSize || DistScheduleChunkSize) && "Chunk size is required");
6271
6272 LLVMContext &Ctx = CLI->getFunction()->getContext();
6273 Value *IV = CLI->getIndVar();
6274 Value *OrigTripCount = CLI->getTripCount();
6275 Type *IVTy = IV->getType();
6276 assert(IVTy->getIntegerBitWidth() <= 64 &&
6277 "Max supported tripcount bitwidth is 64 bits");
6278 Type *InternalIVTy = IVTy->getIntegerBitWidth() <= 32 ? Type::getInt32Ty(C&: Ctx)
6279 : Type::getInt64Ty(C&: Ctx);
6280 Type *I32Type = Type::getInt32Ty(C&: M.getContext());
6281 Constant *Zero = ConstantInt::get(Ty: InternalIVTy, V: 0);
6282 Constant *One = ConstantInt::get(Ty: InternalIVTy, V: 1);
6283
6284 Function *F = CLI->getFunction();
6285 // Blocks must have terminators.
6286 // FIXME: Don't run analyses on incomplete/invalid IR.
6287 SmallVector<Instruction *> UIs;
6288 for (BasicBlock &BB : *F)
6289 if (!BB.hasTerminator())
6290 UIs.push_back(Elt: new UnreachableInst(F->getContext(), &BB));
6291 FunctionAnalysisManager FAM;
6292 FAM.registerPass(PassBuilder: []() { return DominatorTreeAnalysis(); });
6293 FAM.registerPass(PassBuilder: []() { return PassInstrumentationAnalysis(); });
6294 LoopAnalysis LIA;
6295 LoopInfo &&LI = LIA.run(F&: *F, AM&: FAM);
6296 for (Instruction *I : UIs)
6297 I->eraseFromParent();
6298 Loop *L = LI.getLoopFor(BB: CLI->getHeader());
6299 SmallVector<Metadata *> LoopMDList;
6300 if (ChunkSize || DistScheduleChunkSize)
6301 applyParallelAccessesMetadata(CLI, Ctx, Loop: L, LoopInfo&: LI, LoopMDList);
6302 addLoopMetadata(Loop: CLI, Properties: LoopMDList);
6303
6304 // Declare useful OpenMP runtime functions.
6305 FunctionCallee StaticInit =
6306 getKmpcForStaticInitForType(Ty: InternalIVTy, M, OMPBuilder&: *this);
6307 FunctionCallee StaticFini =
6308 getOrCreateRuntimeFunction(M, FnID: omp::OMPRTL___kmpc_for_static_fini);
6309
6310 // Allocate space for computed loop bounds as expected by the "init" function.
6311 Builder.restoreIP(IP: AllocaIP);
6312 Builder.SetCurrentDebugLocation(DL);
6313 Value *PLastIter = Builder.CreateAlloca(Ty: I32Type, ArraySize: nullptr, Name: "p.lastiter");
6314 Value *PLowerBound =
6315 Builder.CreateAlloca(Ty: InternalIVTy, ArraySize: nullptr, Name: "p.lowerbound");
6316 Value *PUpperBound =
6317 Builder.CreateAlloca(Ty: InternalIVTy, ArraySize: nullptr, Name: "p.upperbound");
6318 Value *PStride = Builder.CreateAlloca(Ty: InternalIVTy, ArraySize: nullptr, Name: "p.stride");
6319 CLI->setLastIter(PLastIter);
6320
6321 // Set up the source location value for the OpenMP runtime.
6322 Builder.restoreIP(IP: CLI->getPreheaderIP());
6323 Builder.SetCurrentDebugLocation(DL);
6324
6325 // TODO: Detect overflow in ubsan or max-out with current tripcount.
6326 Value *CastedChunkSize = Builder.CreateZExtOrTrunc(
6327 V: ChunkSize ? ChunkSize : Zero, DestTy: InternalIVTy, Name: "chunksize");
6328 Value *CastedDistScheduleChunkSize = Builder.CreateZExtOrTrunc(
6329 V: DistScheduleChunkSize ? DistScheduleChunkSize : Zero, DestTy: InternalIVTy,
6330 Name: "distschedulechunksize");
6331 Value *CastedTripCount =
6332 Builder.CreateZExt(V: OrigTripCount, DestTy: InternalIVTy, Name: "tripcount");
6333
6334 Constant *SchedulingType =
6335 ConstantInt::get(Ty: I32Type, V: static_cast<int>(SchedType));
6336 Constant *DistSchedulingType =
6337 ConstantInt::get(Ty: I32Type, V: static_cast<int>(DistScheduleSchedType));
6338 Builder.CreateStore(Val: Zero, Ptr: PLowerBound);
6339 Value *OrigUpperBound = Builder.CreateSub(LHS: CastedTripCount, RHS: One);
6340 Value *IsTripCountZero = Builder.CreateICmpEQ(LHS: CastedTripCount, RHS: Zero);
6341 Value *UpperBound =
6342 Builder.CreateSelect(C: IsTripCountZero, True: Zero, False: OrigUpperBound);
6343 Builder.CreateStore(Val: UpperBound, Ptr: PUpperBound);
6344 Builder.CreateStore(Val: One, Ptr: PStride);
6345
6346 // Call the "init" function and update the trip count of the loop with the
6347 // value it produced.
6348 uint32_t SrcLocStrSize;
6349 Constant *SrcLocStr = getOrCreateSrcLocStr(DL, SrcLocStrSize);
6350 IdentFlag Flag = OMP_IDENT_FLAG_WORK_LOOP;
6351 if (DistScheduleSchedType != OMPScheduleType::None) {
6352 Flag |= OMP_IDENT_FLAG_WORK_DISTRIBUTE;
6353 }
6354 Value *SrcLoc = getOrCreateIdent(SrcLocStr, SrcLocStrSize, LocFlags: Flag);
6355 Value *ThreadNum =
6356 getOrCreateThreadID(Ident: getOrCreateIdent(SrcLocStr, SrcLocStrSize));
6357 auto BuildInitCall = [StaticInit, SrcLoc, ThreadNum, PLastIter, PLowerBound,
6358 PUpperBound, PStride, One,
6359 this](Value *SchedulingType, Value *ChunkSize,
6360 auto &Builder) {
6361 createRuntimeFunctionCall(
6362 Callee: StaticInit, Args: {/*loc=*/SrcLoc, /*global_tid=*/ThreadNum,
6363 /*schedtype=*/SchedulingType, /*plastiter=*/PLastIter,
6364 /*plower=*/PLowerBound, /*pupper=*/PUpperBound,
6365 /*pstride=*/PStride, /*incr=*/One,
6366 /*chunk=*/ChunkSize});
6367 };
6368 BuildInitCall(SchedulingType, CastedChunkSize, Builder);
6369 if (DistScheduleSchedType != OMPScheduleType::None &&
6370 SchedType != OMPScheduleType::OrderedDistributeChunked &&
6371 SchedType != OMPScheduleType::OrderedDistribute) {
6372 // We want to emit a second init function call for the dist_schedule clause
6373 // to the Distribute construct. This should only be done however if a
6374 // Workshare Loop is nested within a Distribute Construct
6375 BuildInitCall(DistSchedulingType, CastedDistScheduleChunkSize, Builder);
6376 }
6377
6378 // Load values written by the "init" function.
6379 Value *FirstChunkStart =
6380 Builder.CreateLoad(Ty: InternalIVTy, Ptr: PLowerBound, Name: "omp_firstchunk.lb");
6381 Value *FirstChunkStop =
6382 Builder.CreateLoad(Ty: InternalIVTy, Ptr: PUpperBound, Name: "omp_firstchunk.ub");
6383 Value *FirstChunkEnd = Builder.CreateAdd(LHS: FirstChunkStop, RHS: One);
6384 Value *ChunkRange =
6385 Builder.CreateSub(LHS: FirstChunkEnd, RHS: FirstChunkStart, Name: "omp_chunk.range");
6386 Value *NextChunkStride =
6387 Builder.CreateLoad(Ty: InternalIVTy, Ptr: PStride, Name: "omp_dispatch.stride");
6388
6389 // Create outer "dispatch" loop for enumerating the chunks.
6390 BasicBlock *DispatchEnter = splitBB(Builder, CreateBranch: true);
6391 Value *DispatchCounter;
6392
6393 // It is safe to assume this didn't return an error because the callback
6394 // passed into createCanonicalLoop is the only possible error source, and it
6395 // always returns success.
6396 CanonicalLoopInfo *DispatchCLI = cantFail(ValOrErr: createCanonicalLoop(
6397 Loc: {Builder.saveIP(), DL},
6398 BodyGenCB: [&](InsertPointTy BodyIP, Value *Counter) {
6399 DispatchCounter = Counter;
6400 return Error::success();
6401 },
6402 Start: FirstChunkStart, Stop: CastedTripCount, Step: NextChunkStride,
6403 /*IsSigned=*/false, /*InclusiveStop=*/false, /*ComputeIP=*/{},
6404 Name: "dispatch"));
6405
6406 // Remember the BasicBlocks of the dispatch loop we need, then invalidate to
6407 // not have to preserve the canonical invariant.
6408 BasicBlock *DispatchBody = DispatchCLI->getBody();
6409 BasicBlock *DispatchLatch = DispatchCLI->getLatch();
6410 BasicBlock *DispatchExit = DispatchCLI->getExit();
6411 BasicBlock *DispatchAfter = DispatchCLI->getAfter();
6412 DispatchCLI->invalidate();
6413
6414 // Rewire the original loop to become the chunk loop inside the dispatch loop.
6415 redirectTo(Source: DispatchAfter, Target: CLI->getAfter(), DL);
6416 redirectTo(Source: CLI->getExit(), Target: DispatchLatch, DL);
6417 redirectTo(Source: DispatchBody, Target: DispatchEnter, DL);
6418
6419 // Prepare the prolog of the chunk loop.
6420 Builder.restoreIP(IP: CLI->getPreheaderIP());
6421 Builder.SetCurrentDebugLocation(DL);
6422
6423 // Compute the number of iterations of the chunk loop.
6424 Builder.SetInsertPoint(CLI->getPreheader()->getTerminator());
6425 Value *ChunkEnd = Builder.CreateAdd(LHS: DispatchCounter, RHS: ChunkRange);
6426 Value *IsLastChunk =
6427 Builder.CreateICmpUGE(LHS: ChunkEnd, RHS: CastedTripCount, Name: "omp_chunk.is_last");
6428 Value *CountUntilOrigTripCount =
6429 Builder.CreateSub(LHS: CastedTripCount, RHS: DispatchCounter);
6430 Value *ChunkTripCount = Builder.CreateSelect(
6431 C: IsLastChunk, True: CountUntilOrigTripCount, False: ChunkRange, Name: "omp_chunk.tripcount");
6432 Value *BackcastedChunkTC =
6433 Builder.CreateTrunc(V: ChunkTripCount, DestTy: IVTy, Name: "omp_chunk.tripcount.trunc");
6434 CLI->setTripCount(BackcastedChunkTC);
6435
6436 // Update all uses of the induction variable except the one in the condition
6437 // block that compares it with the actual upper bound, and the increment in
6438 // the latch block.
6439 Value *BackcastedDispatchCounter =
6440 Builder.CreateTrunc(V: DispatchCounter, DestTy: IVTy, Name: "omp_dispatch.iv.trunc");
6441 CLI->mapIndVar(Updater: [&](Instruction *) -> Value * {
6442 Builder.restoreIP(IP: CLI->getBodyIP());
6443 return Builder.CreateAdd(LHS: IV, RHS: BackcastedDispatchCounter);
6444 });
6445
6446 // In the "exit" block, call the "fini" function.
6447 Builder.SetInsertPoint(DispatchExit->getFirstInsertionPt());
6448 createRuntimeFunctionCall(Callee: StaticFini, Args: {SrcLoc, ThreadNum});
6449
6450 // Add the barrier if requested.
6451 if (NeedsBarrier) {
6452 InsertPointOrErrorTy AfterIP =
6453 createBarrier(Loc: LocationDescription(Builder.saveIP(), DL), Kind: OMPD_for,
6454 /*ForceSimpleCall=*/false, /*CheckCancelFlag=*/false);
6455 if (!AfterIP)
6456 return AfterIP.takeError();
6457 }
6458
6459#ifndef NDEBUG
6460 // Even though we currently do not support applying additional methods to it,
6461 // the chunk loop should remain a canonical loop.
6462 CLI->assertOK();
6463#endif
6464
6465 return DispatchAfter->getFirstInsertionPt();
6466}
6467
6468// Returns an LLVM function to call for executing an OpenMP static worksharing
6469// for loop depending on `type`. Only i32 and i64 are supported by the runtime.
6470// Always interpret integers as unsigned similarly to CanonicalLoopInfo.
6471static FunctionCallee
6472getKmpcForStaticLoopForType(Type *Ty, OpenMPIRBuilder *OMPBuilder,
6473 WorksharingLoopType LoopType) {
6474 unsigned Bitwidth = Ty->getIntegerBitWidth();
6475 Module &M = OMPBuilder->M;
6476 switch (LoopType) {
6477 case WorksharingLoopType::ForStaticLoop:
6478 if (Bitwidth == 32)
6479 return OMPBuilder->getOrCreateRuntimeFunction(
6480 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_for_static_loop_4u);
6481 if (Bitwidth == 64)
6482 return OMPBuilder->getOrCreateRuntimeFunction(
6483 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_for_static_loop_8u);
6484 break;
6485 case WorksharingLoopType::DistributeStaticLoop:
6486 if (Bitwidth == 32)
6487 return OMPBuilder->getOrCreateRuntimeFunction(
6488 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_distribute_static_loop_4u);
6489 if (Bitwidth == 64)
6490 return OMPBuilder->getOrCreateRuntimeFunction(
6491 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_distribute_static_loop_8u);
6492 break;
6493 case WorksharingLoopType::DistributeForStaticLoop:
6494 if (Bitwidth == 32)
6495 return OMPBuilder->getOrCreateRuntimeFunction(
6496 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_distribute_for_static_loop_4u);
6497 if (Bitwidth == 64)
6498 return OMPBuilder->getOrCreateRuntimeFunction(
6499 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_distribute_for_static_loop_8u);
6500 break;
6501 }
6502 if (Bitwidth != 32 && Bitwidth != 64) {
6503 llvm_unreachable("Unknown OpenMP loop iterator bitwidth");
6504 }
6505 llvm_unreachable("Unknown type of OpenMP worksharing loop");
6506}
6507
6508// Inserts a call to proper OpenMP Device RTL function which handles
6509// loop worksharing.
6510static void createTargetLoopWorkshareCall(OpenMPIRBuilder *OMPBuilder,
6511 WorksharingLoopType LoopType,
6512 BasicBlock *InsertBlock, Value *Ident,
6513 Value *LoopBodyArg, Value *TripCount,
6514 Function &LoopBodyFn, bool NoLoop) {
6515 Type *TripCountTy = TripCount->getType();
6516 Module &M = OMPBuilder->M;
6517 IRBuilder<> &Builder = OMPBuilder->Builder;
6518 FunctionCallee RTLFn =
6519 getKmpcForStaticLoopForType(Ty: TripCountTy, OMPBuilder, LoopType);
6520 SmallVector<Value *, 8> RealArgs;
6521 RealArgs.push_back(Elt: Ident);
6522 RealArgs.push_back(Elt: &LoopBodyFn);
6523 RealArgs.push_back(Elt: LoopBodyArg);
6524 RealArgs.push_back(Elt: TripCount);
6525 if (LoopType == WorksharingLoopType::DistributeStaticLoop) {
6526 RealArgs.push_back(Elt: ConstantInt::get(Ty: TripCountTy, V: 0));
6527 RealArgs.push_back(Elt: ConstantInt::get(Ty: Builder.getInt8Ty(), V: 0));
6528 Builder.restoreIP(IP: std::prev(x: InsertBlock->end()));
6529 OMPBuilder->createRuntimeFunctionCall(Callee: RTLFn, Args: RealArgs);
6530 return;
6531 }
6532 FunctionCallee RTLNumThreads = OMPBuilder->getOrCreateRuntimeFunction(
6533 M, FnID: omp::RuntimeFunction::OMPRTL_omp_get_num_threads);
6534 Builder.restoreIP(IP: std::prev(x: InsertBlock->end()));
6535 Value *NumThreads = OMPBuilder->createRuntimeFunctionCall(Callee: RTLNumThreads, Args: {});
6536
6537 RealArgs.push_back(
6538 Elt: Builder.CreateZExtOrTrunc(V: NumThreads, DestTy: TripCountTy, Name: "num.threads.cast"));
6539 RealArgs.push_back(Elt: ConstantInt::get(Ty: TripCountTy, V: 0));
6540 if (LoopType == WorksharingLoopType::DistributeForStaticLoop) {
6541 RealArgs.push_back(Elt: ConstantInt::get(Ty: TripCountTy, V: 0));
6542 RealArgs.push_back(Elt: ConstantInt::get(Ty: Builder.getInt8Ty(), V: NoLoop));
6543 } else {
6544 RealArgs.push_back(Elt: ConstantInt::get(Ty: Builder.getInt8Ty(), V: 0));
6545 }
6546
6547 OMPBuilder->createRuntimeFunctionCall(Callee: RTLFn, Args: RealArgs);
6548}
6549
6550static void workshareLoopTargetCallback(
6551 OpenMPIRBuilder *OMPIRBuilder, CanonicalLoopInfo *CLI, Value *Ident,
6552 Function &OutlinedFn, const SmallVector<Instruction *, 4> &ToBeDeleted,
6553 WorksharingLoopType LoopType, bool NoLoop) {
6554 IRBuilder<> &Builder = OMPIRBuilder->Builder;
6555 BasicBlock *Preheader = CLI->getPreheader();
6556 Value *TripCount = CLI->getTripCount();
6557
6558 // After loop body outling, the loop body contains only set up
6559 // of loop body argument structure and the call to the outlined
6560 // loop body function. Firstly, we need to move setup of loop body args
6561 // into loop preheader.
6562 Preheader->splice(ToIt: std::prev(x: Preheader->end()), FromBB: CLI->getBody(),
6563 FromBeginIt: CLI->getBody()->begin(), FromEndIt: std::prev(x: CLI->getBody()->end()));
6564
6565 // The next step is to remove the whole loop. We do not it need anymore.
6566 // That's why make an unconditional branch from loop preheader to loop
6567 // exit block
6568 Builder.restoreIP(IP: Preheader->end());
6569 Builder.SetCurrentDebugLocation(Preheader->getTerminator()->getDebugLoc());
6570 Preheader->getTerminator()->eraseFromParent();
6571 Builder.CreateBr(Dest: CLI->getExit());
6572
6573 // Delete dead loop blocks
6574 OpenMPIRBuilder::OutlineInfo CleanUpInfo;
6575 SmallPtrSet<BasicBlock *, 32> RegionBlockSet;
6576 SmallVector<BasicBlock *, 32> BlocksToBeRemoved;
6577 CleanUpInfo.EntryBB = CLI->getHeader();
6578 CleanUpInfo.ExitBB = CLI->getExit();
6579 CleanUpInfo.collectBlocks(BlockSet&: RegionBlockSet, BlockVector&: BlocksToBeRemoved);
6580 DeleteDeadBlocks(BBs: BlocksToBeRemoved);
6581
6582 // Find the instruction which corresponds to loop body argument structure
6583 // and remove the call to loop body function instruction.
6584 Value *LoopBodyArg;
6585 User *OutlinedFnUser = OutlinedFn.getUniqueUndroppableUser();
6586 assert(OutlinedFnUser &&
6587 "Expected unique undroppable user of outlined function");
6588 CallInst *OutlinedFnCallInstruction = dyn_cast<CallInst>(Val: OutlinedFnUser);
6589 assert(OutlinedFnCallInstruction && "Expected outlined function call");
6590 assert((OutlinedFnCallInstruction->getParent() == Preheader) &&
6591 "Expected outlined function call to be located in loop preheader");
6592 // Check in case no argument structure has been passed.
6593 if (OutlinedFnCallInstruction->arg_size() > 1)
6594 LoopBodyArg = OutlinedFnCallInstruction->getArgOperand(i: 1);
6595 else
6596 LoopBodyArg = Constant::getNullValue(Ty: Builder.getPtrTy());
6597 OutlinedFnCallInstruction->eraseFromParent();
6598
6599 createTargetLoopWorkshareCall(OMPBuilder: OMPIRBuilder, LoopType, InsertBlock: Preheader, Ident,
6600 LoopBodyArg, TripCount, LoopBodyFn&: OutlinedFn, NoLoop);
6601
6602 for (auto &ToBeDeletedItem : ToBeDeleted)
6603 ToBeDeletedItem->eraseFromParent();
6604 CLI->invalidate();
6605}
6606
6607OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::applyWorkshareLoopTarget(
6608 DebugLoc DL, CanonicalLoopInfo *CLI, InsertPointTy AllocaIP,
6609 WorksharingLoopType LoopType, bool NeedsBarrier, bool NoLoop,
6610 bool NeedsLastIter) {
6611 uint32_t SrcLocStrSize;
6612 Constant *SrcLocStr = getOrCreateSrcLocStr(DL, SrcLocStrSize);
6613
6614 // Mirrors host runtime reporting of last iteration by in-body computation.
6615 if (NeedsLastIter) {
6616 Type *I32Type = Type::getInt32Ty(C&: M.getContext());
6617 Builder.restoreIP(IP: AllocaIP);
6618 AllocaInst *PLastIter =
6619 Builder.CreateAlloca(Ty: I32Type, ArraySize: nullptr, Name: "p.lastiter");
6620 CLI->setLastIter(PLastIter);
6621
6622 Builder.SetInsertPoint(CLI->getPreheader()->getTerminator());
6623 Builder.CreateStore(Val: ConstantInt::get(Ty: I32Type, V: 0), Ptr: PLastIter);
6624
6625 Builder.SetInsertPoint(CLI->getBody()->getFirstInsertionPt());
6626 Value *TripCount = CLI->getTripCount();
6627 Value *LastIter =
6628 Builder.CreateSub(LHS: TripCount, RHS: ConstantInt::get(Ty: TripCount->getType(), V: 1));
6629 Value *IsLast =
6630 Builder.CreateICmpEQ(LHS: CLI->getIndVar(), RHS: LastIter, Name: "omp.is_last_iter");
6631 Builder.CreateStore(Val: Builder.CreateZExt(V: IsLast, DestTy: I32Type), Ptr: PLastIter);
6632 }
6633
6634 IdentFlag Flag = IdentFlag(0);
6635 switch (LoopType) {
6636 case WorksharingLoopType::ForStaticLoop:
6637 Flag = OMP_IDENT_FLAG_WORK_LOOP;
6638 break;
6639 case WorksharingLoopType::DistributeStaticLoop:
6640 Flag = OMP_IDENT_FLAG_WORK_DISTRIBUTE;
6641 break;
6642 case WorksharingLoopType::DistributeForStaticLoop:
6643 Flag = OMP_IDENT_FLAG_WORK_DISTRIBUTE | OMP_IDENT_FLAG_WORK_LOOP;
6644 break;
6645 }
6646 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize, LocFlags: Flag);
6647
6648 auto OI = std::make_unique<OutlineInfo>();
6649 OI->OuterAllocBB = CLI->getPreheader();
6650 Function *OuterFn = CLI->getPreheader()->getParent();
6651
6652 // Instructions which need to be deleted at the end of code generation
6653 SmallVector<Instruction *, 4> ToBeDeleted;
6654
6655 OI->OuterAllocBB = AllocaIP.getNodeParent();
6656
6657 // Mark the body loop as region which needs to be extracted
6658 OI->EntryBB = CLI->getBody();
6659 OI->ExitBB = CLI->getLatch()->splitBasicBlockBefore(I: CLI->getLatch()->begin(),
6660 BBName: "omp.prelatch");
6661
6662 // Prepare loop body for extraction
6663 Builder.restoreIP(IP: CLI->getPreheader()->begin());
6664
6665 // Insert new loop counter variable which will be used only in loop
6666 // body.
6667 AllocaInst *NewLoopCnt = Builder.CreateAlloca(Ty: CLI->getIndVarType(), ArraySize: 0, Name: "");
6668 Instruction *NewLoopCntLoad =
6669 Builder.CreateLoad(Ty: CLI->getIndVarType(), Ptr: NewLoopCnt);
6670 // New loop counter instructions are redundant in the loop preheader when
6671 // code generation for workshare loop is finshed. That's why mark them as
6672 // ready for deletion.
6673 ToBeDeleted.push_back(Elt: NewLoopCntLoad);
6674 ToBeDeleted.push_back(Elt: NewLoopCnt);
6675
6676 // Analyse loop body region. Find all input variables which are used inside
6677 // loop body region.
6678 SmallPtrSet<BasicBlock *, 32> ParallelRegionBlockSet;
6679 SmallVector<BasicBlock *, 32> Blocks;
6680 OI->collectBlocks(BlockSet&: ParallelRegionBlockSet, BlockVector&: Blocks);
6681
6682 CodeExtractorAnalysisCache CEAC(*OuterFn);
6683 CodeExtractor Extractor(Blocks,
6684 /* DominatorTree */ nullptr,
6685 /* AggregateArgs */ true,
6686 /* BlockFrequencyInfo */ nullptr,
6687 /* BranchProbabilityInfo */ nullptr,
6688 /* AssumptionCache */ nullptr,
6689 /* AllowVarArgs */ true,
6690 /* AllowAlloca */ true,
6691 /* AllocationBlock */ CLI->getPreheader(),
6692 /* DeallocationBlocks */ {},
6693 /* Suffix */ ".omp_wsloop",
6694 /* AggrArgsIn0AddrSpace */ true);
6695
6696 BasicBlock *CommonExit = nullptr;
6697 SetVector<Value *> SinkingCands, HoistingCands;
6698
6699 // Find allocas outside the loop body region which are used inside loop
6700 // body
6701 Extractor.findAllocas(CEAC, SinkCands&: SinkingCands, HoistCands&: HoistingCands, ExitBlock&: CommonExit);
6702
6703 // We need to model loop body region as the function f(cnt, loop_arg).
6704 // That's why we replace loop induction variable by the new counter
6705 // which will be one of loop body function argument
6706 SmallVector<User *> Users(CLI->getIndVar()->user_begin(),
6707 CLI->getIndVar()->user_end());
6708 for (auto Use : Users) {
6709 if (Instruction *Inst = dyn_cast<Instruction>(Val: Use)) {
6710 if (ParallelRegionBlockSet.count(Ptr: Inst->getParent())) {
6711 Inst->replaceUsesOfWith(From: CLI->getIndVar(), To: NewLoopCntLoad);
6712 }
6713 }
6714 }
6715 // Make sure that loop counter variable is not merged into loop body
6716 // function argument structure and it is passed as separate variable
6717 OI->ExcludeArgsFromAggregate.push_back(Elt: NewLoopCntLoad);
6718
6719 // PostOutline CB is invoked when loop body function is outlined and
6720 // loop body is replaced by call to outlined function. We need to add
6721 // call to OpenMP device rtl inside loop preheader. OpenMP device rtl
6722 // function will handle loop control logic.
6723 //
6724 OI->PostOutlineCB = [=, ToBeDeletedVec =
6725 std::move(ToBeDeleted)](Function &OutlinedFn) {
6726 workshareLoopTargetCallback(OMPIRBuilder: this, CLI, Ident, OutlinedFn, ToBeDeleted: ToBeDeletedVec,
6727 LoopType, NoLoop);
6728 };
6729 addOutlineInfo(OI: std::move(OI));
6730
6731 // Keep the barrier outside the outlined loop body so that every thread
6732 // encounters it, including threads that execute no iterations.
6733 if (NeedsBarrier) {
6734 Builder.SetInsertPoint(CLI->getExit()->getTerminator());
6735 // Standalone distribute loops never request a barrier. For both regular
6736 // worksharing loops and combined distribute/for loops, the barrier is
6737 // associated with the worksharing loop, hence OMPD_for.
6738 InsertPointOrErrorTy BarrierIP =
6739 createBarrier(Loc: LocationDescription(Builder.saveIP(), DL), Kind: OMPD_for,
6740 /*ForceSimpleCall=*/false, /*CheckCancelFlag=*/false);
6741 if (!BarrierIP)
6742 return BarrierIP.takeError();
6743 }
6744 return CLI->getAfterIP();
6745}
6746
6747OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::applyWorkshareLoop(
6748 DebugLoc DL, CanonicalLoopInfo *CLI, InsertPointTy AllocaIP,
6749 bool NeedsBarrier, omp::ScheduleKind SchedKind, Value *ChunkSize,
6750 bool HasSimdModifier, bool HasMonotonicModifier,
6751 bool HasNonmonotonicModifier, bool HasOrderedClause,
6752 WorksharingLoopType LoopType, bool NoLoop, bool HasDistSchedule,
6753 Value *DistScheduleChunkSize, bool NeedsLastIter) {
6754 if (Config.isTargetDevice())
6755 return applyWorkshareLoopTarget(DL, CLI, AllocaIP, LoopType, NeedsBarrier,
6756 NoLoop, NeedsLastIter);
6757 OMPScheduleType EffectiveScheduleType = computeOpenMPScheduleType(
6758 ClauseKind: SchedKind, HasChunks: ChunkSize, HasSimdModifier, HasMonotonicModifier,
6759 HasNonmonotonicModifier, HasOrderedClause, HasDistScheduleChunks: DistScheduleChunkSize);
6760
6761 bool IsOrdered = (EffectiveScheduleType & OMPScheduleType::ModifierOrdered) ==
6762 OMPScheduleType::ModifierOrdered;
6763 OMPScheduleType DistScheduleSchedType = OMPScheduleType::None;
6764 if (HasDistSchedule) {
6765 DistScheduleSchedType = DistScheduleChunkSize
6766 ? OMPScheduleType::OrderedDistributeChunked
6767 : OMPScheduleType::OrderedDistribute;
6768 }
6769 switch (EffectiveScheduleType & ~OMPScheduleType::ModifierMask) {
6770 case OMPScheduleType::BaseStatic:
6771 case OMPScheduleType::BaseDistribute:
6772 assert((!ChunkSize || !DistScheduleChunkSize) &&
6773 "No chunk size with static-chunked schedule");
6774 if (IsOrdered && !HasDistSchedule)
6775 return applyDynamicWorkshareLoop(DL, CLI, AllocaIP, SchedType: EffectiveScheduleType,
6776 NeedsBarrier, Chunk: ChunkSize);
6777 // FIXME: Monotonicity ignored?
6778 if (DistScheduleChunkSize)
6779 return applyStaticChunkedWorkshareLoop(
6780 DL, CLI, AllocaIP, NeedsBarrier, ChunkSize, SchedType: EffectiveScheduleType,
6781 DistScheduleChunkSize, DistScheduleSchedType);
6782 return applyStaticWorkshareLoop(DL, CLI, AllocaIP, LoopType, NeedsBarrier,
6783 HasDistSchedule);
6784
6785 case OMPScheduleType::BaseStaticChunked:
6786 case OMPScheduleType::BaseDistributeChunked:
6787 if (IsOrdered && !HasDistSchedule)
6788 return applyDynamicWorkshareLoop(DL, CLI, AllocaIP, SchedType: EffectiveScheduleType,
6789 NeedsBarrier, Chunk: ChunkSize);
6790 // FIXME: Monotonicity ignored?
6791 return applyStaticChunkedWorkshareLoop(
6792 DL, CLI, AllocaIP, NeedsBarrier, ChunkSize, SchedType: EffectiveScheduleType,
6793 DistScheduleChunkSize, DistScheduleSchedType);
6794
6795 case OMPScheduleType::BaseRuntime:
6796 case OMPScheduleType::BaseAuto:
6797 case OMPScheduleType::BaseGreedy:
6798 case OMPScheduleType::BaseBalanced:
6799 case OMPScheduleType::BaseSteal:
6800 case OMPScheduleType::BaseRuntimeSimd:
6801 assert(!ChunkSize &&
6802 "schedule type does not support user-defined chunk sizes");
6803 [[fallthrough]];
6804 case OMPScheduleType::BaseGuidedSimd:
6805 case OMPScheduleType::BaseDynamicChunked:
6806 case OMPScheduleType::BaseGuidedChunked:
6807 case OMPScheduleType::BaseGuidedIterativeChunked:
6808 case OMPScheduleType::BaseGuidedAnalyticalChunked:
6809 case OMPScheduleType::BaseStaticBalancedChunked:
6810 return applyDynamicWorkshareLoop(DL, CLI, AllocaIP, SchedType: EffectiveScheduleType,
6811 NeedsBarrier, Chunk: ChunkSize);
6812
6813 default:
6814 llvm_unreachable("Unknown/unimplemented schedule kind");
6815 }
6816}
6817
6818/// Returns an LLVM function to call for initializing loop bounds using OpenMP
6819/// dynamic scheduling depending on `type`. Only i32 and i64 are supported by
6820/// the runtime. Always interpret integers as unsigned similarly to
6821/// CanonicalLoopInfo.
6822static FunctionCallee
6823getKmpcForDynamicInitForType(Type *Ty, Module &M, OpenMPIRBuilder &OMPBuilder) {
6824 unsigned Bitwidth = Ty->getIntegerBitWidth();
6825 if (Bitwidth == 32)
6826 return OMPBuilder.getOrCreateRuntimeFunction(
6827 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dispatch_init_4u);
6828 if (Bitwidth == 64)
6829 return OMPBuilder.getOrCreateRuntimeFunction(
6830 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dispatch_init_8u);
6831 llvm_unreachable("unknown OpenMP loop iterator bitwidth");
6832}
6833
6834/// Returns an LLVM function to call for updating the next loop using OpenMP
6835/// dynamic scheduling depending on `type`. Only i32 and i64 are supported by
6836/// the runtime. Always interpret integers as unsigned similarly to
6837/// CanonicalLoopInfo.
6838static FunctionCallee
6839getKmpcForDynamicNextForType(Type *Ty, Module &M, OpenMPIRBuilder &OMPBuilder) {
6840 unsigned Bitwidth = Ty->getIntegerBitWidth();
6841 if (Bitwidth == 32)
6842 return OMPBuilder.getOrCreateRuntimeFunction(
6843 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dispatch_next_4u);
6844 if (Bitwidth == 64)
6845 return OMPBuilder.getOrCreateRuntimeFunction(
6846 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dispatch_next_8u);
6847 llvm_unreachable("unknown OpenMP loop iterator bitwidth");
6848}
6849
6850/// Returns an LLVM function to call for finalizing the dynamic loop using
6851/// depending on `type`. Only i32 and i64 are supported by the runtime. Always
6852/// interpret integers as unsigned similarly to CanonicalLoopInfo.
6853static FunctionCallee
6854getKmpcForDynamicFiniForType(Type *Ty, Module &M, OpenMPIRBuilder &OMPBuilder) {
6855 unsigned Bitwidth = Ty->getIntegerBitWidth();
6856 if (Bitwidth == 32)
6857 return OMPBuilder.getOrCreateRuntimeFunction(
6858 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dispatch_fini_4u);
6859 if (Bitwidth == 64)
6860 return OMPBuilder.getOrCreateRuntimeFunction(
6861 M, FnID: omp::RuntimeFunction::OMPRTL___kmpc_dispatch_fini_8u);
6862 llvm_unreachable("unknown OpenMP loop iterator bitwidth");
6863}
6864
6865OpenMPIRBuilder::InsertPointOrErrorTy
6866OpenMPIRBuilder::applyDynamicWorkshareLoop(DebugLoc DL, CanonicalLoopInfo *CLI,
6867 InsertPointTy AllocaIP,
6868 OMPScheduleType SchedType,
6869 bool NeedsBarrier, Value *Chunk) {
6870 assert(CLI->isValid() && "Requires a valid canonical loop");
6871 assert(!isConflictIP(AllocaIP, CLI->getPreheaderIP()) &&
6872 "Require dedicated allocate IP");
6873 assert(isValidWorkshareLoopScheduleType(SchedType) &&
6874 "Require valid schedule type");
6875
6876 bool Ordered = (SchedType & OMPScheduleType::ModifierOrdered) ==
6877 OMPScheduleType::ModifierOrdered;
6878
6879 // Set up the source location value for OpenMP runtime.
6880 Builder.SetCurrentDebugLocation(DL);
6881
6882 uint32_t SrcLocStrSize;
6883 Constant *SrcLocStr = getOrCreateSrcLocStr(DL, SrcLocStrSize);
6884 Value *SrcLoc =
6885 getOrCreateIdent(SrcLocStr, SrcLocStrSize, LocFlags: OMP_IDENT_FLAG_WORK_LOOP);
6886
6887 // Declare useful OpenMP runtime functions.
6888 Value *IV = CLI->getIndVar();
6889 Type *IVTy = IV->getType();
6890 FunctionCallee DynamicInit = getKmpcForDynamicInitForType(Ty: IVTy, M, OMPBuilder&: *this);
6891 FunctionCallee DynamicNext = getKmpcForDynamicNextForType(Ty: IVTy, M, OMPBuilder&: *this);
6892
6893 // Allocate space for computed loop bounds as expected by the "init" function.
6894 Builder.SetInsertPoint(
6895 AllocaIP.getNodeParent()->getFirstNonPHIOrDbgOrAlloca());
6896 Type *I32Type = Type::getInt32Ty(C&: M.getContext());
6897 Value *PLastIter = Builder.CreateAlloca(Ty: I32Type, ArraySize: nullptr, Name: "p.lastiter");
6898 Value *PLowerBound = Builder.CreateAlloca(Ty: IVTy, ArraySize: nullptr, Name: "p.lowerbound");
6899 Value *PUpperBound = Builder.CreateAlloca(Ty: IVTy, ArraySize: nullptr, Name: "p.upperbound");
6900 Value *PStride = Builder.CreateAlloca(Ty: IVTy, ArraySize: nullptr, Name: "p.stride");
6901 CLI->setLastIter(PLastIter);
6902
6903 // At the end of the preheader, prepare for calling the "init" function by
6904 // storing the current loop bounds into the allocated space. A canonical loop
6905 // always iterates from 0 to trip-count with step 1. Note that "init" expects
6906 // and produces an inclusive upper bound.
6907 BasicBlock *PreHeader = CLI->getPreheader();
6908 Builder.SetInsertPoint(PreHeader->getTerminator());
6909 Constant *One = ConstantInt::get(Ty: IVTy, V: 1);
6910 Builder.CreateStore(Val: One, Ptr: PLowerBound);
6911 Value *UpperBound = CLI->getTripCount();
6912 Builder.CreateStore(Val: UpperBound, Ptr: PUpperBound);
6913 Builder.CreateStore(Val: One, Ptr: PStride);
6914
6915 BasicBlock *Header = CLI->getHeader();
6916 BasicBlock *Exit = CLI->getExit();
6917 BasicBlock *Cond = CLI->getCond();
6918 BasicBlock *Latch = CLI->getLatch();
6919 InsertPointTy AfterIP = CLI->getAfterIP();
6920
6921 // The CLI will be "broken" in the code below, as the loop is no longer
6922 // a valid canonical loop.
6923
6924 if (!Chunk)
6925 Chunk = One;
6926
6927 Value *ThreadNum =
6928 getOrCreateThreadID(Ident: getOrCreateIdent(SrcLocStr, SrcLocStrSize));
6929
6930 Constant *SchedulingType =
6931 ConstantInt::get(Ty: I32Type, V: static_cast<int>(SchedType));
6932
6933 // Call the "init" function.
6934 createRuntimeFunctionCall(Callee: DynamicInit, Args: {SrcLoc, ThreadNum, SchedulingType,
6935 /* LowerBound */ One, UpperBound,
6936 /* step */ One, Chunk});
6937
6938 // An outer loop around the existing one.
6939 BasicBlock *OuterCond = BasicBlock::Create(
6940 Context&: PreHeader->getContext(), Name: Twine(PreHeader->getName()) + ".outer.cond",
6941 Parent: PreHeader->getParent());
6942 // This needs to be 32-bit always, so can't use the IVTy Zero above.
6943 Builder.SetInsertPoint(OuterCond->getFirstInsertionPt());
6944 Value *Res = createRuntimeFunctionCall(
6945 Callee: DynamicNext,
6946 Args: {SrcLoc, ThreadNum, PLastIter, PLowerBound, PUpperBound, PStride});
6947 Constant *Zero32 = ConstantInt::get(Ty: I32Type, V: 0);
6948 Value *MoreWork = Builder.CreateCmp(Pred: CmpInst::ICMP_NE, LHS: Res, RHS: Zero32);
6949 Value *LowerBound =
6950 Builder.CreateSub(LHS: Builder.CreateLoad(Ty: IVTy, Ptr: PLowerBound), RHS: One, Name: "lb");
6951 Builder.CreateCondBr(Cond: MoreWork, True: Header, False: Exit);
6952
6953 // Change PHI-node in loop header to use outer cond rather than preheader,
6954 // and set IV to the LowerBound.
6955 Instruction *Phi = &Header->front();
6956 auto *PI = cast<PHINode>(Val: Phi);
6957 PI->setIncomingBlock(i: 0, BB: OuterCond);
6958 PI->setIncomingValue(i: 0, V: LowerBound);
6959
6960 // Then set the pre-header to jump to the OuterCond
6961 Instruction *Term = PreHeader->getTerminator();
6962 auto *Br = cast<UncondBrInst>(Val: Term);
6963 Br->setSuccessor(OuterCond);
6964
6965 // Modify the inner condition:
6966 // * Use the UpperBound returned from the DynamicNext call.
6967 // * jump to the loop outer loop when done with one of the inner loops.
6968 Builder.SetInsertPoint(Cond->getFirstInsertionPt());
6969 UpperBound = Builder.CreateLoad(Ty: IVTy, Ptr: PUpperBound, Name: "ub");
6970 Instruction *Comp = &*Builder.GetInsertPoint();
6971 auto *CI = cast<CmpInst>(Val: Comp);
6972 CI->setOperand(i_nocapture: 1, Val_nocapture: UpperBound);
6973 // Redirect the inner exit to branch to outer condition.
6974 Instruction *Branch = &Cond->back();
6975 auto *BI = cast<CondBrInst>(Val: Branch);
6976 assert(BI->getSuccessor(1) == Exit);
6977 BI->setSuccessor(idx: 1, NewSucc: OuterCond);
6978
6979 // Call the "fini" function if "ordered" is present in wsloop directive.
6980 if (Ordered) {
6981 Builder.SetInsertPoint(&Latch->back());
6982 FunctionCallee DynamicFini = getKmpcForDynamicFiniForType(Ty: IVTy, M, OMPBuilder&: *this);
6983 createRuntimeFunctionCall(Callee: DynamicFini, Args: {SrcLoc, ThreadNum});
6984 }
6985
6986 // Add the barrier if requested.
6987 if (NeedsBarrier) {
6988 Builder.SetInsertPoint(&Exit->back());
6989 InsertPointOrErrorTy BarrierIP =
6990 createBarrier(Loc: LocationDescription(Builder.saveIP(), DL),
6991 Kind: omp::Directive::OMPD_for, /* ForceSimpleCall */ false,
6992 /* CheckCancelFlag */ false);
6993 if (!BarrierIP)
6994 return BarrierIP.takeError();
6995 }
6996
6997 CLI->invalidate();
6998 return AfterIP;
6999}
7000
7001/// Redirect all edges that branch to \p OldTarget to \p NewTarget. That is,
7002/// after this \p OldTarget will be orphaned.
7003static void redirectAllPredecessorsTo(BasicBlock *OldTarget,
7004 BasicBlock *NewTarget, DebugLoc DL) {
7005 for (BasicBlock *Pred : make_early_inc_range(Range: predecessors(BB: OldTarget)))
7006 redirectTo(Source: Pred, Target: NewTarget, DL);
7007}
7008
7009static void removeUnusedBlocksFromParent(ArrayRef<BasicBlock *> BBs) {
7010 SmallPtrSet<BasicBlock *, 8> InternalBBs(from_range, BBs);
7011 // We add a block to BBsToKeep iff we have proven it has an external use.
7012 SmallPtrSet<BasicBlock *, 8> BBsToKeep;
7013
7014 while (true) {
7015 bool Changed = false;
7016
7017 for (BasicBlock *BB : BBs) {
7018 if (BBsToKeep.contains(Ptr: BB))
7019 continue;
7020
7021 for (Use &U : BB->uses()) {
7022 auto *UseInst = dyn_cast<Instruction>(Val: U.getUser());
7023 if (!UseInst)
7024 continue;
7025 BasicBlock *UseBB = UseInst->getParent();
7026 if (!InternalBBs.contains(Ptr: UseBB) || BBsToKeep.contains(Ptr: UseBB)) {
7027 BBsToKeep.insert(Ptr: BB);
7028 Changed = true;
7029 break;
7030 }
7031 }
7032 }
7033
7034 if (!Changed)
7035 break;
7036 }
7037
7038 SmallVector<BasicBlock *> BBsToDelete = filter_to_vector(
7039 C&: BBs, Pred: [&BBsToKeep](BasicBlock *BB) { return !BBsToKeep.contains(Ptr: BB); });
7040 DeleteDeadBlocks(BBs: BBsToDelete);
7041}
7042
7043CanonicalLoopInfo *
7044OpenMPIRBuilder::collapseLoops(DebugLoc DL, ArrayRef<CanonicalLoopInfo *> Loops,
7045 InsertPointTy ComputeIP) {
7046 assert(Loops.size() >= 1 && "At least one loop required");
7047 size_t NumLoops = Loops.size();
7048
7049 // Nothing to do if there is already just one loop.
7050 if (NumLoops == 1)
7051 return Loops.front();
7052
7053 CanonicalLoopInfo *Outermost = Loops.front();
7054 CanonicalLoopInfo *Innermost = Loops.back();
7055 BasicBlock *OrigPreheader = Outermost->getPreheader();
7056 BasicBlock *OrigAfter = Outermost->getAfter();
7057 Function *F = OrigPreheader->getParent();
7058
7059 // Loop control blocks that may become orphaned later.
7060 SmallVector<BasicBlock *, 12> OldControlBBs;
7061 OldControlBBs.reserve(N: 6 * Loops.size());
7062 for (CanonicalLoopInfo *Loop : Loops)
7063 Loop->collectControlBlocks(BBs&: OldControlBBs);
7064
7065 // Setup the IRBuilder for inserting the trip count computation.
7066 Builder.SetCurrentDebugLocation(DL);
7067 if (ComputeIP.isValid())
7068 Builder.restoreIP(IP: ComputeIP);
7069 else
7070 Builder.restoreIP(IP: Outermost->getPreheaderIP());
7071
7072 // Derive the collapsed' loop trip count.
7073 // TODO: Find common/largest indvar type.
7074 Value *CollapsedTripCount = nullptr;
7075 for (CanonicalLoopInfo *L : Loops) {
7076 assert(L->isValid() &&
7077 "All loops to collapse must be valid canonical loops");
7078 Value *OrigTripCount = L->getTripCount();
7079 if (!CollapsedTripCount) {
7080 CollapsedTripCount = OrigTripCount;
7081 continue;
7082 }
7083
7084 // TODO: Enable UndefinedSanitizer to diagnose an overflow here.
7085 CollapsedTripCount =
7086 Builder.CreateNUWMul(LHS: CollapsedTripCount, RHS: OrigTripCount);
7087 }
7088
7089 // Create the collapsed loop control flow.
7090 CanonicalLoopInfo *Result =
7091 createLoopSkeleton(DL, TripCount: CollapsedTripCount, F,
7092 PreInsertBefore: OrigPreheader->getNextNode(), PostInsertBefore: OrigAfter, Name: "collapsed",
7093 /*IsCollapsed=*/true);
7094
7095 // Build the collapsed loop body code.
7096 // Start with deriving the input loop induction variables from the collapsed
7097 // one, using a divmod scheme. To preserve the original loops' order, the
7098 // innermost loop use the least significant bits.
7099 Builder.restoreIP(IP: Result->getBodyIP());
7100
7101 Value *Leftover = Result->getIndVar();
7102 SmallVector<Value *> NewIndVars;
7103 NewIndVars.resize(N: NumLoops);
7104 for (int i = NumLoops - 1; i >= 1; --i) {
7105 Value *OrigTripCount = Loops[i]->getTripCount();
7106
7107 Value *NewIndVar = Builder.CreateURem(LHS: Leftover, RHS: OrigTripCount);
7108 NewIndVars[i] = NewIndVar;
7109
7110 Leftover = Builder.CreateUDiv(LHS: Leftover, RHS: OrigTripCount);
7111 }
7112 // Outermost loop gets all the remaining bits.
7113 NewIndVars[0] = Leftover;
7114
7115 // Construct the loop body control flow.
7116 // We progressively construct the branch structure following in direction of
7117 // the control flow, from the leading in-between code, the loop nest body, the
7118 // trailing in-between code, and rejoining the collapsed loop's latch.
7119 // ContinueBlock and ContinuePred keep track of the source(s) of next edge. If
7120 // the ContinueBlock is set, continue with that block. If ContinuePred, use
7121 // its predecessors as sources.
7122 BasicBlock *ContinueBlock = Result->getBody();
7123 BasicBlock *ContinuePred = nullptr;
7124 auto ContinueWith = [&ContinueBlock, &ContinuePred, DL](BasicBlock *Dest,
7125 BasicBlock *NextSrc) {
7126 if (ContinueBlock)
7127 redirectTo(Source: ContinueBlock, Target: Dest, DL);
7128 else
7129 redirectAllPredecessorsTo(OldTarget: ContinuePred, NewTarget: Dest, DL);
7130
7131 ContinueBlock = nullptr;
7132 ContinuePred = NextSrc;
7133 };
7134
7135 // The code before the nested loop of each level.
7136 // Because we are sinking it into the nest, it will be executed more often
7137 // that the original loop. More sophisticated schemes could keep track of what
7138 // the in-between code is and instantiate it only once per thread.
7139 for (size_t i = 0; i < NumLoops - 1; ++i)
7140 ContinueWith(Loops[i]->getBody(), Loops[i + 1]->getHeader());
7141
7142 // Connect the loop nest body.
7143 ContinueWith(Innermost->getBody(), Innermost->getLatch());
7144
7145 // The code after the nested loop at each level.
7146 for (size_t i = NumLoops - 1; i > 0; --i)
7147 ContinueWith(Loops[i]->getAfter(), Loops[i - 1]->getLatch());
7148
7149 // Connect the finished loop to the collapsed loop latch.
7150 ContinueWith(Result->getLatch(), nullptr);
7151
7152 // Replace the input loops with the new collapsed loop.
7153 redirectTo(Source: Outermost->getPreheader(), Target: Result->getPreheader(), DL);
7154 redirectTo(Source: Result->getAfter(), Target: Outermost->getAfter(), DL);
7155
7156 // Replace the input loop indvars with the derived ones.
7157 for (size_t i = 0; i < NumLoops; ++i)
7158 Loops[i]->getIndVar()->replaceAllUsesWith(V: NewIndVars[i]);
7159
7160 // Remove unused parts of the input loops.
7161 removeUnusedBlocksFromParent(BBs: OldControlBBs);
7162
7163 for (CanonicalLoopInfo *L : Loops)
7164 L->invalidate();
7165
7166#ifndef NDEBUG
7167 Result->assertOK();
7168#endif
7169 return Result;
7170}
7171
7172std::vector<CanonicalLoopInfo *>
7173OpenMPIRBuilder::tileLoops(DebugLoc DL, ArrayRef<CanonicalLoopInfo *> Loops,
7174 ArrayRef<Value *> TileSizes) {
7175 assert(TileSizes.size() == Loops.size() &&
7176 "Must pass as many tile sizes as there are loops");
7177 int NumLoops = Loops.size();
7178 assert(NumLoops >= 1 && "At least one loop to tile required");
7179
7180 CanonicalLoopInfo *OutermostLoop = Loops.front();
7181 CanonicalLoopInfo *InnermostLoop = Loops.back();
7182 Function *F = OutermostLoop->getBody()->getParent();
7183 BasicBlock *InnerEnter = InnermostLoop->getBody();
7184 BasicBlock *InnerLatch = InnermostLoop->getLatch();
7185
7186 // Loop control blocks that may become orphaned later.
7187 SmallVector<BasicBlock *, 12> OldControlBBs;
7188 OldControlBBs.reserve(N: 6 * Loops.size());
7189 for (CanonicalLoopInfo *Loop : Loops)
7190 Loop->collectControlBlocks(BBs&: OldControlBBs);
7191
7192 // Collect original trip counts and induction variable to be accessible by
7193 // index. Also, the structure of the original loops is not preserved during
7194 // the construction of the tiled loops, so do it before we scavenge the BBs of
7195 // any original CanonicalLoopInfo.
7196 SmallVector<Value *, 4> OrigTripCounts, OrigIndVars;
7197 for (CanonicalLoopInfo *L : Loops) {
7198 assert(L->isValid() && "All input loops must be valid canonical loops");
7199 OrigTripCounts.push_back(Elt: L->getTripCount());
7200 OrigIndVars.push_back(Elt: L->getIndVar());
7201 }
7202
7203 // Collect the code between loop headers. These may contain SSA definitions
7204 // that are used in the loop nest body. To be usable with in the innermost
7205 // body, these BasicBlocks will be sunk into the loop nest body. That is,
7206 // these instructions may be executed more often than before the tiling.
7207 // TODO: It would be sufficient to only sink them into body of the
7208 // corresponding tile loop.
7209 SmallVector<std::pair<BasicBlock *, BasicBlock *>, 4> InbetweenCode;
7210 for (int i = 0; i < NumLoops - 1; ++i) {
7211 CanonicalLoopInfo *Surrounding = Loops[i];
7212 CanonicalLoopInfo *Nested = Loops[i + 1];
7213
7214 BasicBlock *EnterBB = Surrounding->getBody();
7215 BasicBlock *ExitBB = Nested->getHeader();
7216 InbetweenCode.emplace_back(Args&: EnterBB, Args&: ExitBB);
7217 }
7218
7219 // Compute the trip counts of the floor loops.
7220 Builder.SetCurrentDebugLocation(DL);
7221 Builder.restoreIP(IP: OutermostLoop->getPreheaderIP());
7222 SmallVector<Value *, 4> FloorCompleteCount, FloorCount, FloorRems;
7223 for (int i = 0; i < NumLoops; ++i) {
7224 Value *TileSize = TileSizes[i];
7225 Value *OrigTripCount = OrigTripCounts[i];
7226 Type *IVType = OrigTripCount->getType();
7227
7228 Value *FloorCompleteTripCount = Builder.CreateUDiv(LHS: OrigTripCount, RHS: TileSize);
7229 Value *FloorTripRem = Builder.CreateURem(LHS: OrigTripCount, RHS: TileSize);
7230
7231 // 0 if tripcount divides the tilesize, 1 otherwise.
7232 // 1 means we need an additional iteration for a partial tile.
7233 //
7234 // Unfortunately we cannot just use the roundup-formula
7235 // (tripcount + tilesize - 1)/tilesize
7236 // because the summation might overflow. We do not want introduce undefined
7237 // behavior when the untiled loop nest did not.
7238 Value *FloorTripOverflow =
7239 Builder.CreateICmpNE(LHS: FloorTripRem, RHS: ConstantInt::get(Ty: IVType, V: 0));
7240
7241 FloorTripOverflow = Builder.CreateZExt(V: FloorTripOverflow, DestTy: IVType);
7242 Value *FloorTripCount =
7243 Builder.CreateAdd(LHS: FloorCompleteTripCount, RHS: FloorTripOverflow,
7244 Name: "omp_floor" + Twine(i) + ".tripcount", HasNUW: true);
7245
7246 // Remember some values for later use.
7247 FloorCompleteCount.push_back(Elt: FloorCompleteTripCount);
7248 FloorCount.push_back(Elt: FloorTripCount);
7249 FloorRems.push_back(Elt: FloorTripRem);
7250 }
7251
7252 // Generate the new loop nest, from the outermost to the innermost.
7253 std::vector<CanonicalLoopInfo *> Result;
7254 Result.reserve(n: NumLoops * 2);
7255
7256 // The basic block of the surrounding loop that enters the nest generated
7257 // loop.
7258 BasicBlock *Enter = OutermostLoop->getPreheader();
7259
7260 // The basic block of the surrounding loop where the inner code should
7261 // continue.
7262 BasicBlock *Continue = OutermostLoop->getAfter();
7263
7264 // Where the next loop basic block should be inserted.
7265 BasicBlock *OutroInsertBefore = InnermostLoop->getExit();
7266
7267 auto EmbeddNewLoop =
7268 [this, DL, F, InnerEnter, &Enter, &Continue, &OutroInsertBefore](
7269 Value *TripCount, const Twine &Name) -> CanonicalLoopInfo * {
7270 CanonicalLoopInfo *EmbeddedLoop = createLoopSkeleton(
7271 DL, TripCount, F, PreInsertBefore: InnerEnter, PostInsertBefore: OutroInsertBefore, Name);
7272 redirectTo(Source: Enter, Target: EmbeddedLoop->getPreheader(), DL);
7273 redirectTo(Source: EmbeddedLoop->getAfter(), Target: Continue, DL);
7274
7275 // Setup the position where the next embedded loop connects to this loop.
7276 Enter = EmbeddedLoop->getBody();
7277 Continue = EmbeddedLoop->getLatch();
7278 OutroInsertBefore = EmbeddedLoop->getLatch();
7279 return EmbeddedLoop;
7280 };
7281
7282 auto EmbeddNewLoops = [&Result, &EmbeddNewLoop](ArrayRef<Value *> TripCounts,
7283 const Twine &NameBase) {
7284 for (auto P : enumerate(First&: TripCounts)) {
7285 CanonicalLoopInfo *EmbeddedLoop =
7286 EmbeddNewLoop(P.value(), NameBase + Twine(P.index()));
7287 Result.push_back(x: EmbeddedLoop);
7288 }
7289 };
7290
7291 EmbeddNewLoops(FloorCount, "floor");
7292
7293 // Within the innermost floor loop, emit the code that computes the tile
7294 // sizes.
7295 Builder.SetInsertPoint(Enter->getTerminator());
7296 SmallVector<Value *, 4> TileCounts;
7297 for (int i = 0; i < NumLoops; ++i) {
7298 CanonicalLoopInfo *FloorLoop = Result[i];
7299 Value *TileSize = TileSizes[i];
7300
7301 Value *FloorIsEpilogue =
7302 Builder.CreateICmpEQ(LHS: FloorLoop->getIndVar(), RHS: FloorCompleteCount[i]);
7303 Value *TileTripCount =
7304 Builder.CreateSelect(C: FloorIsEpilogue, True: FloorRems[i], False: TileSize);
7305
7306 TileCounts.push_back(Elt: TileTripCount);
7307 }
7308
7309 // Create the tile loops.
7310 EmbeddNewLoops(TileCounts, "tile");
7311
7312 // Insert the inbetween code into the body.
7313 BasicBlock *BodyEnter = Enter;
7314 BasicBlock *BodyEntered = nullptr;
7315 for (std::pair<BasicBlock *, BasicBlock *> P : InbetweenCode) {
7316 BasicBlock *EnterBB = P.first;
7317 BasicBlock *ExitBB = P.second;
7318
7319 if (BodyEnter)
7320 redirectTo(Source: BodyEnter, Target: EnterBB, DL);
7321 else
7322 redirectAllPredecessorsTo(OldTarget: BodyEntered, NewTarget: EnterBB, DL);
7323
7324 BodyEnter = nullptr;
7325 BodyEntered = ExitBB;
7326 }
7327
7328 // Append the original loop nest body into the generated loop nest body.
7329 if (BodyEnter)
7330 redirectTo(Source: BodyEnter, Target: InnerEnter, DL);
7331 else
7332 redirectAllPredecessorsTo(OldTarget: BodyEntered, NewTarget: InnerEnter, DL);
7333 redirectAllPredecessorsTo(OldTarget: InnerLatch, NewTarget: Continue, DL);
7334
7335 // Replace the original induction variable with an induction variable computed
7336 // from the tile and floor induction variables.
7337 Builder.restoreIP(IP: Result.back()->getBodyIP());
7338 for (int i = 0; i < NumLoops; ++i) {
7339 CanonicalLoopInfo *FloorLoop = Result[i];
7340 CanonicalLoopInfo *TileLoop = Result[NumLoops + i];
7341 Value *OrigIndVar = OrigIndVars[i];
7342 Value *Size = TileSizes[i];
7343
7344 Value *Scale =
7345 Builder.CreateMul(LHS: Size, RHS: FloorLoop->getIndVar(), Name: {}, /*HasNUW=*/true);
7346 Value *Shift =
7347 Builder.CreateAdd(LHS: Scale, RHS: TileLoop->getIndVar(), Name: {}, /*HasNUW=*/true);
7348 OrigIndVar->replaceAllUsesWith(V: Shift);
7349 }
7350
7351 // Remove unused parts of the original loops.
7352 removeUnusedBlocksFromParent(BBs: OldControlBBs);
7353
7354 for (CanonicalLoopInfo *L : Loops)
7355 L->invalidate();
7356
7357#ifndef NDEBUG
7358 for (CanonicalLoopInfo *GenL : Result)
7359 GenL->assertOK();
7360#endif
7361 return Result;
7362}
7363
7364/// Attach metadata \p Properties to the basic block described by \p BB. If the
7365/// basic block already has metadata, the basic block properties are appended.
7366static void addBasicBlockMetadata(BasicBlock *BB,
7367 ArrayRef<Metadata *> Properties) {
7368 // Nothing to do if no property to attach.
7369 if (Properties.empty())
7370 return;
7371
7372 LLVMContext &Ctx = BB->getContext();
7373 SmallVector<Metadata *> NewProperties;
7374 NewProperties.push_back(Elt: nullptr);
7375
7376 // If the basic block already has metadata, prepend it to the new metadata.
7377 MDNode *Existing = BB->getTerminator()->getMetadata(KindID: LLVMContext::MD_loop);
7378 if (Existing)
7379 append_range(C&: NewProperties, R: drop_begin(RangeOrContainer: Existing->operands(), N: 1));
7380
7381 append_range(C&: NewProperties, R&: Properties);
7382 MDNode *BasicBlockID = MDNode::getDistinct(Context&: Ctx, MDs: NewProperties);
7383 BasicBlockID->replaceOperandWith(I: 0, New: BasicBlockID);
7384
7385 BB->getTerminator()->setMetadata(KindID: LLVMContext::MD_loop, Node: BasicBlockID);
7386}
7387
7388/// Attach loop metadata \p Properties to the loop described by \p Loop. If the
7389/// loop already has metadata, the loop properties are appended.
7390static void addLoopMetadata(CanonicalLoopInfo *Loop,
7391 ArrayRef<Metadata *> Properties) {
7392 assert(Loop->isValid() && "Expecting a valid CanonicalLoopInfo");
7393
7394 // Attach metadata to the loop's latch
7395 BasicBlock *Latch = Loop->getLatch();
7396 assert(Latch && "A valid CanonicalLoopInfo must have a unique latch");
7397 addBasicBlockMetadata(BB: Latch, Properties);
7398}
7399
7400/// Attach llvm.access.group metadata to the memref instructions of \p Block
7401static void addAccessGroupMetadata(BasicBlock *Block, MDNode *AccessGroup,
7402 LoopInfo &LI) {
7403 for (Instruction &I : *Block) {
7404 if (I.mayReadOrWriteMemory()) {
7405 // TODO: This instruction may already have access group from
7406 // other pragmas e.g. #pragma clang loop vectorize. Append
7407 // so that the existing metadata is not overwritten.
7408 I.setMetadata(KindID: LLVMContext::MD_access_group, Node: AccessGroup);
7409 }
7410 }
7411}
7412
7413CanonicalLoopInfo *
7414OpenMPIRBuilder::fuseLoops(DebugLoc DL, ArrayRef<CanonicalLoopInfo *> Loops) {
7415 CanonicalLoopInfo *firstLoop = Loops.front();
7416 CanonicalLoopInfo *lastLoop = Loops.back();
7417 Function *F = firstLoop->getPreheader()->getParent();
7418
7419 // Loop control blocks that will become orphaned later
7420 SmallVector<BasicBlock *> oldControlBBs;
7421 for (CanonicalLoopInfo *Loop : Loops)
7422 Loop->collectControlBlocks(BBs&: oldControlBBs);
7423
7424 // Collect original trip counts
7425 SmallVector<Value *> origTripCounts;
7426 for (CanonicalLoopInfo *L : Loops) {
7427 assert(L->isValid() && "All input loops must be valid canonical loops");
7428 origTripCounts.push_back(Elt: L->getTripCount());
7429 }
7430
7431 Builder.SetCurrentDebugLocation(DL);
7432
7433 // Compute max trip count.
7434 // The fused loop will be from 0 to max(origTripCounts)
7435 BasicBlock *TCBlock = BasicBlock::Create(Context&: F->getContext(), Name: "omp.fuse.comp.tc",
7436 Parent: F, InsertBefore: firstLoop->getHeader());
7437 Builder.SetInsertPoint(TCBlock);
7438 Value *fusedTripCount = nullptr;
7439 for (CanonicalLoopInfo *L : Loops) {
7440 assert(L->isValid() && "All loops to fuse must be valid canonical loops");
7441 Value *origTripCount = L->getTripCount();
7442 if (!fusedTripCount) {
7443 fusedTripCount = origTripCount;
7444 continue;
7445 }
7446 Value *condTP = Builder.CreateICmpSGT(LHS: fusedTripCount, RHS: origTripCount);
7447 fusedTripCount = Builder.CreateSelect(C: condTP, True: fusedTripCount, False: origTripCount,
7448 Name: ".omp.fuse.tc");
7449 }
7450
7451 // Generate new loop
7452 CanonicalLoopInfo *fused =
7453 createLoopSkeleton(DL, TripCount: fusedTripCount, F, PreInsertBefore: firstLoop->getBody(),
7454 PostInsertBefore: lastLoop->getLatch(), Name: "fused");
7455
7456 // Replace original loops with the fused loop
7457 // Preheader and After are not considered inside the CLI.
7458 // These are used to compute the individual TCs of the loops
7459 // so they have to be put before the resulting fused loop.
7460 // Moving them up for readability.
7461 for (size_t i = 0; i < Loops.size() - 1; ++i) {
7462 Loops[i]->getPreheader()->moveBefore(MovePos: TCBlock);
7463 Loops[i]->getAfter()->moveBefore(MovePos: TCBlock);
7464 }
7465 lastLoop->getPreheader()->moveBefore(MovePos: TCBlock);
7466
7467 for (size_t i = 0; i < Loops.size() - 1; ++i) {
7468 redirectTo(Source: Loops[i]->getPreheader(), Target: Loops[i]->getAfter(), DL);
7469 redirectTo(Source: Loops[i]->getAfter(), Target: Loops[i + 1]->getPreheader(), DL);
7470 }
7471 redirectTo(Source: lastLoop->getPreheader(), Target: TCBlock, DL);
7472 redirectTo(Source: TCBlock, Target: fused->getPreheader(), DL);
7473 redirectTo(Source: fused->getAfter(), Target: lastLoop->getAfter(), DL);
7474
7475 // Build the fused body
7476 // Create new Blocks with conditions that jump to the original loop bodies
7477 SmallVector<BasicBlock *> condBBs;
7478 SmallVector<Value *> condValues;
7479 for (size_t i = 0; i < Loops.size(); ++i) {
7480 BasicBlock *condBlock = BasicBlock::Create(
7481 Context&: F->getContext(), Name: "omp.fused.inner.cond", Parent: F, InsertBefore: Loops[i]->getBody());
7482 Builder.SetInsertPoint(condBlock);
7483 Value *condValue =
7484 Builder.CreateICmpSLT(LHS: fused->getIndVar(), RHS: origTripCounts[i]);
7485 condBBs.push_back(Elt: condBlock);
7486 condValues.push_back(Elt: condValue);
7487 }
7488 // Join the condition blocks with the bodies of the original loops
7489 redirectTo(Source: fused->getBody(), Target: condBBs[0], DL);
7490 for (size_t i = 0; i < Loops.size() - 1; ++i) {
7491 Builder.SetInsertPoint(condBBs[i]);
7492 Builder.CreateCondBr(Cond: condValues[i], True: Loops[i]->getBody(), False: condBBs[i + 1]);
7493 redirectAllPredecessorsTo(OldTarget: Loops[i]->getLatch(), NewTarget: condBBs[i + 1], DL);
7494 // Replace the IV with the fused IV
7495 Loops[i]->getIndVar()->replaceAllUsesWith(V: fused->getIndVar());
7496 }
7497 // Last body jumps to the created end body block
7498 Builder.SetInsertPoint(condBBs.back());
7499 Builder.CreateCondBr(Cond: condValues.back(), True: lastLoop->getBody(),
7500 False: fused->getLatch());
7501 redirectAllPredecessorsTo(OldTarget: lastLoop->getLatch(), NewTarget: fused->getLatch(), DL);
7502 // Replace the IV with the fused IV
7503 lastLoop->getIndVar()->replaceAllUsesWith(V: fused->getIndVar());
7504
7505 // The loop latch must have only one predecessor. Currently it is branched to
7506 // from both the last condition block and the last loop body
7507 fused->getLatch()->splitBasicBlockBefore(I: fused->getLatch()->begin(),
7508 BBName: "omp.fused.pre_latch");
7509
7510 // Remove unused parts
7511 removeUnusedBlocksFromParent(BBs: oldControlBBs);
7512
7513 // Invalidate old CLIs
7514 for (CanonicalLoopInfo *L : Loops)
7515 L->invalidate();
7516
7517#ifndef NDEBUG
7518 fused->assertOK();
7519#endif
7520 return fused;
7521}
7522
7523void OpenMPIRBuilder::unrollLoopFull(DebugLoc, CanonicalLoopInfo *Loop) {
7524 LLVMContext &Ctx = Builder.getContext();
7525 addLoopMetadata(
7526 Loop, Properties: {MDNode::get(Context&: Ctx, MDs: MDString::get(Context&: Ctx, Str: "llvm.loop.unroll.enable")),
7527 MDNode::get(Context&: Ctx, MDs: MDString::get(Context&: Ctx, Str: "llvm.loop.unroll.full"))});
7528}
7529
7530void OpenMPIRBuilder::unrollLoopHeuristic(DebugLoc, CanonicalLoopInfo *Loop) {
7531 LLVMContext &Ctx = Builder.getContext();
7532 addLoopMetadata(
7533 Loop, Properties: {
7534 MDNode::get(Context&: Ctx, MDs: MDString::get(Context&: Ctx, Str: "llvm.loop.unroll.enable")),
7535 });
7536}
7537
7538void OpenMPIRBuilder::createIfVersion(CanonicalLoopInfo *CanonicalLoop,
7539 Value *IfCond, ValueToValueMapTy &VMap,
7540 LoopAnalysis &LIA, LoopInfo &LI, Loop *L,
7541 const Twine &NamePrefix) {
7542 Function *F = CanonicalLoop->getFunction();
7543
7544 // We can't do
7545 // if (cond) {
7546 // simd_loop;
7547 // } else {
7548 // non_simd_loop;
7549 // }
7550 // because then the CanonicalLoopInfo would only point to one of the loops:
7551 // leading to other constructs operating on the same loop to malfunction.
7552 // Instead generate
7553 // while (...) {
7554 // if (cond) {
7555 // simd_body;
7556 // } else {
7557 // not_simd_body;
7558 // }
7559 // }
7560 // At least for simple loops, LLVM seems able to hoist the if out of the loop
7561 // body at -O3
7562
7563 // Define where if branch should be inserted
7564 auto SplitBeforeIt = CanonicalLoop->getBody()->getFirstNonPHIIt();
7565
7566 // Create additional blocks for the if statement
7567 BasicBlock *Cond = SplitBeforeIt->getParent();
7568 llvm::LLVMContext &C = Cond->getContext();
7569 llvm::BasicBlock *ThenBlock = llvm::BasicBlock::Create(
7570 Context&: C, Name: NamePrefix + ".if.then", Parent: Cond->getParent(), InsertBefore: Cond->getNextNode());
7571 llvm::BasicBlock *ElseBlock = llvm::BasicBlock::Create(
7572 Context&: C, Name: NamePrefix + ".if.else", Parent: Cond->getParent(), InsertBefore: CanonicalLoop->getExit());
7573
7574 // Create if condition branch.
7575 Builder.SetInsertPoint(SplitBeforeIt);
7576 Instruction *BrInstr =
7577 Builder.CreateCondBr(Cond: IfCond, True: ThenBlock, /*ifFalse*/ False: ElseBlock);
7578 InsertPointTy IP(++BrInstr->getIterator());
7579 // Then block contains branch to omp loop body which needs to be vectorized
7580 spliceBB(IP, New: ThenBlock, CreateBranch: false, DL: Builder.getCurrentDebugLocation());
7581 ThenBlock->replaceSuccessorsPhiUsesWith(Old: Cond, New: ThenBlock);
7582
7583 Builder.SetInsertPoint(ElseBlock);
7584
7585 // Clone loop for the else branch
7586 SmallVector<BasicBlock *, 8> NewBlocks;
7587
7588 SmallVector<BasicBlock *, 8> ExistingBlocks;
7589 ExistingBlocks.reserve(N: L->getNumBlocks() + 1);
7590 ExistingBlocks.push_back(Elt: ThenBlock);
7591 ExistingBlocks.append(in_start: L->block_begin(), in_end: L->block_end());
7592 // Cond is the block that has the if clause condition
7593 // LoopCond is omp_loop.cond
7594 // LoopHeader is omp_loop.header
7595 BasicBlock *LoopCond = Cond->getUniquePredecessor();
7596 BasicBlock *LoopHeader = LoopCond->getUniquePredecessor();
7597 assert(LoopCond && LoopHeader && "Invalid loop structure");
7598 for (BasicBlock *Block : ExistingBlocks) {
7599 if (Block == L->getLoopPreheader() || Block == L->getLoopLatch() ||
7600 Block == LoopHeader || Block == LoopCond || Block == Cond) {
7601 continue;
7602 }
7603 BasicBlock *NewBB = CloneBasicBlock(BB: Block, VMap, NameSuffix: "", F);
7604
7605 // fix name not to be omp.if.then
7606 if (Block == ThenBlock)
7607 NewBB->setName(NamePrefix + ".if.else");
7608
7609 NewBB->moveBefore(MovePos: CanonicalLoop->getExit());
7610 VMap[Block] = NewBB;
7611 NewBlocks.push_back(Elt: NewBB);
7612 }
7613 remapInstructionsInBlocks(Blocks: NewBlocks, VMap);
7614 Builder.CreateBr(Dest: NewBlocks.front());
7615
7616 // The loop latch must have only one predecessor. Currently it is branched to
7617 // from both the 'then' and 'else' branches.
7618 L->getLoopLatch()->splitBasicBlockBefore(I: L->getLoopLatch()->begin(),
7619 BBName: NamePrefix + ".pre_latch");
7620
7621 // Ensure that the then block is added to the loop so we add the attributes in
7622 // the next step
7623 L->addBasicBlockToLoop(NewBB: ThenBlock, LI);
7624}
7625
7626unsigned
7627OpenMPIRBuilder::getOpenMPDefaultSimdAlign(const Triple &TargetTriple,
7628 const StringMap<bool> &Features) {
7629 if (TargetTriple.isX86()) {
7630 if (Features.lookup(Key: "avx512f"))
7631 return 512;
7632 else if (Features.lookup(Key: "avx"))
7633 return 256;
7634 return 128;
7635 }
7636 if (TargetTriple.isPPC())
7637 return 128;
7638 if (TargetTriple.isWasm())
7639 return 128;
7640 if (TargetTriple.isSystemZ())
7641 return 64;
7642 return 0;
7643}
7644
7645void OpenMPIRBuilder::applySimd(CanonicalLoopInfo *CanonicalLoop,
7646 MapVector<Value *, Value *> AlignedVars,
7647 Value *IfCond, OrderKind Order,
7648 ConstantInt *Simdlen, ConstantInt *Safelen) {
7649 LLVMContext &Ctx = Builder.getContext();
7650
7651 Function *F = CanonicalLoop->getFunction();
7652
7653 // Blocks must have terminators.
7654 // FIXME: Don't run analyses on incomplete/invalid IR.
7655 SmallVector<Instruction *> UIs;
7656 for (BasicBlock &BB : *F)
7657 if (!BB.hasTerminator())
7658 UIs.push_back(Elt: new UnreachableInst(F->getContext(), &BB));
7659
7660 // TODO: We should not rely on pass manager. Currently we use pass manager
7661 // only for getting llvm::Loop which corresponds to given CanonicalLoopInfo
7662 // object. We should have a method which returns all blocks between
7663 // CanonicalLoopInfo::getHeader() and CanonicalLoopInfo::getAfter()
7664 FunctionAnalysisManager FAM;
7665 FAM.registerPass(PassBuilder: []() { return DominatorTreeAnalysis(); });
7666 FAM.registerPass(PassBuilder: []() { return LoopAnalysis(); });
7667 FAM.registerPass(PassBuilder: []() { return PassInstrumentationAnalysis(); });
7668
7669 LoopAnalysis LIA;
7670 LoopInfo &&LI = LIA.run(F&: *F, AM&: FAM);
7671
7672 for (Instruction *I : UIs)
7673 I->eraseFromParent();
7674
7675 Loop *L = LI.getLoopFor(BB: CanonicalLoop->getHeader());
7676 if (AlignedVars.size()) {
7677 InsertPointTy IP = Builder.saveIP();
7678 for (auto &AlignedItem : AlignedVars) {
7679 Value *AlignedPtr = AlignedItem.first;
7680 Value *Alignment = AlignedItem.second;
7681 Instruction *loadInst = dyn_cast<Instruction>(Val: AlignedPtr);
7682 Builder.SetInsertPoint(loadInst->getNextNode());
7683 Builder.CreateAlignmentAssumption(DL: F->getDataLayout(), PtrValue: AlignedPtr,
7684 Alignment);
7685 }
7686 Builder.restoreIP(IP);
7687 }
7688
7689 if (IfCond) {
7690 ValueToValueMapTy VMap;
7691 createIfVersion(CanonicalLoop, IfCond, VMap, LIA, LI, L, NamePrefix: "simd");
7692 }
7693
7694 SmallPtrSet<BasicBlock *, 8> Reachable;
7695
7696 // Get the basic blocks from the loop in which memref instructions
7697 // can be found.
7698 // TODO: Generalize getting all blocks inside a CanonicalizeLoopInfo,
7699 // preferably without running any passes.
7700 for (BasicBlock *Block : L->getBlocks()) {
7701 if (Block == CanonicalLoop->getCond() ||
7702 Block == CanonicalLoop->getHeader())
7703 continue;
7704 Reachable.insert(Ptr: Block);
7705 }
7706
7707 SmallVector<Metadata *> LoopMDList;
7708
7709 // In presence of finite 'safelen', it may be unsafe to mark all
7710 // the memory instructions parallel, because loop-carried
7711 // dependences of 'safelen' iterations are possible.
7712 // If clause order(concurrent) is specified then the memory instructions
7713 // are marked parallel even if 'safelen' is finite.
7714 if ((Safelen == nullptr) || (Order == OrderKind::OMP_ORDER_concurrent))
7715 applyParallelAccessesMetadata(CLI: CanonicalLoop, Ctx, Loop: L, LoopInfo&: LI, LoopMDList);
7716
7717 // FIXME: the IF clause shares a loop backedge for the SIMD and non-SIMD
7718 // versions so we can't add the loop attributes in that case.
7719 if (IfCond) {
7720 // we can still add llvm.loop.parallel_access
7721 addLoopMetadata(Loop: CanonicalLoop, Properties: LoopMDList);
7722 return;
7723 }
7724
7725 // Use the above access group metadata to create loop level
7726 // metadata, which should be distinct for each loop.
7727 LoopMDList.push_back(
7728 Elt: MDNode::get(Context&: Ctx, MDs: {MDString::get(Context&: Ctx, Str: "llvm.loop.vectorize.enable")}));
7729
7730 if (Simdlen || Safelen) {
7731 // If both simdlen and safelen clauses are specified, the value of the
7732 // simdlen parameter must be less than or equal to the value of the safelen
7733 // parameter. Therefore, use safelen only in the absence of simdlen.
7734 ConstantInt *VectorizeWidth = Simdlen == nullptr ? Safelen : Simdlen;
7735 LoopMDList.push_back(
7736 Elt: MDNode::get(Context&: Ctx, MDs: {MDString::get(Context&: Ctx, Str: "llvm.loop.vectorize.width"),
7737 ConstantAsMetadata::get(C: VectorizeWidth)}));
7738 }
7739
7740 addLoopMetadata(Loop: CanonicalLoop, Properties: LoopMDList);
7741}
7742
7743/// Create the TargetMachine object to query the backend for optimization
7744/// preferences.
7745///
7746/// Ideally, this would be passed from the front-end to the OpenMPBuilder, but
7747/// e.g. Clang does not pass it to its CodeGen layer and creates it only when
7748/// needed for the LLVM pass pipline. We use some default options to avoid
7749/// having to pass too many settings from the frontend that probably do not
7750/// matter.
7751///
7752/// Currently, TargetMachine is only used sometimes by the unrollLoopPartial
7753/// method. If we are going to use TargetMachine for more purposes, especially
7754/// those that are sensitive to TargetOptions, RelocModel and CodeModel, it
7755/// might become be worth requiring front-ends to pass on their TargetMachine,
7756/// or at least cache it between methods. Note that while fontends such as Clang
7757/// have just a single main TargetMachine per translation unit, "target-cpu" and
7758/// "target-features" that determine the TargetMachine are per-function and can
7759/// be overrided using __attribute__((target("OPTIONS"))).
7760static std::unique_ptr<TargetMachine>
7761createTargetMachine(Function *F, CodeGenOptLevel OptLevel) {
7762 Module *M = F->getParent();
7763
7764 StringRef CPU = F->getFnAttribute(Kind: "target-cpu").getValueAsString();
7765 StringRef Features = F->getFnAttribute(Kind: "target-features").getValueAsString();
7766 const llvm::Triple &Triple = M->getTargetTriple();
7767
7768 std::string Error;
7769 const llvm::Target *TheTarget = TargetRegistry::lookupTarget(TheTriple: Triple, Error);
7770 if (!TheTarget)
7771 return {};
7772
7773 llvm::TargetOptions Options;
7774 return std::unique_ptr<TargetMachine>(TheTarget->createTargetMachine(
7775 TT: Triple, CPU, Features, Options, /*RelocModel=*/RM: std::nullopt,
7776 /*CodeModel=*/CM: std::nullopt, OL: OptLevel));
7777}
7778
7779/// Heuristically determine the best-performant unroll factor for \p CLI. This
7780/// depends on the target processor. We are re-using the same heuristics as the
7781/// LoopUnrollPass.
7782static int32_t computeHeuristicUnrollFactor(CanonicalLoopInfo *CLI) {
7783 Function *F = CLI->getFunction();
7784
7785 // Assume the user requests the most aggressive unrolling, even if the rest of
7786 // the code is optimized using a lower setting.
7787 CodeGenOptLevel OptLevel = CodeGenOptLevel::Aggressive;
7788 std::unique_ptr<TargetMachine> TM = createTargetMachine(F, OptLevel);
7789
7790 // Blocks must have terminators.
7791 // FIXME: Don't run analyses on incomplete/invalid IR.
7792 SmallVector<Instruction *> UIs;
7793 for (BasicBlock &BB : *F)
7794 if (!BB.hasTerminator())
7795 UIs.push_back(Elt: new UnreachableInst(F->getContext(), &BB));
7796
7797 FunctionAnalysisManager FAM;
7798 FAM.registerPass(PassBuilder: []() { return TargetLibraryAnalysis(); });
7799 FAM.registerPass(PassBuilder: []() { return AssumptionAnalysis(); });
7800 FAM.registerPass(PassBuilder: []() { return DominatorTreeAnalysis(); });
7801 FAM.registerPass(PassBuilder: []() { return LoopAnalysis(); });
7802 FAM.registerPass(PassBuilder: []() { return ScalarEvolutionAnalysis(); });
7803 FAM.registerPass(PassBuilder: []() { return PassInstrumentationAnalysis(); });
7804 TargetIRAnalysis TIRA;
7805 if (TM)
7806 TIRA = TargetIRAnalysis(
7807 [&](const Function &F) { return TM->getTargetTransformInfo(F); });
7808 FAM.registerPass(PassBuilder: [&]() { return TIRA; });
7809
7810 TargetIRAnalysis::Result &&TTI = TIRA.run(F: *F, FAM);
7811 ScalarEvolutionAnalysis SEA;
7812 ScalarEvolution &&SE = SEA.run(F&: *F, AM&: FAM);
7813 DominatorTreeAnalysis DTA;
7814 DominatorTree &&DT = DTA.run(F&: *F, FAM);
7815 LoopAnalysis LIA;
7816 LoopInfo &&LI = LIA.run(F&: *F, AM&: FAM);
7817 AssumptionAnalysis ACT;
7818 AssumptionCache &&AC = ACT.run(F&: *F, FAM);
7819 OptimizationRemarkEmitter ORE{F};
7820
7821 for (Instruction *I : UIs)
7822 I->eraseFromParent();
7823
7824 Loop *L = LI.getLoopFor(BB: CLI->getHeader());
7825 assert(L && "Expecting CanonicalLoopInfo to be recognized as a loop");
7826
7827 TargetTransformInfo::UnrollingPreferences UP = gatherUnrollingPreferences(
7828 L, SE, TTI,
7829 /*BlockFrequencyInfo=*/BFI: nullptr,
7830 /*ProfileSummaryInfo=*/PSI: nullptr, ORE, OptLevel: static_cast<int>(OptLevel),
7831 /*UserThreshold=*/std::nullopt,
7832 /*UserAllowPartial=*/true,
7833 /*UserAllowRuntime=*/UserRuntime: true,
7834 /*UserUpperBound=*/std::nullopt,
7835 /*UserFullUnrollMaxCount=*/std::nullopt);
7836
7837 UP.Force = true;
7838
7839 // Account for additional optimizations taking place before the LoopUnrollPass
7840 // would unroll the loop.
7841 UP.Threshold *= UnrollThresholdFactor;
7842 UP.PartialThreshold *= UnrollThresholdFactor;
7843
7844 // Use normal unroll factors even if the rest of the code is optimized for
7845 // size.
7846 UP.OptSizeThreshold = UP.Threshold;
7847 UP.PartialOptSizeThreshold = UP.PartialThreshold;
7848
7849 LLVM_DEBUG(dbgs() << "Unroll heuristic thresholds:\n"
7850 << " Threshold=" << UP.Threshold << "\n"
7851 << " PartialThreshold=" << UP.PartialThreshold << "\n"
7852 << " OptSizeThreshold=" << UP.OptSizeThreshold << "\n"
7853 << " PartialOptSizeThreshold="
7854 << UP.PartialOptSizeThreshold << "\n");
7855
7856 // Disable peeling.
7857 TargetTransformInfo::PeelingPreferences PP =
7858 gatherPeelingPreferences(L, SE, TTI,
7859 /*UserAllowPeeling=*/false,
7860 /*UserAllowProfileBasedPeeling=*/false,
7861 /*UnrollingSpecficValues=*/false);
7862
7863 SmallPtrSet<const Value *, 32> EphValues;
7864 CodeMetrics::collectEphemeralValues(L, AC: &AC, EphValues);
7865
7866 // Assume that reads and writes to stack variables can be eliminated by
7867 // Mem2Reg, SROA or LICM. That is, don't count them towards the loop body's
7868 // size.
7869 for (BasicBlock *BB : L->blocks()) {
7870 for (Instruction &I : *BB) {
7871 Value *Ptr;
7872 if (auto *Load = dyn_cast<LoadInst>(Val: &I)) {
7873 Ptr = Load->getPointerOperand();
7874 } else if (auto *Store = dyn_cast<StoreInst>(Val: &I)) {
7875 Ptr = Store->getPointerOperand();
7876 } else
7877 continue;
7878
7879 Ptr = Ptr->stripPointerCasts();
7880
7881 if (auto *Alloca = dyn_cast<AllocaInst>(Val: Ptr)) {
7882 if (Alloca->getParent() == &F->getEntryBlock())
7883 EphValues.insert(Ptr: &I);
7884 }
7885 }
7886 }
7887
7888 UnrollCostEstimator UCE(L, TTI, EphValues, UP.BEInsns);
7889
7890 // Loop is not unrollable if the loop contains certain instructions.
7891 if (!UCE.canUnroll()) {
7892 LLVM_DEBUG(dbgs() << "Loop not considered unrollable\n");
7893 return 1;
7894 }
7895
7896 LLVM_DEBUG(dbgs() << "Estimated loop size is " << UCE.getRolledLoopSize()
7897 << "\n");
7898
7899 // TODO: Determine trip count of \p CLI if constant, computeUnrollCount might
7900 // be able to use it.
7901 int TripCount = 0;
7902 int MaxTripCount = 0;
7903 bool MaxOrZero = false;
7904 unsigned TripMultiple = 0;
7905
7906 unsigned Factor =
7907 computeUnrollCount(L, TTI, DT, LI: &LI, AC: &AC, SE, EphValues, ORE: &ORE, TripCount,
7908 MaxTripCount, MaxOrZero, TripMultiple, UCE, UP, PP);
7909 LLVM_DEBUG(dbgs() << "Suggesting unroll factor of " << Factor << "\n");
7910
7911 // This function returns 1 to signal to not unroll a loop.
7912 if (Factor == 0)
7913 return 1;
7914 return Factor;
7915}
7916
7917void OpenMPIRBuilder::unrollLoopPartial(DebugLoc DL, CanonicalLoopInfo *Loop,
7918 int32_t Factor,
7919 CanonicalLoopInfo **UnrolledCLI) {
7920 assert(Factor >= 0 && "Unroll factor must not be negative");
7921
7922 Function *F = Loop->getFunction();
7923 LLVMContext &Ctx = F->getContext();
7924
7925 // If the unrolled loop is not used for another loop-associated directive, it
7926 // is sufficient to add metadata for the LoopUnrollPass.
7927 if (!UnrolledCLI) {
7928 SmallVector<Metadata *, 2> LoopMetadata;
7929 LoopMetadata.push_back(
7930 Elt: MDNode::get(Context&: Ctx, MDs: MDString::get(Context&: Ctx, Str: "llvm.loop.unroll.enable")));
7931
7932 if (Factor >= 1) {
7933 ConstantAsMetadata *FactorConst = ConstantAsMetadata::get(
7934 C: ConstantInt::get(Ty: Type::getInt32Ty(C&: Ctx), V: APInt(32, Factor)));
7935 LoopMetadata.push_back(Elt: MDNode::get(
7936 Context&: Ctx, MDs: {MDString::get(Context&: Ctx, Str: "llvm.loop.unroll.count"), FactorConst}));
7937 }
7938
7939 addLoopMetadata(Loop, Properties: LoopMetadata);
7940 return;
7941 }
7942
7943 // Heuristically determine the unroll factor.
7944 if (Factor == 0)
7945 Factor = computeHeuristicUnrollFactor(CLI: Loop);
7946
7947 // No change required with unroll factor 1.
7948 if (Factor == 1) {
7949 *UnrolledCLI = Loop;
7950 return;
7951 }
7952
7953 assert(Factor >= 2 &&
7954 "unrolling only makes sense with a factor of 2 or larger");
7955
7956 Type *IndVarTy = Loop->getIndVarType();
7957
7958 // Apply partial unrolling by tiling the loop by the unroll-factor, then fully
7959 // unroll the inner loop.
7960 Value *FactorVal =
7961 ConstantInt::get(Ty: IndVarTy, V: APInt(IndVarTy->getIntegerBitWidth(), Factor,
7962 /*isSigned=*/false));
7963 std::vector<CanonicalLoopInfo *> LoopNest =
7964 tileLoops(DL, Loops: {Loop}, TileSizes: {FactorVal});
7965 assert(LoopNest.size() == 2 && "Expect 2 loops after tiling");
7966 *UnrolledCLI = LoopNest[0];
7967 CanonicalLoopInfo *InnerLoop = LoopNest[1];
7968
7969 // LoopUnrollPass can only fully unroll loops with constant trip count.
7970 // Unroll by the unroll factor with a fallback epilog for the remainder
7971 // iterations if necessary.
7972 ConstantAsMetadata *FactorConst = ConstantAsMetadata::get(
7973 C: ConstantInt::get(Ty: Type::getInt32Ty(C&: Ctx), V: APInt(32, Factor)));
7974 addLoopMetadata(
7975 Loop: InnerLoop,
7976 Properties: {MDNode::get(Context&: Ctx, MDs: MDString::get(Context&: Ctx, Str: "llvm.loop.unroll.enable")),
7977 MDNode::get(
7978 Context&: Ctx, MDs: {MDString::get(Context&: Ctx, Str: "llvm.loop.unroll.count"), FactorConst})});
7979
7980#ifndef NDEBUG
7981 (*UnrolledCLI)->assertOK();
7982#endif
7983}
7984
7985OpenMPIRBuilder::InsertPointTy
7986OpenMPIRBuilder::createCopyPrivate(const LocationDescription &Loc,
7987 llvm::Value *BufSize, llvm::Value *CpyBuf,
7988 llvm::Value *CpyFn, llvm::Value *DidIt) {
7989 if (!updateToLocation(Loc))
7990 return Loc.IP;
7991
7992 uint32_t SrcLocStrSize;
7993 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
7994 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
7995 Value *ThreadId = getOrCreateThreadID(Ident);
7996
7997 llvm::Value *DidItLD = Builder.CreateLoad(Ty: Builder.getInt32Ty(), Ptr: DidIt);
7998
7999 Value *Args[] = {Ident, ThreadId, BufSize, CpyBuf, CpyFn, DidItLD};
8000
8001 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_copyprivate);
8002 createRuntimeFunctionCall(Callee: Fn, Args);
8003
8004 return Builder.saveIP();
8005}
8006
8007OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createSingle(
8008 const LocationDescription &Loc, BodyGenCallbackTy BodyGenCB,
8009 FinalizeCallbackTy FiniCB, bool IsNowait, ArrayRef<llvm::Value *> CPVars,
8010 ArrayRef<llvm::Function *> CPFuncs) {
8011
8012 if (!updateToLocation(Loc))
8013 return Loc.IP;
8014
8015 // If needed allocate and initialize `DidIt` with 0.
8016 // DidIt: flag variable: 1=single thread; 0=not single thread.
8017 llvm::Value *DidIt = nullptr;
8018 if (!CPVars.empty()) {
8019 DidIt = Builder.CreateAlloca(Ty: llvm::Type::getInt32Ty(C&: Builder.getContext()));
8020 Builder.CreateStore(Val: Builder.getInt32(C: 0), Ptr: DidIt);
8021 }
8022
8023 Directive OMPD = Directive::OMPD_single;
8024 uint32_t SrcLocStrSize;
8025 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8026 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8027 Value *ThreadId = getOrCreateThreadID(Ident);
8028 Value *Args[] = {Ident, ThreadId};
8029
8030 Function *EntryRTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_single);
8031 Instruction *EntryCall = createRuntimeFunctionCall(Callee: EntryRTLFn, Args);
8032
8033 Function *ExitRTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_end_single);
8034 Instruction *ExitCall = createRuntimeFunctionCall(Callee: ExitRTLFn, Args);
8035
8036 auto FiniCBWrapper = [&](InsertPointTy IP) -> Error {
8037 if (Error Err = FiniCB(IP))
8038 return Err;
8039
8040 // The thread that executes the single region must set `DidIt` to 1.
8041 // This is used by __kmpc_copyprivate, to know if the caller is the
8042 // single thread or not.
8043 if (DidIt)
8044 Builder.CreateStore(Val: Builder.getInt32(C: 1), Ptr: DidIt);
8045
8046 return Error::success();
8047 };
8048
8049 // generates the following:
8050 // if (__kmpc_single()) {
8051 // .... single region ...
8052 // __kmpc_end_single
8053 // }
8054 // __kmpc_copyprivate
8055 // __kmpc_barrier
8056
8057 InsertPointOrErrorTy AfterIP =
8058 EmitOMPInlinedRegion(OMPD, EntryCall, ExitCall, BodyGenCB, FiniCB: FiniCBWrapper,
8059 /*Conditional*/ true,
8060 /*hasFinalize*/ HasFinalize: true);
8061 if (!AfterIP)
8062 return AfterIP.takeError();
8063
8064 if (DidIt) {
8065 for (size_t I = 0, E = CPVars.size(); I < E; ++I)
8066 // NOTE BufSize is currently unused, so just pass 0.
8067 createCopyPrivate(Loc: LocationDescription(Builder.saveIP(), Loc.DL),
8068 /*BufSize=*/ConstantInt::get(Ty: Int64, V: 0), CpyBuf: CPVars[I],
8069 CpyFn: CPFuncs[I], DidIt);
8070 // NOTE __kmpc_copyprivate already inserts a barrier
8071 } else if (!IsNowait) {
8072 InsertPointOrErrorTy AfterIP =
8073 createBarrier(Loc: LocationDescription(Builder.saveIP(), Loc.DL),
8074 Kind: omp::Directive::OMPD_unknown, /* ForceSimpleCall */ false,
8075 /* CheckCancelFlag */ false);
8076 if (!AfterIP)
8077 return AfterIP.takeError();
8078 }
8079 return Builder.saveIP();
8080}
8081
8082OpenMPIRBuilder::InsertPointOrErrorTy
8083OpenMPIRBuilder::createScope(const LocationDescription &Loc,
8084 BodyGenCallbackTy BodyGenCB,
8085 FinalizeCallbackTy FiniCB, bool IsNowait) {
8086
8087 if (!updateToLocation(Loc))
8088 return Loc.IP;
8089
8090 // All threads execute the scope body — no conditional entry.
8091 InsertPointOrErrorTy AfterIP = EmitOMPInlinedRegion(
8092 OMPD: Directive::OMPD_scope, /*EntryCall=*/nullptr, /*ExitCall=*/nullptr,
8093 BodyGenCB, FiniCB, /*Conditional=*/false, /*HasFinalize=*/true,
8094 /*IsCancellable=*/false);
8095 if (!AfterIP)
8096 return AfterIP.takeError();
8097
8098 Builder.restoreIP(IP: *AfterIP);
8099 if (!IsNowait) {
8100 AfterIP = createBarrier(Loc: LocationDescription(Builder.saveIP(), Loc.DL),
8101 Kind: omp::Directive::OMPD_unknown,
8102 /*ForceSimpleCall=*/false,
8103 /*CheckCancelFlag=*/false);
8104 if (!AfterIP)
8105 return AfterIP.takeError();
8106 }
8107 return Builder.saveIP();
8108}
8109
8110OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createCritical(
8111 const LocationDescription &Loc, BodyGenCallbackTy BodyGenCB,
8112 FinalizeCallbackTy FiniCB, StringRef CriticalName, Value *HintInst) {
8113
8114 if (!updateToLocation(Loc))
8115 return Loc.IP;
8116
8117 Directive OMPD = Directive::OMPD_critical;
8118 uint32_t SrcLocStrSize;
8119 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8120 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8121 Value *ThreadId = getOrCreateThreadID(Ident);
8122 Value *LockVar = getOMPCriticalRegionLock(CriticalName);
8123 Value *Args[] = {Ident, ThreadId, LockVar};
8124
8125 SmallVector<llvm::Value *, 4> EnterArgs(std::begin(arr&: Args), std::end(arr&: Args));
8126 Function *RTFn = nullptr;
8127 if (HintInst) {
8128 // Add Hint to entry Args and create call
8129 EnterArgs.push_back(Elt: HintInst);
8130 RTFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_critical_with_hint);
8131 } else {
8132 RTFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_critical);
8133 }
8134 Instruction *EntryCall = createRuntimeFunctionCall(Callee: RTFn, Args: EnterArgs);
8135
8136 Function *ExitRTLFn =
8137 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_end_critical);
8138 Instruction *ExitCall = createRuntimeFunctionCall(Callee: ExitRTLFn, Args);
8139
8140 return EmitOMPInlinedRegion(OMPD, EntryCall, ExitCall, BodyGenCB, FiniCB,
8141 /*Conditional*/ false, /*hasFinalize*/ HasFinalize: true);
8142}
8143
8144OpenMPIRBuilder::InsertPointTy
8145OpenMPIRBuilder::createOrderedDepend(const LocationDescription &Loc,
8146 InsertPointTy AllocaIP, unsigned NumLoops,
8147 ArrayRef<llvm::Value *> StoreValues,
8148 const Twine &Name, bool IsDependSource) {
8149 assert(
8150 llvm::all_of(StoreValues,
8151 [](Value *SV) { return SV->getType()->isIntegerTy(64); }) &&
8152 "OpenMP runtime requires depend vec with i64 type");
8153
8154 if (!updateToLocation(Loc))
8155 return Loc.IP;
8156
8157 // Allocate space for vector and generate alloc instruction.
8158 auto *ArrI64Ty = ArrayType::get(ElementType: Int64, NumElements: NumLoops);
8159 Builder.restoreIP(IP: AllocaIP);
8160 AllocaInst *ArgsBase = Builder.CreateAlloca(Ty: ArrI64Ty, ArraySize: nullptr, Name);
8161 ArgsBase->setAlignment(Align(8));
8162 updateToLocation(Loc);
8163
8164 // Store the index value with offset in depend vector.
8165 for (unsigned I = 0; I < NumLoops; ++I) {
8166 Value *DependAddrGEPIter = Builder.CreateInBoundsGEP(
8167 Ty: ArrI64Ty, Ptr: ArgsBase, IdxList: {Builder.getInt64(C: 0), Builder.getInt64(C: I)});
8168 StoreInst *STInst = Builder.CreateStore(Val: StoreValues[I], Ptr: DependAddrGEPIter);
8169 STInst->setAlignment(Align(8));
8170 }
8171
8172 Value *DependBaseAddrGEP = Builder.CreateInBoundsGEP(
8173 Ty: ArrI64Ty, Ptr: ArgsBase, IdxList: {Builder.getInt64(C: 0), Builder.getInt64(C: 0)});
8174
8175 uint32_t SrcLocStrSize;
8176 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8177 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8178 Value *ThreadId = getOrCreateThreadID(Ident);
8179 Value *Args[] = {Ident, ThreadId, DependBaseAddrGEP};
8180
8181 Function *RTLFn = nullptr;
8182 if (IsDependSource)
8183 RTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_doacross_post);
8184 else
8185 RTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_doacross_wait);
8186 createRuntimeFunctionCall(Callee: RTLFn, Args);
8187
8188 return Builder.saveIP();
8189}
8190
8191OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createOrderedThreadsSimd(
8192 const LocationDescription &Loc, BodyGenCallbackTy BodyGenCB,
8193 FinalizeCallbackTy FiniCB, bool IsThreads) {
8194 if (!updateToLocation(Loc))
8195 return Loc.IP;
8196
8197 Directive OMPD = Directive::OMPD_ordered_blockassoc;
8198 Instruction *EntryCall = nullptr;
8199 Instruction *ExitCall = nullptr;
8200
8201 if (IsThreads) {
8202 uint32_t SrcLocStrSize;
8203 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8204 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8205 Value *ThreadId = getOrCreateThreadID(Ident);
8206 Value *Args[] = {Ident, ThreadId};
8207
8208 Function *EntryRTLFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_ordered);
8209 EntryCall = createRuntimeFunctionCall(Callee: EntryRTLFn, Args);
8210
8211 Function *ExitRTLFn =
8212 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_end_ordered);
8213 ExitCall = createRuntimeFunctionCall(Callee: ExitRTLFn, Args);
8214 }
8215
8216 return EmitOMPInlinedRegion(OMPD, EntryCall, ExitCall, BodyGenCB, FiniCB,
8217 /*Conditional*/ false, /*hasFinalize*/ HasFinalize: true);
8218}
8219
8220OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::EmitOMPInlinedRegion(
8221 Directive OMPD, Instruction *EntryCall, Instruction *ExitCall,
8222 BodyGenCallbackTy BodyGenCB, FinalizeCallbackTy FiniCB, bool Conditional,
8223 bool HasFinalize, bool IsCancellable) {
8224
8225 if (HasFinalize)
8226 FinalizationStack.push_back(Elt: {FiniCB, OMPD, IsCancellable});
8227
8228 // Create inlined region's entry and body blocks, in preparation
8229 // for conditional creation
8230 BasicBlock *EntryBB = Builder.GetInsertBlock();
8231 Instruction *SplitPos = EntryBB->getTerminatorOrNull();
8232 if (!isa_and_nonnull<UncondBrInst, CondBrInst>(Val: SplitPos))
8233 SplitPos = new UnreachableInst(Builder.getContext(), EntryBB);
8234 BasicBlock *ExitBB = EntryBB->splitBasicBlock(I: SplitPos, BBName: "omp_region.end");
8235 BasicBlock *FiniBB =
8236 EntryBB->splitBasicBlock(I: EntryBB->getTerminator(), BBName: "omp_region.finalize");
8237
8238 Builder.SetInsertPoint(EntryBB->getTerminator());
8239 emitCommonDirectiveEntry(OMPD, EntryCall, ExitBB, Conditional);
8240
8241 // generate body
8242 if (Error Err =
8243 BodyGenCB(/* AllocaIP */ InsertPointTy(),
8244 /* CodeGenIP */ Builder.saveIP(), /* DeallocBlocks */ {}))
8245 return Err;
8246
8247 // emit exit call and do any needed finalization.
8248 auto FinIP = FiniBB->getFirstInsertionPt();
8249 assert(FiniBB->getTerminator()->getNumSuccessors() == 1 &&
8250 FiniBB->getTerminator()->getSuccessor(0) == ExitBB &&
8251 "Unexpected control flow graph state!!");
8252 InsertPointOrErrorTy AfterIP =
8253 emitCommonDirectiveExit(OMPD, FinIP, ExitCall, HasFinalize);
8254 if (!AfterIP)
8255 return AfterIP.takeError();
8256
8257 // If we are skipping the region of a non conditional, remove the exit
8258 // block, and clear the builder's insertion point.
8259 assert(SplitPos->getParent() == ExitBB &&
8260 "Unexpected Insertion point location!");
8261 auto merged = MergeBlockIntoPredecessor(BB: ExitBB);
8262 BasicBlock *ExitPredBB = SplitPos->getParent();
8263 auto InsertBB = merged ? ExitPredBB : ExitBB;
8264 if (!isa_and_nonnull<UncondBrInst, CondBrInst>(Val: SplitPos))
8265 SplitPos->eraseFromParent();
8266 Builder.SetInsertPoint(InsertBB);
8267
8268 return Builder.saveIP();
8269}
8270
8271OpenMPIRBuilder::InsertPointTy OpenMPIRBuilder::emitCommonDirectiveEntry(
8272 Directive OMPD, Value *EntryCall, BasicBlock *ExitBB, bool Conditional) {
8273 // if nothing to do, Return current insertion point.
8274 if (!Conditional || !EntryCall)
8275 return Builder.saveIP();
8276
8277 BasicBlock *EntryBB = Builder.GetInsertBlock();
8278 Value *CallBool = Builder.CreateIsNotNull(Arg: EntryCall);
8279 auto *ThenBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp_region.body");
8280 auto *UI = new UnreachableInst(Builder.getContext(), ThenBB);
8281
8282 // Emit thenBB and set the Builder's insertion point there for
8283 // body generation next. Place the block after the current block.
8284 Function *CurFn = EntryBB->getParent();
8285 CurFn->insert(Position: std::next(x: EntryBB->getIterator()), BB: ThenBB);
8286
8287 // Move Entry branch to end of ThenBB, and replace with conditional
8288 // branch (If-stmt)
8289 Instruction *EntryBBTI = EntryBB->getTerminator();
8290 Builder.CreateCondBr(Cond: CallBool, True: ThenBB, False: ExitBB);
8291 EntryBBTI->removeFromParent();
8292 Builder.SetInsertPoint(UI);
8293 Builder.Insert(I: EntryBBTI);
8294 UI->eraseFromParent();
8295 Builder.SetInsertPoint(ThenBB->getTerminator());
8296
8297 // return an insertion point to ExitBB.
8298 return ExitBB->getFirstInsertionPt();
8299}
8300
8301OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::emitCommonDirectiveExit(
8302 omp::Directive OMPD, InsertPointTy FinIP, Instruction *ExitCall,
8303 bool HasFinalize) {
8304
8305 Builder.restoreIP(IP: FinIP);
8306
8307 // If there is finalization to do, emit it before the exit call
8308 if (HasFinalize) {
8309 assert(!FinalizationStack.empty() &&
8310 "Unexpected finalization stack state!");
8311
8312 FinalizationInfo Fi = FinalizationStack.pop_back_val();
8313 assert(Fi.DK == OMPD && "Unexpected Directive for Finalization call!");
8314
8315 BasicBlock *FinBB = FinIP.getNodeParent();
8316 if (Error Err = Fi.mergeFiniBB(Builder, OtherFiniBB: FinBB))
8317 return std::move(Err);
8318
8319 // Exit condition: insertion point is before the terminator of the new Fini
8320 // block
8321 Builder.SetInsertPoint(FinBB->getTerminator());
8322 }
8323
8324 if (!ExitCall)
8325 return Builder.saveIP();
8326
8327 // place the Exitcall as last instruction before Finalization block terminator
8328 ExitCall->removeFromParent();
8329 Builder.Insert(I: ExitCall);
8330
8331 return ExitCall->getIterator();
8332}
8333
8334OpenMPIRBuilder::InsertPointTy OpenMPIRBuilder::createCopyinClauseBlocks(
8335 InsertPointTy IP, Value *MasterAddr, Value *PrivateAddr,
8336 llvm::IntegerType *IntPtrTy, bool BranchtoEnd) {
8337 if (!IP.isValid())
8338 return IP;
8339
8340 IRBuilder<>::InsertPointGuard IPG(Builder);
8341
8342 // creates the following CFG structure
8343 // OMP_Entry : (MasterAddr != PrivateAddr)?
8344 // F T
8345 // | \
8346 // | copin.not.master
8347 // | /
8348 // v /
8349 // copyin.not.master.end
8350 // |
8351 // v
8352 // OMP.Entry.Next
8353
8354 BasicBlock *OMP_Entry = IP.getNodeParent();
8355 Function *CurFn = OMP_Entry->getParent();
8356 BasicBlock *CopyBegin =
8357 BasicBlock::Create(Context&: M.getContext(), Name: "copyin.not.master", Parent: CurFn);
8358 BasicBlock *CopyEnd = nullptr;
8359
8360 // If entry block is terminated, split to preserve the branch to following
8361 // basic block (i.e. OMP.Entry.Next), otherwise, leave everything as is.
8362 if (isa_and_nonnull<CondBrInst>(Val: OMP_Entry->getTerminatorOrNull())) {
8363 CopyEnd = OMP_Entry->splitBasicBlock(I: OMP_Entry->getTerminator(),
8364 BBName: "copyin.not.master.end");
8365 OMP_Entry->getTerminator()->eraseFromParent();
8366 } else {
8367 CopyEnd =
8368 BasicBlock::Create(Context&: M.getContext(), Name: "copyin.not.master.end", Parent: CurFn);
8369 }
8370
8371 Builder.SetInsertPoint(OMP_Entry);
8372 Value *MasterPtr = Builder.CreatePtrToInt(V: MasterAddr, DestTy: IntPtrTy);
8373 Value *PrivatePtr = Builder.CreatePtrToInt(V: PrivateAddr, DestTy: IntPtrTy);
8374 Value *cmp = Builder.CreateICmpNE(LHS: MasterPtr, RHS: PrivatePtr);
8375 Builder.CreateCondBr(Cond: cmp, True: CopyBegin, False: CopyEnd);
8376
8377 Builder.SetInsertPoint(CopyBegin);
8378 if (BranchtoEnd)
8379 Builder.SetInsertPoint(Builder.CreateBr(Dest: CopyEnd));
8380
8381 return Builder.saveIP();
8382}
8383
8384CallInst *OpenMPIRBuilder::createOMPAlloc(const LocationDescription &Loc,
8385 Value *Size, Value *Allocator,
8386 std::string Name) {
8387 IRBuilder<>::InsertPointGuard IPG(Builder);
8388 if (!updateToLocation(Loc))
8389 return nullptr;
8390
8391 uint32_t SrcLocStrSize;
8392 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8393 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8394 Value *ThreadId = getOrCreateThreadID(Ident);
8395 Value *Args[] = {ThreadId, Size, Allocator};
8396
8397 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_alloc);
8398
8399 return createRuntimeFunctionCall(Callee: Fn, Args, Name);
8400}
8401
8402CallInst *OpenMPIRBuilder::createOMPAlignedAlloc(const LocationDescription &Loc,
8403 Value *Align, Value *Size,
8404 Value *Allocator,
8405 std::string Name) {
8406 IRBuilder<>::InsertPointGuard IPG(Builder);
8407 if (!updateToLocation(Loc))
8408 return nullptr;
8409
8410 uint32_t SrcLocStrSize;
8411 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8412 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8413 Value *ThreadId = getOrCreateThreadID(Ident);
8414 Value *Args[] = {ThreadId, Align, Size, Allocator};
8415
8416 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_aligned_alloc);
8417
8418 return Builder.CreateCall(Callee: Fn, Args, Name);
8419}
8420
8421CallInst *OpenMPIRBuilder::createOMPFree(const LocationDescription &Loc,
8422 Value *Addr, Value *Allocator,
8423 std::string Name) {
8424 IRBuilder<>::InsertPointGuard IPG(Builder);
8425 if (!updateToLocation(Loc))
8426 return nullptr;
8427
8428 uint32_t SrcLocStrSize;
8429 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8430 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8431 Value *ThreadId = getOrCreateThreadID(Ident);
8432 Value *Args[] = {ThreadId, Addr, Allocator};
8433 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_free);
8434 return createRuntimeFunctionCall(Callee: Fn, Args, Name);
8435}
8436
8437CallInst *OpenMPIRBuilder::createOMPAllocShared(const LocationDescription &Loc,
8438 Value *Size,
8439 const Twine &Name) {
8440 IRBuilder<>::InsertPointGuard IPG(Builder);
8441 updateToLocation(Loc);
8442
8443 Value *Args[] = {Size};
8444 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_alloc_shared);
8445 CallInst *Call = Builder.CreateCall(Callee: Fn, Args, Name);
8446 Call->addRetAttr(Attr: Attribute::getWithAlignment(
8447 Context&: M.getContext(), Alignment: M.getDataLayout().getPrefTypeAlign(Ty: Int64)));
8448 return Call;
8449}
8450
8451CallInst *OpenMPIRBuilder::createOMPAllocShared(const LocationDescription &Loc,
8452 Type *VarType,
8453 const Twine &Name) {
8454 return createOMPAllocShared(
8455 Loc, Size: Builder.getInt64(C: M.getDataLayout().getTypeAllocSize(Ty: VarType)), Name);
8456}
8457
8458CallInst *OpenMPIRBuilder::createOMPFreeShared(const LocationDescription &Loc,
8459 Value *Addr, Value *Size,
8460 const Twine &Name) {
8461 IRBuilder<>::InsertPointGuard IPG(Builder);
8462 updateToLocation(Loc);
8463
8464 Value *Args[] = {Addr, Size};
8465 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_free_shared);
8466 return Builder.CreateCall(Callee: Fn, Args, Name);
8467}
8468
8469CallInst *OpenMPIRBuilder::createOMPFreeShared(const LocationDescription &Loc,
8470 Value *Addr, Type *VarType,
8471 const Twine &Name) {
8472 return createOMPFreeShared(
8473 Loc, Addr, Size: Builder.getInt64(C: M.getDataLayout().getTypeAllocSize(Ty: VarType)),
8474 Name);
8475}
8476
8477CallInst *OpenMPIRBuilder::createOMPInteropInit(
8478 const LocationDescription &Loc, Value *InteropVar,
8479 omp::OMPInteropType InteropType, Value *Device, Value *NumDependences,
8480 Value *DependenceAddress, bool HaveNowaitClause) {
8481 IRBuilder<>::InsertPointGuard IPG(Builder);
8482 updateToLocation(Loc);
8483
8484 uint32_t SrcLocStrSize;
8485 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8486 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8487 Value *ThreadId = getOrCreateThreadID(Ident);
8488 if (Device == nullptr)
8489 Device = Constant::getAllOnesValue(Ty: Int32);
8490 else if (Device->getType() != Int32)
8491 Device = Builder.CreateIntCast(V: Device, DestTy: Int32, /*isSigned=*/true);
8492 Constant *InteropTypeVal = ConstantInt::get(Ty: Int32, V: (int)InteropType);
8493 if (NumDependences == nullptr) {
8494 NumDependences = ConstantInt::get(Ty: Int32, V: 0);
8495 PointerType *PointerTypeVar = PointerType::getUnqual(C&: M.getContext());
8496 DependenceAddress = ConstantPointerNull::get(T: PointerTypeVar);
8497 }
8498 Value *HaveNowaitClauseVal = ConstantInt::get(Ty: Int32, V: HaveNowaitClause);
8499 Value *Args[] = {
8500 Ident, ThreadId, InteropVar, InteropTypeVal,
8501 Device, NumDependences, DependenceAddress, HaveNowaitClauseVal};
8502
8503 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___tgt_interop_init);
8504
8505 return createRuntimeFunctionCall(Callee: Fn, Args);
8506}
8507
8508CallInst *OpenMPIRBuilder::createOMPInteropDestroy(
8509 const LocationDescription &Loc, Value *InteropVar, Value *Device,
8510 Value *NumDependences, Value *DependenceAddress, bool HaveNowaitClause) {
8511 IRBuilder<>::InsertPointGuard IPG(Builder);
8512 updateToLocation(Loc);
8513
8514 uint32_t SrcLocStrSize;
8515 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8516 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8517 Value *ThreadId = getOrCreateThreadID(Ident);
8518 if (Device == nullptr)
8519 Device = Constant::getAllOnesValue(Ty: Int32);
8520 else if (Device->getType() != Int32)
8521 Device = Builder.CreateIntCast(V: Device, DestTy: Int32, /*isSigned=*/true);
8522 if (NumDependences == nullptr) {
8523 NumDependences = ConstantInt::get(Ty: Int32, V: 0);
8524 PointerType *PointerTypeVar = PointerType::getUnqual(C&: M.getContext());
8525 DependenceAddress = ConstantPointerNull::get(T: PointerTypeVar);
8526 }
8527 Value *HaveNowaitClauseVal = ConstantInt::get(Ty: Int32, V: HaveNowaitClause);
8528 Value *Args[] = {
8529 Ident, ThreadId, InteropVar, Device,
8530 NumDependences, DependenceAddress, HaveNowaitClauseVal};
8531
8532 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___tgt_interop_destroy);
8533
8534 return createRuntimeFunctionCall(Callee: Fn, Args);
8535}
8536
8537CallInst *OpenMPIRBuilder::createOMPInteropUse(const LocationDescription &Loc,
8538 Value *InteropVar, Value *Device,
8539 Value *NumDependences,
8540 Value *DependenceAddress,
8541 bool HaveNowaitClause) {
8542 IRBuilder<>::InsertPointGuard IPG(Builder);
8543 updateToLocation(Loc);
8544 uint32_t SrcLocStrSize;
8545 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8546 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8547 Value *ThreadId = getOrCreateThreadID(Ident);
8548 if (Device == nullptr)
8549 Device = Constant::getAllOnesValue(Ty: Int32);
8550 else if (Device->getType() != Int32)
8551 Device = Builder.CreateIntCast(V: Device, DestTy: Int32, /*isSigned=*/true);
8552 if (NumDependences == nullptr) {
8553 NumDependences = ConstantInt::get(Ty: Int32, V: 0);
8554 PointerType *PointerTypeVar = PointerType::getUnqual(C&: M.getContext());
8555 DependenceAddress = ConstantPointerNull::get(T: PointerTypeVar);
8556 }
8557 Value *HaveNowaitClauseVal = ConstantInt::get(Ty: Int32, V: HaveNowaitClause);
8558 Value *Args[] = {
8559 Ident, ThreadId, InteropVar, Device,
8560 NumDependences, DependenceAddress, HaveNowaitClauseVal};
8561
8562 Function *Fn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___tgt_interop_use);
8563
8564 return createRuntimeFunctionCall(Callee: Fn, Args);
8565}
8566
8567CallInst *OpenMPIRBuilder::createCachedThreadPrivate(
8568 const LocationDescription &Loc, llvm::Value *Pointer,
8569 llvm::ConstantInt *Size, const llvm::Twine &Name) {
8570 IRBuilder<>::InsertPointGuard IPG(Builder);
8571 updateToLocation(Loc);
8572
8573 uint32_t SrcLocStrSize;
8574 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8575 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8576 Value *ThreadId = getOrCreateThreadID(Ident);
8577 Constant *ThreadPrivateCache =
8578 getOrCreateInternalVariable(Ty: Int8PtrPtr, Name: Name.str());
8579 llvm::Value *Args[] = {Ident, ThreadId, Pointer, Size, ThreadPrivateCache};
8580
8581 Function *Fn =
8582 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_threadprivate_cached);
8583
8584 return createRuntimeFunctionCall(Callee: Fn, Args);
8585}
8586
8587Constant *OpenMPIRBuilder::emitKernelEnvironment(
8588 const LocationDescription &Loc,
8589 const llvm::OpenMPIRBuilder::TargetKernelDefaultAttrs &Attrs) {
8590 assert(!Attrs.MaxThreads.empty() && !Attrs.MaxTeams.empty() &&
8591 "expected num_threads and num_teams to be specified");
8592
8593 if (!updateToLocation(Loc))
8594 return nullptr;
8595
8596 uint32_t SrcLocStrSize;
8597 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8598 Constant *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8599 Constant *IsSPMDVal = ConstantInt::getSigned(Ty: Int8, V: Attrs.ExecFlags);
8600 Constant *UseGenericStateMachineVal = ConstantInt::getSigned(
8601 Ty: Int8, V: Attrs.ExecFlags != omp::OMP_TGT_EXEC_MODE_SPMD &&
8602 Attrs.ExecFlags != omp::OMP_TGT_EXEC_MODE_SPMD_NO_LOOP);
8603 Constant *MayUseNestedParallelismVal = ConstantInt::getSigned(Ty: Int8, V: true);
8604 Constant *DebugIndentionLevelVal = ConstantInt::getSigned(Ty: Int16, V: 0);
8605
8606 Function *DebugKernelWrapper = Builder.GetInsertBlock()->getParent();
8607 Function *Kernel = DebugKernelWrapper;
8608
8609 // We need to strip the debug prefix to get the correct kernel name.
8610 StringRef KernelName = Kernel->getName();
8611 const std::string DebugPrefix = "_debug__";
8612 if (KernelName.ends_with(Suffix: DebugPrefix)) {
8613 KernelName = KernelName.drop_back(N: DebugPrefix.length());
8614 Kernel = M.getFunction(Name: KernelName);
8615 assert(Kernel && "Expected the real kernel to exist");
8616 }
8617
8618 // Manifest the launch configuration in the metadata matching the kernel
8619 // environment.
8620 if (Attrs.MinTeams.front() > 1 || Attrs.MaxTeams.front() > 0)
8621 writeTeamsForKernel(T, Kernel&: *Kernel, LB: Attrs.MinTeams.front(),
8622 UB: Attrs.MaxTeams.front());
8623
8624 // Don't derive or write thread bounds for Bare kernels.
8625 int32_t MaxThreadsVal = Attrs.MaxThreads.front();
8626 if (Attrs.ExecFlags != omp::OMP_TGT_EXEC_MODE_BARE) {
8627 // If MaxThreads is not set and needs adjustment, select the maximum
8628 // between the default workgroup size and the MinThreads value. This is
8629 // only meaningful for targets with a known grid value (i.e. GPUs); for
8630 // other targets (e.g. host kernels) leave it unset so the runtime falls
8631 // back to its own device-specific default.
8632 if (MaxThreadsVal < 0 && UseDefaultMaxThreads && hasGridValue(T))
8633 MaxThreadsVal =
8634 std::max(a: int32_t(getGridValue(T, Kernel).GV_Default_WG_Size),
8635 b: Attrs.MinThreads.front());
8636
8637 // Generic mode runs the main thread on a warp of its own, past
8638 // thread_limit. Reserve the widest warp any target has. Not on SPIR-V,
8639 // causes problems with Level Zero.
8640 if (MaxThreadsVal > 0 &&
8641 Attrs.ExecFlags == omp::OMP_TGT_EXEC_MODE_GENERIC && hasGridValue(T) &&
8642 !T.isSPIRV())
8643 MaxThreadsVal = int32_t(
8644 std::min<int64_t>(a: int64_t(MaxThreadsVal) + 64,
8645 b: int64_t(getGridValue(T, Kernel).GV_Max_WG_Size)));
8646
8647 if (MaxThreadsVal > 0)
8648 writeThreadBoundsForKernel(T, Kernel&: *Kernel, LB: Attrs.MinThreads.front(),
8649 UB: MaxThreadsVal);
8650 }
8651
8652 Constant *MinThreads =
8653 ConstantInt::getSigned(Ty: Int32, V: Attrs.MinThreads.front());
8654 Constant *MaxThreads = ConstantInt::getSigned(Ty: Int32, V: MaxThreadsVal);
8655 Constant *MinTeams = ConstantInt::getSigned(Ty: Int32, V: Attrs.MinTeams.front());
8656 Constant *MaxTeams = ConstantInt::getSigned(Ty: Int32, V: Attrs.MaxTeams.front());
8657 Constant *ReductionDataSize =
8658 ConstantInt::getSigned(Ty: Int32, V: Attrs.ReductionDataSize);
8659
8660 const DataLayout &DL = M.getDataLayout();
8661
8662 Twine DynamicEnvironmentName = KernelName + "_dynamic_environment";
8663 Constant *DynamicEnvironmentInitializer =
8664 ConstantStruct::get(T: DynamicEnvironment, V: {DebugIndentionLevelVal});
8665 GlobalVariable *DynamicEnvironmentGV = new GlobalVariable(
8666 M, DynamicEnvironment, /*IsConstant=*/false, GlobalValue::WeakODRLinkage,
8667 DynamicEnvironmentInitializer, DynamicEnvironmentName,
8668 /*InsertBefore=*/nullptr, GlobalValue::NotThreadLocal,
8669 DL.getDefaultGlobalsAddressSpace());
8670 DynamicEnvironmentGV->setVisibility(GlobalValue::ProtectedVisibility);
8671
8672 Constant *DynamicEnvironment =
8673 DynamicEnvironmentGV->getType() == DynamicEnvironmentPtr
8674 ? DynamicEnvironmentGV
8675 : ConstantExpr::getAddrSpaceCast(C: DynamicEnvironmentGV,
8676 Ty: DynamicEnvironmentPtr);
8677
8678 Constant *ConfigurationEnvironmentInitializer = ConstantStruct::get(
8679 T: ConfigurationEnvironment, V: {
8680 UseGenericStateMachineVal,
8681 MayUseNestedParallelismVal,
8682 IsSPMDVal,
8683 MinThreads,
8684 MaxThreads,
8685 MinTeams,
8686 MaxTeams,
8687 ReductionDataSize,
8688 });
8689 Constant *KernelEnvironmentInitializer = ConstantStruct::get(
8690 T: KernelEnvironment, V: {
8691 ConfigurationEnvironmentInitializer,
8692 Ident,
8693 DynamicEnvironment,
8694 });
8695 std::string KernelEnvironmentName =
8696 (KernelName + "_kernel_environment").str();
8697 GlobalVariable *KernelEnvironmentGV = new GlobalVariable(
8698 M, KernelEnvironment, /*IsConstant=*/true, GlobalValue::WeakODRLinkage,
8699 KernelEnvironmentInitializer, KernelEnvironmentName,
8700 /*InsertBefore=*/nullptr, GlobalValue::NotThreadLocal,
8701 DL.getDefaultGlobalsAddressSpace());
8702 KernelEnvironmentGV->setVisibility(GlobalValue::ProtectedVisibility);
8703
8704 return KernelEnvironmentGV->getType() == KernelEnvironmentPtr
8705 ? KernelEnvironmentGV
8706 : ConstantExpr::getAddrSpaceCast(C: KernelEnvironmentGV,
8707 Ty: KernelEnvironmentPtr);
8708}
8709
8710OpenMPIRBuilder::InsertPointTy OpenMPIRBuilder::createTargetInit(
8711 const LocationDescription &Loc,
8712 const llvm::OpenMPIRBuilder::TargetKernelDefaultAttrs &Attrs) {
8713 Constant *KernelEnvironment = emitKernelEnvironment(Loc, Attrs);
8714 if (!KernelEnvironment)
8715 return Loc.IP;
8716
8717 if (!updateToLocation(Loc))
8718 return Loc.IP;
8719
8720 Function *DebugKernelWrapper = Builder.GetInsertBlock()->getParent();
8721 Function *Fn = getOrCreateRuntimeFunctionPtr(
8722 FnID: omp::RuntimeFunction::OMPRTL___kmpc_target_init);
8723
8724 Value *KernelLaunchEnvironment =
8725 DebugKernelWrapper->getArg(i: DebugKernelWrapper->arg_size() - 1);
8726 Type *KernelLaunchEnvParamTy = Fn->getFunctionType()->getParamType(i: 1);
8727 KernelLaunchEnvironment =
8728 KernelLaunchEnvironment->getType() == KernelLaunchEnvParamTy
8729 ? KernelLaunchEnvironment
8730 : Builder.CreateAddrSpaceCast(V: KernelLaunchEnvironment,
8731 DestTy: KernelLaunchEnvParamTy);
8732 CallInst *ThreadKind = createRuntimeFunctionCall(
8733 Callee: Fn, Args: {KernelEnvironment, KernelLaunchEnvironment});
8734
8735 Value *ExecUserCode = Builder.CreateICmpEQ(
8736 LHS: ThreadKind, RHS: Constant::getAllOnesValue(Ty: ThreadKind->getType()),
8737 Name: "exec_user_code");
8738
8739 // ThreadKind = __kmpc_target_init(...)
8740 // if (ThreadKind == -1)
8741 // user_code
8742 // else
8743 // return;
8744
8745 auto *UI = Builder.CreateUnreachable();
8746 BasicBlock *CheckBB = UI->getParent();
8747 BasicBlock *UserCodeEntryBB = CheckBB->splitBasicBlock(I: UI, BBName: "user_code.entry");
8748
8749 BasicBlock *WorkerExitBB = BasicBlock::Create(
8750 Context&: CheckBB->getContext(), Name: "worker.exit", Parent: CheckBB->getParent());
8751 Builder.SetInsertPoint(WorkerExitBB);
8752 Builder.CreateRetVoid();
8753
8754 auto *CheckBBTI = CheckBB->getTerminator();
8755 Builder.SetInsertPoint(CheckBBTI);
8756 Builder.CreateCondBr(Cond: ExecUserCode, True: UI->getParent(), False: WorkerExitBB);
8757
8758 CheckBBTI->eraseFromParent();
8759 UI->eraseFromParent();
8760
8761 // Continue in the "user_code" block, see diagram above and in
8762 // openmp/libomptarget/deviceRTLs/common/include/target.h .
8763 return UserCodeEntryBB->getFirstInsertionPt();
8764}
8765
8766void OpenMPIRBuilder::createTargetDeinit(const LocationDescription &Loc,
8767 int32_t TeamsReductionDataSize) {
8768 if (!updateToLocation(Loc))
8769 return;
8770
8771 Function *Fn = getOrCreateRuntimeFunctionPtr(
8772 FnID: omp::RuntimeFunction::OMPRTL___kmpc_target_deinit);
8773
8774 createRuntimeFunctionCall(Callee: Fn, Args: {});
8775
8776 if (!TeamsReductionDataSize)
8777 return;
8778
8779 Function *Kernel = Builder.GetInsertBlock()->getParent();
8780 // We need to strip the debug prefix to get the correct kernel name.
8781 StringRef KernelName = Kernel->getName();
8782 const std::string DebugPrefix = "_debug__";
8783 if (KernelName.ends_with(Suffix: DebugPrefix))
8784 KernelName = KernelName.drop_back(N: DebugPrefix.length());
8785 auto *KernelEnvironmentGV =
8786 M.getNamedGlobal(Name: (KernelName + "_kernel_environment").str());
8787 assert(KernelEnvironmentGV && "Expected kernel environment global\n");
8788 auto *KernelEnvironmentInitializer = KernelEnvironmentGV->getInitializer();
8789 auto *NewInitializer = ConstantFoldInsertValueInstruction(
8790 Agg: KernelEnvironmentInitializer,
8791 Val: ConstantInt::get(Ty: Int32, V: TeamsReductionDataSize), Idxs: {0, 7});
8792 KernelEnvironmentGV->setInitializer(NewInitializer);
8793}
8794
8795static void updateNVPTXAttr(Function &Kernel, StringRef Name, int32_t Value,
8796 bool Min) {
8797 if (Kernel.hasFnAttribute(Kind: Name)) {
8798 int32_t OldLimit = Kernel.getFnAttributeAsParsedInteger(Kind: Name);
8799 Value = Min ? std::min(a: OldLimit, b: Value) : std::max(a: OldLimit, b: Value);
8800 }
8801 Kernel.addFnAttr(Kind: Name, Val: llvm::utostr(X: Value));
8802}
8803
8804std::pair<int32_t, int32_t>
8805OpenMPIRBuilder::readThreadBoundsForKernel(const Triple &T, Function &Kernel) {
8806 int32_t ThreadLimit =
8807 Kernel.getFnAttributeAsParsedInteger(Kind: "omp_target_thread_limit");
8808
8809 if (T.isAMDGPU()) {
8810 const auto &Attr = Kernel.getFnAttribute(Kind: "amdgpu-flat-work-group-size");
8811 if (!Attr.isValid() || !Attr.isStringAttribute())
8812 return {0, ThreadLimit};
8813 auto [LBStr, UBStr] = Attr.getValueAsString().split(Separator: ',');
8814 int32_t LB, UB;
8815 if (!llvm::to_integer(S: UBStr, Num&: UB, Base: 10))
8816 return {0, ThreadLimit};
8817 UB = ThreadLimit ? std::min(a: ThreadLimit, b: UB) : UB;
8818 if (!llvm::to_integer(S: LBStr, Num&: LB, Base: 10))
8819 return {0, UB};
8820 return {LB, UB};
8821 }
8822
8823 if (Kernel.hasFnAttribute(Kind: NVVMAttr::MaxNTID)) {
8824 int32_t UB = Kernel.getFnAttributeAsParsedInteger(Kind: NVVMAttr::MaxNTID);
8825 return {0, ThreadLimit ? std::min(a: ThreadLimit, b: UB) : UB};
8826 }
8827 return {0, ThreadLimit};
8828}
8829
8830void OpenMPIRBuilder::writeThreadBoundsForKernel(const Triple &T,
8831 Function &Kernel, int32_t LB,
8832 int32_t UB) {
8833 Kernel.addFnAttr(Kind: "omp_target_thread_limit", Val: std::to_string(val: UB));
8834
8835 if (T.isAMDGPU()) {
8836 Kernel.addFnAttr(Kind: "amdgpu-flat-work-group-size",
8837 Val: llvm::utostr(X: LB) + "," + llvm::utostr(X: UB));
8838 return;
8839 }
8840
8841 updateNVPTXAttr(Kernel, Name: NVVMAttr::MaxNTID, Value: UB, Min: true);
8842}
8843
8844std::pair<int32_t, int32_t>
8845OpenMPIRBuilder::readTeamBoundsForKernel(const Triple &, Function &Kernel) {
8846 // TODO: Read from backend annotations if available.
8847 return {0, Kernel.getFnAttributeAsParsedInteger(Kind: "omp_target_num_teams")};
8848}
8849
8850void OpenMPIRBuilder::writeTeamsForKernel(const Triple &T, Function &Kernel,
8851 int32_t LB, int32_t UB) {
8852 if (UB > 0) {
8853 if (T.isNVPTX())
8854 Kernel.addFnAttr(Kind: NVVMAttr::MaxClusterRank, Val: llvm::utostr(X: UB));
8855 if (T.isAMDGPU())
8856 Kernel.addFnAttr(Kind: "amdgpu-max-num-workgroups", Val: llvm::utostr(X: UB) + ",1,1");
8857 }
8858
8859 Kernel.addFnAttr(Kind: "omp_target_num_teams", Val: std::to_string(val: LB));
8860}
8861
8862void OpenMPIRBuilder::setOutlinedTargetRegionFunctionAttributes(
8863 Function *OutlinedFn) {
8864 if (Config.isTargetDevice()) {
8865 OutlinedFn->setLinkage(GlobalValue::WeakODRLinkage);
8866 // TODO: Determine if DSO local can be set to true.
8867 OutlinedFn->setDSOLocal(false);
8868 OutlinedFn->setVisibility(GlobalValue::ProtectedVisibility);
8869 if (T.isAMDGCN())
8870 OutlinedFn->setCallingConv(CallingConv::AMDGPU_KERNEL);
8871 else if (T.isNVPTX())
8872 OutlinedFn->setCallingConv(CallingConv::PTX_Kernel);
8873 else if (T.isSPIRV())
8874 OutlinedFn->setCallingConv(CallingConv::SPIR_KERNEL);
8875 }
8876}
8877
8878Constant *OpenMPIRBuilder::createOutlinedFunctionID(Function *OutlinedFn,
8879 StringRef EntryFnIDName) {
8880 if (Config.isTargetDevice()) {
8881 assert(OutlinedFn && "The outlined function must exist if embedded");
8882 return OutlinedFn;
8883 }
8884
8885 return new GlobalVariable(
8886 M, Builder.getInt8Ty(), /*isConstant=*/true, GlobalValue::WeakAnyLinkage,
8887 Constant::getNullValue(Ty: Builder.getInt8Ty()), EntryFnIDName);
8888}
8889
8890Constant *OpenMPIRBuilder::createTargetRegionEntryAddr(Function *OutlinedFn,
8891 StringRef EntryFnName) {
8892 if (OutlinedFn)
8893 return OutlinedFn;
8894
8895 assert(!M.getGlobalVariable(EntryFnName, true) &&
8896 "Named kernel already exists?");
8897 return new GlobalVariable(
8898 M, Builder.getInt8Ty(), /*isConstant=*/true, GlobalValue::InternalLinkage,
8899 Constant::getNullValue(Ty: Builder.getInt8Ty()), EntryFnName);
8900}
8901
8902Error OpenMPIRBuilder::emitTargetRegionFunction(
8903 TargetRegionEntryInfo &EntryInfo,
8904 FunctionGenCallback &GenerateFunctionCallback, bool IsOffloadEntry,
8905 Function *&OutlinedFn, Constant *&OutlinedFnID) {
8906
8907 SmallString<64> EntryFnName;
8908 OffloadInfoManager.getTargetRegionEntryFnName(Name&: EntryFnName, EntryInfo);
8909
8910 if (Config.isTargetDevice() || !Config.openMPOffloadMandatory()) {
8911 Expected<Function *> CBResult = GenerateFunctionCallback(EntryFnName);
8912 if (!CBResult)
8913 return CBResult.takeError();
8914 OutlinedFn = *CBResult;
8915 } else {
8916 OutlinedFn = nullptr;
8917 }
8918
8919 // If this target outline function is not an offload entry, we don't need to
8920 // register it. This may be in the case of a false if clause, or if there are
8921 // no OpenMP targets.
8922 if (!IsOffloadEntry)
8923 return Error::success();
8924
8925 std::string EntryFnIDName =
8926 Config.isTargetDevice()
8927 ? std::string(EntryFnName)
8928 : createPlatformSpecificName(Parts: {EntryFnName, "region_id"});
8929
8930 OutlinedFnID = registerTargetRegionFunction(EntryInfo, OutlinedFunction: OutlinedFn,
8931 EntryFnName, EntryFnIDName);
8932 return Error::success();
8933}
8934
8935Constant *OpenMPIRBuilder::registerTargetRegionFunction(
8936 TargetRegionEntryInfo &EntryInfo, Function *OutlinedFn,
8937 StringRef EntryFnName, StringRef EntryFnIDName) {
8938 if (OutlinedFn)
8939 setOutlinedTargetRegionFunctionAttributes(OutlinedFn);
8940 auto OutlinedFnID = createOutlinedFunctionID(OutlinedFn, EntryFnIDName);
8941 auto EntryAddr = createTargetRegionEntryAddr(OutlinedFn, EntryFnName);
8942 OffloadInfoManager.registerTargetRegionEntryInfo(
8943 EntryInfo, Addr: EntryAddr, ID: OutlinedFnID,
8944 Flags: OffloadEntriesInfoManager::OMPTargetRegionEntryTargetRegion);
8945 return OutlinedFnID;
8946}
8947
8948OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createTargetData(
8949 const LocationDescription &Loc, InsertPointTy AllocaIP,
8950 InsertPointTy CodeGenIP, ArrayRef<BasicBlock *> DeallocBlocks,
8951 Value *DeviceID, Value *IfCond, TargetDataInfo &Info,
8952 GenMapInfoCallbackTy GenMapInfoCB, CustomMapperCallbackTy CustomMapperCB,
8953 omp::RuntimeFunction *MapperFunc,
8954 function_ref<InsertPointOrErrorTy(InsertPointTy CodeGenIP,
8955 BodyGenTy BodyGenType)>
8956 BodyGenCB,
8957 function_ref<void(unsigned int, Value *)> DeviceAddrCB, Value *SrcLocInfo) {
8958 if (!updateToLocation(Loc))
8959 return InsertPointTy();
8960
8961 Builder.restoreIP(IP: CodeGenIP);
8962
8963 bool IsStandAlone = !BodyGenCB;
8964 MapInfosTy *MapInfo;
8965 // Generate the code for the opening of the data environment. Capture all the
8966 // arguments of the runtime call by reference because they are used in the
8967 // closing of the region.
8968 auto BeginThenGen = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
8969 ArrayRef<BasicBlock *> DeallocBlocks) -> Error {
8970 MapInfo = &GenMapInfoCB(Builder.saveIP());
8971 if (Error Err = emitOffloadingArrays(
8972 AllocaIP, CodeGenIP: Builder.saveIP(), CombinedInfo&: *MapInfo, Info, CustomMapperCB,
8973 /*IsNonContiguous=*/true, DeviceAddrCB))
8974 return Err;
8975
8976 TargetDataRTArgs RTArgs;
8977 emitOffloadingArraysArgument(Builder, RTArgs, Info);
8978
8979 // Emit the number of elements in the offloading arrays.
8980 Value *PointerNum = Builder.getInt32(C: Info.NumberOfPtrs);
8981
8982 // Source location for the ident struct
8983 if (!SrcLocInfo) {
8984 uint32_t SrcLocStrSize;
8985 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
8986 SrcLocInfo = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
8987 }
8988
8989 SmallVector<llvm::Value *, 13> OffloadingArgs = {
8990 SrcLocInfo, DeviceID,
8991 PointerNum, RTArgs.BasePointersArray,
8992 RTArgs.PointersArray, RTArgs.SizesArray,
8993 RTArgs.MapTypesArray, RTArgs.MapNamesArray,
8994 RTArgs.MappersArray};
8995
8996 if (IsStandAlone) {
8997 assert(MapperFunc && "MapperFunc missing for standalone target data");
8998
8999 auto TaskBodyCB = [&](Value *, Value *,
9000 IRBuilderBase::InsertPoint) -> Error {
9001 if (Info.HasNoWait) {
9002 OffloadingArgs.append(IL: {llvm::Constant::getNullValue(Ty: Int32),
9003 llvm::Constant::getNullValue(Ty: VoidPtr),
9004 llvm::Constant::getNullValue(Ty: Int32),
9005 llvm::Constant::getNullValue(Ty: VoidPtr)});
9006 }
9007
9008 createRuntimeFunctionCall(Callee: getOrCreateRuntimeFunctionPtr(FnID: *MapperFunc),
9009 Args: OffloadingArgs);
9010
9011 if (Info.HasNoWait) {
9012 BasicBlock *OffloadContBlock =
9013 BasicBlock::Create(Context&: Builder.getContext(), Name: "omp_offload.cont");
9014 Function *CurFn = Builder.GetInsertBlock()->getParent();
9015 emitBlock(BB: OffloadContBlock, CurFn, /*IsFinished=*/true);
9016 Builder.restoreIP(IP: Builder.saveIP());
9017 }
9018 return Error::success();
9019 };
9020
9021 bool RequiresOuterTargetTask = Info.HasNoWait;
9022 if (!RequiresOuterTargetTask)
9023 cantFail(Err: TaskBodyCB(/*DeviceID=*/nullptr, /*RTLoc=*/nullptr,
9024 /*TargetTaskAllocaIP=*/{}));
9025 else
9026 cantFail(ValOrErr: emitTargetTask(TaskBodyCB, DeviceID, RTLoc: SrcLocInfo, AllocaIP,
9027 /*Dependencies=*/{}, RTArgs, HasNoWait: Info.HasNoWait));
9028 } else {
9029 Function *BeginMapperFunc = getOrCreateRuntimeFunctionPtr(
9030 FnID: omp::OMPRTL___tgt_target_data_begin_mapper);
9031
9032 createRuntimeFunctionCall(Callee: BeginMapperFunc, Args: OffloadingArgs);
9033
9034 for (auto DeviceMap : Info.DevicePtrInfoMap) {
9035 if (isa<AllocaInst>(Val: DeviceMap.second.second)) {
9036 auto *LI =
9037 Builder.CreateLoad(Ty: Builder.getPtrTy(), Ptr: DeviceMap.second.first);
9038 Builder.CreateStore(Val: LI, Ptr: DeviceMap.second.second);
9039 }
9040 }
9041
9042 // If device pointer privatization is required, emit the body of the
9043 // region here. It will have to be duplicated: with and without
9044 // privatization.
9045 InsertPointOrErrorTy AfterIP =
9046 BodyGenCB(Builder.saveIP(), BodyGenTy::Priv);
9047 if (!AfterIP)
9048 return AfterIP.takeError();
9049 Builder.restoreIP(IP: *AfterIP);
9050 }
9051 return Error::success();
9052 };
9053
9054 // If we need device pointer privatization, we need to emit the body of the
9055 // region with no privatization in the 'else' branch of the conditional.
9056 // Otherwise, we don't have to do anything.
9057 auto BeginElseGen = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
9058 ArrayRef<BasicBlock *> DeallocBlocks) -> Error {
9059 InsertPointOrErrorTy AfterIP =
9060 BodyGenCB(Builder.saveIP(), BodyGenTy::DupNoPriv);
9061 if (!AfterIP)
9062 return AfterIP.takeError();
9063 Builder.restoreIP(IP: *AfterIP);
9064 return Error::success();
9065 };
9066
9067 // Generate code for the closing of the data region.
9068 auto EndThenGen = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
9069 ArrayRef<BasicBlock *> DeallocBlocks) {
9070 TargetDataRTArgs RTArgs;
9071 Info.EmitDebug = !MapInfo->Names.empty();
9072 emitOffloadingArraysArgument(Builder, RTArgs, Info, /*ForEndCall=*/true);
9073
9074 // Emit the number of elements in the offloading arrays.
9075 Value *PointerNum = Builder.getInt32(C: Info.NumberOfPtrs);
9076
9077 // Source location for the ident struct
9078 if (!SrcLocInfo) {
9079 uint32_t SrcLocStrSize;
9080 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
9081 SrcLocInfo = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
9082 }
9083
9084 Value *OffloadingArgs[] = {SrcLocInfo, DeviceID,
9085 PointerNum, RTArgs.BasePointersArray,
9086 RTArgs.PointersArray, RTArgs.SizesArray,
9087 RTArgs.MapTypesArray, RTArgs.MapNamesArray,
9088 RTArgs.MappersArray};
9089 Function *EndMapperFunc =
9090 getOrCreateRuntimeFunctionPtr(FnID: omp::OMPRTL___tgt_target_data_end_mapper);
9091
9092 createRuntimeFunctionCall(Callee: EndMapperFunc, Args: OffloadingArgs);
9093 return Error::success();
9094 };
9095
9096 // We don't have to do anything to close the region if the if clause evaluates
9097 // to false.
9098 auto EndElseGen = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
9099 ArrayRef<BasicBlock *> DeallocBlocks) {
9100 return Error::success();
9101 };
9102
9103 Error Err = [&]() -> Error {
9104 if (BodyGenCB) {
9105 Error Err = [&]() {
9106 if (IfCond)
9107 return emitIfClause(Cond: IfCond, ThenGen: BeginThenGen, ElseGen: BeginElseGen, AllocaIP);
9108 return BeginThenGen(AllocaIP, Builder.saveIP(), DeallocBlocks);
9109 }();
9110
9111 if (Err)
9112 return Err;
9113
9114 // If we don't require privatization of device pointers, we emit the body
9115 // in between the runtime calls. This avoids duplicating the body code.
9116 InsertPointOrErrorTy AfterIP =
9117 BodyGenCB(Builder.saveIP(), BodyGenTy::NoPriv);
9118 if (!AfterIP)
9119 return AfterIP.takeError();
9120 restoreIPandDebugLoc(Builder, IP: *AfterIP);
9121
9122 if (IfCond)
9123 return emitIfClause(Cond: IfCond, ThenGen: EndThenGen, ElseGen: EndElseGen, AllocaIP);
9124 return EndThenGen(AllocaIP, Builder.saveIP(), DeallocBlocks);
9125 }
9126 if (IfCond)
9127 return emitIfClause(Cond: IfCond, ThenGen: BeginThenGen, ElseGen: EndElseGen, AllocaIP);
9128 return BeginThenGen(AllocaIP, Builder.saveIP(), DeallocBlocks);
9129 }();
9130
9131 if (Err)
9132 return Err;
9133
9134 return Builder.saveIP();
9135}
9136
9137FunctionCallee
9138OpenMPIRBuilder::createForStaticInitFunction(unsigned IVSize, bool IVSigned,
9139 bool IsGPUDistribute) {
9140 assert((IVSize == 32 || IVSize == 64) &&
9141 "IV size is not compatible with the omp runtime");
9142 RuntimeFunction Name;
9143 if (IsGPUDistribute)
9144 Name = IVSize == 32
9145 ? (IVSigned ? omp::OMPRTL___kmpc_distribute_static_init_4
9146 : omp::OMPRTL___kmpc_distribute_static_init_4u)
9147 : (IVSigned ? omp::OMPRTL___kmpc_distribute_static_init_8
9148 : omp::OMPRTL___kmpc_distribute_static_init_8u);
9149 else
9150 Name = IVSize == 32 ? (IVSigned ? omp::OMPRTL___kmpc_for_static_init_4
9151 : omp::OMPRTL___kmpc_for_static_init_4u)
9152 : (IVSigned ? omp::OMPRTL___kmpc_for_static_init_8
9153 : omp::OMPRTL___kmpc_for_static_init_8u);
9154
9155 return getOrCreateRuntimeFunction(M, FnID: Name);
9156}
9157
9158FunctionCallee OpenMPIRBuilder::createDispatchInitFunction(unsigned IVSize,
9159 bool IVSigned) {
9160 assert((IVSize == 32 || IVSize == 64) &&
9161 "IV size is not compatible with the omp runtime");
9162 RuntimeFunction Name = IVSize == 32
9163 ? (IVSigned ? omp::OMPRTL___kmpc_dispatch_init_4
9164 : omp::OMPRTL___kmpc_dispatch_init_4u)
9165 : (IVSigned ? omp::OMPRTL___kmpc_dispatch_init_8
9166 : omp::OMPRTL___kmpc_dispatch_init_8u);
9167
9168 return getOrCreateRuntimeFunction(M, FnID: Name);
9169}
9170
9171FunctionCallee OpenMPIRBuilder::createDispatchNextFunction(unsigned IVSize,
9172 bool IVSigned) {
9173 assert((IVSize == 32 || IVSize == 64) &&
9174 "IV size is not compatible with the omp runtime");
9175 RuntimeFunction Name = IVSize == 32
9176 ? (IVSigned ? omp::OMPRTL___kmpc_dispatch_next_4
9177 : omp::OMPRTL___kmpc_dispatch_next_4u)
9178 : (IVSigned ? omp::OMPRTL___kmpc_dispatch_next_8
9179 : omp::OMPRTL___kmpc_dispatch_next_8u);
9180
9181 return getOrCreateRuntimeFunction(M, FnID: Name);
9182}
9183
9184FunctionCallee OpenMPIRBuilder::createDispatchFiniFunction(unsigned IVSize,
9185 bool IVSigned) {
9186 assert((IVSize == 32 || IVSize == 64) &&
9187 "IV size is not compatible with the omp runtime");
9188 RuntimeFunction Name = IVSize == 32
9189 ? (IVSigned ? omp::OMPRTL___kmpc_dispatch_fini_4
9190 : omp::OMPRTL___kmpc_dispatch_fini_4u)
9191 : (IVSigned ? omp::OMPRTL___kmpc_dispatch_fini_8
9192 : omp::OMPRTL___kmpc_dispatch_fini_8u);
9193
9194 return getOrCreateRuntimeFunction(M, FnID: Name);
9195}
9196
9197FunctionCallee OpenMPIRBuilder::createDispatchDeinitFunction() {
9198 return getOrCreateRuntimeFunction(M, FnID: omp::OMPRTL___kmpc_dispatch_deinit);
9199}
9200
9201static void FixupDebugInfoForOutlinedFunction(
9202 OpenMPIRBuilder &OMPBuilder, IRBuilderBase &Builder, Function *Func,
9203 DenseMap<Value *, std::tuple<Value *, unsigned>> &ValueReplacementMap) {
9204
9205 DISubprogram *NewSP = Func->getSubprogram();
9206 if (!NewSP)
9207 return;
9208
9209 SmallDenseMap<DILocalVariable *, DILocalVariable *> RemappedVariables;
9210
9211 auto GetUpdatedDIVariable = [&](DILocalVariable *OldVar, unsigned arg) {
9212 DILocalVariable *&NewVar = RemappedVariables[OldVar];
9213 // Only use cached variable if the arg number matches. This is important
9214 // so that DIVariable created for privatized variables are not discarded.
9215 if (NewVar && (arg == NewVar->getArg()))
9216 return NewVar;
9217
9218 NewVar = llvm::DILocalVariable::get(
9219 Context&: Builder.getContext(), Scope: OldVar->getScope(), Name: OldVar->getName(),
9220 File: OldVar->getFile(), Line: OldVar->getLine(), Type: OldVar->getType(), Arg: arg,
9221 Flags: OldVar->getFlags(), AlignInBits: OldVar->getAlignInBits(), Annotations: OldVar->getAnnotations());
9222 return NewVar;
9223 };
9224
9225 auto UpdateDebugRecord = [&](auto *DR) {
9226 DILocalVariable *OldVar = DR->getVariable();
9227 unsigned ArgNo = 0;
9228 for (auto Loc : DR->location_ops()) {
9229 auto Iter = ValueReplacementMap.find(Loc);
9230 if (Iter != ValueReplacementMap.end()) {
9231 DR->replaceVariableLocationOp(Loc, std::get<0>(Iter->second));
9232 ArgNo = std::get<1>(Iter->second) + 1;
9233 }
9234 }
9235 if (ArgNo != 0)
9236 DR->setVariable(GetUpdatedDIVariable(OldVar, ArgNo));
9237 };
9238
9239 SmallVector<DbgVariableRecord *, 4> DVRsToDelete;
9240 auto MoveDebugRecordToCorrectBlock = [&](DbgVariableRecord *DVR) {
9241 if (DVR->getNumVariableLocationOps() != 1u) {
9242 DVR->setKillLocation();
9243 return;
9244 }
9245 Value *Loc = DVR->getVariableLocationOp(OpIdx: 0u);
9246 BasicBlock *CurBB = DVR->getParent();
9247 BasicBlock *RequiredBB = nullptr;
9248
9249 if (Instruction *LocInst = dyn_cast<Instruction>(Val: Loc))
9250 RequiredBB = LocInst->getParent();
9251 else if (isa<llvm::Argument>(Val: Loc))
9252 RequiredBB = &DVR->getFunction()->getEntryBlock();
9253
9254 if (RequiredBB && RequiredBB != CurBB) {
9255 assert(!RequiredBB->empty());
9256 RequiredBB->insertDbgRecordBefore(DR: DVR->clone(),
9257 Here: RequiredBB->back().getIterator());
9258 DVRsToDelete.push_back(Elt: DVR);
9259 }
9260 };
9261
9262 // The location and scope of variable intrinsics and records still point to
9263 // the parent function of the target region. Update them.
9264 for (Instruction &I : instructions(F: Func)) {
9265 assert(!isa<llvm::DbgVariableIntrinsic>(&I) &&
9266 "Unexpected debug intrinsic");
9267 for (DbgVariableRecord &DVR : filterDbgVars(R: I.getDbgRecordRange())) {
9268 UpdateDebugRecord(&DVR);
9269 MoveDebugRecordToCorrectBlock(&DVR);
9270 }
9271 }
9272 for (auto *DVR : DVRsToDelete)
9273 DVR->getMarker()->MarkedInstr->dropOneDbgRecord(I: DVR);
9274 // An extra argument is passed to the device. Create the debug data for it.
9275 if (OMPBuilder.Config.isTargetDevice()) {
9276 DICompileUnit *CU = NewSP->getUnit();
9277 Module *M = Func->getParent();
9278 DIBuilder DB(*M, true, CU);
9279 DIType *VoidPtrTy =
9280 DB.createQualifiedType(Tag: dwarf::DW_TAG_pointer_type, FromTy: nullptr);
9281 unsigned ArgNo = Func->arg_size();
9282 DILocalVariable *Var = DB.createParameterVariable(
9283 Scope: NewSP, Name: "dyn_ptr", ArgNo, File: NewSP->getFile(), /*LineNo=*/0, Ty: VoidPtrTy,
9284 /*AlwaysPreserve=*/false, Flags: DINode::DIFlags::FlagArtificial);
9285 auto Loc = DILocation::get(Context&: Func->getContext(), Line: 0, Column: 0, Scope: NewSP, InlinedAt: 0);
9286 Argument *LastArg = Func->getArg(i: Func->arg_size() - 1);
9287 DB.insertDeclare(Storage: LastArg, VarInfo: Var, Expr: DB.createExpression(), DL: Loc,
9288 InsertAtEnd: &(*Func->begin()));
9289 }
9290}
9291
9292static Value *removeASCastIfPresent(Value *V) {
9293 if (Operator::getOpcode(V) == Instruction::AddrSpaceCast)
9294 return cast<Operator>(Val: V)->getOperand(i: 0);
9295 return V;
9296}
9297
9298static Expected<Function *> createOutlinedFunction(
9299 OpenMPIRBuilder &OMPBuilder, IRBuilderBase &Builder,
9300 const OpenMPIRBuilder::TargetKernelDefaultAttrs &DefaultAttrs,
9301 StringRef FuncName, SmallVectorImpl<Value *> &Inputs,
9302 OpenMPIRBuilder::TargetBodyGenCallbackTy &CBFunc,
9303 OpenMPIRBuilder::TargetGenArgAccessorsCallbackTy &ArgAccessorFuncCB,
9304 DebugLoc OutlinedFnLoc) {
9305 SmallVector<Type *> ParameterTypes;
9306 if (OMPBuilder.Config.isTargetDevice()) {
9307 // All parameters to target devices are passed as pointers
9308 // or i64. This assumes 64-bit address spaces/pointers.
9309 for (auto &Arg : Inputs)
9310 ParameterTypes.push_back(Elt: Arg->getType()->isPointerTy()
9311 ? Arg->getType()
9312 : Type::getInt64Ty(C&: Builder.getContext()));
9313 } else {
9314 for (auto &Arg : Inputs)
9315 ParameterTypes.push_back(Elt: Arg->getType());
9316 }
9317
9318 // The implicit dyn_ptr argument is always the last parameter on both host
9319 // and device so the argument counts match without runtime manipulation.
9320 auto *PtrTy = PointerType::getUnqual(C&: Builder.getContext());
9321 ParameterTypes.push_back(Elt: PtrTy);
9322
9323 auto BB = Builder.GetInsertBlock();
9324 auto M = BB->getModule();
9325 auto FuncType = FunctionType::get(Result: Builder.getVoidTy(), Params: ParameterTypes,
9326 /*isVarArg*/ false);
9327 auto Func =
9328 Function::Create(Ty: FuncType, Linkage: GlobalValue::InternalLinkage, N: FuncName, M);
9329
9330 // Forward target-cpu and target-features function attributes from the
9331 // original function to the new outlined function.
9332 Function *ParentFn = Builder.GetInsertBlock()->getParent();
9333
9334 auto TargetCpuAttr = ParentFn->getFnAttribute(Kind: "target-cpu");
9335 if (TargetCpuAttr.isStringAttribute())
9336 Func->addFnAttr(Attr: TargetCpuAttr);
9337
9338 auto TargetFeaturesAttr = ParentFn->getFnAttribute(Kind: "target-features");
9339 if (TargetFeaturesAttr.isStringAttribute())
9340 Func->addFnAttr(Attr: TargetFeaturesAttr);
9341
9342 if (OMPBuilder.Config.isTargetDevice()) {
9343 Value *ExecMode =
9344 OMPBuilder.emitKernelExecutionMode(KernelName: FuncName, Mode: DefaultAttrs.ExecFlags);
9345 OMPBuilder.emitUsed(Name: "llvm.compiler.used", List: {ExecMode});
9346 }
9347
9348 // Save insert point.
9349 IRBuilder<>::InsertPointGuard IPG(Builder);
9350 // We will generate the entries in the outlined function but the debug
9351 // location is still pointing to the parent function, which is the wrong
9352 // scope. OutlinedFnLoc, when the caller provides one, is the same source
9353 // position scoped to the subprogram that will be attached to the outlined
9354 // function, so it is what everything emitted below needs.
9355 Builder.SetCurrentDebugLocation(OutlinedFnLoc);
9356
9357 // Generate the region into the function.
9358 BasicBlock *EntryBB = BasicBlock::Create(Context&: Builder.getContext(), Name: "entry", Parent: Func);
9359 Builder.SetInsertPoint(EntryBB);
9360
9361 // Insert target init call in the device compilation pass. On the host
9362 // (e.g. a non-GPU offload target), there is no runtime init/deinit
9363 // sequence, but the runtime still needs a '<kernel>_kernel_environment'
9364 // global to know how the kernel was configured, so emit it directly.
9365 if (OMPBuilder.Config.isTargetDevice())
9366 Builder.restoreIP(IP: OMPBuilder.createTargetInit(Loc: Builder, Attrs: DefaultAttrs));
9367 else
9368 OMPBuilder.emitKernelEnvironment(Loc: Builder, Attrs: DefaultAttrs);
9369
9370 BasicBlock *UserCodeEntryBB = Builder.GetInsertBlock();
9371
9372 // As we embed the user code in the middle of our target region after we
9373 // generate entry code, we must move what allocas we can into the entry
9374 // block to avoid possible breaking optimisations for device
9375 if (OMPBuilder.Config.isTargetDevice())
9376 OMPBuilder.ConstantAllocaRaiseCandidates.emplace_back(Args&: Func);
9377
9378 BasicBlock *ExitBB = splitBB(Builder, /*CreateBranch=*/true, Name: "target.exit");
9379 BasicBlock *OutlinedBodyBB =
9380 splitBB(Builder, /*CreateBranch=*/true, Name: "outlined.body");
9381 llvm::OpenMPIRBuilder::InsertPointOrErrorTy AfterIP =
9382 CBFunc(Builder.saveIP(), OutlinedBodyBB->begin(), ExitBB);
9383 if (!AfterIP)
9384 return AfterIP.takeError();
9385 Builder.SetInsertPoint(ExitBB);
9386 // The body callback builds the body with its own IRBuilder and cannot reach
9387 // this one directly. But a body holding another OpenMP construct, a nested
9388 // parallel say, calls OpenMPIRBuilder::createParallel, and that can leave
9389 // this Builder pointing at the wrong debug location, or at none at all. The
9390 // epilogue below belongs to the target construct rather than to whatever the
9391 // body emitted last, so re-establish the location the prologue was emitted
9392 // with.
9393 Builder.SetCurrentDebugLocation(OutlinedFnLoc);
9394
9395 // Insert target deinit call in the device compilation pass.
9396 if (OMPBuilder.Config.isTargetDevice())
9397 OMPBuilder.createTargetDeinit(Loc: Builder);
9398
9399 // Insert return instruction.
9400 Builder.CreateRetVoid();
9401
9402 // New Alloca IP at entry point of created device function.
9403 Builder.SetInsertPoint(EntryBB->getFirstNonPHIIt());
9404 auto AllocaIP = Builder.saveIP();
9405
9406 Builder.SetInsertPoint(UserCodeEntryBB->getFirstNonPHIOrDbg());
9407
9408 // Do not include the artificial dyn_ptr argument.
9409 const auto &ArgRange = make_range(x: Func->arg_begin(), y: Func->arg_end() - 1);
9410
9411 DenseMap<Value *, std::tuple<Value *, unsigned>> ValueReplacementMap;
9412
9413 auto ReplaceValue = [](Value *Input, Value *InputCopy, Function *Func) {
9414 // Things like GEP's can come in the form of Constants. Constants and
9415 // ConstantExpr's do not have access to the knowledge of what they're
9416 // contained in, so we must dig a little to find an instruction so we
9417 // can tell if they're used inside of the function we're outlining. We
9418 // also replace the original constant expression with a new instruction
9419 // equivalent; an instruction as it allows easy modification in the
9420 // following loop, as we can now know the constant (instruction) is
9421 // owned by our target function and replaceUsesOfWith can now be invoked
9422 // on it (cannot do this with constants it seems). A brand new one also
9423 // allows us to be cautious as it is perhaps possible the old expression
9424 // was used inside of the function but exists and is used externally
9425 // (unlikely by the nature of a Constant, but still).
9426 // NOTE: We cannot remove dead constants that have been rewritten to
9427 // instructions at this stage, we run the risk of breaking later lowering
9428 // by doing so as we could still be in the process of lowering the module
9429 // from MLIR to LLVM-IR and the MLIR lowering may still require the original
9430 // constants we have created rewritten versions of.
9431 if (auto *Const = dyn_cast<Constant>(Val: Input))
9432 convertUsersOfConstantsToInstructions(Consts: Const, RestrictToFunc: Func, RemoveDeadConstants: false);
9433
9434 // Collect users before iterating over them to avoid invalidating the
9435 // iteration in case a user uses Input more than once (e.g. a call
9436 // instruction).
9437 SetVector<User *> Users(Input->users().begin(), Input->users().end());
9438 // Collect all the instructions
9439 for (User *User : make_early_inc_range(Range&: Users))
9440 if (auto *Instr = dyn_cast<Instruction>(Val: User))
9441 if (Instr->getFunction() == Func)
9442 Instr->replaceUsesOfWith(From: Input, To: InputCopy);
9443 };
9444
9445 SmallVector<std::pair<Value *, Value *>> DeferredReplacement;
9446
9447 // Rewrite uses of input valus to parameters.
9448 for (auto InArg : zip(t&: Inputs, u: ArgRange)) {
9449 Value *Input = std::get<0>(t&: InArg);
9450 Argument &Arg = std::get<1>(t&: InArg);
9451 Value *InputCopy = nullptr;
9452
9453 llvm::OpenMPIRBuilder::InsertPointOrErrorTy AfterIP = ArgAccessorFuncCB(
9454 Arg, Input, InputCopy, AllocaIP, Builder.saveIP(), ExitBB->begin());
9455 if (!AfterIP)
9456 return AfterIP.takeError();
9457 Builder.restoreIP(IP: *AfterIP);
9458 ValueReplacementMap[Input] = std::make_tuple(args&: InputCopy, args: Arg.getArgNo());
9459
9460 // In certain cases a Global may be set up for replacement, however, this
9461 // Global may be used in multiple arguments to the kernel, just segmented
9462 // apart, for example, if we have a global array, that is sectioned into
9463 // multiple mappings (technically not legal in OpenMP, but there is a case
9464 // in Fortran for Common Blocks where this is neccesary), we will end up
9465 // with GEP's into this array inside the kernel, that refer to the Global
9466 // but are technically separate arguments to the kernel for all intents and
9467 // purposes. If we have mapped a segment that requires a GEP into the 0-th
9468 // index, it will fold into an referal to the Global, if we then encounter
9469 // this folded GEP during replacement all of the references to the
9470 // Global in the kernel will be replaced with the argument we have generated
9471 // that corresponds to it, including any other GEP's that refer to the
9472 // Global that may be other arguments. This will invalidate all of the other
9473 // preceding mapped arguments that refer to the same global that may be
9474 // separate segments. To prevent this, we defer global processing until all
9475 // other processing has been performed.
9476 if (llvm::isa<llvm::GlobalValue, llvm::GlobalObject, llvm::GlobalVariable>(
9477 Val: removeASCastIfPresent(V: Input))) {
9478 DeferredReplacement.push_back(Elt: std::make_pair(x&: Input, y&: InputCopy));
9479 continue;
9480 }
9481
9482 if (isa<ConstantData>(Val: Input))
9483 continue;
9484
9485 ReplaceValue(Input, InputCopy, Func);
9486 }
9487
9488 // Replace all of our deferred Input values, currently just Globals.
9489 for (auto Deferred : DeferredReplacement)
9490 ReplaceValue(std::get<0>(in&: Deferred), std::get<1>(in&: Deferred), Func);
9491
9492 FixupDebugInfoForOutlinedFunction(OMPBuilder, Builder, Func,
9493 ValueReplacementMap);
9494 return Func;
9495}
9496/// Given a task descriptor, TaskWithPrivates, return the pointer to the block
9497/// of pointers containing shared data between the parent task and the created
9498/// task.
9499static LoadInst *loadSharedDataFromTaskDescriptor(OpenMPIRBuilder &OMPIRBuilder,
9500 IRBuilderBase &Builder,
9501 Value *TaskWithPrivates,
9502 Type *TaskWithPrivatesTy) {
9503
9504 Type *TaskTy = OMPIRBuilder.Task;
9505 LLVMContext &Ctx = Builder.getContext();
9506 Value *TaskT =
9507 Builder.CreateStructGEP(Ty: TaskWithPrivatesTy, Ptr: TaskWithPrivates, Idx: 0);
9508 Value *Shareds = TaskT;
9509 // TaskWithPrivatesTy can be one of the following
9510 // 1. %struct.task_with_privates = type { %struct.kmp_task_ompbuilder_t,
9511 // %struct.privates }
9512 // 2. %struct.kmp_task_ompbuilder_t ;; This is simply TaskTy
9513 //
9514 // In the former case, that is when TaskWithPrivatesTy != TaskTy,
9515 // its first member has to be the task descriptor. TaskTy is the type of the
9516 // task descriptor. TaskT is the pointer to the task descriptor. Loading the
9517 // first member of TaskT, gives us the pointer to shared data.
9518 if (TaskWithPrivatesTy != TaskTy)
9519 Shareds = Builder.CreateStructGEP(Ty: TaskTy, Ptr: TaskT, Idx: 0);
9520 return Builder.CreateLoad(Ty: PointerType::getUnqual(C&: Ctx), Ptr: Shareds);
9521}
9522/// Create an entry point for a target task with the following.
9523/// It'll have the following signature
9524/// void @.omp_target_task_proxy_func(i32 %thread.id, ptr %task)
9525/// This function is called from emitTargetTask once the
9526/// code to launch the target kernel has been outlined already.
9527/// NumOffloadingArrays is the number of offloading arrays that we need to copy
9528/// into the task structure so that the deferred target task can access this
9529/// data even after the stack frame of the generating task has been rolled
9530/// back. Offloading arrays contain base pointers, pointers, sizes etc
9531/// of the data that the target kernel will access. These in effect are the
9532/// non-empty arrays of pointers held by OpenMPIRBuilder::TargetDataRTArgs.
9533static Function *emitTargetTaskProxyFunction(
9534 OpenMPIRBuilder &OMPBuilder, IRBuilderBase &Builder, CallInst *StaleCI,
9535 StructType *PrivatesTy, StructType *TaskWithPrivatesTy,
9536 const size_t NumOffloadingArrays, const int SharedArgsOperandNo) {
9537
9538 // If NumOffloadingArrays is non-zero, PrivatesTy better not be nullptr.
9539 // This is because PrivatesTy is the type of the structure in which
9540 // we pass the offloading arrays to the deferred target task.
9541 assert((!NumOffloadingArrays || PrivatesTy) &&
9542 "PrivatesTy cannot be nullptr when there are offloadingArrays"
9543 "to privatize");
9544
9545 Module &M = OMPBuilder.M;
9546 // KernelLaunchFunction is the target launch function, i.e.
9547 // the function that sets up kernel arguments and calls
9548 // __tgt_target_kernel to launch the kernel on the device.
9549 //
9550 Function *KernelLaunchFunction = StaleCI->getCalledFunction();
9551
9552 // StaleCI is the CallInst which is the call to the outlined
9553 // target kernel launch function. If there are local live-in values
9554 // that the outlined function uses then these are aggregated into a structure
9555 // which is passed as the second argument. If there are no local live-in
9556 // values or if all values used by the outlined kernel are global variables,
9557 // then there's only one argument, the threadID. So, StaleCI can be
9558 //
9559 // %structArg = alloca { ptr, ptr }, align 8
9560 // %gep_ = getelementptr { ptr, ptr }, ptr %structArg, i32 0, i32 0
9561 // store ptr %20, ptr %gep_, align 8
9562 // %gep_8 = getelementptr { ptr, ptr }, ptr %structArg, i32 0, i32 1
9563 // store ptr %21, ptr %gep_8, align 8
9564 // call void @_QQmain..omp_par.1(i32 %global.tid.val6, ptr %structArg)
9565 //
9566 // OR
9567 //
9568 // call void @_QQmain..omp_par.1(i32 %global.tid.val6)
9569 LLVMContext &Ctx = StaleCI->getParent()->getContext();
9570
9571 Type *ThreadIDTy = Type::getInt32Ty(C&: Ctx);
9572 Type *TaskPtrTy = OMPBuilder.TaskPtr;
9573 [[maybe_unused]] Type *TaskTy = OMPBuilder.Task;
9574
9575 auto ProxyFnTy =
9576 FunctionType::get(Result: Builder.getVoidTy(), Params: {ThreadIDTy, TaskPtrTy},
9577 /* isVarArg */ false);
9578 auto ProxyFn = Function::Create(Ty: ProxyFnTy, Linkage: GlobalValue::InternalLinkage,
9579 N: ".omp_target_task_proxy_func", M);
9580 Value *ThreadId = ProxyFn->getArg(i: 0);
9581 Value *TaskWithPrivates = ProxyFn->getArg(i: 1);
9582 ThreadId->setName("thread.id");
9583 TaskWithPrivates->setName("task");
9584
9585 bool HasShareds = SharedArgsOperandNo > 0;
9586 bool HasOffloadingArrays = NumOffloadingArrays > 0;
9587 IRBuilder<>::InsertPointGuard IPG(Builder);
9588 BasicBlock *EntryBB =
9589 BasicBlock::Create(Context&: Builder.getContext(), Name: "entry", Parent: ProxyFn);
9590 Builder.SetInsertPoint(EntryBB);
9591 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
9592
9593 SmallVector<Value *> KernelLaunchArgs;
9594 KernelLaunchArgs.reserve(N: StaleCI->arg_size());
9595 KernelLaunchArgs.push_back(Elt: ThreadId);
9596
9597 if (HasOffloadingArrays) {
9598 assert(TaskTy != TaskWithPrivatesTy &&
9599 "If there are offloading arrays to pass to the target"
9600 "TaskTy cannot be the same as TaskWithPrivatesTy");
9601 (void)TaskTy;
9602 Value *Privates =
9603 Builder.CreateStructGEP(Ty: TaskWithPrivatesTy, Ptr: TaskWithPrivates, Idx: 1);
9604 for (unsigned int i = 0; i < NumOffloadingArrays; ++i)
9605 KernelLaunchArgs.push_back(
9606 Elt: Builder.CreateStructGEP(Ty: PrivatesTy, Ptr: Privates, Idx: i));
9607 }
9608
9609 if (HasShareds) {
9610 auto *ArgStructAlloca =
9611 dyn_cast<AllocaInst>(Val: StaleCI->getArgOperand(i: SharedArgsOperandNo));
9612 assert(ArgStructAlloca &&
9613 "Unable to find the alloca instruction corresponding to arguments "
9614 "for extracted function");
9615 auto *ArgStructType = cast<StructType>(Val: ArgStructAlloca->getAllocatedType());
9616 std::optional<TypeSize> ArgAllocSize =
9617 ArgStructAlloca->getAllocationSize(DL: M.getDataLayout());
9618 assert(ArgStructType && ArgAllocSize &&
9619 "Unable to determine size of arguments for extracted function");
9620 uint64_t StructSize = ArgAllocSize->getFixedValue();
9621
9622 AllocaInst *NewArgStructAlloca =
9623 Builder.CreateAlloca(Ty: ArgStructType, ArraySize: nullptr, Name: "structArg");
9624
9625 Value *SharedsSize = Builder.getInt64(C: StructSize);
9626
9627 LoadInst *LoadShared = loadSharedDataFromTaskDescriptor(
9628 OMPIRBuilder&: OMPBuilder, Builder, TaskWithPrivates, TaskWithPrivatesTy);
9629
9630 Builder.CreateMemCpy(
9631 Dst: NewArgStructAlloca, DstAlign: NewArgStructAlloca->getAlign(), Src: LoadShared,
9632 SrcAlign: LoadShared->getPointerAlignment(DL: M.getDataLayout()), Size: SharedsSize);
9633 KernelLaunchArgs.push_back(Elt: NewArgStructAlloca);
9634 }
9635 OMPBuilder.createRuntimeFunctionCall(Callee: KernelLaunchFunction, Args: KernelLaunchArgs);
9636 Builder.CreateRetVoid();
9637 return ProxyFn;
9638}
9639static Type *getOffloadingArrayType(Value *V) {
9640
9641 if (auto *GEP = dyn_cast<GetElementPtrInst>(Val: V))
9642 return GEP->getSourceElementType();
9643 if (auto *Alloca = dyn_cast<AllocaInst>(Val: V))
9644 return Alloca->getAllocatedType();
9645
9646 llvm_unreachable("Unhandled Instruction type");
9647 return nullptr;
9648}
9649// This function returns a struct that has at most two members.
9650// The first member is always %struct.kmp_task_ompbuilder_t, that is the task
9651// descriptor. The second member, if needed, is a struct containing arrays
9652// that need to be passed to the offloaded target kernel. For example,
9653// if .offload_baseptrs, .offload_ptrs and .offload_sizes have to be passed to
9654// the target kernel and their types are [3 x ptr], [3 x ptr] and [3 x i64]
9655// respectively, then the types created by this function are
9656//
9657// %struct.privates = type { [3 x ptr], [3 x ptr], [3 x i64] }
9658// %struct.task_with_privates = type { %struct.kmp_task_ompbuilder_t,
9659// %struct.privates }
9660// %struct.task_with_privates is returned by this function.
9661// If there aren't any offloading arrays to pass to the target kernel,
9662// %struct.kmp_task_ompbuilder_t is returned.
9663static StructType *
9664createTaskWithPrivatesTy(OpenMPIRBuilder &OMPIRBuilder,
9665 ArrayRef<Value *> OffloadingArraysToPrivatize) {
9666
9667 if (OffloadingArraysToPrivatize.empty())
9668 return OMPIRBuilder.Task;
9669
9670 SmallVector<Type *, 4> StructFieldTypes;
9671 for (Value *V : OffloadingArraysToPrivatize) {
9672 assert(V->getType()->isPointerTy() &&
9673 "Expected pointer to array to privatize. Got a non-pointer value "
9674 "instead");
9675 Type *ArrayTy = getOffloadingArrayType(V);
9676 assert(ArrayTy && "ArrayType cannot be nullptr");
9677 StructFieldTypes.push_back(Elt: ArrayTy);
9678 }
9679 StructType *PrivatesStructTy =
9680 StructType::create(Elements: StructFieldTypes, Name: "struct.privates");
9681 return StructType::create(Elements: {OMPIRBuilder.Task, PrivatesStructTy},
9682 Name: "struct.task_with_privates");
9683}
9684static Error emitTargetOutlinedFunction(
9685 OpenMPIRBuilder &OMPBuilder, IRBuilderBase &Builder, bool IsOffloadEntry,
9686 TargetRegionEntryInfo &EntryInfo,
9687 const OpenMPIRBuilder::TargetKernelDefaultAttrs &DefaultAttrs,
9688 Function *&OutlinedFn, Constant *&OutlinedFnID,
9689 SmallVectorImpl<Value *> &Inputs,
9690 OpenMPIRBuilder::TargetBodyGenCallbackTy &CBFunc,
9691 OpenMPIRBuilder::TargetGenArgAccessorsCallbackTy &ArgAccessorFuncCB,
9692 DebugLoc OutlinedFnLoc) {
9693
9694 OpenMPIRBuilder::FunctionGenCallback &&GenerateOutlinedFunction =
9695 [&](StringRef EntryFnName) {
9696 return createOutlinedFunction(OMPBuilder, Builder, DefaultAttrs,
9697 FuncName: EntryFnName, Inputs, CBFunc,
9698 ArgAccessorFuncCB, OutlinedFnLoc);
9699 };
9700
9701 return OMPBuilder.emitTargetRegionFunction(
9702 EntryInfo, GenerateFunctionCallback&: GenerateOutlinedFunction, IsOffloadEntry, OutlinedFn,
9703 OutlinedFnID);
9704}
9705
9706OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::emitTargetTask(
9707 TargetTaskBodyCallbackTy TaskBodyCB, Value *DeviceID, Value *RTLoc,
9708 OpenMPIRBuilder::InsertPointTy AllocaIP,
9709 const DependenciesInfo &Dependencies, const TargetDataRTArgs &RTArgs,
9710 bool HasNoWait) {
9711
9712 // The following explains the code-gen scenario for the `target` directive. A
9713 // similar scneario is followed for other device-related directives (e.g.
9714 // `target enter data`) but in similar fashion since we only need to emit task
9715 // that encapsulates the proper runtime call.
9716 //
9717 // When we arrive at this function, the target region itself has been
9718 // outlined into the function OutlinedFn.
9719 // So at ths point, for
9720 // --------------------------------------------------------------
9721 // void user_code_that_offloads(...) {
9722 // omp target depend(..) map(from:a) map(to:b) private(i)
9723 // do i = 1, 10
9724 // a(i) = b(i) + n
9725 // }
9726 //
9727 // --------------------------------------------------------------
9728 //
9729 // we have
9730 //
9731 // --------------------------------------------------------------
9732 //
9733 // void user_code_that_offloads(...) {
9734 // %.offload_baseptrs = alloca [2 x ptr], align 8
9735 // %.offload_ptrs = alloca [2 x ptr], align 8
9736 // %.offload_mappers = alloca [2 x ptr], align 8
9737 // ;; target region has been outlined and now we need to
9738 // ;; offload to it via a target task.
9739 // }
9740 // void outlined_device_function(ptr a, ptr b, ptr n) {
9741 // n = *n_ptr;
9742 // do i = 1, 10
9743 // a(i) = b(i) + n
9744 // }
9745 //
9746 // We have to now do the following
9747 // (i) Make an offloading call to outlined_device_function using the OpenMP
9748 // RTL. See 'kernel_launch_function' in the pseudo code below. This is
9749 // emitted by emitKernelLaunch
9750 // (ii) Create a task entry point function that calls kernel_launch_function
9751 // and is the entry point for the target task. See
9752 // '@.omp_target_task_proxy_func in the pseudocode below.
9753 // (iii) Create a task with the task entry point created in (ii)
9754 //
9755 // That is we create the following
9756 // struct task_with_privates {
9757 // struct kmp_task_ompbuilder_t task_struct;
9758 // struct privates {
9759 // [2 x ptr] ; baseptrs
9760 // [2 x ptr] ; ptrs
9761 // [2 x i64] ; sizes
9762 // }
9763 // }
9764 // void user_code_that_offloads(...) {
9765 // %.offload_baseptrs = alloca [2 x ptr], align 8
9766 // %.offload_ptrs = alloca [2 x ptr], align 8
9767 // %.offload_sizes = alloca [2 x i64], align 8
9768 //
9769 // %structArg = alloca { ptr, ptr, ptr }, align 8
9770 // %strucArg[0] = a
9771 // %strucArg[1] = b
9772 // %strucArg[2] = &n
9773 //
9774 // target_task_with_privates = @__kmpc_omp_target_task_alloc(...,
9775 // sizeof(kmp_task_ompbuilder_t),
9776 // sizeof(structArg),
9777 // @.omp_target_task_proxy_func,
9778 // ...)
9779 // memcpy(target_task_with_privates->task_struct->shareds, %structArg,
9780 // sizeof(structArg))
9781 // memcpy(target_task_with_privates->privates->baseptrs,
9782 // offload_baseptrs, sizeof(offload_baseptrs)
9783 // memcpy(target_task_with_privates->privates->ptrs,
9784 // offload_ptrs, sizeof(offload_ptrs)
9785 // memcpy(target_task_with_privates->privates->sizes,
9786 // offload_sizes, sizeof(offload_sizes)
9787 // dependencies_array = ...
9788 // ;; if nowait not present
9789 // call @__kmpc_omp_wait_deps(..., dependencies_array)
9790 // call @__kmpc_omp_task_begin_if0(...)
9791 // call @ @.omp_target_task_proxy_func(i32 thread_id, ptr
9792 // %target_task_with_privates)
9793 // call @__kmpc_omp_task_complete_if0(...)
9794 // }
9795 //
9796 // define internal void @.omp_target_task_proxy_func(i32 %thread.id,
9797 // ptr %task) {
9798 // %structArg = alloca {ptr, ptr, ptr}
9799 // %task_ptr = getelementptr(%task, 0, 0)
9800 // %shared_data = load (getelementptr %task_ptr, 0, 0)
9801 // mempcy(%structArg, %shared_data, sizeof(%structArg))
9802 //
9803 // %offloading_arrays = getelementptr(%task, 0, 1)
9804 // %offload_baseptrs = getelementptr(%offloading_arrays, 0, 0)
9805 // %offload_ptrs = getelementptr(%offloading_arrays, 0, 1)
9806 // %offload_sizes = getelementptr(%offloading_arrays, 0, 2)
9807 // kernel_launch_function(%thread.id, %offload_baseptrs, %offload_ptrs,
9808 // %offload_sizes, %structArg)
9809 // }
9810 //
9811 // We need the proxy function because the signature of the task entry point
9812 // expected by kmpc_omp_task is always the same and will be different from
9813 // that of the kernel_launch function.
9814 //
9815 // kernel_launch_function is generated by emitKernelLaunch and has the
9816 // always_inline attribute. For this example, it'll look like so:
9817 // void kernel_launch_function(%thread_id, %offload_baseptrs, %offload_ptrs,
9818 // %offload_sizes, %structArg) alwaysinline {
9819 // %kernel_args = alloca %struct.__tgt_kernel_arguments, align 8
9820 // ; load aggregated data from %structArg
9821 // ; setup kernel_args using offload_baseptrs, offload_ptrs and
9822 // ; offload_sizes
9823 // call i32 @__tgt_target_kernel(...,
9824 // outlined_device_function,
9825 // ptr %kernel_args)
9826 // }
9827 // void outlined_device_function(ptr a, ptr b, ptr n) {
9828 // n = *n_ptr;
9829 // do i = 1, 10
9830 // a(i) = b(i) + n
9831 // }
9832 //
9833 BasicBlock *TargetTaskBodyBB =
9834 splitBB(Builder, /*CreateBranch=*/true, Name: "target.task.body");
9835 BasicBlock *TargetTaskAllocaBB =
9836 splitBB(Builder, /*CreateBranch=*/true, Name: "target.task.alloca");
9837
9838 InsertPointTy TargetTaskAllocaIP(TargetTaskAllocaBB->begin());
9839 InsertPointTy TargetTaskBodyIP(TargetTaskBodyBB->begin());
9840
9841 auto OI = std::make_unique<OutlineInfo>();
9842 OI->EntryBB = TargetTaskAllocaBB;
9843 OI->OuterAllocBB = AllocaIP.getNodeParent();
9844
9845 // Add the thread ID argument.
9846 SmallVector<Instruction *, 4> ToBeDeleted;
9847 OI->ExcludeArgsFromAggregate.push_back(Elt: createFakeIntVal(
9848 Builder, OuterAllocaIP: AllocaIP, ToBeDeleted, InnerAllocaIP: TargetTaskAllocaIP, Name: "global.tid", AsPtr: false));
9849
9850 // Generate the task body which will subsequently be outlined.
9851 Builder.restoreIP(IP: TargetTaskBodyIP);
9852 if (Error Err = TaskBodyCB(DeviceID, RTLoc, TargetTaskAllocaIP))
9853 return Err;
9854
9855 // The outliner (CodeExtractor) extract a sequence or vector of blocks that
9856 // it is given. These blocks are enumerated by
9857 // OpenMPIRBuilder::OutlineInfo::collectBlocks which expects the OI.ExitBlock
9858 // to be outside the region. In other words, OI.ExitBlock is expected to be
9859 // the start of the region after the outlining. We used to set OI.ExitBlock
9860 // to the InsertBlock after TaskBodyCB is done. This is fine in most cases
9861 // except when the task body is a single basic block. In that case,
9862 // OI.ExitBlock is set to the single task body block and will get left out of
9863 // the outlining process. So, simply create a new empty block to which we
9864 // uncoditionally branch from where TaskBodyCB left off
9865 OI->ExitBB = BasicBlock::Create(Context&: Builder.getContext(), Name: "target.task.cont");
9866 emitBlock(BB: OI->ExitBB, CurFn: Builder.GetInsertBlock()->getParent(),
9867 /*IsFinished=*/true);
9868
9869 SmallVector<Value *, 2> OffloadingArraysToPrivatize;
9870 bool NeedsTargetTask = HasNoWait && DeviceID;
9871 if (NeedsTargetTask) {
9872 for (auto *V :
9873 {RTArgs.BasePointersArray, RTArgs.PointersArray, RTArgs.MappersArray,
9874 RTArgs.MapNamesArray, RTArgs.MapTypesArray, RTArgs.MapTypesArrayEnd,
9875 RTArgs.SizesArray}) {
9876 if (V && !isa<ConstantPointerNull, GlobalVariable>(Val: V)) {
9877 OffloadingArraysToPrivatize.push_back(Elt: V);
9878 OI->ExcludeArgsFromAggregate.push_back(Elt: V);
9879 }
9880 }
9881 }
9882 OI->PostOutlineCB = [this, ToBeDeleted, Dependencies, NeedsTargetTask,
9883 DeviceID, OffloadingArraysToPrivatize](
9884 Function &OutlinedFn) mutable {
9885 assert(OutlinedFn.hasOneUse() &&
9886 "there must be a single user for the outlined function");
9887
9888 CallInst *StaleCI = cast<CallInst>(Val: OutlinedFn.user_back());
9889
9890 // The first argument of StaleCI is always the thread id.
9891 // The next few arguments are the pointers to offloading arrays
9892 // if any. (see OffloadingArraysToPrivatize)
9893 // Finally, all other local values that are live-in into the outlined region
9894 // end up in a structure whose pointer is passed as the last argument. This
9895 // piece of data is passed in the "shared" field of the task structure. So,
9896 // we know we have to pass shareds to the task if the number of arguments is
9897 // greater than OffloadingArraysToPrivatize.size() + 1 The 1 is for the
9898 // thread id. Further, for safety, we assert that the number of arguments of
9899 // StaleCI is exactly OffloadingArraysToPrivatize.size() + 2
9900 const unsigned int NumStaleCIArgs = StaleCI->arg_size();
9901 bool HasShareds = NumStaleCIArgs > OffloadingArraysToPrivatize.size() + 1;
9902 assert((!HasShareds ||
9903 NumStaleCIArgs == (OffloadingArraysToPrivatize.size() + 2)) &&
9904 "Wrong number of arguments for StaleCI when shareds are present");
9905 int SharedArgOperandNo =
9906 HasShareds ? OffloadingArraysToPrivatize.size() + 1 : 0;
9907
9908 StructType *TaskWithPrivatesTy =
9909 createTaskWithPrivatesTy(OMPIRBuilder&: *this, OffloadingArraysToPrivatize);
9910 StructType *PrivatesTy = nullptr;
9911
9912 if (!OffloadingArraysToPrivatize.empty())
9913 PrivatesTy =
9914 static_cast<StructType *>(TaskWithPrivatesTy->getElementType(N: 1));
9915
9916 Function *ProxyFn = emitTargetTaskProxyFunction(
9917 OMPBuilder&: *this, Builder, StaleCI, PrivatesTy, TaskWithPrivatesTy,
9918 NumOffloadingArrays: OffloadingArraysToPrivatize.size(), SharedArgsOperandNo: SharedArgOperandNo);
9919
9920 LLVM_DEBUG(dbgs() << "Proxy task entry function created: " << *ProxyFn
9921 << "\n");
9922
9923 Builder.SetInsertPoint(StaleCI);
9924
9925 // Gather the arguments for emitting the runtime call.
9926 uint32_t SrcLocStrSize;
9927 Constant *SrcLocStr =
9928 getOrCreateSrcLocStr(Loc: LocationDescription(Builder), SrcLocStrSize);
9929 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
9930
9931 // @__kmpc_omp_task_alloc or @__kmpc_omp_target_task_alloc
9932 //
9933 // If `HasNoWait == true`, we call @__kmpc_omp_target_task_alloc to provide
9934 // the DeviceID to the deferred task and also since
9935 // @__kmpc_omp_target_task_alloc creates an untied/async task.
9936 Function *TaskAllocFn =
9937 !NeedsTargetTask
9938 ? getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_alloc)
9939 : getOrCreateRuntimeFunctionPtr(
9940 FnID: OMPRTL___kmpc_omp_target_task_alloc);
9941
9942 // Arguments - `loc_ref` (Ident) and `gtid` (ThreadID)
9943 // call.
9944 Value *ThreadID = getOrCreateThreadID(Ident);
9945
9946 // Argument - `sizeof_kmp_task_t` (TaskSize)
9947 // Tasksize refers to the size in bytes of kmp_task_t data structure
9948 // plus any other data to be passed to the target task, if any, which
9949 // is packed into a struct. kmp_task_t and the struct so created are
9950 // packed into a wrapper struct whose type is TaskWithPrivatesTy.
9951 Value *TaskSize = Builder.getInt64(
9952 C: M.getDataLayout().getTypeStoreSize(Ty: TaskWithPrivatesTy));
9953
9954 // Argument - `sizeof_shareds` (SharedsSize)
9955 // SharedsSize refers to the shareds array size in the kmp_task_t data
9956 // structure.
9957 Value *SharedsSize = Builder.getInt64(C: 0);
9958 if (HasShareds) {
9959 auto *ArgStructAlloca =
9960 dyn_cast<AllocaInst>(Val: StaleCI->getArgOperand(i: SharedArgOperandNo));
9961 assert(ArgStructAlloca &&
9962 "Unable to find the alloca instruction corresponding to arguments "
9963 "for extracted function");
9964 std::optional<TypeSize> ArgAllocSize =
9965 ArgStructAlloca->getAllocationSize(DL: M.getDataLayout());
9966 assert(ArgAllocSize &&
9967 "Unable to determine size of arguments for extracted function");
9968 SharedsSize = Builder.getInt64(C: ArgAllocSize->getFixedValue());
9969 }
9970
9971 // Argument - `flags`
9972 // Task is tied iff (Flags & 1) == 1.
9973 // Task is untied iff (Flags & 1) == 0.
9974 // Task is final iff (Flags & 2) == 2.
9975 // Task is not final iff (Flags & 2) == 0.
9976 // A target task is not final and is untied.
9977 Value *Flags = Builder.getInt32(C: 0);
9978
9979 // Emit the @__kmpc_omp_task_alloc runtime call
9980 // The runtime call returns a pointer to an area where the task captured
9981 // variables must be copied before the task is run (TaskData)
9982 CallInst *TaskData = nullptr;
9983
9984 SmallVector<llvm::Value *> TaskAllocArgs = {
9985 /*loc_ref=*/Ident, /*gtid=*/ThreadID,
9986 /*flags=*/Flags,
9987 /*sizeof_task=*/TaskSize, /*sizeof_shared=*/SharedsSize,
9988 /*task_func=*/ProxyFn};
9989
9990 if (NeedsTargetTask) {
9991 assert(DeviceID && "Expected non-empty device ID.");
9992 TaskAllocArgs.push_back(Elt: DeviceID);
9993 }
9994
9995 TaskData = createRuntimeFunctionCall(Callee: TaskAllocFn, Args: TaskAllocArgs);
9996
9997 Align Alignment = TaskData->getPointerAlignment(DL: M.getDataLayout());
9998 if (HasShareds) {
9999 Value *Shareds = StaleCI->getArgOperand(i: SharedArgOperandNo);
10000 Value *TaskShareds = loadSharedDataFromTaskDescriptor(
10001 OMPIRBuilder&: *this, Builder, TaskWithPrivates: TaskData, TaskWithPrivatesTy);
10002 Builder.CreateMemCpy(Dst: TaskShareds, DstAlign: Alignment, Src: Shareds, SrcAlign: Alignment,
10003 Size: SharedsSize);
10004 }
10005 if (!OffloadingArraysToPrivatize.empty()) {
10006 Value *Privates =
10007 Builder.CreateStructGEP(Ty: TaskWithPrivatesTy, Ptr: TaskData, Idx: 1);
10008 for (unsigned int i = 0; i < OffloadingArraysToPrivatize.size(); ++i) {
10009 Value *PtrToPrivatize = OffloadingArraysToPrivatize[i];
10010 [[maybe_unused]] Type *ArrayType =
10011 getOffloadingArrayType(V: PtrToPrivatize);
10012 assert(ArrayType && "ArrayType cannot be nullptr");
10013
10014 Type *ElementType = PrivatesTy->getElementType(N: i);
10015 assert(ElementType == ArrayType &&
10016 "ElementType should match ArrayType");
10017 (void)ArrayType;
10018
10019 Value *Dst = Builder.CreateStructGEP(Ty: PrivatesTy, Ptr: Privates, Idx: i);
10020 Builder.CreateMemCpy(
10021 Dst, DstAlign: Alignment, Src: PtrToPrivatize, SrcAlign: Alignment,
10022 Size: Builder.getInt64(C: M.getDataLayout().getTypeStoreSize(Ty: ElementType)));
10023 }
10024 }
10025
10026 Value *DepArray = nullptr;
10027 Value *NumDeps = nullptr;
10028 if (Dependencies.DepArray) {
10029 DepArray = Dependencies.DepArray;
10030 NumDeps = Dependencies.NumDeps;
10031 } else if (!Dependencies.Deps.empty()) {
10032 DepArray = emitTaskDependencies(OMPBuilder&: *this, Dependencies: Dependencies.Deps);
10033 NumDeps = Builder.getInt32(C: Dependencies.Deps.size());
10034 }
10035
10036 // ---------------------------------------------------------------
10037 // V5.2 13.8 target construct
10038 // If the nowait clause is present, execution of the target task
10039 // may be deferred. If the nowait clause is not present, the target task is
10040 // an included task.
10041 // ---------------------------------------------------------------
10042 // The above means that the lack of a nowait on the target construct
10043 // translates to '#pragma omp task if(0)'
10044 if (!NeedsTargetTask) {
10045 if (DepArray) {
10046 Function *TaskWaitFn =
10047 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_wait_deps);
10048 createRuntimeFunctionCall(
10049 Callee: TaskWaitFn,
10050 Args: {/*loc_ref=*/Ident, /*gtid=*/ThreadID,
10051 /*ndeps=*/NumDeps,
10052 /*dep_list=*/DepArray,
10053 /*ndeps_noalias=*/ConstantInt::get(Ty: Builder.getInt32Ty(), V: 0),
10054 /*noalias_dep_list=*/
10055 ConstantPointerNull::get(T: PointerType::getUnqual(C&: M.getContext()))});
10056 }
10057 // Included task.
10058 Function *TaskBeginFn =
10059 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_begin_if0);
10060 Function *TaskCompleteFn =
10061 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_complete_if0);
10062 createRuntimeFunctionCall(Callee: TaskBeginFn, Args: {Ident, ThreadID, TaskData});
10063 CallInst *CI = createRuntimeFunctionCall(Callee: ProxyFn, Args: {ThreadID, TaskData});
10064 CI->setDebugLoc(StaleCI->getDebugLoc());
10065 createRuntimeFunctionCall(Callee: TaskCompleteFn, Args: {Ident, ThreadID, TaskData});
10066 } else if (DepArray) {
10067 // HasNoWait - meaning the task may be deferred. Call
10068 // __kmpc_omp_task_with_deps if there are dependencies,
10069 // else call __kmpc_omp_task
10070 Function *TaskFn =
10071 getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task_with_deps);
10072 createRuntimeFunctionCall(
10073 Callee: TaskFn,
10074 Args: {Ident, ThreadID, TaskData, NumDeps, DepArray,
10075 ConstantInt::get(Ty: Builder.getInt32Ty(), V: 0),
10076 ConstantPointerNull::get(T: PointerType::getUnqual(C&: M.getContext()))});
10077 } else {
10078 // Emit the @__kmpc_omp_task runtime call to spawn the task
10079 Function *TaskFn = getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_omp_task);
10080 createRuntimeFunctionCall(Callee: TaskFn, Args: {Ident, ThreadID, TaskData});
10081 }
10082
10083 Builder.ClearInsertionPoint();
10084 StaleCI->eraseFromParent();
10085 for (Instruction *I : llvm::reverse(C&: ToBeDeleted))
10086 I->eraseFromParent();
10087 };
10088 addOutlineInfo(OI: std::move(OI));
10089
10090 LLVM_DEBUG(dbgs() << "Insert block after emitKernelLaunch = \n"
10091 << *(Builder.GetInsertBlock()) << "\n");
10092 LLVM_DEBUG(dbgs() << "Module after emitKernelLaunch = \n"
10093 << *(Builder.getModule()) << "\n");
10094 return Builder.saveIP();
10095}
10096
10097Error OpenMPIRBuilder::emitOffloadingArraysAndArgs(
10098 InsertPointTy AllocaIP, InsertPointTy CodeGenIP, TargetDataInfo &Info,
10099 TargetDataRTArgs &RTArgs, MapInfosTy &CombinedInfo,
10100 CustomMapperCallbackTy CustomMapperCB, bool IsNonContiguous,
10101 bool ForEndCall, function_ref<void(unsigned int, Value *)> DeviceAddrCB) {
10102 if (Error Err =
10103 emitOffloadingArrays(AllocaIP, CodeGenIP, CombinedInfo, Info,
10104 CustomMapperCB, IsNonContiguous, DeviceAddrCB))
10105 return Err;
10106 emitOffloadingArraysArgument(Builder, RTArgs, Info, ForEndCall);
10107 return Error::success();
10108}
10109
10110static void emitTargetCall(
10111 OpenMPIRBuilder &OMPBuilder, IRBuilderBase &Builder, Value *RTLocOverride,
10112 OpenMPIRBuilder::InsertPointTy AllocaIP,
10113 ArrayRef<BasicBlock *> DeallocBlocks, OpenMPIRBuilder::TargetDataInfo &Info,
10114 const OpenMPIRBuilder::TargetKernelDefaultAttrs &DefaultAttrs,
10115 const OpenMPIRBuilder::TargetKernelRuntimeAttrs &RuntimeAttrs,
10116 Value *IfCond, Function *OutlinedFn, Constant *OutlinedFnID,
10117 SmallVectorImpl<Value *> &Args,
10118 OpenMPIRBuilder::GenMapInfoCallbackTy GenMapInfoCB,
10119 OpenMPIRBuilder::CustomMapperCallbackTy CustomMapperCB,
10120 const OpenMPIRBuilder::DependenciesInfo &Dependencies, bool HasNoWait,
10121 Value *DynCGroupMem, OMPDynGroupprivateFallbackType DynCGroupMemFallback) {
10122 // Generate a function call to the host fallback implementation of the target
10123 // region. This is called by the host when no offload entry was generated for
10124 // the target region and when the offloading call fails at runtime.
10125 auto &&EmitTargetCallFallbackCB = [&](OpenMPIRBuilder::InsertPointTy IP)
10126 -> OpenMPIRBuilder::InsertPointOrErrorTy {
10127 Builder.restoreIP(IP);
10128 // Ensure the host fallback has the same dyn_ptr ABI as the device.
10129 SmallVector<Value *> FallbackArgs(Args.begin(), Args.end());
10130 FallbackArgs.push_back(
10131 Elt: Constant::getNullValue(Ty: PointerType::getUnqual(C&: Builder.getContext())));
10132 OMPBuilder.createRuntimeFunctionCall(Callee: OutlinedFn, Args: FallbackArgs);
10133 return Builder.saveIP();
10134 };
10135
10136 bool HasDependencies = !Dependencies.empty();
10137 bool RequiresOuterTargetTask = HasNoWait || HasDependencies;
10138
10139 OpenMPIRBuilder::TargetKernelArgs KArgs;
10140
10141 auto TaskBodyCB =
10142 [&](Value *DeviceID, Value *RTLoc,
10143 IRBuilderBase::InsertPoint TargetTaskAllocaIP) -> Error {
10144 // Assume no error was returned because EmitTargetCallFallbackCB doesn't
10145 // produce any.
10146 llvm::OpenMPIRBuilder::InsertPointTy AfterIP = cantFail(ValOrErr: [&]() {
10147 // emitKernelLaunch makes the necessary runtime call to offload the
10148 // kernel. We then outline all that code into a separate function
10149 // ('kernel_launch_function' in the pseudo code above). This function is
10150 // then called by the target task proxy function (see
10151 // '@.omp_target_task_proxy_func' in the pseudo code above)
10152 // "@.omp_target_task_proxy_func' is generated by
10153 // emitTargetTaskProxyFunction.
10154 if (OutlinedFnID && DeviceID)
10155 return OMPBuilder.emitKernelLaunch(Loc: Builder, OutlinedFnID,
10156 EmitTargetCallFallbackCB, Args&: KArgs,
10157 DeviceID, RTLoc, AllocaIP: TargetTaskAllocaIP);
10158
10159 // We only need to do the outlining if `DeviceID` is set to avoid calling
10160 // `emitKernelLaunch` if we want to code-gen for the host; e.g. if we are
10161 // generating the `else` branch of an `if` clause.
10162 //
10163 // When OutlinedFnID is set to nullptr, then it's not an offloading call.
10164 // In this case, we execute the host implementation directly.
10165 return EmitTargetCallFallbackCB(OMPBuilder.Builder.saveIP());
10166 }());
10167
10168 OMPBuilder.Builder.restoreIP(IP: AfterIP);
10169 return Error::success();
10170 };
10171
10172 auto &&EmitTargetCallElse =
10173 [&](OpenMPIRBuilder::InsertPointTy AllocaIP,
10174 OpenMPIRBuilder::InsertPointTy CodeGenIP,
10175 ArrayRef<BasicBlock *> DeallocBlocks) -> Error {
10176 // Assume no error was returned because EmitTargetCallFallbackCB doesn't
10177 // produce any.
10178 OpenMPIRBuilder::InsertPointTy AfterIP = cantFail(ValOrErr: [&]() {
10179 if (RequiresOuterTargetTask) {
10180 // Arguments that are intended to be directly forwarded to an
10181 // emitKernelLaunch call are pased as nullptr, since
10182 // OutlinedFnID=nullptr results in that call not being done.
10183 OpenMPIRBuilder::TargetDataRTArgs EmptyRTArgs;
10184 return OMPBuilder.emitTargetTask(TaskBodyCB, /*DeviceID=*/nullptr,
10185 /*RTLoc=*/nullptr, AllocaIP,
10186 Dependencies, RTArgs: EmptyRTArgs, HasNoWait);
10187 }
10188 return EmitTargetCallFallbackCB(Builder.saveIP());
10189 }());
10190
10191 Builder.restoreIP(IP: AfterIP);
10192 return Error::success();
10193 };
10194
10195 auto &&EmitTargetCallThen =
10196 [&](OpenMPIRBuilder::InsertPointTy AllocaIP,
10197 OpenMPIRBuilder::InsertPointTy CodeGenIP,
10198 ArrayRef<BasicBlock *> DeallocBlocks) -> Error {
10199 Info.HasNoWait = HasNoWait;
10200 OpenMPIRBuilder::MapInfosTy &MapInfo = GenMapInfoCB(Builder.saveIP());
10201
10202 OpenMPIRBuilder::TargetDataRTArgs RTArgs;
10203 if (Error Err = OMPBuilder.emitOffloadingArraysAndArgs(
10204 AllocaIP, CodeGenIP: Builder.saveIP(), Info, RTArgs, CombinedInfo&: MapInfo, CustomMapperCB,
10205 /*IsNonContiguous=*/true,
10206 /*ForEndCall=*/false))
10207 return Err;
10208
10209 SmallVector<Value *, 3> NumTeamsC;
10210 for (auto [DefaultVal, RuntimeVal] :
10211 zip_equal(t: DefaultAttrs.MaxTeams, u: RuntimeAttrs.MaxTeams))
10212 NumTeamsC.push_back(Elt: RuntimeVal ? RuntimeVal
10213 : Builder.getInt32(C: DefaultVal));
10214
10215 // Calculate number of threads: 0 if no clauses specified, otherwise it is
10216 // the minimum between optional THREAD_LIMIT and NUM_THREADS clauses.
10217 auto InitMaxThreadsClause = [&Builder](Value *Clause) {
10218 if (Clause)
10219 Clause = Builder.CreateIntCast(V: Clause, DestTy: Builder.getInt32Ty(),
10220 /*isSigned=*/false);
10221 return Clause;
10222 };
10223 auto CombineMaxThreadsClauses = [&Builder](Value *Clause, Value *&Result) {
10224 if (Clause)
10225 Result =
10226 Result ? Builder.CreateSelect(C: Builder.CreateICmpULT(LHS: Result, RHS: Clause),
10227 True: Result, False: Clause)
10228 : Clause;
10229 };
10230
10231 // If a multi-dimensional THREAD_LIMIT is set, it is the OMPX_BARE case, so
10232 // the NUM_THREADS clause is overriden by THREAD_LIMIT.
10233 SmallVector<Value *, 3> NumThreadsC;
10234 Value *MaxThreadsClause =
10235 RuntimeAttrs.TeamsThreadLimit.size() == 1
10236 ? InitMaxThreadsClause(RuntimeAttrs.MaxThreads.front())
10237 : nullptr;
10238
10239 for (auto [TeamsVal, TargetVal] : zip_equal(
10240 t: RuntimeAttrs.TeamsThreadLimit, u: RuntimeAttrs.TargetThreadLimit)) {
10241 Value *TeamsThreadLimitClause = InitMaxThreadsClause(TeamsVal);
10242 Value *NumThreads = InitMaxThreadsClause(TargetVal);
10243
10244 CombineMaxThreadsClauses(TeamsThreadLimitClause, NumThreads);
10245 CombineMaxThreadsClauses(MaxThreadsClause, NumThreads);
10246
10247 NumThreadsC.push_back(Elt: NumThreads ? NumThreads : Builder.getInt32(C: 0));
10248 }
10249
10250 unsigned NumTargetItems = Info.NumberOfPtrs;
10251 Value *RTLoc = RTLocOverride;
10252 if (!RTLoc) {
10253 uint32_t SrcLocStrSize;
10254 Constant *SrcLocStr =
10255 OMPBuilder.getOrCreateDefaultSrcLocStr(SrcLocStrSize);
10256 RTLoc = OMPBuilder.getOrCreateIdent(SrcLocStr, SrcLocStrSize,
10257 LocFlags: llvm::omp::IdentFlag(0), Reserve2Flags: 0);
10258 }
10259
10260 Value *TripCount = RuntimeAttrs.LoopTripCount
10261 ? Builder.CreateIntCast(V: RuntimeAttrs.LoopTripCount,
10262 DestTy: Builder.getInt64Ty(),
10263 /*isSigned=*/false)
10264 : Builder.getInt64(C: 0);
10265
10266 // Request zero groupprivate bytes by default.
10267 if (!DynCGroupMem)
10268 DynCGroupMem = Builder.getInt32(C: 0);
10269
10270 KArgs = OpenMPIRBuilder::TargetKernelArgs(
10271 NumTargetItems, RTArgs, TripCount, NumTeamsC, NumThreadsC, DynCGroupMem,
10272 HasNoWait, /*StrictBlocks=*/false, /*StrictThreads=*/false,
10273 DynCGroupMemFallback);
10274
10275 // Assume no error was returned because TaskBodyCB and
10276 // EmitTargetCallFallbackCB don't produce any.
10277 OpenMPIRBuilder::InsertPointTy AfterIP = cantFail(ValOrErr: [&]() {
10278 // The presence of certain clauses on the target directive require the
10279 // explicit generation of the target task.
10280 if (RequiresOuterTargetTask)
10281 return OMPBuilder.emitTargetTask(TaskBodyCB, DeviceID: RuntimeAttrs.DeviceID,
10282 RTLoc, AllocaIP, Dependencies,
10283 RTArgs: KArgs.RTArgs, HasNoWait: Info.HasNoWait);
10284
10285 return OMPBuilder.emitKernelLaunch(
10286 Loc: Builder, OutlinedFnID, EmitTargetCallFallbackCB, Args&: KArgs,
10287 DeviceID: RuntimeAttrs.DeviceID, RTLoc, AllocaIP);
10288 }());
10289
10290 Builder.restoreIP(IP: AfterIP);
10291 return Error::success();
10292 };
10293
10294 // If we don't have an ID for the target region, it means an offload entry
10295 // wasn't created. In this case we just run the host fallback directly and
10296 // ignore any potential 'if' clauses.
10297 if (!OutlinedFnID) {
10298 cantFail(Err: EmitTargetCallElse(AllocaIP, Builder.saveIP(), DeallocBlocks));
10299 return;
10300 }
10301
10302 // If there's no 'if' clause, only generate the kernel launch code path.
10303 if (!IfCond) {
10304 cantFail(Err: EmitTargetCallThen(AllocaIP, Builder.saveIP(), DeallocBlocks));
10305 return;
10306 }
10307
10308 cantFail(Err: OMPBuilder.emitIfClause(Cond: IfCond, ThenGen: EmitTargetCallThen,
10309 ElseGen: EmitTargetCallElse, AllocaIP));
10310}
10311
10312OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createTarget(
10313 const LocationDescription &Loc, bool IsOffloadEntry, InsertPointTy AllocaIP,
10314 InsertPointTy CodeGenIP, ArrayRef<BasicBlock *> DeallocBlocks,
10315 TargetDataInfo &Info, TargetRegionEntryInfo &EntryInfo,
10316 const TargetKernelDefaultAttrs &DefaultAttrs,
10317 const TargetKernelRuntimeAttrs &RuntimeAttrs, Value *IfCond,
10318 SmallVectorImpl<Value *> &Inputs, GenMapInfoCallbackTy GenMapInfoCB,
10319 OpenMPIRBuilder::TargetBodyGenCallbackTy CBFunc,
10320 OpenMPIRBuilder::TargetGenArgAccessorsCallbackTy ArgAccessorFuncCB,
10321 CustomMapperCallbackTy CustomMapperCB, const DependenciesInfo &Dependencies,
10322 bool HasNowait, Value *DynCGroupMem,
10323 OMPDynGroupprivateFallbackType DynCGroupMemFallback, DebugLoc OutlinedFnLoc,
10324 Value *RTLocOverride) {
10325
10326 if (!updateToLocation(Loc))
10327 return InsertPointTy();
10328
10329 Builder.restoreIP(IP: CodeGenIP);
10330
10331 Function *OutlinedFn;
10332 Constant *OutlinedFnID = nullptr;
10333 // The target region is outlined into its own function. The LLVM IR for
10334 // the target region itself is generated using the callbacks CBFunc
10335 // and ArgAccessorFuncCB
10336 if (Error Err = emitTargetOutlinedFunction(
10337 OMPBuilder&: *this, Builder, IsOffloadEntry, EntryInfo, DefaultAttrs, OutlinedFn,
10338 OutlinedFnID, Inputs, CBFunc, ArgAccessorFuncCB, OutlinedFnLoc))
10339 return Err;
10340
10341 // If we are not on the target device, then we need to generate code
10342 // to make a remote call (offload) to the previously outlined function
10343 // that represents the target region. Do that now.
10344 if (!Config.isTargetDevice())
10345 emitTargetCall(OMPBuilder&: *this, Builder, RTLocOverride, AllocaIP, DeallocBlocks, Info,
10346 DefaultAttrs, RuntimeAttrs, IfCond, OutlinedFn, OutlinedFnID,
10347 Args&: Inputs, GenMapInfoCB, CustomMapperCB, Dependencies,
10348 HasNoWait: HasNowait, DynCGroupMem, DynCGroupMemFallback);
10349 return Builder.saveIP();
10350}
10351
10352std::string OpenMPIRBuilder::getNameWithSeparators(ArrayRef<StringRef> Parts,
10353 StringRef FirstSeparator,
10354 StringRef Separator) {
10355 SmallString<128> Buffer;
10356 llvm::raw_svector_ostream OS(Buffer);
10357 StringRef Sep = FirstSeparator;
10358 for (StringRef Part : Parts) {
10359 OS << Sep << Part;
10360 Sep = Separator;
10361 }
10362 return OS.str().str();
10363}
10364
10365std::string
10366OpenMPIRBuilder::createPlatformSpecificName(ArrayRef<StringRef> Parts) const {
10367 return OpenMPIRBuilder::getNameWithSeparators(Parts, FirstSeparator: Config.firstSeparator(),
10368 Separator: Config.separator());
10369}
10370
10371GlobalVariable *OpenMPIRBuilder::getOrCreateInternalVariable(
10372 Type *Ty, const StringRef &Name, std::optional<unsigned> AddressSpace) {
10373 auto &Elem = *InternalVars.try_emplace(Key: Name, Args: nullptr).first;
10374 if (Elem.second) {
10375 assert(Elem.second->getValueType() == Ty &&
10376 "OMP internal variable has different type than requested");
10377 } else {
10378 // TODO: investigate the appropriate linkage type used for the global
10379 // variable for possibly changing that to internal or private, or maybe
10380 // create different versions of the function for different OMP internal
10381 // variables.
10382 const DataLayout &DL = M.getDataLayout();
10383 // TODO: Investigate why AMDGPU expects AS 0 for globals even though the
10384 // default global AS is 1.
10385 // See double-target-call-with-declare-target.f90 and
10386 // declare-target-vars-in-target-region.f90 libomptarget
10387 // tests.
10388 unsigned AddressSpaceVal = AddressSpace ? *AddressSpace
10389 : M.getTargetTriple().isAMDGPU()
10390 ? 0
10391 : DL.getDefaultGlobalsAddressSpace();
10392 auto Linkage = this->M.getTargetTriple().isWasm()
10393 ? GlobalValue::InternalLinkage
10394 : GlobalValue::CommonLinkage;
10395 auto *GV = new GlobalVariable(M, Ty, /*IsConstant=*/false, Linkage,
10396 Constant::getNullValue(Ty), Elem.first(),
10397 /*InsertBefore=*/nullptr,
10398 GlobalValue::NotThreadLocal, AddressSpaceVal);
10399 const llvm::Align TypeAlign = DL.getABITypeAlign(Ty);
10400 const llvm::Align PtrAlign = DL.getPointerABIAlignment(AS: AddressSpaceVal);
10401 GV->setAlignment(std::max(a: TypeAlign, b: PtrAlign));
10402 Elem.second = GV;
10403 }
10404
10405 return Elem.second;
10406}
10407
10408Value *OpenMPIRBuilder::getOMPCriticalRegionLock(StringRef CriticalName) {
10409 std::string Prefix = Twine("gomp_critical_user_", CriticalName).str();
10410 std::string Name = getNameWithSeparators(Parts: {Prefix, "var"}, FirstSeparator: ".", Separator: ".");
10411 return getOrCreateInternalVariable(Ty: KmpCriticalNameTy, Name);
10412}
10413
10414Value *OpenMPIRBuilder::getSizeInBytes(Value *BasePtr) {
10415 LLVMContext &Ctx = Builder.getContext();
10416 Value *Null =
10417 Constant::getNullValue(Ty: PointerType::getUnqual(C&: BasePtr->getContext()));
10418 Value *SizeGep =
10419 Builder.CreateGEP(Ty: BasePtr->getType(), Ptr: Null, IdxList: Builder.getInt32(C: 1));
10420 Value *SizePtrToInt = Builder.CreatePtrToInt(V: SizeGep, DestTy: Type::getInt64Ty(C&: Ctx));
10421 return SizePtrToInt;
10422}
10423
10424GlobalVariable *
10425OpenMPIRBuilder::createOffloadMaptypes(SmallVectorImpl<uint64_t> &Mappings,
10426 std::string VarName) {
10427 llvm::Constant *MaptypesArrayInit =
10428 llvm::ConstantDataArray::get(Context&: M.getContext(), Elts&: Mappings);
10429 auto *MaptypesArrayGlobal = new llvm::GlobalVariable(
10430 M, MaptypesArrayInit->getType(),
10431 /*isConstant=*/true, llvm::GlobalValue::PrivateLinkage, MaptypesArrayInit,
10432 VarName);
10433 MaptypesArrayGlobal->setUnnamedAddr(llvm::GlobalValue::UnnamedAddr::Global);
10434 return MaptypesArrayGlobal;
10435}
10436
10437void OpenMPIRBuilder::createMapperAllocas(const LocationDescription &Loc,
10438 InsertPointTy AllocaIP,
10439 unsigned NumOperands,
10440 struct MapperAllocas &MapperAllocas) {
10441 if (!updateToLocation(Loc))
10442 return;
10443
10444 auto *ArrI8PtrTy = ArrayType::get(ElementType: Int8Ptr, NumElements: NumOperands);
10445 auto *ArrI64Ty = ArrayType::get(ElementType: Int64, NumElements: NumOperands);
10446 Builder.restoreIP(IP: AllocaIP);
10447 AllocaInst *ArgsBase = Builder.CreateAlloca(
10448 Ty: ArrI8PtrTy, /* ArraySize = */ nullptr, Name: ".offload_baseptrs");
10449 AllocaInst *Args = Builder.CreateAlloca(Ty: ArrI8PtrTy, /* ArraySize = */ nullptr,
10450 Name: ".offload_ptrs");
10451 AllocaInst *ArgSizes = Builder.CreateAlloca(
10452 Ty: ArrI64Ty, /* ArraySize = */ nullptr, Name: ".offload_sizes");
10453 updateToLocation(Loc);
10454 MapperAllocas.ArgsBase = ArgsBase;
10455 MapperAllocas.Args = Args;
10456 MapperAllocas.ArgSizes = ArgSizes;
10457}
10458
10459void OpenMPIRBuilder::emitMapperCall(const LocationDescription &Loc,
10460 Function *MapperFunc, Value *SrcLocInfo,
10461 Value *MaptypesArg, Value *MapnamesArg,
10462 struct MapperAllocas &MapperAllocas,
10463 int64_t DeviceID, unsigned NumOperands) {
10464 if (!updateToLocation(Loc))
10465 return;
10466
10467 auto *ArrI8PtrTy = ArrayType::get(ElementType: Int8Ptr, NumElements: NumOperands);
10468 auto *ArrI64Ty = ArrayType::get(ElementType: Int64, NumElements: NumOperands);
10469 Value *ArgsBaseGEP =
10470 Builder.CreateInBoundsGEP(Ty: ArrI8PtrTy, Ptr: MapperAllocas.ArgsBase,
10471 IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: 0)});
10472 Value *ArgsGEP =
10473 Builder.CreateInBoundsGEP(Ty: ArrI8PtrTy, Ptr: MapperAllocas.Args,
10474 IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: 0)});
10475 Value *ArgSizesGEP =
10476 Builder.CreateInBoundsGEP(Ty: ArrI64Ty, Ptr: MapperAllocas.ArgSizes,
10477 IdxList: {Builder.getInt32(C: 0), Builder.getInt32(C: 0)});
10478 Value *NullPtr =
10479 Constant::getNullValue(Ty: PointerType::getUnqual(C&: Int8Ptr->getContext()));
10480 createRuntimeFunctionCall(Callee: MapperFunc, Args: {SrcLocInfo, Builder.getInt64(C: DeviceID),
10481 Builder.getInt32(C: NumOperands),
10482 ArgsBaseGEP, ArgsGEP, ArgSizesGEP,
10483 MaptypesArg, MapnamesArg, NullPtr});
10484}
10485
10486void OpenMPIRBuilder::emitOffloadingArraysArgument(IRBuilderBase &Builder,
10487 TargetDataRTArgs &RTArgs,
10488 TargetDataInfo &Info,
10489 bool ForEndCall) {
10490 assert((!ForEndCall || Info.separateBeginEndCalls()) &&
10491 "expected region end call to runtime only when end call is separate");
10492 auto UnqualPtrTy = PointerType::getUnqual(C&: M.getContext());
10493 auto VoidPtrTy = UnqualPtrTy;
10494 auto VoidPtrPtrTy = UnqualPtrTy;
10495 auto Int64Ty = Type::getInt64Ty(C&: M.getContext());
10496 auto Int64PtrTy = UnqualPtrTy;
10497
10498 if (!Info.NumberOfPtrs) {
10499 RTArgs.BasePointersArray = ConstantPointerNull::get(T: VoidPtrPtrTy);
10500 RTArgs.PointersArray = ConstantPointerNull::get(T: VoidPtrPtrTy);
10501 RTArgs.SizesArray = ConstantPointerNull::get(T: Int64PtrTy);
10502 RTArgs.MapTypesArray = ConstantPointerNull::get(T: Int64PtrTy);
10503 RTArgs.MapNamesArray = ConstantPointerNull::get(T: VoidPtrPtrTy);
10504 RTArgs.MappersArray = ConstantPointerNull::get(T: VoidPtrPtrTy);
10505 return;
10506 }
10507
10508 RTArgs.BasePointersArray = Builder.CreateConstInBoundsGEP2_32(
10509 Ty: ArrayType::get(ElementType: VoidPtrTy, NumElements: Info.NumberOfPtrs),
10510 Ptr: Info.RTArgs.BasePointersArray,
10511 /*Idx0=*/0, /*Idx1=*/0);
10512 RTArgs.PointersArray = Builder.CreateConstInBoundsGEP2_32(
10513 Ty: ArrayType::get(ElementType: VoidPtrTy, NumElements: Info.NumberOfPtrs), Ptr: Info.RTArgs.PointersArray,
10514 /*Idx0=*/0,
10515 /*Idx1=*/0);
10516 RTArgs.SizesArray = Builder.CreateConstInBoundsGEP2_32(
10517 Ty: ArrayType::get(ElementType: Int64Ty, NumElements: Info.NumberOfPtrs), Ptr: Info.RTArgs.SizesArray,
10518 /*Idx0=*/0, /*Idx1=*/0);
10519 RTArgs.MapTypesArray = Builder.CreateConstInBoundsGEP2_32(
10520 Ty: ArrayType::get(ElementType: Int64Ty, NumElements: Info.NumberOfPtrs),
10521 Ptr: ForEndCall && Info.RTArgs.MapTypesArrayEnd ? Info.RTArgs.MapTypesArrayEnd
10522 : Info.RTArgs.MapTypesArray,
10523 /*Idx0=*/0,
10524 /*Idx1=*/0);
10525
10526 // Only emit the mapper information arrays if debug information is
10527 // requested.
10528 if (!Info.EmitDebug)
10529 RTArgs.MapNamesArray = ConstantPointerNull::get(T: VoidPtrPtrTy);
10530 else
10531 RTArgs.MapNamesArray = Builder.CreateConstInBoundsGEP2_32(
10532 Ty: ArrayType::get(ElementType: VoidPtrTy, NumElements: Info.NumberOfPtrs), Ptr: Info.RTArgs.MapNamesArray,
10533 /*Idx0=*/0,
10534 /*Idx1=*/0);
10535 // If there is no user-defined mapper, set the mapper array to nullptr to
10536 // avoid an unnecessary data privatization
10537 if (!Info.HasMapper)
10538 RTArgs.MappersArray = ConstantPointerNull::get(T: VoidPtrPtrTy);
10539 else
10540 RTArgs.MappersArray =
10541 Builder.CreatePointerCast(V: Info.RTArgs.MappersArray, DestTy: VoidPtrPtrTy);
10542}
10543
10544void OpenMPIRBuilder::emitNonContiguousDescriptor(InsertPointTy AllocaIP,
10545 InsertPointTy CodeGenIP,
10546 MapInfosTy &CombinedInfo,
10547 TargetDataInfo &Info) {
10548 MapInfosTy::StructNonContiguousInfo &NonContigInfo =
10549 CombinedInfo.NonContigInfo;
10550
10551 // Build an array of struct descriptor_dim and then assign it to
10552 // offload_args.
10553 //
10554 // struct descriptor_dim {
10555 // uint64_t offset;
10556 // uint64_t count;
10557 // uint64_t stride
10558 // };
10559 Type *Int64Ty = Builder.getInt64Ty();
10560 StructType *DimTy = StructType::create(
10561 Context&: M.getContext(), Elements: ArrayRef<Type *>({Int64Ty, Int64Ty, Int64Ty}),
10562 Name: "struct.descriptor_dim");
10563
10564 enum { OffsetFD = 0, CountFD, StrideFD };
10565 // We need two index variable here since the size of "Dims" is the same as
10566 // the size of Components, however, the size of offset, count, and stride is
10567 // equal to the size of base declaration that is non-contiguous.
10568 for (unsigned I = 0, L = 0, E = NonContigInfo.Dims.size(); I < E; ++I) {
10569 // Skip emitting ir if dimension size is 1 since it cannot be
10570 // non-contiguous.
10571 if (NonContigInfo.Dims[I] == 1)
10572 continue;
10573 Builder.restoreIP(IP: AllocaIP);
10574 ArrayType *ArrayTy = ArrayType::get(ElementType: DimTy, NumElements: NonContigInfo.Dims[I]);
10575 AllocaInst *DimsAddr =
10576 Builder.CreateAlloca(Ty: ArrayTy, /* ArraySize = */ nullptr, Name: "dims");
10577 Builder.restoreIP(IP: CodeGenIP);
10578 for (unsigned II = 0, EE = NonContigInfo.Dims[I]; II < EE; ++II) {
10579 unsigned RevIdx = EE - II - 1;
10580 Value *DimsLVal = Builder.CreateInBoundsGEP(
10581 Ty: ArrayTy, Ptr: DimsAddr, IdxList: {Builder.getInt64(C: 0), Builder.getInt64(C: II)});
10582 // Offset
10583 Value *OffsetLVal = Builder.CreateStructGEP(Ty: DimTy, Ptr: DimsLVal, Idx: OffsetFD);
10584 Builder.CreateAlignedStore(
10585 Val: NonContigInfo.Offsets[L][RevIdx], Ptr: OffsetLVal,
10586 Align: M.getDataLayout().getPrefTypeAlign(Ty: OffsetLVal->getType()));
10587 // Count
10588 Value *CountLVal = Builder.CreateStructGEP(Ty: DimTy, Ptr: DimsLVal, Idx: CountFD);
10589 Builder.CreateAlignedStore(
10590 Val: NonContigInfo.Counts[L][RevIdx], Ptr: CountLVal,
10591 Align: M.getDataLayout().getPrefTypeAlign(Ty: CountLVal->getType()));
10592 // Stride
10593 Value *StrideLVal = Builder.CreateStructGEP(Ty: DimTy, Ptr: DimsLVal, Idx: StrideFD);
10594 Builder.CreateAlignedStore(
10595 Val: NonContigInfo.Strides[L][RevIdx], Ptr: StrideLVal,
10596 Align: M.getDataLayout().getPrefTypeAlign(Ty: CountLVal->getType()));
10597 }
10598 // args[I] = &dims
10599 Builder.restoreIP(IP: CodeGenIP);
10600 Value *DAddr = Builder.CreatePointerBitCastOrAddrSpaceCast(
10601 V: DimsAddr, DestTy: Builder.getPtrTy());
10602 Value *P = Builder.CreateConstInBoundsGEP2_32(
10603 Ty: ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: Info.NumberOfPtrs),
10604 Ptr: Info.RTArgs.PointersArray, Idx0: 0, Idx1: I);
10605 Builder.CreateAlignedStore(
10606 Val: DAddr, Ptr: P, Align: M.getDataLayout().getPrefTypeAlign(Ty: Builder.getPtrTy()));
10607 ++L;
10608 }
10609}
10610
10611void OpenMPIRBuilder::emitUDMapperArrayInitOrDel(
10612 Function *MapperFn, Value *MapperHandle, Value *Base, Value *Begin,
10613 Value *Size, Value *MapType, Value *MapName, TypeSize ElementSize,
10614 BasicBlock *ExitBB, bool IsInit) {
10615 StringRef Prefix = IsInit ? ".init" : ".del";
10616
10617 // Evaluate if this is an array section.
10618 BasicBlock *BodyBB = BasicBlock::Create(
10619 Context&: M.getContext(), Name: createPlatformSpecificName(Parts: {"omp.array", Prefix}));
10620 Value *IsArray =
10621 Builder.CreateICmpSGT(LHS: Size, RHS: Builder.getInt64(C: 1), Name: "omp.arrayinit.isarray");
10622 Value *DeleteBit = Builder.CreateAnd(
10623 LHS: MapType,
10624 RHS: Builder.getInt64(
10625 C: static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10626 OpenMPOffloadMappingFlags::OMP_MAP_DELETE)));
10627 Value *DeleteCond;
10628 Value *Cond;
10629 if (IsInit) {
10630 // base != begin?
10631 Value *BaseIsBegin = Builder.CreateICmpNE(LHS: Base, RHS: Begin);
10632 Cond = Builder.CreateOr(LHS: IsArray, RHS: BaseIsBegin);
10633 DeleteCond = Builder.CreateIsNull(
10634 Arg: DeleteBit,
10635 Name: createPlatformSpecificName(Parts: {"omp.array", Prefix, ".delete"}));
10636 } else {
10637 Cond = IsArray;
10638 DeleteCond = Builder.CreateIsNotNull(
10639 Arg: DeleteBit,
10640 Name: createPlatformSpecificName(Parts: {"omp.array", Prefix, ".delete"}));
10641 }
10642 Cond = Builder.CreateAnd(LHS: Cond, RHS: DeleteCond);
10643 Builder.CreateCondBr(Cond, True: BodyBB, False: ExitBB);
10644
10645 emitBlock(BB: BodyBB, CurFn: MapperFn);
10646 // Get the array size by multiplying element size and element number (i.e., \p
10647 // Size).
10648 Value *ArraySize = Builder.CreateNUWMul(LHS: Size, RHS: Builder.getInt64(C: ElementSize));
10649 // Remove OMP_MAP_TO and OMP_MAP_FROM from the map type, so that it achieves
10650 // memory allocation/deletion purpose only.
10651 Value *MapTypeArg = Builder.CreateAnd(
10652 LHS: MapType,
10653 RHS: Builder.getInt64(
10654 C: ~static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10655 OpenMPOffloadMappingFlags::OMP_MAP_TO |
10656 OpenMPOffloadMappingFlags::OMP_MAP_FROM)));
10657 MapTypeArg = Builder.CreateOr(
10658 LHS: MapTypeArg,
10659 RHS: Builder.getInt64(
10660 C: static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10661 OpenMPOffloadMappingFlags::OMP_MAP_IMPLICIT)));
10662
10663 // Call the runtime API __tgt_push_mapper_component to fill up the runtime
10664 // data structure.
10665 Value *OffloadingArgs[] = {MapperHandle, Base, Begin,
10666 ArraySize, MapTypeArg, MapName};
10667 createRuntimeFunctionCall(
10668 Callee: getOrCreateRuntimeFunction(M, FnID: OMPRTL___tgt_push_mapper_component),
10669 Args: OffloadingArgs);
10670}
10671
10672Expected<Function *> OpenMPIRBuilder::emitUserDefinedMapper(
10673 function_ref<MapInfosOrErrorTy(InsertPointTy CodeGenIP, llvm::Value *PtrPHI,
10674 llvm::Value *BeginArg)>
10675 GenMapInfoCB,
10676 Type *ElemTy, StringRef FuncName, CustomMapperCallbackTy CustomMapperCB,
10677 bool PreserveMemberOfFlags, bool PropagatePresentToPointee) {
10678 SmallVector<Type *> Params;
10679 Params.emplace_back(Args: Builder.getPtrTy());
10680 Params.emplace_back(Args: Builder.getPtrTy());
10681 Params.emplace_back(Args: Builder.getPtrTy());
10682 Params.emplace_back(Args: Builder.getInt64Ty());
10683 Params.emplace_back(Args: Builder.getInt64Ty());
10684 Params.emplace_back(Args: Builder.getPtrTy());
10685
10686 auto *FnTy =
10687 FunctionType::get(Result: Builder.getVoidTy(), Params, /* IsVarArg */ isVarArg: false);
10688
10689 SmallString<64> TyStr;
10690 raw_svector_ostream Out(TyStr);
10691 Function *MapperFn =
10692 Function::Create(Ty: FnTy, Linkage: GlobalValue::InternalLinkage, N: FuncName, M);
10693 MapperFn->addFnAttr(Kind: Attribute::NoInline);
10694 MapperFn->addFnAttr(Kind: Attribute::NoUnwind);
10695 MapperFn->addParamAttr(ArgNo: 0, Kind: Attribute::NoUndef);
10696 MapperFn->addParamAttr(ArgNo: 1, Kind: Attribute::NoUndef);
10697 MapperFn->addParamAttr(ArgNo: 2, Kind: Attribute::NoUndef);
10698 MapperFn->addParamAttr(ArgNo: 3, Kind: Attribute::NoUndef);
10699 MapperFn->addParamAttr(ArgNo: 4, Kind: Attribute::NoUndef);
10700 MapperFn->addParamAttr(ArgNo: 5, Kind: Attribute::NoUndef);
10701
10702 // Start the mapper function code generation.
10703 BasicBlock *EntryBB = BasicBlock::Create(Context&: M.getContext(), Name: "entry", Parent: MapperFn);
10704 IRBuilder<>::InsertPointGuard IPG(Builder);
10705 Builder.SetInsertPoint(EntryBB);
10706 Builder.SetCurrentDebugLocation(llvm::DebugLoc());
10707
10708 Value *MapperHandle = MapperFn->getArg(i: 0);
10709 Value *BaseIn = MapperFn->getArg(i: 1);
10710 Value *BeginIn = MapperFn->getArg(i: 2);
10711 Value *Size = MapperFn->getArg(i: 3);
10712 Value *MapType = MapperFn->getArg(i: 4);
10713 Value *MapName = MapperFn->getArg(i: 5);
10714
10715 // Compute the starting and end addresses of array elements.
10716 // Prepare common arguments for array initiation and deletion.
10717 // Convert the size in bytes into the number of array elements.
10718 TypeSize ElementSize = M.getDataLayout().getTypeStoreSize(Ty: ElemTy);
10719 Size = Builder.CreateExactUDiv(LHS: Size, RHS: Builder.getInt64(C: ElementSize));
10720 Value *PtrBegin = BeginIn;
10721 Value *PtrEnd = Builder.CreateGEP(Ty: ElemTy, Ptr: PtrBegin, IdxList: Size);
10722
10723 // Emit array initiation if this is an array section and \p MapType indicates
10724 // that memory allocation is required.
10725 BasicBlock *HeadBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.arraymap.head");
10726 emitUDMapperArrayInitOrDel(MapperFn, MapperHandle, Base: BaseIn, Begin: BeginIn, Size,
10727 MapType, MapName, ElementSize, ExitBB: HeadBB,
10728 /*IsInit=*/true);
10729
10730 // Emit a for loop to iterate through SizeArg of elements and map all of them.
10731
10732 // Emit the loop header block.
10733 emitBlock(BB: HeadBB, CurFn: MapperFn);
10734 BasicBlock *BodyBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.arraymap.body");
10735 BasicBlock *DoneBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.done");
10736 // Evaluate whether the initial condition is satisfied.
10737 Value *IsEmpty =
10738 Builder.CreateICmpEQ(LHS: PtrBegin, RHS: PtrEnd, Name: "omp.arraymap.isempty");
10739 Builder.CreateCondBr(Cond: IsEmpty, True: DoneBB, False: BodyBB);
10740
10741 // Emit the loop body block.
10742 emitBlock(BB: BodyBB, CurFn: MapperFn);
10743 BasicBlock *LastBB = BodyBB;
10744 PHINode *PtrPHI =
10745 Builder.CreatePHI(Ty: PtrBegin->getType(), NumReservedValues: 2, Name: "omp.arraymap.ptrcurrent");
10746 PtrPHI->addIncoming(V: PtrBegin, BB: HeadBB);
10747
10748 // Get map clause information. Fill up the arrays with all mapped variables.
10749 MapInfosOrErrorTy Info = GenMapInfoCB(Builder.saveIP(), PtrPHI, BeginIn);
10750 if (!Info)
10751 return Info.takeError();
10752
10753 // Call the runtime API __tgt_mapper_num_components to get the number of
10754 // pre-existing components.
10755 Value *OffloadingArgs[] = {MapperHandle};
10756 Value *PreviousSize = createRuntimeFunctionCall(
10757 Callee: getOrCreateRuntimeFunction(M, FnID: OMPRTL___tgt_mapper_num_components),
10758 Args: OffloadingArgs);
10759 Value *ShiftedPreviousSize =
10760 Builder.CreateShl(LHS: PreviousSize, RHS: Builder.getInt64(C: getFlagMemberOffset()));
10761
10762 // Fill up the runtime mapper handle for all components.
10763 for (unsigned I = 0; I < Info->BasePointers.size(); ++I) {
10764 Value *CurBaseArg = Info->BasePointers[I];
10765 Value *CurBeginArg = Info->Pointers[I];
10766 Value *CurSizeArg = Info->Sizes[I];
10767 Value *CurNameArg = Info->Names.size()
10768 ? Info->Names[I]
10769 : Constant::getNullValue(Ty: Builder.getPtrTy());
10770
10771 Value *OriMapType = Builder.getInt64(
10772 C: static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10773 Info->Types[I]));
10774 auto RawType =
10775 static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10776 Info->Types[I]);
10777 constexpr uint64_t MemberOfMask =
10778 static_cast<uint64_t>(OpenMPOffloadMappingFlags::OMP_MAP_MEMBER_OF);
10779 constexpr uint64_t AttachBit =
10780 static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10781 OpenMPOffloadMappingFlags::OMP_MAP_ATTACH);
10782
10783 // Add MEMBER_OF (ShiftedPreviousSize) to group this sub-map with the
10784 // current array element (N = __tgt_mapper_num_components() at loop body
10785 // start).
10786 //
10787 // Example 1:
10788 // struct S { int x; int *p; };
10789 //
10790 // mapper: #pragma omp declare mapper(id: S s) map(s.x, s.p[0:10])
10791 // use: S arr[2]; ... map(arr)
10792 // entries per element:
10793 //
10794 // &arr[i], &arr[i].x, sizeof(int), MEMBER_OF(N)|TO|FROM
10795 // &arr[i].p[0], &arr[i].p[0], 10*sizeof(int), TO|FROM (*)
10796 // &arr[i].p, &arr[i].p[0], sizeof(int*), ATTACH (**)
10797 //
10798 // Example 2:
10799 // struct S1 { int x; int y; };
10800 // struct S2 { int z; S1 *s1p; };
10801 //
10802 // mapper: #pragma omp declare mapper(S2 s2) map(s2.z, s2.s1p->x,
10803 // s2.s1p->y)
10804 // use: S2 arr[2]; ... map(arr)
10805 // entries per element:
10806 //
10807 // &arr[i], &arr[i].z, sizeof(int), MEMBER_OF(N)|TO|FROM
10808 // &arr[i].s1p[0], &arr[i].s1p->x, sizeof(s1p->x..y), ALLOC (*)
10809 // &arr[i].s1p[0], &arr[i].s1p->x, 4, MEMBER_OF(N+2)|TO|FROM (*)(***)
10810 // &arr[i].s1p[0], &arr[i].s1p->y, 4, MEMBER_OF(N+2)|TO|FROM (*)(***)
10811 // &arr[i].s1p, &arr[i].s1p->x, sizeof(ptr), ATTACH (**)
10812 //
10813 // x/y carry inner MEMBER_OF(2)
10814 // which is shifted by N to become MEMBER_OF(N+2).
10815 //
10816 // HasAttachPtr is set on all of the s1p entries except the ATTACH one:
10817 // the combined ALLOC entry for the s1p->x..y block, and the individual
10818 // x/y entries that are MEMBER_OF that block, all describe storage
10819 // reached through the attach ptr arr[i].s1p.
10820 //
10821 // Entries of the following kinds do NOT receive a new outer MEMBER_OF
10822 // linking them to the parent struct:
10823 //
10824 // * (*) Entries with HasAttachPtr: they represent pointee data that
10825 // occupies a different storage block than the struct being mapped, so
10826 // they are not a member of it. They may still be MEMBER_OF an entry
10827 // within that pointee block, in which case those pre-existing bits are
10828 // shifted -- see (***).
10829 // * (**) ATTACH entries: they are not a member of anything — they just
10830 // link a ptr to its ptee.
10831 // * All entries when PreserveMemberOfFlags is set (the Flang/MLIR path):
10832 // its pre-shaped entries already carry their final MEMBER_OF bits.
10833 // TODO: set HasAttachPtr from Flang for entries whose storage is the
10834 // pointee's (e.g. s%p(0:10)) and drop PreserveMemberOfFlags in favor of
10835 // it.
10836 //
10837 // (***) If such an entry already has its own MEMBER_OF bits (e.g. the
10838 // s1p->x/y entries above), those bits are still shifted by N.
10839 Value *MemberMapType;
10840 if (PreserveMemberOfFlags || (RawType & AttachBit) ||
10841 Info->HasAttachPtr[I]) {
10842 if (RawType & MemberOfMask)
10843 MemberMapType = Builder.CreateNUWAdd(LHS: OriMapType, RHS: ShiftedPreviousSize);
10844 else
10845 MemberMapType = OriMapType;
10846 } else {
10847 MemberMapType = Builder.CreateNUWAdd(LHS: OriMapType, RHS: ShiftedPreviousSize);
10848 }
10849
10850 // Combine the map type inherited from user-defined mapper with that
10851 // specified in the program. According to the OMP_MAP_TO and OMP_MAP_FROM
10852 // bits of the \a MapType, which is the input argument of the mapper
10853 // function, the following code will set the OMP_MAP_TO and OMP_MAP_FROM
10854 // bits of MemberMapType.
10855 // [OpenMP 5.0], 1.2.6. map-type decay.
10856 // | alloc | to | from | tofrom | release | delete
10857 // ----------------------------------------------------------
10858 // alloc | alloc | alloc | alloc | alloc | release | delete
10859 // to | alloc | to | alloc | to | release | delete
10860 // from | alloc | alloc | from | from | release | delete
10861 // tofrom | alloc | to | from | tofrom | release | delete
10862 Value *LeftToFrom = Builder.CreateAnd(
10863 LHS: MapType,
10864 RHS: Builder.getInt64(
10865 C: static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10866 OpenMPOffloadMappingFlags::OMP_MAP_TO |
10867 OpenMPOffloadMappingFlags::OMP_MAP_FROM)));
10868 BasicBlock *AllocBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.type.alloc");
10869 BasicBlock *AllocElseBB =
10870 BasicBlock::Create(Context&: M.getContext(), Name: "omp.type.alloc.else");
10871 BasicBlock *ToBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.type.to");
10872 BasicBlock *ToElseBB =
10873 BasicBlock::Create(Context&: M.getContext(), Name: "omp.type.to.else");
10874 BasicBlock *FromBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.type.from");
10875 BasicBlock *EndBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.type.end");
10876 Value *IsAlloc = Builder.CreateIsNull(Arg: LeftToFrom);
10877 Builder.CreateCondBr(Cond: IsAlloc, True: AllocBB, False: AllocElseBB);
10878 // In case of alloc, clear OMP_MAP_TO and OMP_MAP_FROM.
10879 emitBlock(BB: AllocBB, CurFn: MapperFn);
10880 Value *AllocMapType = Builder.CreateAnd(
10881 LHS: MemberMapType,
10882 RHS: Builder.getInt64(
10883 C: ~static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10884 OpenMPOffloadMappingFlags::OMP_MAP_TO |
10885 OpenMPOffloadMappingFlags::OMP_MAP_FROM)));
10886 Builder.CreateBr(Dest: EndBB);
10887 emitBlock(BB: AllocElseBB, CurFn: MapperFn);
10888 Value *IsTo = Builder.CreateICmpEQ(
10889 LHS: LeftToFrom,
10890 RHS: Builder.getInt64(
10891 C: static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10892 OpenMPOffloadMappingFlags::OMP_MAP_TO)));
10893 Builder.CreateCondBr(Cond: IsTo, True: ToBB, False: ToElseBB);
10894 // In case of to, clear OMP_MAP_FROM.
10895 emitBlock(BB: ToBB, CurFn: MapperFn);
10896 Value *ToMapType = Builder.CreateAnd(
10897 LHS: MemberMapType,
10898 RHS: Builder.getInt64(
10899 C: ~static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10900 OpenMPOffloadMappingFlags::OMP_MAP_FROM)));
10901 Builder.CreateBr(Dest: EndBB);
10902 emitBlock(BB: ToElseBB, CurFn: MapperFn);
10903 Value *IsFrom = Builder.CreateICmpEQ(
10904 LHS: LeftToFrom,
10905 RHS: Builder.getInt64(
10906 C: static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10907 OpenMPOffloadMappingFlags::OMP_MAP_FROM)));
10908 Builder.CreateCondBr(Cond: IsFrom, True: FromBB, False: EndBB);
10909 // In case of from, clear OMP_MAP_TO.
10910 emitBlock(BB: FromBB, CurFn: MapperFn);
10911 Value *FromMapType = Builder.CreateAnd(
10912 LHS: MemberMapType,
10913 RHS: Builder.getInt64(
10914 C: ~static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10915 OpenMPOffloadMappingFlags::OMP_MAP_TO)));
10916 // In case of tofrom, do nothing.
10917 emitBlock(BB: EndBB, CurFn: MapperFn);
10918 LastBB = EndBB;
10919 PHINode *CurMapType =
10920 Builder.CreatePHI(Ty: Builder.getInt64Ty(), NumReservedValues: 4, Name: "omp.maptype");
10921 CurMapType->addIncoming(V: AllocMapType, BB: AllocBB);
10922 CurMapType->addIncoming(V: ToMapType, BB: ToBB);
10923 CurMapType->addIncoming(V: FromMapType, BB: FromBB);
10924 CurMapType->addIncoming(V: MemberMapType, BB: ToElseBB);
10925
10926 // Propagate map-type-modifying bits from the outer map clause to each map
10927 // inserted by the mapper.
10928 //
10929 // OpenMP 6.0:281:34: The effect of the mapper modifier is to remove the
10930 // list item from the map clause and to apply the clauses specified in the
10931 // declared mapper to the construct on which the map clause appears...
10932 // If any modifier with the map-type-modifying property appears in the map
10933 // clause then the effect is as if that modifier appears in each map clause
10934 // specified in the declared mapper.
10935 //
10936 // Map-type-modifying bits: ALWAYS, DELETE, CLOSE, PRESENT.
10937 //
10938 // ALWAYS/DELETE/CLOSE are propagated to every (non-ATTACH) entry.
10939 //
10940 // PRESENT is propagated only to entries that have an attach ptr
10941 // (HasAttachPtr): the pointee data, which occupies a different storage
10942 // block than the struct being mapped and so is not covered by the
10943 // present-check on the struct's own storage. A present modifier on the
10944 // outer clause must still require that pointee to be present on the device.
10945 //
10946 // This is gated on \p PropagatePresentToPointee (set by callers only for
10947 // OpenMP >= 6.0). Before 6.0 the present modifier is treated as not
10948 // applying to the pointee: the spec committee confirmed the divergence
10949 // between the present "motion" modifier (to/from) and the present map-type
10950 // modifier (map) was unintentional, to be fixed as an OpenMP 6.0 erratum,
10951 // so for 5.2 present is ignored for the pointee for both map and to/from.
10952 //
10953 // TODO: PRESENT should also be propagated to the struct's own members
10954 // (e.g. the s.x, s.y of map(present, mapper(id): s)) so that an absent
10955 // member triggers the present-check. We cannot do that yet: while pointer
10956 // members are mapped with PTR_AND_OBJ, a single combined entry allocates
10957 // the whole struct (including the pointer's storage), so propagating
10958 // PRESENT to it would wrongly require the pointer's pointee to be present.
10959 // Enable member propagation once Clang stops emitting PTR_AND_OBJ and uses
10960 // attach-style maps throughout.
10961 uint64_t ModifierBits =
10962 static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10963 OpenMPOffloadMappingFlags::OMP_MAP_ALWAYS |
10964 OpenMPOffloadMappingFlags::OMP_MAP_DELETE |
10965 OpenMPOffloadMappingFlags::OMP_MAP_CLOSE);
10966 if (PropagatePresentToPointee && Info->HasAttachPtr[I])
10967 ModifierBits |=
10968 static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
10969 OpenMPOffloadMappingFlags::OMP_MAP_PRESENT);
10970 Value *ImportedModifierBits =
10971 Builder.CreateAnd(LHS: MapType, RHS: Builder.getInt64(C: ModifierBits));
10972 Value *CurMapTypeWithModifiers = Builder.CreateOr(
10973 LHS: CurMapType, RHS: ImportedModifierBits, Name: "omp.maptype.with.modifiers");
10974
10975 // ATTACH entries must not receive map-type-modifying bits: ATTACH|ALWAYS is
10976 // reserved for the attach(always) map-type modifier, and other modifier
10977 // bits (DELETE, CLOSE) have no meaning for an ATTACH entry.
10978 Value *FinalMapType =
10979 (RawType & AttachBit) ? CurMapType : CurMapTypeWithModifiers;
10980
10981 Value *OffloadingArgs[] = {MapperHandle, CurBaseArg, CurBeginArg,
10982 CurSizeArg, FinalMapType, CurNameArg};
10983
10984 auto ChildMapperFn = CustomMapperCB(I);
10985 if (!ChildMapperFn)
10986 return ChildMapperFn.takeError();
10987 if (*ChildMapperFn) {
10988 // Call the corresponding mapper function.
10989 createRuntimeFunctionCall(Callee: *ChildMapperFn, Args: OffloadingArgs)
10990 ->setDoesNotThrow();
10991 } else {
10992 // Call the runtime API __tgt_push_mapper_component to fill up the runtime
10993 // data structure.
10994 createRuntimeFunctionCall(
10995 Callee: getOrCreateRuntimeFunction(M, FnID: OMPRTL___tgt_push_mapper_component),
10996 Args: OffloadingArgs);
10997 }
10998 }
10999
11000 // Update the pointer to point to the next element that needs to be mapped,
11001 // and check whether we have mapped all elements.
11002 Value *PtrNext = Builder.CreateConstGEP1_32(Ty: ElemTy, Ptr: PtrPHI, /*Idx0=*/1,
11003 Name: "omp.arraymap.next");
11004 PtrPHI->addIncoming(V: PtrNext, BB: LastBB);
11005 Value *IsDone = Builder.CreateICmpEQ(LHS: PtrNext, RHS: PtrEnd, Name: "omp.arraymap.isdone");
11006 BasicBlock *ExitBB = BasicBlock::Create(Context&: M.getContext(), Name: "omp.arraymap.exit");
11007 Builder.CreateCondBr(Cond: IsDone, True: ExitBB, False: BodyBB);
11008
11009 emitBlock(BB: ExitBB, CurFn: MapperFn);
11010 // Emit array deletion if this is an array section and \p MapType indicates
11011 // that deletion is required.
11012 emitUDMapperArrayInitOrDel(MapperFn, MapperHandle, Base: BaseIn, Begin: BeginIn, Size,
11013 MapType, MapName, ElementSize, ExitBB: DoneBB,
11014 /*IsInit=*/false);
11015
11016 // Emit the function exit block.
11017 emitBlock(BB: DoneBB, CurFn: MapperFn, /*IsFinished=*/true);
11018
11019 Builder.CreateRetVoid();
11020 return MapperFn;
11021}
11022
11023Error OpenMPIRBuilder::emitOffloadingArrays(
11024 InsertPointTy AllocaIP, InsertPointTy CodeGenIP, MapInfosTy &CombinedInfo,
11025 TargetDataInfo &Info, CustomMapperCallbackTy CustomMapperCB,
11026 bool IsNonContiguous,
11027 function_ref<void(unsigned int, Value *)> DeviceAddrCB) {
11028
11029 // Reset the array information.
11030 Info.clearArrayInfo();
11031 Info.NumberOfPtrs = CombinedInfo.BasePointers.size();
11032
11033 if (Info.NumberOfPtrs == 0)
11034 return Error::success();
11035
11036 Builder.restoreIP(IP: AllocaIP);
11037 // Detect if we have any capture size requiring runtime evaluation of the
11038 // size so that a constant array could be eventually used.
11039 ArrayType *PointerArrayType =
11040 ArrayType::get(ElementType: Builder.getPtrTy(), NumElements: Info.NumberOfPtrs);
11041
11042 Info.RTArgs.BasePointersArray = Builder.CreateAlloca(
11043 Ty: PointerArrayType, /* ArraySize = */ nullptr, Name: ".offload_baseptrs");
11044
11045 Info.RTArgs.PointersArray = Builder.CreateAlloca(
11046 Ty: PointerArrayType, /* ArraySize = */ nullptr, Name: ".offload_ptrs");
11047 AllocaInst *MappersArray = Builder.CreateAlloca(
11048 Ty: PointerArrayType, /* ArraySize = */ nullptr, Name: ".offload_mappers");
11049 Info.RTArgs.MappersArray = MappersArray;
11050
11051 // If we don't have any VLA types or other types that require runtime
11052 // evaluation, we can use a constant array for the map sizes, otherwise we
11053 // need to fill up the arrays as we do for the pointers.
11054 Type *Int64Ty = Builder.getInt64Ty();
11055 SmallVector<Constant *> ConstSizes(CombinedInfo.Sizes.size(),
11056 ConstantInt::get(Ty: Int64Ty, V: 0));
11057 SmallBitVector RuntimeSizes(CombinedInfo.Sizes.size());
11058 for (unsigned I = 0, E = CombinedInfo.Sizes.size(); I < E; ++I) {
11059 bool IsNonContigEntry =
11060 IsNonContiguous &&
11061 (static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
11062 CombinedInfo.Types[I] &
11063 OpenMPOffloadMappingFlags::OMP_MAP_NON_CONTIG) != 0);
11064 // For NON_CONTIG entries, ArgSizes stores the dimension count (number of
11065 // descriptor_dim records), not the byte size.
11066 if (IsNonContigEntry) {
11067 assert(I < CombinedInfo.NonContigInfo.Dims.size() &&
11068 "Index must be in-bounds for NON_CONTIG Dims array");
11069 const uint64_t DimCount = CombinedInfo.NonContigInfo.Dims[I];
11070 assert(DimCount > 0 && "NON_CONTIG DimCount must be > 0");
11071 ConstSizes[I] = ConstantInt::get(Ty: Int64Ty, V: DimCount);
11072 continue;
11073 }
11074 if (auto *CI = dyn_cast<Constant>(Val: CombinedInfo.Sizes[I])) {
11075 if (!isa<ConstantExpr>(Val: CI) && !isa<GlobalValue>(Val: CI)) {
11076 ConstSizes[I] = CI;
11077 continue;
11078 }
11079 }
11080 RuntimeSizes.set(I);
11081 }
11082
11083 if (RuntimeSizes.all()) {
11084 ArrayType *SizeArrayType = ArrayType::get(ElementType: Int64Ty, NumElements: Info.NumberOfPtrs);
11085 Info.RTArgs.SizesArray = Builder.CreateAlloca(
11086 Ty: SizeArrayType, /* ArraySize = */ nullptr, Name: ".offload_sizes");
11087 restoreIPandDebugLoc(Builder, IP: CodeGenIP);
11088 } else {
11089 auto *SizesArrayInit = ConstantArray::get(
11090 T: ArrayType::get(ElementType: Int64Ty, NumElements: ConstSizes.size()), V: ConstSizes);
11091 std::string Name = createPlatformSpecificName(Parts: {"offload_sizes"});
11092 auto *SizesArrayGbl =
11093 new GlobalVariable(M, SizesArrayInit->getType(), /*isConstant=*/true,
11094 GlobalValue::PrivateLinkage, SizesArrayInit, Name);
11095 SizesArrayGbl->setUnnamedAddr(GlobalValue::UnnamedAddr::Global);
11096
11097 if (!RuntimeSizes.any()) {
11098 Info.RTArgs.SizesArray = SizesArrayGbl;
11099 } else {
11100 unsigned IndexSize = M.getDataLayout().getIndexSizeInBits(AS: 0);
11101 Align OffloadSizeAlign = M.getDataLayout().getABIIntegerTypeAlignment(BitWidth: 64);
11102 ArrayType *SizeArrayType = ArrayType::get(ElementType: Int64Ty, NumElements: Info.NumberOfPtrs);
11103 AllocaInst *Buffer = Builder.CreateAlloca(
11104 Ty: SizeArrayType, /* ArraySize = */ nullptr, Name: ".offload_sizes");
11105 Buffer->setAlignment(OffloadSizeAlign);
11106 restoreIPandDebugLoc(Builder, IP: CodeGenIP);
11107 Builder.CreateMemCpy(
11108 Dst: Buffer, DstAlign: M.getDataLayout().getPrefTypeAlign(Ty: Buffer->getType()),
11109 Src: SizesArrayGbl, SrcAlign: OffloadSizeAlign,
11110 Size: Builder.getIntN(
11111 N: IndexSize,
11112 C: Buffer->getAllocationSize(DL: M.getDataLayout())->getFixedValue()));
11113
11114 Info.RTArgs.SizesArray = Buffer;
11115 }
11116 restoreIPandDebugLoc(Builder, IP: CodeGenIP);
11117 }
11118
11119 // The map types are always constant so we don't need to generate code to
11120 // fill arrays. Instead, we create an array constant.
11121 SmallVector<uint64_t, 4> Mapping;
11122 for (auto mapFlag : CombinedInfo.Types)
11123 Mapping.push_back(
11124 Elt: static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
11125 mapFlag));
11126 std::string MaptypesName = createPlatformSpecificName(Parts: {"offload_maptypes"});
11127 auto *MapTypesArrayGbl = createOffloadMaptypes(Mappings&: Mapping, VarName: MaptypesName);
11128 Info.RTArgs.MapTypesArray = MapTypesArrayGbl;
11129
11130 // The information types are only built if provided.
11131 if (!CombinedInfo.Names.empty()) {
11132 auto *MapNamesArrayGbl = createOffloadMapnames(
11133 Names&: CombinedInfo.Names, VarName: createPlatformSpecificName(Parts: {"offload_mapnames"}));
11134 Info.RTArgs.MapNamesArray = MapNamesArrayGbl;
11135 Info.EmitDebug = true;
11136 } else {
11137 Info.RTArgs.MapNamesArray =
11138 Constant::getNullValue(Ty: PointerType::getUnqual(C&: Builder.getContext()));
11139 Info.EmitDebug = false;
11140 }
11141
11142 // If there's a present map type modifier, it must not be applied to the end
11143 // of a region, so generate a separate map type array in that case.
11144 if (Info.separateBeginEndCalls()) {
11145 bool EndMapTypesDiffer = false;
11146 for (uint64_t &Type : Mapping) {
11147 if (Type & static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
11148 OpenMPOffloadMappingFlags::OMP_MAP_PRESENT)) {
11149 Type &= ~static_cast<std::underlying_type_t<OpenMPOffloadMappingFlags>>(
11150 OpenMPOffloadMappingFlags::OMP_MAP_PRESENT);
11151 EndMapTypesDiffer = true;
11152 }
11153 }
11154 if (EndMapTypesDiffer) {
11155 MapTypesArrayGbl = createOffloadMaptypes(Mappings&: Mapping, VarName: MaptypesName);
11156 Info.RTArgs.MapTypesArrayEnd = MapTypesArrayGbl;
11157 }
11158 }
11159
11160 PointerType *PtrTy = Builder.getPtrTy();
11161 for (unsigned I = 0; I < Info.NumberOfPtrs; ++I) {
11162 Value *BPVal = CombinedInfo.BasePointers[I];
11163 Value *BP = Builder.CreateConstInBoundsGEP2_32(
11164 Ty: ArrayType::get(ElementType: PtrTy, NumElements: Info.NumberOfPtrs), Ptr: Info.RTArgs.BasePointersArray,
11165 Idx0: 0, Idx1: I);
11166 Builder.CreateAlignedStore(Val: BPVal, Ptr: BP,
11167 Align: M.getDataLayout().getPrefTypeAlign(Ty: PtrTy));
11168
11169 if (Info.requiresDevicePointerInfo()) {
11170 if (CombinedInfo.DevicePointers[I] == DeviceInfoTy::Pointer) {
11171 CodeGenIP = Builder.saveIP();
11172 Builder.restoreIP(IP: AllocaIP);
11173 Info.DevicePtrInfoMap[BPVal] = {BP, Builder.CreateAlloca(Ty: PtrTy)};
11174 restoreIPandDebugLoc(Builder, IP: CodeGenIP);
11175 if (DeviceAddrCB)
11176 DeviceAddrCB(I, Info.DevicePtrInfoMap[BPVal].second);
11177 } else if (CombinedInfo.DevicePointers[I] == DeviceInfoTy::Address) {
11178 Info.DevicePtrInfoMap[BPVal] = {BP, BP};
11179 if (DeviceAddrCB)
11180 DeviceAddrCB(I, BP);
11181 }
11182 }
11183
11184 Value *PVal = CombinedInfo.Pointers[I];
11185 Value *P = Builder.CreateConstInBoundsGEP2_32(
11186 Ty: ArrayType::get(ElementType: PtrTy, NumElements: Info.NumberOfPtrs), Ptr: Info.RTArgs.PointersArray, Idx0: 0,
11187 Idx1: I);
11188 // TODO: Check alignment correct.
11189 Builder.CreateAlignedStore(Val: PVal, Ptr: P,
11190 Align: M.getDataLayout().getPrefTypeAlign(Ty: PtrTy));
11191
11192 if (RuntimeSizes.test(Idx: I)) {
11193 Value *S = Builder.CreateConstInBoundsGEP2_32(
11194 Ty: ArrayType::get(ElementType: Int64Ty, NumElements: Info.NumberOfPtrs), Ptr: Info.RTArgs.SizesArray,
11195 /*Idx0=*/0,
11196 /*Idx1=*/I);
11197 Builder.CreateAlignedStore(Val: Builder.CreateIntCast(V: CombinedInfo.Sizes[I],
11198 DestTy: Int64Ty,
11199 /*isSigned=*/true),
11200 Ptr: S, Align: M.getDataLayout().getPrefTypeAlign(Ty: PtrTy));
11201 }
11202 // Fill up the mapper array.
11203 unsigned IndexSize = M.getDataLayout().getIndexSizeInBits(AS: 0);
11204 Value *MFunc = ConstantPointerNull::get(T: PtrTy);
11205
11206 auto CustomMFunc = CustomMapperCB(I);
11207 if (!CustomMFunc)
11208 return CustomMFunc.takeError();
11209 if (*CustomMFunc)
11210 MFunc = Builder.CreatePointerCast(V: *CustomMFunc, DestTy: PtrTy);
11211
11212 Value *MAddr = Builder.CreateInBoundsGEP(
11213 Ty: PointerArrayType, Ptr: MappersArray,
11214 IdxList: {Builder.getIntN(N: IndexSize, C: 0), Builder.getIntN(N: IndexSize, C: I)});
11215 Builder.CreateAlignedStore(
11216 Val: MFunc, Ptr: MAddr, Align: M.getDataLayout().getPrefTypeAlign(Ty: MAddr->getType()));
11217 }
11218
11219 if (!IsNonContiguous || CombinedInfo.NonContigInfo.Offsets.empty() ||
11220 Info.NumberOfPtrs == 0)
11221 return Error::success();
11222 emitNonContiguousDescriptor(AllocaIP, CodeGenIP, CombinedInfo, Info);
11223 return Error::success();
11224}
11225
11226void OpenMPIRBuilder::emitBranch(BasicBlock *Target) {
11227 BasicBlock *CurBB = Builder.GetInsertBlock();
11228
11229 if (!CurBB || CurBB->hasTerminator()) {
11230 // If there is no insert point or the previous block is already
11231 // terminated, don't touch it.
11232 } else {
11233 // Otherwise, create a fall-through branch.
11234 Builder.CreateBr(Dest: Target);
11235 }
11236
11237 Builder.ClearInsertionPoint();
11238}
11239
11240void OpenMPIRBuilder::emitBlock(BasicBlock *BB, Function *CurFn,
11241 bool IsFinished) {
11242 BasicBlock *CurBB = Builder.GetInsertBlock();
11243
11244 // Fall out of the current block (if necessary).
11245 emitBranch(Target: BB);
11246
11247 if (IsFinished && BB->use_empty()) {
11248 BB->eraseFromParent();
11249 return;
11250 }
11251
11252 // Place the block after the current block, if possible, or else at
11253 // the end of the function.
11254 if (CurBB && CurBB->getParent())
11255 CurFn->insert(Position: std::next(x: CurBB->getIterator()), BB);
11256 else
11257 CurFn->insert(Position: CurFn->end(), BB);
11258 Builder.SetInsertPoint(BB);
11259}
11260
11261Error OpenMPIRBuilder::emitIfClause(Value *Cond, BodyGenCallbackTy ThenGen,
11262 BodyGenCallbackTy ElseGen,
11263 InsertPointTy AllocaIP,
11264 ArrayRef<BasicBlock *> DeallocBlocks) {
11265 // If the condition constant folds and can be elided, try to avoid emitting
11266 // the condition and the dead arm of the if/else.
11267 if (auto *CI = dyn_cast<ConstantInt>(Val: Cond)) {
11268 auto CondConstant = CI->getSExtValue();
11269 if (CondConstant)
11270 return ThenGen(AllocaIP, Builder.saveIP(), DeallocBlocks);
11271
11272 return ElseGen(AllocaIP, Builder.saveIP(), DeallocBlocks);
11273 }
11274
11275 Function *CurFn = Builder.GetInsertBlock()->getParent();
11276
11277 // Otherwise, the condition did not fold, or we couldn't elide it. Just
11278 // emit the conditional branch.
11279 BasicBlock *ThenBlock = BasicBlock::Create(Context&: M.getContext(), Name: "omp_if.then");
11280 BasicBlock *ElseBlock = BasicBlock::Create(Context&: M.getContext(), Name: "omp_if.else");
11281 BasicBlock *ContBlock = BasicBlock::Create(Context&: M.getContext(), Name: "omp_if.end");
11282 Builder.CreateCondBr(Cond, True: ThenBlock, False: ElseBlock);
11283 // Emit the 'then' code.
11284 emitBlock(BB: ThenBlock, CurFn);
11285 if (Error Err = ThenGen(AllocaIP, Builder.saveIP(), DeallocBlocks))
11286 return Err;
11287 emitBranch(Target: ContBlock);
11288 // Emit the 'else' code if present.
11289 // There is no need to emit line number for unconditional branch.
11290 emitBlock(BB: ElseBlock, CurFn);
11291 if (Error Err = ElseGen(AllocaIP, Builder.saveIP(), DeallocBlocks))
11292 return Err;
11293 // There is no need to emit line number for unconditional branch.
11294 emitBranch(Target: ContBlock);
11295 // Emit the continuation block for code after the if.
11296 emitBlock(BB: ContBlock, CurFn, /*IsFinished=*/true);
11297 return Error::success();
11298}
11299
11300bool OpenMPIRBuilder::checkAndEmitFlushAfterAtomic(
11301 const LocationDescription &Loc, llvm::AtomicOrdering AO, AtomicKind AK) {
11302 assert(!(AO == AtomicOrdering::NotAtomic ||
11303 AO == llvm::AtomicOrdering::Unordered) &&
11304 "Unexpected Atomic Ordering.");
11305
11306 bool Flush = false;
11307 llvm::AtomicOrdering FlushAO = AtomicOrdering::Monotonic;
11308
11309 switch (AK) {
11310 case Read:
11311 if (AO == AtomicOrdering::Acquire || AO == AtomicOrdering::AcquireRelease ||
11312 AO == AtomicOrdering::SequentiallyConsistent) {
11313 FlushAO = AtomicOrdering::Acquire;
11314 Flush = true;
11315 }
11316 break;
11317 case Write:
11318 case Compare:
11319 case Update:
11320 if (AO == AtomicOrdering::Release || AO == AtomicOrdering::AcquireRelease ||
11321 AO == AtomicOrdering::SequentiallyConsistent) {
11322 FlushAO = AtomicOrdering::Release;
11323 Flush = true;
11324 }
11325 break;
11326 case Capture:
11327 switch (AO) {
11328 case AtomicOrdering::Acquire:
11329 FlushAO = AtomicOrdering::Acquire;
11330 Flush = true;
11331 break;
11332 case AtomicOrdering::Release:
11333 FlushAO = AtomicOrdering::Release;
11334 Flush = true;
11335 break;
11336 case AtomicOrdering::AcquireRelease:
11337 case AtomicOrdering::SequentiallyConsistent:
11338 FlushAO = AtomicOrdering::AcquireRelease;
11339 Flush = true;
11340 break;
11341 default:
11342 // do nothing - leave silently.
11343 break;
11344 }
11345 }
11346
11347 if (Flush) {
11348 // Currently Flush RT call still doesn't take memory_ordering, so for when
11349 // that happens, this tries to do the resolution of which atomic ordering
11350 // to use with but issue the flush call
11351 // TODO: pass `FlushAO` after memory ordering support is added
11352 (void)FlushAO;
11353 emitFlush(Loc);
11354 }
11355
11356 // for AO == AtomicOrdering::Monotonic and all other case combinations
11357 // do nothing
11358 return Flush;
11359}
11360
11361OpenMPIRBuilder::InsertPointTy
11362OpenMPIRBuilder::createAtomicRead(const LocationDescription &Loc,
11363 AtomicOpValue &X, AtomicOpValue &V,
11364 AtomicOrdering AO, InsertPointTy AllocaIP) {
11365 if (!updateToLocation(Loc))
11366 return Loc.IP;
11367
11368 assert(X.Var->getType()->isPointerTy() &&
11369 "OMP Atomic expects a pointer to target memory");
11370 Type *XElemTy = X.ElemTy;
11371 assert((XElemTy->isFloatingPointTy() || XElemTy->isIntegerTy() ||
11372 XElemTy->isPointerTy() || XElemTy->isStructTy()) &&
11373 "OMP atomic read expected a scalar type");
11374
11375 Value *XRead = nullptr;
11376
11377 if (XElemTy->isIntegerTy()) {
11378 LoadInst *XLD =
11379 Builder.CreateLoad(Ty: XElemTy, Ptr: X.Var, isVolatile: X.IsVolatile, Name: "omp.atomic.read");
11380 XLD->setAtomic(Ordering: AO);
11381 XRead = cast<Value>(Val: XLD);
11382 } else if (XElemTy->isStructTy()) {
11383 // FIXME: Add checks to ensure __atomic_load is emitted iff the
11384 // target does not support `atomicrmw` of the size of the struct
11385 LoadInst *OldVal = Builder.CreateLoad(Ty: XElemTy, Ptr: X.Var, Name: "omp.atomic.read");
11386 OldVal->setAtomic(Ordering: AO);
11387 const DataLayout &DL = OldVal->getDataLayout();
11388 unsigned LoadSize = DL.getTypeStoreSize(Ty: XElemTy);
11389 OpenMPIRBuilder::AtomicInfo atomicInfo(
11390 &Builder, XElemTy, LoadSize * 8, LoadSize * 8, OldVal->getAlign(),
11391 OldVal->getAlign(), true /* UseLibcall */, AllocaIP, X.Var);
11392 auto AtomicLoadRes = atomicInfo.EmitAtomicLoadLibcall(AO);
11393 XRead = AtomicLoadRes.first;
11394 OldVal->eraseFromParent();
11395 } else {
11396 // We need to perform atomic op as integer
11397 IntegerType *IntCastTy =
11398 IntegerType::get(C&: M.getContext(), NumBits: XElemTy->getScalarSizeInBits());
11399 LoadInst *XLoad =
11400 Builder.CreateLoad(Ty: IntCastTy, Ptr: X.Var, isVolatile: X.IsVolatile, Name: "omp.atomic.load");
11401 XLoad->setAtomic(Ordering: AO);
11402 if (XElemTy->isFloatingPointTy()) {
11403 XRead = Builder.CreateBitCast(V: XLoad, DestTy: XElemTy, Name: "atomic.flt.cast");
11404 } else {
11405 XRead = Builder.CreateIntToPtr(V: XLoad, DestTy: XElemTy, Name: "atomic.ptr.cast");
11406 }
11407 }
11408 checkAndEmitFlushAfterAtomic(Loc, AO, AK: AtomicKind::Read);
11409 Builder.CreateStore(Val: XRead, Ptr: V.Var, isVolatile: V.IsVolatile);
11410 return Builder.saveIP();
11411}
11412
11413OpenMPIRBuilder::InsertPointTy
11414OpenMPIRBuilder::createAtomicWrite(const LocationDescription &Loc,
11415 AtomicOpValue &X, Value *Expr,
11416 AtomicOrdering AO, InsertPointTy AllocaIP) {
11417 if (!updateToLocation(Loc))
11418 return Loc.IP;
11419
11420 assert(X.Var->getType()->isPointerTy() &&
11421 "OMP Atomic expects a pointer to target memory");
11422 Type *XElemTy = X.ElemTy;
11423 assert((XElemTy->isFloatingPointTy() || XElemTy->isIntegerTy() ||
11424 XElemTy->isPointerTy() || XElemTy->isStructTy()) &&
11425 "OMP atomic write expected a scalar type");
11426
11427 if (XElemTy->isIntegerTy()) {
11428 StoreInst *XSt = Builder.CreateStore(Val: Expr, Ptr: X.Var, isVolatile: X.IsVolatile);
11429 XSt->setAtomic(Ordering: AO);
11430 } else if (XElemTy->isStructTy()) {
11431 LoadInst *OldVal = Builder.CreateLoad(Ty: XElemTy, Ptr: X.Var, Name: "omp.atomic.read");
11432 const DataLayout &DL = OldVal->getDataLayout();
11433 unsigned LoadSize = DL.getTypeStoreSize(Ty: XElemTy);
11434 OpenMPIRBuilder::AtomicInfo atomicInfo(
11435 &Builder, XElemTy, LoadSize * 8, LoadSize * 8, OldVal->getAlign(),
11436 OldVal->getAlign(), true /* UseLibcall */, AllocaIP, X.Var);
11437 atomicInfo.EmitAtomicStoreLibcall(AO, Source: Expr);
11438 OldVal->eraseFromParent();
11439 } else {
11440 // We need to bitcast and perform atomic op as integers
11441 IntegerType *IntCastTy =
11442 IntegerType::get(C&: M.getContext(), NumBits: XElemTy->getScalarSizeInBits());
11443 Value *ExprCast =
11444 Builder.CreateBitCast(V: Expr, DestTy: IntCastTy, Name: "atomic.src.int.cast");
11445 StoreInst *XSt = Builder.CreateStore(Val: ExprCast, Ptr: X.Var, isVolatile: X.IsVolatile);
11446 XSt->setAtomic(Ordering: AO);
11447 }
11448
11449 checkAndEmitFlushAfterAtomic(Loc, AO, AK: AtomicKind::Write);
11450 return Builder.saveIP();
11451}
11452
11453OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createAtomicUpdate(
11454 const LocationDescription &Loc, InsertPointTy AllocaIP, AtomicOpValue &X,
11455 Value *Expr, AtomicOrdering AO, AtomicRMWInst::BinOp RMWOp,
11456 AtomicUpdateCallbackTy &UpdateOp, bool IsXBinopExpr,
11457 bool IsIgnoreDenormalMode, bool IsFineGrainedMemory, bool IsRemoteMemory) {
11458 assert(!isConflictIP(Loc.IP, AllocaIP) && "IPs must not be ambiguous");
11459 if (!updateToLocation(Loc))
11460 return Loc.IP;
11461
11462 LLVM_DEBUG({
11463 Type *XTy = X.Var->getType();
11464 assert(XTy->isPointerTy() &&
11465 "OMP Atomic expects a pointer to target memory");
11466 Type *XElemTy = X.ElemTy;
11467 assert((XElemTy->isFloatingPointTy() || XElemTy->isIntegerTy() ||
11468 XElemTy->isPointerTy() || XElemTy->isStructTy()) &&
11469 "OMP atomic update expected a scalar or struct type");
11470 assert((RMWOp != AtomicRMWInst::Max) && (RMWOp != AtomicRMWInst::Min) &&
11471 (RMWOp != AtomicRMWInst::UMax) && (RMWOp != AtomicRMWInst::UMin) &&
11472 "OpenMP atomic does not support LT or GT operations");
11473 });
11474
11475 Expected<std::pair<Value *, Value *>> AtomicResult = emitAtomicUpdate(
11476 AllocaIP, X: X.Var, XElemTy: X.ElemTy, Expr, AO, RMWOp, UpdateOp, VolatileX: X.IsVolatile,
11477 IsXBinopExpr, IsIgnoreDenormalMode, IsFineGrainedMemory, IsRemoteMemory);
11478 if (!AtomicResult)
11479 return AtomicResult.takeError();
11480 checkAndEmitFlushAfterAtomic(Loc, AO, AK: AtomicKind::Update);
11481 return Builder.saveIP();
11482}
11483
11484// FIXME: Duplicating AtomicExpand
11485Value *OpenMPIRBuilder::emitRMWOpAsInstruction(Value *Src1, Value *Src2,
11486 AtomicRMWInst::BinOp RMWOp) {
11487 switch (RMWOp) {
11488 case AtomicRMWInst::Add:
11489 return Builder.CreateAdd(LHS: Src1, RHS: Src2);
11490 case AtomicRMWInst::Sub:
11491 return Builder.CreateSub(LHS: Src1, RHS: Src2);
11492 case AtomicRMWInst::And:
11493 return Builder.CreateAnd(LHS: Src1, RHS: Src2);
11494 case AtomicRMWInst::Nand:
11495 return Builder.CreateNeg(V: Builder.CreateAnd(LHS: Src1, RHS: Src2));
11496 case AtomicRMWInst::Or:
11497 return Builder.CreateOr(LHS: Src1, RHS: Src2);
11498 case AtomicRMWInst::Xor:
11499 return Builder.CreateXor(LHS: Src1, RHS: Src2);
11500 case AtomicRMWInst::Xchg:
11501 case AtomicRMWInst::FAdd:
11502 case AtomicRMWInst::FSub:
11503 case AtomicRMWInst::BAD_BINOP:
11504 case AtomicRMWInst::Max:
11505 case AtomicRMWInst::Min:
11506 case AtomicRMWInst::UMax:
11507 case AtomicRMWInst::UMin:
11508 case AtomicRMWInst::FMax:
11509 case AtomicRMWInst::FMin:
11510 case AtomicRMWInst::FMaximum:
11511 case AtomicRMWInst::FMinimum:
11512 case AtomicRMWInst::FMaximumNum:
11513 case AtomicRMWInst::FMinimumNum:
11514 case AtomicRMWInst::UIncWrap:
11515 case AtomicRMWInst::UDecWrap:
11516 case AtomicRMWInst::USubCond:
11517 case AtomicRMWInst::USubSat:
11518 llvm_unreachable("Unsupported atomic update operation");
11519 }
11520 llvm_unreachable("Unsupported atomic update operation");
11521}
11522
11523static AtomicOrdering TransformReleaseAcquireRelease(AtomicOrdering AO) {
11524 // Loads cannot use Release or AcquireRelease ordering. This load is
11525 // just the initial value for the cmpxchg loop; the cmpxchg itself
11526 // retains the original ordering.
11527 AtomicOrdering LoadAO = AO;
11528
11529 if (AO == AtomicOrdering::Release) {
11530 LoadAO = AtomicOrdering::Monotonic;
11531 } else if (AO == AtomicOrdering::AcquireRelease) {
11532 LoadAO = AtomicOrdering::Acquire;
11533 }
11534
11535 return LoadAO;
11536}
11537
11538Expected<std::pair<Value *, Value *>> OpenMPIRBuilder::emitAtomicUpdate(
11539 InsertPointTy AllocaIP, Value *X, Type *XElemTy, Value *Expr,
11540 AtomicOrdering AO, AtomicRMWInst::BinOp RMWOp,
11541 AtomicUpdateCallbackTy &UpdateOp, bool VolatileX, bool IsXBinopExpr,
11542 bool IsIgnoreDenormalMode, bool IsFineGrainedMemory, bool IsRemoteMemory) {
11543 // TODO: handle the case where XElemTy is not byte-sized or not a power of 2.
11544 bool emitRMWOp = false;
11545 switch (RMWOp) {
11546 case AtomicRMWInst::Add:
11547 case AtomicRMWInst::And:
11548 case AtomicRMWInst::Nand:
11549 case AtomicRMWInst::Or:
11550 case AtomicRMWInst::Xor:
11551 case AtomicRMWInst::Xchg:
11552 emitRMWOp = XElemTy;
11553 break;
11554 case AtomicRMWInst::Sub:
11555 emitRMWOp = (IsXBinopExpr && XElemTy);
11556 break;
11557 default:
11558 emitRMWOp = false;
11559 }
11560 emitRMWOp &= XElemTy->isIntegerTy();
11561
11562 std::pair<Value *, Value *> Res;
11563 if (emitRMWOp) {
11564 AtomicRMWInst *RMWInst =
11565 Builder.CreateAtomicRMW(Op: RMWOp, Ptr: X, Val: Expr, Align: llvm::MaybeAlign(), Ordering: AO);
11566 if (IsIgnoreDenormalMode)
11567 RMWInst->setMetadata(KindID: llvm::LLVMContext::MD_atomic_ignore_denormal_mode,
11568 Node: llvm::MDNode::get(Context&: Builder.getContext(), MDs: {}));
11569 if (T.isAMDGPU()) {
11570 if (!IsFineGrainedMemory)
11571 RMWInst->setMetadata(Kind: "amdgpu.no.fine.grained.memory",
11572 Node: llvm::MDNode::get(Context&: Builder.getContext(), MDs: {}));
11573 if (!IsRemoteMemory)
11574 RMWInst->setMetadata(Kind: "amdgpu.no.remote.memory",
11575 Node: llvm::MDNode::get(Context&: Builder.getContext(), MDs: {}));
11576 }
11577 Res.first = RMWInst;
11578 // not needed except in case of postfix captures. Generate anyway for
11579 // consistency with the else part. Will be removed with any DCE pass.
11580 // AtomicRMWInst::Xchg does not have a coressponding instruction.
11581 if (RMWOp == AtomicRMWInst::Xchg)
11582 Res.second = Res.first;
11583 else
11584 Res.second = emitRMWOpAsInstruction(Src1: Res.first, Src2: Expr, RMWOp);
11585 } else if (XElemTy->isStructTy()) {
11586 LoadInst *OldVal =
11587 Builder.CreateLoad(Ty: XElemTy, Ptr: X, Name: X->getName() + ".atomic.load");
11588 AtomicOrdering LoadAO = TransformReleaseAcquireRelease(AO);
11589 OldVal->setAtomic(Ordering: LoadAO);
11590 const DataLayout &LoadDL = OldVal->getDataLayout();
11591 unsigned LoadSize = LoadDL.getTypeStoreSize(Ty: XElemTy);
11592
11593 OpenMPIRBuilder::AtomicInfo atomicInfo(
11594 &Builder, XElemTy, LoadSize * 8, LoadSize * 8, OldVal->getAlign(),
11595 OldVal->getAlign(), true /* UseLibcall */, AllocaIP, X);
11596 auto AtomicLoadRes = atomicInfo.EmitAtomicLoadLibcall(AO);
11597 BasicBlock *CurBB = Builder.GetInsertBlock();
11598 Instruction *CurBBTI = CurBB->getTerminatorOrNull();
11599 CurBBTI = CurBBTI ? CurBBTI : Builder.CreateUnreachable();
11600 BasicBlock *ExitBB =
11601 CurBB->splitBasicBlock(I: CurBBTI, BBName: X->getName() + ".atomic.exit");
11602 BasicBlock *ContBB = CurBB->splitBasicBlock(I: CurBB->getTerminator(),
11603 BBName: X->getName() + ".atomic.cont");
11604 ContBB->getTerminator()->eraseFromParent();
11605 Builder.restoreIP(IP: AllocaIP);
11606 AllocaInst *NewAtomicAddr = Builder.CreateAlloca(Ty: XElemTy);
11607 NewAtomicAddr->setName(X->getName() + "x.new.val");
11608 Builder.SetInsertPoint(ContBB);
11609 llvm::PHINode *PHI = Builder.CreatePHI(Ty: OldVal->getType(), NumReservedValues: 2);
11610 PHI->addIncoming(V: AtomicLoadRes.first, BB: CurBB);
11611 Value *OldExprVal = PHI;
11612 Expected<Value *> CBResult = UpdateOp(OldExprVal, Builder);
11613 if (!CBResult)
11614 return CBResult.takeError();
11615 Value *Upd = *CBResult;
11616 Builder.CreateStore(Val: Upd, Ptr: NewAtomicAddr);
11617 AtomicOrdering Failure =
11618 llvm::AtomicCmpXchgInst::getStrongestFailureOrdering(SuccessOrdering: AO);
11619 auto Result = atomicInfo.EmitAtomicCompareExchangeLibcall(
11620 ExpectedVal: AtomicLoadRes.second, DesiredVal: NewAtomicAddr, Success: AO, Failure);
11621 LoadInst *PHILoad = Builder.CreateLoad(Ty: XElemTy, Ptr: Result.first);
11622 PHI->addIncoming(V: PHILoad, BB: Builder.GetInsertBlock());
11623 Builder.CreateCondBr(Cond: Result.second, True: ExitBB, False: ContBB);
11624 OldVal->eraseFromParent();
11625 Res.first = OldExprVal;
11626 Res.second = Upd;
11627
11628 if (UnreachableInst *ExitTI =
11629 dyn_cast<UnreachableInst>(Val: ExitBB->getTerminator())) {
11630 CurBBTI->eraseFromParent();
11631 Builder.SetInsertPoint(ExitBB);
11632 } else {
11633 Builder.SetInsertPoint(ExitTI);
11634 }
11635 } else {
11636 IntegerType *IntCastTy =
11637 IntegerType::get(C&: M.getContext(), NumBits: XElemTy->getScalarSizeInBits());
11638 LoadInst *OldVal =
11639 Builder.CreateLoad(Ty: IntCastTy, Ptr: X, Name: X->getName() + ".atomic.load");
11640 AtomicOrdering LoadAO = TransformReleaseAcquireRelease(AO);
11641 OldVal->setAtomic(Ordering: LoadAO);
11642 // CurBB
11643 // | /---\
11644 // ContBB |
11645 // | \---/
11646 // ExitBB
11647 BasicBlock *CurBB = Builder.GetInsertBlock();
11648 Instruction *CurBBTI = CurBB->getTerminatorOrNull();
11649 CurBBTI = CurBBTI ? CurBBTI : Builder.CreateUnreachable();
11650 BasicBlock *ExitBB =
11651 CurBB->splitBasicBlock(I: CurBBTI, BBName: X->getName() + ".atomic.exit");
11652 BasicBlock *ContBB = CurBB->splitBasicBlock(I: CurBB->getTerminator(),
11653 BBName: X->getName() + ".atomic.cont");
11654 ContBB->getTerminator()->eraseFromParent();
11655 Builder.restoreIP(IP: AllocaIP);
11656 AllocaInst *NewAtomicAddr = Builder.CreateAlloca(Ty: XElemTy);
11657 NewAtomicAddr->setName(X->getName() + "x.new.val");
11658 Builder.SetInsertPoint(ContBB);
11659 llvm::PHINode *PHI = Builder.CreatePHI(Ty: OldVal->getType(), NumReservedValues: 2);
11660 PHI->addIncoming(V: OldVal, BB: CurBB);
11661 bool IsIntTy = XElemTy->isIntegerTy();
11662 Value *OldExprVal = PHI;
11663 if (!IsIntTy) {
11664 if (XElemTy->isFloatingPointTy()) {
11665 OldExprVal = Builder.CreateBitCast(V: PHI, DestTy: XElemTy,
11666 Name: X->getName() + ".atomic.fltCast");
11667 } else {
11668 OldExprVal = Builder.CreateIntToPtr(V: PHI, DestTy: XElemTy,
11669 Name: X->getName() + ".atomic.ptrCast");
11670 }
11671 }
11672
11673 Expected<Value *> CBResult = UpdateOp(OldExprVal, Builder);
11674 if (!CBResult)
11675 return CBResult.takeError();
11676 Value *Upd = *CBResult;
11677 Builder.CreateStore(Val: Upd, Ptr: NewAtomicAddr);
11678 LoadInst *DesiredVal = Builder.CreateLoad(Ty: IntCastTy, Ptr: NewAtomicAddr);
11679 AtomicOrdering Failure =
11680 llvm::AtomicCmpXchgInst::getStrongestFailureOrdering(SuccessOrdering: AO);
11681 AtomicCmpXchgInst *Result = Builder.CreateAtomicCmpXchg(
11682 Ptr: X, Cmp: PHI, New: DesiredVal, Align: llvm::MaybeAlign(), SuccessOrdering: AO, FailureOrdering: Failure);
11683 Result->setVolatile(VolatileX);
11684 Value *PreviousVal = Builder.CreateExtractValue(Agg: Result, /*Idxs=*/0);
11685 Value *SuccessFailureVal = Builder.CreateExtractValue(Agg: Result, /*Idxs=*/1);
11686 PHI->addIncoming(V: PreviousVal, BB: Builder.GetInsertBlock());
11687 Builder.CreateCondBr(Cond: SuccessFailureVal, True: ExitBB, False: ContBB);
11688
11689 Res.first = OldExprVal;
11690 Res.second = Upd;
11691
11692 // set Insertion point in exit block
11693 if (UnreachableInst *ExitTI =
11694 dyn_cast<UnreachableInst>(Val: ExitBB->getTerminator())) {
11695 CurBBTI->eraseFromParent();
11696 Builder.SetInsertPoint(ExitBB);
11697 } else {
11698 Builder.SetInsertPoint(ExitTI);
11699 }
11700 }
11701
11702 return Res;
11703}
11704
11705OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createAtomicCapture(
11706 const LocationDescription &Loc, InsertPointTy AllocaIP, AtomicOpValue &X,
11707 AtomicOpValue &V, Value *Expr, AtomicOrdering AO,
11708 AtomicRMWInst::BinOp RMWOp, AtomicUpdateCallbackTy &UpdateOp,
11709 bool UpdateExpr, bool IsPostfixUpdate, bool IsXBinopExpr,
11710 bool IsIgnoreDenormalMode, bool IsFineGrainedMemory, bool IsRemoteMemory) {
11711 if (!updateToLocation(Loc))
11712 return Loc.IP;
11713
11714 LLVM_DEBUG({
11715 Type *XTy = X.Var->getType();
11716 assert(XTy->isPointerTy() &&
11717 "OMP Atomic expects a pointer to target memory");
11718 Type *XElemTy = X.ElemTy;
11719 assert((XElemTy->isFloatingPointTy() || XElemTy->isIntegerTy() ||
11720 XElemTy->isPointerTy() || XElemTy->isStructTy()) &&
11721 "OMP atomic capture expected a scalar or struct type");
11722 assert((RMWOp != AtomicRMWInst::Max) && (RMWOp != AtomicRMWInst::Min) &&
11723 "OpenMP atomic does not support LT or GT operations");
11724 });
11725
11726 // If UpdateExpr is 'x' updated with some `expr` not based on 'x',
11727 // 'x' is simply atomically rewritten with 'expr'.
11728 AtomicRMWInst::BinOp AtomicOp = (UpdateExpr ? RMWOp : AtomicRMWInst::Xchg);
11729 Expected<std::pair<Value *, Value *>> AtomicResult = emitAtomicUpdate(
11730 AllocaIP, X: X.Var, XElemTy: X.ElemTy, Expr, AO, RMWOp: AtomicOp, UpdateOp, VolatileX: X.IsVolatile,
11731 IsXBinopExpr, IsIgnoreDenormalMode, IsFineGrainedMemory, IsRemoteMemory);
11732 if (!AtomicResult)
11733 return AtomicResult.takeError();
11734 Value *CapturedVal =
11735 (IsPostfixUpdate ? AtomicResult->first : AtomicResult->second);
11736 Builder.CreateStore(Val: CapturedVal, Ptr: V.Var, isVolatile: V.IsVolatile);
11737
11738 checkAndEmitFlushAfterAtomic(Loc, AO, AK: AtomicKind::Capture);
11739 return Builder.saveIP();
11740}
11741
11742OpenMPIRBuilder::InsertPointTy OpenMPIRBuilder::createAtomicCompare(
11743 const LocationDescription &Loc, AtomicOpValue &X, AtomicOpValue &V,
11744 AtomicOpValue &R, Value *E, Value *D, AtomicOrdering AO,
11745 omp::OMPAtomicCompareOp Op, bool IsXBinopExpr, bool IsPostfixUpdate,
11746 bool IsFailOnly, bool IsWeak) {
11747
11748 AtomicOrdering Failure = AtomicCmpXchgInst::getStrongestFailureOrdering(SuccessOrdering: AO);
11749 return createAtomicCompare(Loc, X, V, R, E, D, AO, Op, IsXBinopExpr,
11750 IsPostfixUpdate, IsFailOnly, Failure, IsWeak);
11751}
11752
11753OpenMPIRBuilder::InsertPointTy OpenMPIRBuilder::createAtomicCompare(
11754 const LocationDescription &Loc, AtomicOpValue &X, AtomicOpValue &V,
11755 AtomicOpValue &R, Value *E, Value *D, AtomicOrdering AO,
11756 omp::OMPAtomicCompareOp Op, bool IsXBinopExpr, bool IsPostfixUpdate,
11757 bool IsFailOnly, AtomicOrdering Failure, bool IsWeak) {
11758
11759 if (!updateToLocation(Loc))
11760 return Loc.IP;
11761
11762 assert(X.Var->getType()->isPointerTy() &&
11763 "OMP atomic expects a pointer to target memory");
11764 // compare capture
11765 if (V.Var) {
11766 assert(V.Var->getType()->isPointerTy() && "v.var must be of pointer type");
11767 assert(V.ElemTy == X.ElemTy && "x and v must be of same type");
11768 }
11769
11770 bool IsInteger = E->getType()->isIntegerTy();
11771
11772 if (Op == OMPAtomicCompareOp::EQ) {
11773 // OldValue and SuccessOrFail are set below and used in the shared V.Var /
11774 // R.Var handling.
11775 Value *OldValue = nullptr;
11776 Value *SuccessOrFail = nullptr;
11777
11778 if (!IsInteger && HandleFPNegZero) {
11779 // IEEE 754 special cases for cmpxchg (which is bitwise):
11780 // 1. -0.0 == +0.0 but they have different bit patterns.
11781 // 2. NaN != NaN but identical NaN bit patterns would match.
11782 //
11783 // CurBB:
11784 // %e_int = bitcast E to intN
11785 // %d_int = bitcast D to intN
11786 // %x_curr = load atomic intN, X
11787 // %x_fp = bitcast %x_curr to FP
11788 // %e_is_nan = fcmp uno E, E
11789 // %x_is_nan = fcmp uno %x_fp, %x_fp
11790 // %either_nan = or %e_is_nan, %x_is_nan
11791 // br %either_nan, NaNBB, NotNaNBB
11792 // NaNBB: ; NaN == anything is always false
11793 // br ExitBB
11794 // NotNaNBB:
11795 // %x_is_zero = fcmp oeq %x_fp, 0.0
11796 // %e_is_zero = fcmp oeq E, 0.0
11797 // %both_zero = and %x_is_zero, %e_is_zero
11798 // br %both_zero, ZeroBB, NormalBB
11799 // ZeroBB: ; both ±0.0 → x = d
11800 // cmpxchg X, %x_curr, %d_int
11801 // br ExitBB
11802 // NormalBB: ; original path
11803 // cmpxchg X, %e_int, %d_int
11804 // br ExitBB
11805 // ExitBB:
11806 // phi merge
11807 IntegerType *IntCastTy =
11808 IntegerType::get(C&: M.getContext(), NumBits: X.ElemTy->getScalarSizeInBits());
11809 Value *EBCast = Builder.CreateBitCast(V: E, DestTy: IntCastTy);
11810 Value *DBCast = Builder.CreateBitCast(V: D, DestTy: IntCastTy);
11811
11812 // Load X atomically.
11813 LoadInst *XCurr = Builder.CreateLoad(Ty: IntCastTy, Ptr: X.Var,
11814 Name: X.Var->getName() + ".atomic.load");
11815 XCurr->setAtomic(Ordering: AtomicOrdering::Monotonic);
11816 Value *XFP = Builder.CreateBitCast(V: XCurr, DestTy: X.ElemTy);
11817
11818 // IEEE 754: NaN != NaN, but cmpxchg would succeed if E and X have
11819 // the same NaN bit pattern. Skip cmpxchg when either is NaN.
11820 Value *EIsNaN = Builder.CreateFCmpUNO(LHS: E, RHS: E, Name: "atomic.e.isnan");
11821 Value *XIsNaN = Builder.CreateFCmpUNO(LHS: XFP, RHS: XFP, Name: "atomic.x.isnan");
11822 Value *EitherNaN = Builder.CreateOr(LHS: EIsNaN, RHS: XIsNaN, Name: "atomic.either.nan");
11823
11824 BasicBlock *CurBB = Builder.GetInsertBlock();
11825 Function *F = CurBB->getParent();
11826 Instruction *CurBBTI = CurBB->getTerminatorOrNull();
11827 CurBBTI = CurBBTI ? CurBBTI : Builder.CreateUnreachable();
11828 BasicBlock *ExitBB =
11829 CurBB->splitBasicBlock(I: CurBBTI, BBName: X.Var->getName() + ".atomic.exit");
11830 BasicBlock *NaNBB = BasicBlock::Create(
11831 Context&: M.getContext(), Name: X.Var->getName() + ".atomic.nan", Parent: F, InsertBefore: ExitBB);
11832 BasicBlock *NotNaNBB = BasicBlock::Create(
11833 Context&: M.getContext(), Name: X.Var->getName() + ".atomic.notnan", Parent: F, InsertBefore: ExitBB);
11834 BasicBlock *ZeroBB = BasicBlock::Create(
11835 Context&: M.getContext(), Name: X.Var->getName() + ".atomic.zero", Parent: F, InsertBefore: ExitBB);
11836 BasicBlock *NormalBB = BasicBlock::Create(
11837 Context&: M.getContext(), Name: X.Var->getName() + ".atomic.normal", Parent: F, InsertBefore: ExitBB);
11838
11839 // If either E or X is NaN → NaNBB (always fails), else check for ±0.0.
11840 CurBB->getTerminator()->eraseFromParent();
11841 Builder.SetInsertPoint(CurBB);
11842 Builder.CreateCondBr(Cond: EitherNaN, True: NaNBB, False: NotNaNBB);
11843
11844 // NaNBB: NaN == anything is always false; skip cmpxchg.
11845 Builder.SetInsertPoint(NaNBB);
11846 Builder.CreateBr(Dest: ExitBB);
11847
11848 // NotNaNBB: check both X and E for ±0.0.
11849 Builder.SetInsertPoint(NotNaNBB);
11850 Value *XIsZero =
11851 Builder.CreateFCmpOEQ(LHS: XFP, RHS: ConstantFP::getZero(Ty: X.ElemTy),
11852 Name: X.Var->getName() + ".atomic.xiszero");
11853 Value *EIsZero = Builder.CreateFCmpOEQ(LHS: E, RHS: ConstantFP::getZero(Ty: X.ElemTy),
11854 Name: "atomic.e.iszero");
11855 Value *BothZero = Builder.CreateAnd(LHS: XIsZero, RHS: EIsZero, Name: "atomic.both.zero");
11856 Builder.CreateCondBr(Cond: BothZero, True: ZeroBB, False: NormalBB);
11857
11858 // ZeroBB: cmpxchg with X's loaded bit-pattern.
11859 Builder.SetInsertPoint(ZeroBB);
11860 AtomicCmpXchgInst *ResZero = Builder.CreateAtomicCmpXchg(
11861 Ptr: X.Var, Cmp: XCurr, New: DBCast, Align: MaybeAlign(), SuccessOrdering: AO, FailureOrdering: Failure);
11862 ResZero->setWeak(IsWeak);
11863 Value *OldZero = Builder.CreateExtractValue(Agg: ResZero, /*Idxs=*/0);
11864 Value *OkZero = Builder.CreateExtractValue(Agg: ResZero, /*Idxs=*/1);
11865 Builder.CreateBr(Dest: ExitBB);
11866
11867 // NormalBB: original bitwise cmpxchg.
11868 Builder.SetInsertPoint(NormalBB);
11869 AtomicCmpXchgInst *ResNormal = Builder.CreateAtomicCmpXchg(
11870 Ptr: X.Var, Cmp: EBCast, New: DBCast, Align: MaybeAlign(), SuccessOrdering: AO, FailureOrdering: Failure);
11871 ResNormal->setWeak(IsWeak);
11872 Value *OldNormal = Builder.CreateExtractValue(Agg: ResNormal, /*Idxs=*/0);
11873 Value *OkNormal = Builder.CreateExtractValue(Agg: ResNormal, /*Idxs=*/1);
11874 Builder.CreateBr(Dest: ExitBB);
11875
11876 // ExitBB: merge results from NaN, Zero, and Normal paths.
11877 Builder.SetInsertPoint(ExitBB->begin());
11878 PHINode *OldIntPHI =
11879 Builder.CreatePHI(Ty: IntCastTy, NumReservedValues: 3, Name: X.Var->getName() + ".atomic.old");
11880 OldIntPHI->addIncoming(V: XCurr, BB: NaNBB);
11881 OldIntPHI->addIncoming(V: OldZero, BB: ZeroBB);
11882 OldIntPHI->addIncoming(V: OldNormal, BB: NormalBB);
11883 PHINode *SuccessPHI = Builder.CreatePHI(Ty: Builder.getInt1Ty(), NumReservedValues: 3,
11884 Name: X.Var->getName() + ".atomic.ok");
11885 SuccessPHI->addIncoming(V: Builder.getFalse(), BB: NaNBB);
11886 SuccessPHI->addIncoming(V: OkZero, BB: ZeroBB);
11887 SuccessPHI->addIncoming(V: OkNormal, BB: NormalBB);
11888
11889 if (isa<UnreachableInst>(Val: ExitBB->getTerminator())) {
11890 CurBBTI->eraseFromParent();
11891 Builder.SetInsertPoint(ExitBB);
11892 } else {
11893 Builder.SetInsertPoint(&*ExitBB->getFirstNonPHIIt());
11894 }
11895
11896 OldValue = Builder.CreateBitCast(V: OldIntPHI, DestTy: X.ElemTy,
11897 Name: X.Var->getName() + ".atomic.old.fp");
11898 SuccessOrFail = SuccessPHI;
11899 } else {
11900 AtomicCmpXchgInst *Result = nullptr;
11901 if (!IsInteger) {
11902 IntegerType *IntCastTy =
11903 IntegerType::get(C&: M.getContext(), NumBits: X.ElemTy->getScalarSizeInBits());
11904 Value *EBCast = Builder.CreateBitCast(V: E, DestTy: IntCastTy);
11905 Value *DBCast = Builder.CreateBitCast(V: D, DestTy: IntCastTy);
11906 Result = Builder.CreateAtomicCmpXchg(Ptr: X.Var, Cmp: EBCast, New: DBCast,
11907 Align: MaybeAlign(), SuccessOrdering: AO, FailureOrdering: Failure);
11908 } else {
11909 Result =
11910 Builder.CreateAtomicCmpXchg(Ptr: X.Var, Cmp: E, New: D, Align: MaybeAlign(), SuccessOrdering: AO, FailureOrdering: Failure);
11911 }
11912 Result->setWeak(IsWeak);
11913
11914 if (V.Var) {
11915 OldValue = Builder.CreateExtractValue(Agg: Result, /*Idxs=*/0);
11916 if (!IsInteger)
11917 OldValue = Builder.CreateBitCast(V: OldValue, DestTy: X.ElemTy);
11918 assert(OldValue->getType() == V.ElemTy &&
11919 "OldValue and V must be of same type");
11920 if (IsPostfixUpdate) {
11921 Builder.CreateStore(Val: OldValue, Ptr: V.Var, isVolatile: V.IsVolatile);
11922 } else {
11923 SuccessOrFail = Builder.CreateExtractValue(Agg: Result, /*Idxs=*/1);
11924 if (IsFailOnly) {
11925 BasicBlock *CurBB = Builder.GetInsertBlock();
11926 Instruction *CurBBTI = CurBB->getTerminatorOrNull();
11927 CurBBTI = CurBBTI ? CurBBTI : Builder.CreateUnreachable();
11928 BasicBlock *ExitBB = CurBB->splitBasicBlock(
11929 I: CurBBTI, BBName: X.Var->getName() + ".atomic.exit");
11930 BasicBlock *ContBB = CurBB->splitBasicBlock(
11931 I: CurBB->getTerminator(), BBName: X.Var->getName() + ".atomic.cont");
11932 ContBB->getTerminator()->eraseFromParent();
11933 CurBB->getTerminator()->eraseFromParent();
11934
11935 Builder.CreateCondBr(Cond: SuccessOrFail, True: ExitBB, False: ContBB);
11936
11937 Builder.SetInsertPoint(ContBB);
11938 Builder.CreateStore(Val: OldValue, Ptr: V.Var);
11939 Builder.CreateBr(Dest: ExitBB);
11940
11941 if (UnreachableInst *ExitTI =
11942 dyn_cast<UnreachableInst>(Val: ExitBB->getTerminator())) {
11943 CurBBTI->eraseFromParent();
11944 Builder.SetInsertPoint(ExitBB);
11945 } else {
11946 Builder.SetInsertPoint(ExitTI);
11947 }
11948 } else {
11949 Value *CapturedValue =
11950 Builder.CreateSelect(C: SuccessOrFail, True: E, False: OldValue);
11951 Builder.CreateStore(Val: CapturedValue, Ptr: V.Var, isVolatile: V.IsVolatile);
11952 }
11953 }
11954 }
11955 // The comparison result has to be stored.
11956 if (R.Var) {
11957 assert(R.Var->getType()->isPointerTy() &&
11958 "r.var must be of pointer type");
11959 assert(R.ElemTy->isIntegerTy() && "r must be of integral type");
11960
11961 Value *SuccessFailureVal =
11962 Builder.CreateExtractValue(Agg: Result, /*Idxs=*/1);
11963 Value *ResultCast =
11964 R.IsSigned ? Builder.CreateSExt(V: SuccessFailureVal, DestTy: R.ElemTy)
11965 : Builder.CreateZExt(V: SuccessFailureVal, DestTy: R.ElemTy);
11966 Builder.CreateStore(Val: ResultCast, Ptr: R.Var, isVolatile: R.IsVolatile);
11967 }
11968 }
11969
11970 // For the HandleFPNegZero path, handle V.Var and R.Var using the
11971 // pre-computed OldValue and SuccessOrFail.
11972 if (HandleFPNegZero && !IsInteger) {
11973 if (V.Var) {
11974 assert(OldValue->getType() == V.ElemTy &&
11975 "OldValue and V must be of same type");
11976 if (IsPostfixUpdate) {
11977 Builder.CreateStore(Val: OldValue, Ptr: V.Var, isVolatile: V.IsVolatile);
11978 } else {
11979 if (IsFailOnly) {
11980 BasicBlock *CurBB = Builder.GetInsertBlock();
11981 Instruction *CurBBTI = CurBB->getTerminatorOrNull();
11982 CurBBTI = CurBBTI ? CurBBTI : Builder.CreateUnreachable();
11983 BasicBlock *ExitBB = CurBB->splitBasicBlock(
11984 I: CurBBTI, BBName: X.Var->getName() + ".atomic.exit");
11985 BasicBlock *ContBB = CurBB->splitBasicBlock(
11986 I: CurBB->getTerminator(), BBName: X.Var->getName() + ".atomic.cont");
11987 ContBB->getTerminator()->eraseFromParent();
11988 CurBB->getTerminator()->eraseFromParent();
11989
11990 Builder.CreateCondBr(Cond: SuccessOrFail, True: ExitBB, False: ContBB);
11991
11992 Builder.SetInsertPoint(ContBB);
11993 Builder.CreateStore(Val: OldValue, Ptr: V.Var);
11994 Builder.CreateBr(Dest: ExitBB);
11995
11996 if (UnreachableInst *ExitTI =
11997 dyn_cast<UnreachableInst>(Val: ExitBB->getTerminator())) {
11998 CurBBTI->eraseFromParent();
11999 Builder.SetInsertPoint(ExitBB);
12000 } else {
12001 Builder.SetInsertPoint(ExitTI);
12002 }
12003 } else {
12004 Value *CapturedValue =
12005 Builder.CreateSelect(C: SuccessOrFail, True: E, False: OldValue);
12006 Builder.CreateStore(Val: CapturedValue, Ptr: V.Var, isVolatile: V.IsVolatile);
12007 }
12008 }
12009 }
12010 // The comparison result has to be stored.
12011 if (R.Var) {
12012 assert(R.Var->getType()->isPointerTy() &&
12013 "r.var must be of pointer type");
12014 assert(R.ElemTy->isIntegerTy() && "r must be of integral type");
12015
12016 Value *ResultCast = R.IsSigned
12017 ? Builder.CreateSExt(V: SuccessOrFail, DestTy: R.ElemTy)
12018 : Builder.CreateZExt(V: SuccessOrFail, DestTy: R.ElemTy);
12019 Builder.CreateStore(Val: ResultCast, Ptr: R.Var, isVolatile: R.IsVolatile);
12020 }
12021 }
12022 } else {
12023 assert((Op == OMPAtomicCompareOp::MAX || Op == OMPAtomicCompareOp::MIN) &&
12024 "Op should be either max or min at this point");
12025 assert(!IsFailOnly && "IsFailOnly is only valid when the comparison is ==");
12026
12027 // Reverse the ordop as the OpenMP forms are different from LLVM forms.
12028 // Let's take max as example.
12029 // OpenMP form:
12030 // x = x > expr ? expr : x;
12031 // LLVM form:
12032 // *ptr = *ptr > val ? *ptr : val;
12033 // We need to transform to LLVM form.
12034 // x = x <= expr ? x : expr;
12035 AtomicRMWInst::BinOp NewOp;
12036 if (IsXBinopExpr) {
12037 if (IsInteger) {
12038 if (X.IsSigned)
12039 NewOp = Op == OMPAtomicCompareOp::MAX ? AtomicRMWInst::Min
12040 : AtomicRMWInst::Max;
12041 else
12042 NewOp = Op == OMPAtomicCompareOp::MAX ? AtomicRMWInst::UMin
12043 : AtomicRMWInst::UMax;
12044 } else {
12045 NewOp = Op == OMPAtomicCompareOp::MAX ? AtomicRMWInst::FMin
12046 : AtomicRMWInst::FMax;
12047 }
12048 } else {
12049 if (IsInteger) {
12050 if (X.IsSigned)
12051 NewOp = Op == OMPAtomicCompareOp::MAX ? AtomicRMWInst::Max
12052 : AtomicRMWInst::Min;
12053 else
12054 NewOp = Op == OMPAtomicCompareOp::MAX ? AtomicRMWInst::UMax
12055 : AtomicRMWInst::UMin;
12056 } else {
12057 NewOp = Op == OMPAtomicCompareOp::MAX ? AtomicRMWInst::FMax
12058 : AtomicRMWInst::FMin;
12059 }
12060 }
12061
12062 AtomicRMWInst *OldValue =
12063 Builder.CreateAtomicRMW(Op: NewOp, Ptr: X.Var, Val: E, Align: MaybeAlign(), Ordering: AO);
12064 if (V.Var) {
12065 Value *CapturedValue = nullptr;
12066 if (IsPostfixUpdate) {
12067 CapturedValue = OldValue;
12068 } else {
12069 CmpInst::Predicate Pred;
12070 switch (NewOp) {
12071 case AtomicRMWInst::Max:
12072 Pred = CmpInst::ICMP_SGT;
12073 break;
12074 case AtomicRMWInst::UMax:
12075 Pred = CmpInst::ICMP_UGT;
12076 break;
12077 case AtomicRMWInst::FMax:
12078 Pred = CmpInst::FCMP_OGT;
12079 break;
12080 case AtomicRMWInst::Min:
12081 Pred = CmpInst::ICMP_SLT;
12082 break;
12083 case AtomicRMWInst::UMin:
12084 Pred = CmpInst::ICMP_ULT;
12085 break;
12086 case AtomicRMWInst::FMin:
12087 Pred = CmpInst::FCMP_OLT;
12088 break;
12089 default:
12090 llvm_unreachable("unexpected comparison op");
12091 }
12092 Value *NonAtomicCmp = Builder.CreateCmp(Pred, LHS: OldValue, RHS: E);
12093 CapturedValue = Builder.CreateSelect(C: NonAtomicCmp, True: E, False: OldValue);
12094 }
12095 Builder.CreateStore(Val: CapturedValue, Ptr: V.Var, isVolatile: V.IsVolatile);
12096 }
12097 }
12098
12099 checkAndEmitFlushAfterAtomic(Loc, AO, AK: AtomicKind::Compare);
12100
12101 return Builder.saveIP();
12102}
12103
12104OpenMPIRBuilder::InsertPointOrErrorTy
12105OpenMPIRBuilder::createTeams(const LocationDescription &Loc,
12106 BodyGenCallbackTy BodyGenCB, Value *NumTeamsLower,
12107 Value *NumTeamsUpper, Value *ThreadLimit,
12108 Value *IfExpr) {
12109 if (!updateToLocation(Loc))
12110 return InsertPointTy();
12111
12112 uint32_t SrcLocStrSize;
12113 Constant *SrcLocStr = getOrCreateSrcLocStr(Loc, SrcLocStrSize);
12114 Value *Ident = getOrCreateIdent(SrcLocStr, SrcLocStrSize);
12115 Function *CurrentFunction = Builder.GetInsertBlock()->getParent();
12116
12117 // Outer allocation basicblock is the entry block of the current function.
12118 BasicBlock &OuterAllocaBB = CurrentFunction->getEntryBlock();
12119 if (&OuterAllocaBB == Builder.GetInsertBlock()) {
12120 BasicBlock *BodyBB = splitBB(Builder, /*CreateBranch=*/true, Name: "teams.entry");
12121 Builder.SetInsertPoint(BodyBB->begin());
12122 }
12123
12124 // The current basic block is split into four basic blocks. After outlining,
12125 // they will be mapped as follows:
12126 // ```
12127 // def current_fn() {
12128 // current_basic_block:
12129 // br label %teams.exit
12130 // teams.exit:
12131 // ; instructions after teams
12132 // }
12133 //
12134 // def outlined_fn() {
12135 // teams.alloca:
12136 // br label %teams.body
12137 // teams.body:
12138 // ; instructions within teams body
12139 // }
12140 // ```
12141 BasicBlock *ExitBB = splitBB(Builder, /*CreateBranch=*/true, Name: "teams.exit");
12142 BasicBlock *BodyBB = splitBB(Builder, /*CreateBranch=*/true, Name: "teams.body");
12143 BasicBlock *AllocaBB =
12144 splitBB(Builder, /*CreateBranch=*/true, Name: "teams.alloca");
12145
12146 bool SubClausesPresent =
12147 (NumTeamsLower || NumTeamsUpper || ThreadLimit || IfExpr);
12148 // Push num_teams
12149 if (!Config.isTargetDevice() && SubClausesPresent) {
12150 assert((NumTeamsLower == nullptr || NumTeamsUpper != nullptr) &&
12151 "if lowerbound is non-null, then upperbound must also be non-null "
12152 "for bounds on num_teams");
12153
12154 if (NumTeamsUpper == nullptr)
12155 NumTeamsUpper = Builder.getInt32(C: 0);
12156
12157 if (NumTeamsLower == nullptr)
12158 NumTeamsLower = NumTeamsUpper;
12159
12160 if (IfExpr) {
12161 assert(IfExpr->getType()->isIntegerTy() &&
12162 "argument to if clause must be an integer value");
12163
12164 // upper = ifexpr ? upper : 1
12165 if (IfExpr->getType() != Int1)
12166 IfExpr = Builder.CreateICmpNE(LHS: IfExpr,
12167 RHS: ConstantInt::get(Ty: IfExpr->getType(), V: 0));
12168 NumTeamsUpper = Builder.CreateSelect(
12169 C: IfExpr, True: NumTeamsUpper, False: Builder.getInt32(C: 1), Name: "numTeamsUpper");
12170
12171 // lower = ifexpr ? lower : 1
12172 NumTeamsLower = Builder.CreateSelect(
12173 C: IfExpr, True: NumTeamsLower, False: Builder.getInt32(C: 1), Name: "numTeamsLower");
12174 }
12175
12176 if (ThreadLimit == nullptr)
12177 ThreadLimit = Builder.getInt32(C: 0);
12178
12179 // The __kmpc_push_num_teams_51 function expects int32 as the arguments. So,
12180 // truncate or sign extend the passed values to match the int32 parameters.
12181 Value *NumTeamsLowerInt32 =
12182 Builder.CreateSExtOrTrunc(V: NumTeamsLower, DestTy: Builder.getInt32Ty());
12183 Value *NumTeamsUpperInt32 =
12184 Builder.CreateSExtOrTrunc(V: NumTeamsUpper, DestTy: Builder.getInt32Ty());
12185 Value *ThreadLimitInt32 =
12186 Builder.CreateSExtOrTrunc(V: ThreadLimit, DestTy: Builder.getInt32Ty());
12187
12188 Value *ThreadNum = getOrCreateThreadID(Ident);
12189
12190 createRuntimeFunctionCall(
12191 Callee: getOrCreateRuntimeFunctionPtr(FnID: OMPRTL___kmpc_push_num_teams_51),
12192 Args: {Ident, ThreadNum, NumTeamsLowerInt32, NumTeamsUpperInt32,
12193 ThreadLimitInt32});
12194 }
12195 // Generate the body of teams.
12196 InsertPointTy AllocaIP(AllocaBB->begin());
12197 InsertPointTy CodeGenIP(BodyBB->begin());
12198 if (Error Err = BodyGenCB(AllocaIP, CodeGenIP, ExitBB))
12199 return Err;
12200
12201 auto OI = std::make_unique<OutlineInfo>();
12202 OI->EntryBB = AllocaBB;
12203 OI->ExitBB = ExitBB;
12204 OI->OuterAllocBB = &OuterAllocaBB;
12205
12206 // Insert fake values for global tid and bound tid.
12207 SmallVector<Instruction *, 8> ToBeDeleted;
12208 InsertPointTy OuterAllocaIP(OuterAllocaBB.begin());
12209 OI->ExcludeArgsFromAggregate.push_back(Elt: createFakeIntVal(
12210 Builder, OuterAllocaIP, ToBeDeleted, InnerAllocaIP: AllocaIP, Name: "gid", AsPtr: true));
12211 OI->ExcludeArgsFromAggregate.push_back(Elt: createFakeIntVal(
12212 Builder, OuterAllocaIP, ToBeDeleted, InnerAllocaIP: AllocaIP, Name: "tid", AsPtr: true));
12213
12214 auto HostPostOutlineCB = [this, Ident,
12215 ToBeDeleted](Function &OutlinedFn) mutable {
12216 // The stale call instruction will be replaced with a new call instruction
12217 // for runtime call with the outlined function.
12218
12219 assert(OutlinedFn.hasOneUse() &&
12220 "there must be a single user for the outlined function");
12221 CallInst *StaleCI = cast<CallInst>(Val: OutlinedFn.user_back());
12222 ToBeDeleted.push_back(Elt: StaleCI);
12223
12224 assert((OutlinedFn.arg_size() == 2 || OutlinedFn.arg_size() == 3) &&
12225 "Outlined function must have two or three arguments only");
12226
12227 bool HasShared = OutlinedFn.arg_size() == 3;
12228
12229 OutlinedFn.getArg(i: 0)->setName("global.tid.ptr");
12230 OutlinedFn.getArg(i: 1)->setName("bound.tid.ptr");
12231 if (HasShared)
12232 OutlinedFn.getArg(i: 2)->setName("data");
12233
12234 // Call to the runtime function for teams in the current function.
12235 assert(StaleCI && "Error while outlining - no CallInst user found for the "
12236 "outlined function.");
12237 Builder.SetInsertPoint(StaleCI);
12238 SmallVector<Value *> Args = {
12239 Ident, Builder.getInt32(C: StaleCI->arg_size() - 2), &OutlinedFn};
12240 if (HasShared)
12241 Args.push_back(Elt: StaleCI->getArgOperand(i: 2));
12242 createRuntimeFunctionCall(
12243 Callee: getOrCreateRuntimeFunctionPtr(
12244 FnID: omp::RuntimeFunction::OMPRTL___kmpc_fork_teams),
12245 Args);
12246
12247 Builder.ClearInsertionPoint();
12248 for (Instruction *I : llvm::reverse(C&: ToBeDeleted))
12249 I->eraseFromParent();
12250 };
12251
12252 if (!Config.isTargetDevice())
12253 OI->PostOutlineCB = HostPostOutlineCB;
12254
12255 addOutlineInfo(OI: std::move(OI));
12256
12257 Builder.SetInsertPoint(ExitBB);
12258
12259 return Builder.saveIP();
12260}
12261
12262OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createDistribute(
12263 const LocationDescription &Loc, InsertPointTy OuterAllocIP,
12264 ArrayRef<BasicBlock *> OuterDeallocBlocks, BodyGenCallbackTy BodyGenCB) {
12265 if (!updateToLocation(Loc))
12266 return InsertPointTy();
12267
12268 BasicBlock *OuterAllocaBB = OuterAllocIP.getNodeParent();
12269
12270 if (OuterAllocaBB == Builder.GetInsertBlock()) {
12271 BasicBlock *BodyBB =
12272 splitBB(Builder, /*CreateBranch=*/true, Name: "distribute.entry");
12273 Builder.SetInsertPoint(BodyBB->begin());
12274 }
12275 BasicBlock *ExitBB =
12276 splitBB(Builder, /*CreateBranch=*/true, Name: "distribute.exit");
12277 BasicBlock *BodyBB =
12278 splitBB(Builder, /*CreateBranch=*/true, Name: "distribute.body");
12279 BasicBlock *AllocaBB =
12280 splitBB(Builder, /*CreateBranch=*/true, Name: "distribute.alloca");
12281
12282 // Generate the body of distribute clause
12283 InsertPointTy AllocaIP(AllocaBB->begin());
12284 InsertPointTy CodeGenIP(BodyBB->begin());
12285 if (Error Err = BodyGenCB(AllocaIP, CodeGenIP, ExitBB))
12286 return Err;
12287
12288 // When using target we use different runtime functions which require a
12289 // callback.
12290 if (Config.isTargetDevice()) {
12291 auto OI = std::make_unique<OutlineInfo>();
12292 OI->OuterAllocBB = OuterAllocIP.getNodeParent();
12293 OI->EntryBB = AllocaBB;
12294 OI->ExitBB = ExitBB;
12295 OI->OuterDeallocBBs.reserve(N: OuterDeallocBlocks.size());
12296 copy(Range&: OuterDeallocBlocks, Out: OI->OuterDeallocBBs.end());
12297
12298 addOutlineInfo(OI: std::move(OI));
12299 }
12300 Builder.SetInsertPoint(ExitBB);
12301
12302 return Builder.saveIP();
12303}
12304
12305GlobalVariable *
12306OpenMPIRBuilder::createOffloadMapnames(SmallVectorImpl<llvm::Constant *> &Names,
12307 std::string VarName) {
12308 llvm::Constant *MapNamesArrayInit = llvm::ConstantArray::get(
12309 T: llvm::ArrayType::get(ElementType: llvm::PointerType::getUnqual(C&: M.getContext()),
12310 NumElements: Names.size()),
12311 V: Names);
12312 auto *MapNamesArrayGlobal = new llvm::GlobalVariable(
12313 M, MapNamesArrayInit->getType(),
12314 /*isConstant=*/true, llvm::GlobalValue::PrivateLinkage, MapNamesArrayInit,
12315 VarName);
12316 return MapNamesArrayGlobal;
12317}
12318
12319// Create all simple and struct types exposed by the runtime and remember
12320// the llvm::PointerTypes of them for easy access later.
12321void OpenMPIRBuilder::initializeTypes(Module &M) {
12322 LLVMContext &Ctx = M.getContext();
12323 StructType *T;
12324 unsigned DefaultTargetAS = Config.getDefaultTargetAS();
12325 unsigned ProgramAS = M.getDataLayout().getProgramAddressSpace();
12326#define OMP_TYPE(VarName, InitValue) VarName = InitValue;
12327#define OMP_ARRAY_TYPE(VarName, ElemTy, ArraySize) \
12328 VarName##Ty = ArrayType::get(ElemTy, ArraySize); \
12329 VarName##PtrTy = PointerType::get(Ctx, DefaultTargetAS);
12330#define OMP_FUNCTION_TYPE(VarName, IsVarArg, ReturnType, ...) \
12331 VarName = FunctionType::get(ReturnType, {__VA_ARGS__}, IsVarArg); \
12332 VarName##Ptr = PointerType::get(Ctx, ProgramAS);
12333#define OMP_STRUCT_TYPE(VarName, StructName, Packed, ...) \
12334 T = StructType::getTypeByName(Ctx, StructName); \
12335 if (!T) \
12336 T = StructType::create(Ctx, {__VA_ARGS__}, StructName, Packed); \
12337 VarName = T; \
12338 VarName##Ptr = PointerType::get(Ctx, DefaultTargetAS);
12339#include "llvm/Frontend/OpenMP/OMPKinds.def"
12340}
12341
12342void OpenMPIRBuilder::OutlineInfo::collectBlocks(
12343 SmallPtrSetImpl<BasicBlock *> &BlockSet,
12344 SmallVectorImpl<BasicBlock *> &BlockVector) {
12345 SmallVector<BasicBlock *, 32> Worklist;
12346 BlockSet.insert(Ptr: EntryBB);
12347 BlockSet.insert(Ptr: ExitBB);
12348
12349 Worklist.push_back(Elt: EntryBB);
12350 while (!Worklist.empty()) {
12351 BasicBlock *BB = Worklist.pop_back_val();
12352 BlockVector.push_back(Elt: BB);
12353 for (BasicBlock *SuccBB : successors(BB))
12354 if (BlockSet.insert(Ptr: SuccBB).second)
12355 Worklist.push_back(Elt: SuccBB);
12356 }
12357}
12358
12359std::unique_ptr<CodeExtractor>
12360OpenMPIRBuilder::OutlineInfo::createCodeExtractor(ArrayRef<BasicBlock *> Blocks,
12361 bool ArgsInZeroAddressSpace,
12362 Twine Suffix) {
12363 return std::make_unique<CodeExtractor>(
12364 args&: Blocks, /* DominatorTree */ args: nullptr,
12365 /* AggregateArgs */ args: true,
12366 /* BlockFrequencyInfo */ args: nullptr,
12367 /* BranchProbabilityInfo */ args: nullptr,
12368 /* AssumptionCache */ args: nullptr,
12369 /* AllowVarArgs */ args: true,
12370 /* AllowAlloca */ args: true,
12371 /* AllocationBlock*/ args&: OuterAllocBB,
12372 /* DeallocationBlocks */ args: ArrayRef<BasicBlock *>(),
12373 /* Suffix */ args: Suffix.str(), args&: ArgsInZeroAddressSpace);
12374}
12375
12376std::unique_ptr<CodeExtractor> DeviceSharedMemOutlineInfo::createCodeExtractor(
12377 ArrayRef<BasicBlock *> Blocks, bool ArgsInZeroAddressSpace, Twine Suffix) {
12378 return std::make_unique<DeviceSharedMemCodeExtractor>(
12379 args&: OMPBuilder, args&: Blocks, /* DominatorTree */ args: nullptr,
12380 /* AggregateArgs */ args: true,
12381 /* BlockFrequencyInfo */ args: nullptr,
12382 /* BranchProbabilityInfo */ args: nullptr,
12383 /* AssumptionCache */ args: nullptr,
12384 /* AllowVarArgs */ args: true,
12385 /* AllowAlloca */ args: true,
12386 /* AllocationBlock*/ args&: OuterAllocBB,
12387 /* DeallocationBlocks */ args: OuterDeallocBBs.empty()
12388 ? SmallVector<BasicBlock *>{ExitBB}
12389 : OuterDeallocBBs,
12390 /* Suffix */ args: Suffix.str(), args&: ArgsInZeroAddressSpace);
12391}
12392
12393void OpenMPIRBuilder::createOffloadEntry(Constant *ID, Constant *Addr,
12394 uint64_t Size, int32_t Flags,
12395 GlobalValue::LinkageTypes,
12396 StringRef Name) {
12397 if (!Config.isGPU()) {
12398 llvm::offloading::emitOffloadingEntry(
12399 M, Kind: object::OffloadKind::OFK_OpenMP, Addr: ID,
12400 Name: Name.empty() ? Addr->getName() : Name, Size, Flags, /*Data=*/0);
12401 return;
12402 }
12403 // TODO: Add support for global variables on the device after declare target
12404 // support.
12405 Function *Fn = dyn_cast<Function>(Val: Addr);
12406 if (!Fn)
12407 return;
12408
12409 // Add a function attribute for the kernel.
12410 Fn->addFnAttr(Kind: "kernel");
12411 if (T.isAMDGCN())
12412 Fn->addFnAttr(Kind: "uniform-work-group-size");
12413 Fn->addFnAttr(Kind: Attribute::MustProgress);
12414}
12415
12416// We only generate metadata for function that contain target regions.
12417void OpenMPIRBuilder::createOffloadEntriesAndInfoMetadata(
12418 EmitMetadataErrorReportFunctionTy &ErrorFn) {
12419
12420 // If there are no entries, we don't need to do anything.
12421 if (OffloadInfoManager.empty())
12422 return;
12423
12424 LLVMContext &C = M.getContext();
12425 SmallVector<std::pair<const OffloadEntriesInfoManager::OffloadEntryInfo *,
12426 TargetRegionEntryInfo>,
12427 16>
12428 OrderedEntries(OffloadInfoManager.size());
12429
12430 // Auxiliary methods to create metadata values and strings.
12431 auto &&GetMDInt = [this](unsigned V) {
12432 return ConstantAsMetadata::get(C: ConstantInt::get(Ty: Builder.getInt32Ty(), V));
12433 };
12434
12435 auto &&GetMDString = [&C](StringRef V) { return MDString::get(Context&: C, Str: V); };
12436
12437 // Create the offloading info metadata node.
12438 NamedMDNode *MD = M.getOrInsertNamedMetadata(Name: "omp_offload.info");
12439 auto &&TargetRegionMetadataEmitter =
12440 [&C, MD, &OrderedEntries, &GetMDInt, &GetMDString](
12441 const TargetRegionEntryInfo &EntryInfo,
12442 const OffloadEntriesInfoManager::OffloadEntryInfoTargetRegion &E) {
12443 // Generate metadata for target regions. Each entry of this metadata
12444 // contains:
12445 // - Entry 0 -> Kind of this type of metadata (0).
12446 // - Entry 1 -> Device ID of the file where the entry was identified.
12447 // - Entry 2 -> File ID of the file where the entry was identified.
12448 // - Entry 3 -> Mangled name of the function where the entry was
12449 // identified.
12450 // - Entry 4 -> Line in the file where the entry was identified.
12451 // - Entry 5 -> Count of regions at this DeviceID/FilesID/Line.
12452 // - Entry 6 -> Order the entry was created.
12453 // The first element of the metadata node is the kind.
12454 Metadata *Ops[] = {
12455 GetMDInt(E.getKind()), GetMDInt(EntryInfo.DeviceID),
12456 GetMDInt(EntryInfo.FileID), GetMDString(EntryInfo.ParentName),
12457 GetMDInt(EntryInfo.Line), GetMDInt(EntryInfo.Count),
12458 GetMDInt(E.getOrder())};
12459
12460 // Save this entry in the right position of the ordered entries array.
12461 OrderedEntries[E.getOrder()] = std::make_pair(x: &E, y: EntryInfo);
12462
12463 // Add metadata to the named metadata node.
12464 MD->addOperand(M: MDNode::get(Context&: C, MDs: Ops));
12465 };
12466
12467 OffloadInfoManager.actOnTargetRegionEntriesInfo(Action: TargetRegionMetadataEmitter);
12468
12469 // Create function that emits metadata for each device global variable entry;
12470 auto &&DeviceGlobalVarMetadataEmitter =
12471 [&C, &OrderedEntries, &GetMDInt, &GetMDString, MD](
12472 StringRef MangledName,
12473 const OffloadEntriesInfoManager::OffloadEntryInfoDeviceGlobalVar &E) {
12474 // Generate metadata for global variables. Each entry of this metadata
12475 // contains:
12476 // - Entry 0 -> Kind of this type of metadata (1).
12477 // - Entry 1 -> Mangled name of the variable.
12478 // - Entry 2 -> Declare target kind.
12479 // - Entry 3 -> Order the entry was created.
12480 // The first element of the metadata node is the kind.
12481 Metadata *Ops[] = {GetMDInt(E.getKind()), GetMDString(MangledName),
12482 GetMDInt(E.getFlags()), GetMDInt(E.getOrder())};
12483
12484 // Save this entry in the right position of the ordered entries array.
12485 TargetRegionEntryInfo varInfo(MangledName, 0, 0, 0);
12486 OrderedEntries[E.getOrder()] = std::make_pair(x: &E, y&: varInfo);
12487
12488 // Add metadata to the named metadata node.
12489 MD->addOperand(M: MDNode::get(Context&: C, MDs: Ops));
12490 };
12491
12492 OffloadInfoManager.actOnDeviceGlobalVarEntriesInfo(
12493 Action: DeviceGlobalVarMetadataEmitter);
12494
12495 for (const auto &E : OrderedEntries) {
12496 assert(E.first && "All ordered entries must exist!");
12497 if (const auto *CE =
12498 dyn_cast<OffloadEntriesInfoManager::OffloadEntryInfoTargetRegion>(
12499 Val: E.first)) {
12500 if (!CE->getID() || !CE->getAddress()) {
12501 // Do not blame the entry if the parent funtion is not emitted.
12502 TargetRegionEntryInfo EntryInfo = E.second;
12503 StringRef FnName = EntryInfo.ParentName;
12504 if (!M.getNamedValue(Name: FnName))
12505 continue;
12506 ErrorFn(EMIT_MD_TARGET_REGION_ERROR, EntryInfo);
12507 continue;
12508 }
12509 createOffloadEntry(ID: CE->getID(), Addr: CE->getAddress(),
12510 /*Size=*/0, Flags: CE->getFlags(),
12511 GlobalValue::WeakAnyLinkage);
12512 } else if (const auto *CE = dyn_cast<
12513 OffloadEntriesInfoManager::OffloadEntryInfoDeviceGlobalVar>(
12514 Val: E.first)) {
12515 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind Flags =
12516 static_cast<OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind>(
12517 CE->getFlags());
12518 switch (Flags) {
12519 case OffloadEntriesInfoManager::OMPTargetGlobalVarEntryEnter:
12520 case OffloadEntriesInfoManager::OMPTargetGlobalVarEntryTo:
12521 if (Config.isTargetDevice() && Config.hasRequiresUnifiedSharedMemory())
12522 continue;
12523 if (!CE->getAddress()) {
12524 ErrorFn(EMIT_MD_DECLARE_TARGET_ERROR, E.second);
12525 continue;
12526 }
12527 // The vaiable has no definition - no need to add the entry.
12528 if (CE->getVarSize() == 0)
12529 continue;
12530 break;
12531 case OffloadEntriesInfoManager::OMPTargetGlobalVarEntryLink:
12532 assert(((Config.isTargetDevice() && !CE->getAddress()) ||
12533 (!Config.isTargetDevice() && CE->getAddress())) &&
12534 "Declaret target link address is set.");
12535 if (Config.isTargetDevice())
12536 continue;
12537 if (!CE->getAddress()) {
12538 ErrorFn(EMIT_MD_GLOBAL_VAR_LINK_ERROR, TargetRegionEntryInfo());
12539 continue;
12540 }
12541 break;
12542 case OffloadEntriesInfoManager::OMPTargetGlobalVarEntryIndirect:
12543 case OffloadEntriesInfoManager::OMPTargetGlobalVarEntryIndirectVTable:
12544 if (!CE->getAddress()) {
12545 ErrorFn(EMIT_MD_GLOBAL_VAR_INDIRECT_ERROR, E.second);
12546 continue;
12547 }
12548 break;
12549 default:
12550 break;
12551 }
12552
12553 // Hidden or internal symbols on the device are not externally visible.
12554 // We should not attempt to register them by creating an offloading
12555 // entry. Indirect variables are handled separately on the device.
12556 if (auto *GV = dyn_cast<GlobalValue>(Val: CE->getAddress()))
12557 if ((GV->hasLocalLinkage() || GV->hasHiddenVisibility()) &&
12558 (Flags !=
12559 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryIndirect &&
12560 Flags != OffloadEntriesInfoManager::
12561 OMPTargetGlobalVarEntryIndirectVTable))
12562 continue;
12563
12564 // Indirect globals need to use a special name that doesn't match the name
12565 // of the associated host global.
12566 if (Flags == OffloadEntriesInfoManager::OMPTargetGlobalVarEntryIndirect ||
12567 Flags ==
12568 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryIndirectVTable)
12569 createOffloadEntry(ID: CE->getAddress(), Addr: CE->getAddress(), Size: CE->getVarSize(),
12570 Flags, CE->getLinkage(), Name: CE->getVarName());
12571 else
12572 createOffloadEntry(ID: CE->getAddress(), Addr: CE->getAddress(), Size: CE->getVarSize(),
12573 Flags, CE->getLinkage());
12574
12575 } else {
12576 llvm_unreachable("Unsupported entry kind.");
12577 }
12578 }
12579
12580 // Emit requires directive globals to a special entry so the runtime can
12581 // register them when the device image is loaded.
12582 // TODO: This reduces the offloading entries to a 32-bit integer. Offloading
12583 // entries should be redesigned to better suit this use-case.
12584 if (Config.hasRequiresFlags() && !Config.isTargetDevice())
12585 offloading::emitOffloadingEntry(
12586 M, Kind: object::OffloadKind::OFK_OpenMP,
12587 Addr: Constant::getNullValue(Ty: PointerType::getUnqual(C&: M.getContext())),
12588 Name: ".requires", /*Size=*/0,
12589 Flags: OffloadEntriesInfoManager::OMPTargetGlobalRegisterRequires,
12590 Data: Config.getRequiresFlags());
12591}
12592
12593void TargetRegionEntryInfo::getTargetRegionEntryFnName(
12594 SmallVectorImpl<char> &Name, StringRef ParentName, unsigned DeviceID,
12595 unsigned FileID, unsigned Line, unsigned Count) {
12596 raw_svector_ostream OS(Name);
12597 OS << KernelNamePrefix << llvm::format(Fmt: "%x", Vals: DeviceID)
12598 << llvm::format(Fmt: "_%x_", Vals: FileID) << ParentName << "_l" << Line;
12599 if (Count)
12600 OS << "_" << Count;
12601}
12602
12603void OffloadEntriesInfoManager::getTargetRegionEntryFnName(
12604 SmallVectorImpl<char> &Name, const TargetRegionEntryInfo &EntryInfo) {
12605 unsigned NewCount = getTargetRegionEntryInfoCount(EntryInfo);
12606 TargetRegionEntryInfo::getTargetRegionEntryFnName(
12607 Name, ParentName: EntryInfo.ParentName, DeviceID: EntryInfo.DeviceID, FileID: EntryInfo.FileID,
12608 Line: EntryInfo.Line, Count: NewCount);
12609}
12610
12611TargetRegionEntryInfo
12612OpenMPIRBuilder::getTargetEntryUniqueInfo(FileIdentifierInfoCallbackTy CallBack,
12613 vfs::FileSystem &VFS,
12614 StringRef ParentName) {
12615 sys::fs::UniqueID ID(0xdeadf17e, 0);
12616 auto FileIDInfo = CallBack();
12617 uint64_t FileID = 0;
12618 if (ErrorOr<vfs::Status> Status = VFS.status(Path: std::get<0>(t&: FileIDInfo))) {
12619 ID = Status->getUniqueID();
12620 FileID = Status->getUniqueID().getFile();
12621 } else {
12622 // If the inode ID could not be determined, create a hash value
12623 // the current file name and use that as an ID.
12624 FileID = hash_value(arg: std::get<0>(t&: FileIDInfo));
12625 }
12626
12627 return TargetRegionEntryInfo(ParentName, ID.getDevice(), FileID,
12628 std::get<1>(t&: FileIDInfo));
12629}
12630
12631unsigned OpenMPIRBuilder::getFlagMemberOffset() {
12632 unsigned Offset = 0;
12633 for (uint64_t Remain =
12634 static_cast<std::underlying_type_t<omp::OpenMPOffloadMappingFlags>>(
12635 omp::OpenMPOffloadMappingFlags::OMP_MAP_MEMBER_OF);
12636 !(Remain & 1); Remain = Remain >> 1)
12637 Offset++;
12638 return Offset;
12639}
12640
12641omp::OpenMPOffloadMappingFlags
12642OpenMPIRBuilder::getMemberOfFlag(unsigned Position) {
12643 // Rotate by getFlagMemberOffset() bits.
12644 return static_cast<omp::OpenMPOffloadMappingFlags>(((uint64_t)Position + 1)
12645 << getFlagMemberOffset());
12646}
12647
12648void OpenMPIRBuilder::setCorrectMemberOfFlag(
12649 omp::OpenMPOffloadMappingFlags &Flags,
12650 omp::OpenMPOffloadMappingFlags MemberOfFlag) {
12651 // If the entry is PTR_AND_OBJ but has not been marked with the special
12652 // placeholder value 0xFFFF in the MEMBER_OF field, then it should not be
12653 // marked as MEMBER_OF.
12654 if (static_cast<std::underlying_type_t<omp::OpenMPOffloadMappingFlags>>(
12655 Flags & omp::OpenMPOffloadMappingFlags::OMP_MAP_PTR_AND_OBJ) &&
12656 static_cast<std::underlying_type_t<omp::OpenMPOffloadMappingFlags>>(
12657 (Flags & omp::OpenMPOffloadMappingFlags::OMP_MAP_MEMBER_OF) !=
12658 omp::OpenMPOffloadMappingFlags::OMP_MAP_MEMBER_OF))
12659 return;
12660
12661 // Entries with ATTACH are not members-of anything. They are handled
12662 // separately by the runtime after other maps have been handled.
12663 if (static_cast<std::underlying_type_t<omp::OpenMPOffloadMappingFlags>>(
12664 Flags & omp::OpenMPOffloadMappingFlags::OMP_MAP_ATTACH))
12665 return;
12666
12667 // Reset the placeholder value to prepare the flag for the assignment of the
12668 // proper MEMBER_OF value.
12669 Flags &= ~omp::OpenMPOffloadMappingFlags::OMP_MAP_MEMBER_OF;
12670 Flags |= MemberOfFlag;
12671}
12672
12673Constant *OpenMPIRBuilder::getAddrOfDeclareTargetVar(
12674 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind CaptureClause,
12675 OffloadEntriesInfoManager::OMPTargetDeviceClauseKind DeviceClause,
12676 bool IsDeclaration, bool IsExternallyVisible,
12677 TargetRegionEntryInfo EntryInfo, StringRef MangledName,
12678 std::vector<GlobalVariable *> &GeneratedRefs, bool OpenMPSIMD,
12679 std::vector<Triple> TargetTriple, Type *LlvmPtrTy,
12680 std::function<Constant *()> GlobalInitializer,
12681 std::function<GlobalValue::LinkageTypes()> VariableLinkage) {
12682 // TODO: convert this to utilise the IRBuilder Config rather than
12683 // a passed down argument.
12684 if (OpenMPSIMD)
12685 return nullptr;
12686
12687 if (CaptureClause == OffloadEntriesInfoManager::OMPTargetGlobalVarEntryLink ||
12688 ((CaptureClause == OffloadEntriesInfoManager::OMPTargetGlobalVarEntryTo ||
12689 CaptureClause ==
12690 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryEnter) &&
12691 Config.hasRequiresUnifiedSharedMemory())) {
12692 SmallString<64> PtrName;
12693 {
12694 raw_svector_ostream OS(PtrName);
12695 OS << MangledName;
12696 if (!IsExternallyVisible)
12697 OS << format(Fmt: "_%x", Vals: EntryInfo.FileID);
12698 OS << "_decl_tgt_ref_ptr";
12699 }
12700
12701 Value *Ptr = M.getNamedValue(Name: PtrName);
12702
12703 if (!Ptr) {
12704 GlobalValue *GlobalValue = M.getNamedValue(Name: MangledName);
12705 Ptr = getOrCreateInternalVariable(Ty: LlvmPtrTy, Name: PtrName);
12706
12707 auto *GV = cast<GlobalVariable>(Val: Ptr);
12708 GV->setLinkage(GlobalValue::WeakAnyLinkage);
12709
12710 if (!Config.isTargetDevice()) {
12711 if (GlobalInitializer)
12712 GV->setInitializer(GlobalInitializer());
12713 else
12714 GV->setInitializer(GlobalValue);
12715 }
12716
12717 registerTargetGlobalVariable(
12718 CaptureClause, DeviceClause, IsDeclaration, IsExternallyVisible,
12719 EntryInfo, MangledName, GeneratedRefs, OpenMPSIMD, TargetTriple,
12720 GlobalInitializer, VariableLinkage, LlvmPtrTy, Addr: cast<Constant>(Val: Ptr));
12721 }
12722
12723 return cast<Constant>(Val: Ptr);
12724 }
12725
12726 return nullptr;
12727}
12728
12729void OpenMPIRBuilder::registerTargetGlobalVariable(
12730 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind CaptureClause,
12731 OffloadEntriesInfoManager::OMPTargetDeviceClauseKind DeviceClause,
12732 bool IsDeclaration, bool IsExternallyVisible,
12733 TargetRegionEntryInfo EntryInfo, StringRef MangledName,
12734 std::vector<GlobalVariable *> &GeneratedRefs, bool OpenMPSIMD,
12735 std::vector<Triple> TargetTriple,
12736 std::function<Constant *()> GlobalInitializer,
12737 std::function<GlobalValue::LinkageTypes()> VariableLinkage, Type *LlvmPtrTy,
12738 Constant *Addr) {
12739 if (DeviceClause != OffloadEntriesInfoManager::OMPTargetDeviceClauseAny ||
12740 (TargetTriple.empty() && !Config.isTargetDevice()))
12741 return;
12742
12743 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind Flags;
12744 StringRef VarName;
12745 int64_t VarSize;
12746 GlobalValue::LinkageTypes Linkage;
12747
12748 if ((CaptureClause == OffloadEntriesInfoManager::OMPTargetGlobalVarEntryTo ||
12749 CaptureClause ==
12750 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryEnter) &&
12751 !Config.hasRequiresUnifiedSharedMemory()) {
12752 Flags = OffloadEntriesInfoManager::OMPTargetGlobalVarEntryTo;
12753 VarName = MangledName;
12754 GlobalValue *LlvmVal = M.getNamedValue(Name: VarName);
12755
12756 if (!IsDeclaration)
12757 VarSize = divideCeil(
12758 Numerator: M.getDataLayout().getTypeSizeInBits(Ty: LlvmVal->getValueType()), Denominator: 8);
12759 else
12760 VarSize = 0;
12761 Linkage = (VariableLinkage) ? VariableLinkage() : LlvmVal->getLinkage();
12762
12763 // This is a workaround carried over from Clang which prevents undesired
12764 // optimisation of internal variables.
12765 if (Config.isTargetDevice() &&
12766 (!IsExternallyVisible || Linkage == GlobalValue::LinkOnceODRLinkage)) {
12767 // Do not create a "ref-variable" if the original is not also available
12768 // on the host.
12769 if (!OffloadInfoManager.hasDeviceGlobalVarEntryInfo(VarName))
12770 return;
12771
12772 std::string RefName = createPlatformSpecificName(Parts: {VarName, "ref"});
12773
12774 if (!M.getNamedValue(Name: RefName)) {
12775 Constant *AddrRef =
12776 getOrCreateInternalVariable(Ty: Addr->getType(), Name: RefName);
12777 auto *GvAddrRef = cast<GlobalVariable>(Val: AddrRef);
12778 GvAddrRef->setConstant(true);
12779 GvAddrRef->setLinkage(GlobalValue::InternalLinkage);
12780 GvAddrRef->setInitializer(Addr);
12781 GeneratedRefs.push_back(x: GvAddrRef);
12782 }
12783 }
12784 } else {
12785 if (CaptureClause == OffloadEntriesInfoManager::OMPTargetGlobalVarEntryLink)
12786 Flags = OffloadEntriesInfoManager::OMPTargetGlobalVarEntryLink;
12787 else
12788 Flags = OffloadEntriesInfoManager::OMPTargetGlobalVarEntryTo;
12789
12790 if (Config.isTargetDevice()) {
12791 VarName = (Addr) ? Addr->getName() : "";
12792 Addr = nullptr;
12793 } else {
12794 Addr = getAddrOfDeclareTargetVar(
12795 CaptureClause, DeviceClause, IsDeclaration, IsExternallyVisible,
12796 EntryInfo, MangledName, GeneratedRefs, OpenMPSIMD, TargetTriple,
12797 LlvmPtrTy, GlobalInitializer, VariableLinkage);
12798 VarName = (Addr) ? Addr->getName() : "";
12799 }
12800 VarSize = M.getDataLayout().getPointerSize();
12801 Linkage = GlobalValue::WeakAnyLinkage;
12802 }
12803
12804 OffloadInfoManager.registerDeviceGlobalVarEntryInfo(VarName, Addr, VarSize,
12805 Flags, Linkage);
12806}
12807
12808/// Loads all the offload entries information from the host IR
12809/// metadata.
12810void OpenMPIRBuilder::loadOffloadInfoMetadata(Module &M) {
12811 // If we are in target mode, load the metadata from the host IR. This code has
12812 // to match the metadata creation in createOffloadEntriesAndInfoMetadata().
12813
12814 NamedMDNode *MD = M.getNamedMetadata(Name: ompOffloadInfoName);
12815 if (!MD)
12816 return;
12817
12818 for (MDNode *MN : MD->operands()) {
12819 auto &&GetMDInt = [MN](unsigned Idx) {
12820 auto *V = cast<ConstantAsMetadata>(Val: MN->getOperand(I: Idx));
12821 return cast<ConstantInt>(Val: V->getValue())->getZExtValue();
12822 };
12823
12824 auto &&GetMDString = [MN](unsigned Idx) {
12825 auto *V = cast<MDString>(Val: MN->getOperand(I: Idx));
12826 return V->getString();
12827 };
12828
12829 switch (GetMDInt(0)) {
12830 default:
12831 llvm_unreachable("Unexpected metadata!");
12832 break;
12833 case OffloadEntriesInfoManager::OffloadEntryInfo::
12834 OffloadingEntryInfoTargetRegion: {
12835 TargetRegionEntryInfo EntryInfo(/*ParentName=*/GetMDString(3),
12836 /*DeviceID=*/GetMDInt(1),
12837 /*FileID=*/GetMDInt(2),
12838 /*Line=*/GetMDInt(4),
12839 /*Count=*/GetMDInt(5));
12840 OffloadInfoManager.initializeTargetRegionEntryInfo(EntryInfo,
12841 /*Order=*/GetMDInt(6));
12842 break;
12843 }
12844 case OffloadEntriesInfoManager::OffloadEntryInfo::
12845 OffloadingEntryInfoDeviceGlobalVar:
12846 OffloadInfoManager.initializeDeviceGlobalVarEntryInfo(
12847 /*MangledName=*/Name: GetMDString(1),
12848 Flags: static_cast<OffloadEntriesInfoManager::OMPTargetGlobalVarEntryKind>(
12849 /*Flags=*/GetMDInt(2)),
12850 /*Order=*/GetMDInt(3));
12851 break;
12852 }
12853 }
12854}
12855
12856void OpenMPIRBuilder::loadOffloadInfoMetadata(vfs::FileSystem &VFS,
12857 StringRef HostFilePath) {
12858 if (HostFilePath.empty())
12859 return;
12860
12861 auto Buf = VFS.getBufferForFile(Name: HostFilePath);
12862 if (std::error_code Err = Buf.getError()) {
12863 report_fatal_error(reason: ("error opening host file from host file path inside of "
12864 "OpenMPIRBuilder: " +
12865 Err.message())
12866 .c_str());
12867 }
12868
12869 LLVMContext Ctx;
12870 auto M = expectedToErrorOrAndEmitErrors(
12871 Ctx, Val: parseBitcodeFile(Buffer: Buf.get()->getMemBufferRef(), Context&: Ctx));
12872 if (std::error_code Err = M.getError()) {
12873 report_fatal_error(
12874 reason: ("error parsing host file inside of OpenMPIRBuilder: " + Err.message())
12875 .c_str());
12876 }
12877
12878 loadOffloadInfoMetadata(M&: *M.get());
12879}
12880
12881OpenMPIRBuilder::InsertPointOrErrorTy OpenMPIRBuilder::createIteratorLoop(
12882 LocationDescription Loc, llvm::Value *TripCount, IteratorBodyGenTy BodyGen,
12883 llvm::StringRef Name) {
12884 Builder.restoreIP(IP: Loc.IP);
12885
12886 BasicBlock *CurBB = Builder.GetInsertBlock();
12887 assert(CurBB &&
12888 "expected a valid insertion block for creating an iterator loop");
12889 Function *F = CurBB->getParent();
12890
12891 InsertPointTy SplitIP = Builder.saveIP();
12892 if (SplitIP == CurBB->end())
12893 if (Instruction *Terminator = CurBB->getTerminatorOrNull())
12894 SplitIP = Terminator->getIterator();
12895
12896 BasicBlock *ContBB =
12897 splitBB(IP: SplitIP, /*CreateBranch=*/false,
12898 DL: Builder.getCurrentDebugLocation(), Name: "omp.it.cont");
12899
12900 CanonicalLoopInfo *CLI =
12901 createLoopSkeleton(DL: Builder.getCurrentDebugLocation(), TripCount, F,
12902 /*PreInsertBefore=*/ContBB,
12903 /*PostInsertBefore=*/ContBB, Name);
12904
12905 // Enter loop from original block.
12906 redirectTo(Source: CurBB, Target: CLI->getPreheader(), DL: Builder.getCurrentDebugLocation());
12907
12908 // Remove the unconditional branch inserted by createLoopSkeleton in the body
12909 if (Instruction *T = CLI->getBody()->getTerminatorOrNull())
12910 T->eraseFromParent();
12911
12912 InsertPointTy BodyIP = CLI->getBodyIP();
12913 // Blocks that exist before BodyGen, other than the body and latch, are
12914 // outside the loop body.
12915 SmallPtrSet<BasicBlock *, 32> ExistingBlocks;
12916 for (BasicBlock &Block : *F)
12917 ExistingBlocks.insert(Ptr: &Block);
12918 if (llvm::Error Err = BodyGen(BodyIP, CLI->getIndVar()))
12919 return Err;
12920
12921 // The body may span several blocks. Branch its single unterminated block to
12922 // the latch; otherwise some block must already branch there.
12923 BasicBlock *Latch = CLI->getLatch();
12924 BasicBlock *OpenBB = nullptr;
12925 bool ReachesLatch = false;
12926 SmallVector<BasicBlock *> Worklist{CLI->getBody()};
12927 SmallPtrSet<BasicBlock *, 8> Visited{CLI->getBody()};
12928 while (!Worklist.empty()) {
12929 BasicBlock *BB = Worklist.pop_back_val();
12930 if (!BB->hasTerminator()) {
12931 if (OpenBB)
12932 return make_error<StringError>(
12933 Args: "iterator bodygen must leave at most one unterminated block",
12934 Args: inconvertibleErrorCode());
12935 OpenBB = BB;
12936 continue;
12937 }
12938 for (BasicBlock *Succ : successors(BB)) {
12939 if (Succ == Latch) {
12940 ReachesLatch = true;
12941 continue;
12942 }
12943 if (Succ != CLI->getBody() && ExistingBlocks.contains(Ptr: Succ))
12944 return make_error<StringError>(
12945 Args: "iterator bodygen must not branch out of the loop body",
12946 Args: inconvertibleErrorCode());
12947 if (Visited.insert(Ptr: Succ).second)
12948 Worklist.push_back(Elt: Succ);
12949 }
12950 }
12951
12952 if (OpenBB) {
12953 Builder.SetInsertPoint(OpenBB);
12954 Builder.CreateBr(Dest: Latch);
12955 } else if (!ReachesLatch) {
12956 return make_error<StringError>(Args: "iterator bodygen must reach the loop latch",
12957 Args: inconvertibleErrorCode());
12958 }
12959
12960 // Link After -> ContBB
12961 Builder.SetInsertPoint(CLI->getAfter()->begin());
12962 if (!CLI->getAfter()->hasTerminator())
12963 Builder.CreateBr(Dest: ContBB);
12964
12965 return ContBB->begin();
12966}
12967
12968/// Mangle the parameter part of the vector function name according to
12969/// their OpenMP classification. The mangling function is defined in
12970/// section 4.5 of the AAVFABI(2021Q1).
12971static std::string mangleVectorParameters(
12972 ArrayRef<llvm::OpenMPIRBuilder::DeclareSimdAttrTy> ParamAttrs) {
12973 SmallString<256> Buffer;
12974 llvm::raw_svector_ostream Out(Buffer);
12975 for (const auto &ParamAttr : ParamAttrs) {
12976 switch (ParamAttr.Kind) {
12977 case llvm::OpenMPIRBuilder::DeclareSimdKindTy::Linear:
12978 Out << 'l';
12979 break;
12980 case llvm::OpenMPIRBuilder::DeclareSimdKindTy::LinearRef:
12981 Out << 'R';
12982 break;
12983 case llvm::OpenMPIRBuilder::DeclareSimdKindTy::LinearUVal:
12984 Out << 'U';
12985 break;
12986 case llvm::OpenMPIRBuilder::DeclareSimdKindTy::LinearVal:
12987 Out << 'L';
12988 break;
12989 case llvm::OpenMPIRBuilder::DeclareSimdKindTy::Uniform:
12990 Out << 'u';
12991 break;
12992 case llvm::OpenMPIRBuilder::DeclareSimdKindTy::Vector:
12993 Out << 'v';
12994 break;
12995 }
12996 if (ParamAttr.HasVarStride)
12997 Out << "s" << ParamAttr.StrideOrArg;
12998 else if (ParamAttr.Kind ==
12999 llvm::OpenMPIRBuilder::DeclareSimdKindTy::Linear ||
13000 ParamAttr.Kind ==
13001 llvm::OpenMPIRBuilder::DeclareSimdKindTy::LinearRef ||
13002 ParamAttr.Kind ==
13003 llvm::OpenMPIRBuilder::DeclareSimdKindTy::LinearUVal ||
13004 ParamAttr.Kind ==
13005 llvm::OpenMPIRBuilder::DeclareSimdKindTy::LinearVal) {
13006 // Don't print the step value if it is not present or if it is
13007 // equal to 1.
13008 if (ParamAttr.StrideOrArg < 0)
13009 Out << 'n' << -ParamAttr.StrideOrArg;
13010 else if (ParamAttr.StrideOrArg != 1)
13011 Out << ParamAttr.StrideOrArg;
13012 }
13013
13014 if (!!ParamAttr.Alignment)
13015 Out << 'a' << ParamAttr.Alignment;
13016 }
13017
13018 return std::string(Out.str());
13019}
13020
13021void OpenMPIRBuilder::emitX86DeclareSimdFunction(
13022 llvm::Function *Fn, unsigned NumElts, const llvm::APSInt &VLENVal,
13023 llvm::ArrayRef<DeclareSimdAttrTy> ParamAttrs, DeclareSimdBranch Branch) {
13024 struct ISADataTy {
13025 char ISA;
13026 unsigned VecRegSize;
13027 };
13028 ISADataTy ISAData[] = {
13029 {.ISA: 'b', .VecRegSize: 128}, // SSE
13030 {.ISA: 'c', .VecRegSize: 256}, // AVX
13031 {.ISA: 'd', .VecRegSize: 256}, // AVX2
13032 {.ISA: 'e', .VecRegSize: 512}, // AVX512
13033 };
13034 llvm::SmallVector<char, 2> Masked;
13035 switch (Branch) {
13036 case DeclareSimdBranch::Undefined:
13037 Masked.push_back(Elt: 'N');
13038 Masked.push_back(Elt: 'M');
13039 break;
13040 case DeclareSimdBranch::Notinbranch:
13041 Masked.push_back(Elt: 'N');
13042 break;
13043 case DeclareSimdBranch::Inbranch:
13044 Masked.push_back(Elt: 'M');
13045 break;
13046 }
13047 for (char Mask : Masked) {
13048 for (const ISADataTy &Data : ISAData) {
13049 llvm::SmallString<256> Buffer;
13050 llvm::raw_svector_ostream Out(Buffer);
13051 Out << "_ZGV" << Data.ISA << Mask;
13052 if (!VLENVal) {
13053 assert(NumElts && "Non-zero simdlen/cdtsize expected");
13054 Out << llvm::APSInt::getUnsigned(X: Data.VecRegSize / NumElts);
13055 } else {
13056 Out << VLENVal;
13057 }
13058 Out << mangleVectorParameters(ParamAttrs);
13059 Out << '_' << Fn->getName();
13060 Fn->addFnAttr(Kind: Out.str());
13061 }
13062 }
13063}
13064
13065// Function used to add the attribute. The parameter `VLEN` is templated to
13066// allow the use of `x` when targeting scalable functions for SVE.
13067template <typename T>
13068static void addAArch64VectorName(T VLEN, StringRef LMask, StringRef Prefix,
13069 char ISA, StringRef ParSeq,
13070 StringRef MangledName, bool OutputBecomesInput,
13071 llvm::Function *Fn) {
13072 SmallString<256> Buffer;
13073 llvm::raw_svector_ostream Out(Buffer);
13074 Out << Prefix << ISA << LMask << VLEN;
13075 if (OutputBecomesInput)
13076 Out << 'v';
13077 Out << ParSeq << '_' << MangledName;
13078 Fn->addFnAttr(Kind: Out.str());
13079}
13080
13081// Helper function to generate the Advanced SIMD names depending on the value
13082// of the NDS when simdlen is not present.
13083static void addAArch64AdvSIMDNDSNames(unsigned NDS, StringRef Mask,
13084 StringRef Prefix, char ISA,
13085 StringRef ParSeq, StringRef MangledName,
13086 bool OutputBecomesInput,
13087 llvm::Function *Fn) {
13088 switch (NDS) {
13089 case 8:
13090 addAArch64VectorName(VLEN: 8, LMask: Mask, Prefix, ISA, ParSeq, MangledName,
13091 OutputBecomesInput, Fn);
13092 addAArch64VectorName(VLEN: 16, LMask: Mask, Prefix, ISA, ParSeq, MangledName,
13093 OutputBecomesInput, Fn);
13094 break;
13095 case 16:
13096 addAArch64VectorName(VLEN: 4, LMask: Mask, Prefix, ISA, ParSeq, MangledName,
13097 OutputBecomesInput, Fn);
13098 addAArch64VectorName(VLEN: 8, LMask: Mask, Prefix, ISA, ParSeq, MangledName,
13099 OutputBecomesInput, Fn);
13100 break;
13101 case 32:
13102 addAArch64VectorName(VLEN: 2, LMask: Mask, Prefix, ISA, ParSeq, MangledName,
13103 OutputBecomesInput, Fn);
13104 addAArch64VectorName(VLEN: 4, LMask: Mask, Prefix, ISA, ParSeq, MangledName,
13105 OutputBecomesInput, Fn);
13106 break;
13107 case 64:
13108 case 128:
13109 addAArch64VectorName(VLEN: 2, LMask: Mask, Prefix, ISA, ParSeq, MangledName,
13110 OutputBecomesInput, Fn);
13111 break;
13112 default:
13113 llvm_unreachable("Scalar type is too wide.");
13114 }
13115}
13116
13117/// Emit vector function attributes for AArch64, as defined in the AAVFABI.
13118void OpenMPIRBuilder::emitAArch64DeclareSimdFunction(
13119 llvm::Function *Fn, unsigned UserVLEN,
13120 llvm::ArrayRef<DeclareSimdAttrTy> ParamAttrs, DeclareSimdBranch Branch,
13121 char ISA, unsigned NarrowestDataSize, bool OutputBecomesInput) {
13122 assert((ISA == 'n' || ISA == 's') && "Expected ISA either 's' or 'n'.");
13123
13124 // Sort out parameter sequence.
13125 const std::string ParSeq = mangleVectorParameters(ParamAttrs);
13126 StringRef Prefix = "_ZGV";
13127 StringRef MangledName = Fn->getName();
13128
13129 // Generate simdlen from user input (if any).
13130 if (UserVLEN) {
13131 if (ISA == 's') {
13132 // SVE generates only a masked function.
13133 addAArch64VectorName(VLEN: UserVLEN, LMask: "M", Prefix, ISA, ParSeq, MangledName,
13134 OutputBecomesInput, Fn);
13135 return;
13136 }
13137
13138 switch (Branch) {
13139 case DeclareSimdBranch::Undefined:
13140 addAArch64VectorName(VLEN: UserVLEN, LMask: "N", Prefix, ISA, ParSeq, MangledName,
13141 OutputBecomesInput, Fn);
13142 addAArch64VectorName(VLEN: UserVLEN, LMask: "M", Prefix, ISA, ParSeq, MangledName,
13143 OutputBecomesInput, Fn);
13144 break;
13145 case DeclareSimdBranch::Inbranch:
13146 addAArch64VectorName(VLEN: UserVLEN, LMask: "M", Prefix, ISA, ParSeq, MangledName,
13147 OutputBecomesInput, Fn);
13148 break;
13149 case DeclareSimdBranch::Notinbranch:
13150 addAArch64VectorName(VLEN: UserVLEN, LMask: "N", Prefix, ISA, ParSeq, MangledName,
13151 OutputBecomesInput, Fn);
13152 break;
13153 }
13154 return;
13155 }
13156
13157 if (ISA == 's') {
13158 // SVE, section 3.4.1, item 1.
13159 addAArch64VectorName(VLEN: "x", LMask: "M", Prefix, ISA, ParSeq, MangledName,
13160 OutputBecomesInput, Fn);
13161 return;
13162 }
13163
13164 switch (Branch) {
13165 case DeclareSimdBranch::Undefined:
13166 addAArch64AdvSIMDNDSNames(NDS: NarrowestDataSize, Mask: "N", Prefix, ISA, ParSeq,
13167 MangledName, OutputBecomesInput, Fn);
13168 addAArch64AdvSIMDNDSNames(NDS: NarrowestDataSize, Mask: "M", Prefix, ISA, ParSeq,
13169 MangledName, OutputBecomesInput, Fn);
13170 break;
13171 case DeclareSimdBranch::Inbranch:
13172 addAArch64AdvSIMDNDSNames(NDS: NarrowestDataSize, Mask: "M", Prefix, ISA, ParSeq,
13173 MangledName, OutputBecomesInput, Fn);
13174 break;
13175 case DeclareSimdBranch::Notinbranch:
13176 addAArch64AdvSIMDNDSNames(NDS: NarrowestDataSize, Mask: "N", Prefix, ISA, ParSeq,
13177 MangledName, OutputBecomesInput, Fn);
13178 break;
13179 }
13180}
13181
13182//===----------------------------------------------------------------------===//
13183// OffloadEntriesInfoManager
13184//===----------------------------------------------------------------------===//
13185
13186bool OffloadEntriesInfoManager::empty() const {
13187 return OffloadEntriesTargetRegion.empty() &&
13188 OffloadEntriesDeviceGlobalVar.empty();
13189}
13190
13191unsigned OffloadEntriesInfoManager::getTargetRegionEntryInfoCount(
13192 const TargetRegionEntryInfo &EntryInfo) const {
13193 auto It = OffloadEntriesTargetRegionCount.find(
13194 x: getTargetRegionEntryCountKey(EntryInfo));
13195 if (It == OffloadEntriesTargetRegionCount.end())
13196 return 0;
13197 return It->second;
13198}
13199
13200void OffloadEntriesInfoManager::incrementTargetRegionEntryInfoCount(
13201 const TargetRegionEntryInfo &EntryInfo) {
13202 OffloadEntriesTargetRegionCount[getTargetRegionEntryCountKey(EntryInfo)] =
13203 EntryInfo.Count + 1;
13204}
13205
13206/// Initialize target region entry.
13207void OffloadEntriesInfoManager::initializeTargetRegionEntryInfo(
13208 const TargetRegionEntryInfo &EntryInfo, unsigned Order) {
13209 OffloadEntriesTargetRegion[EntryInfo] =
13210 OffloadEntryInfoTargetRegion(Order, /*Addr=*/nullptr, /*ID=*/nullptr,
13211 OMPTargetRegionEntryTargetRegion);
13212 ++OffloadingEntriesNum;
13213}
13214
13215void OffloadEntriesInfoManager::registerTargetRegionEntryInfo(
13216 TargetRegionEntryInfo EntryInfo, Constant *Addr, Constant *ID,
13217 OMPTargetRegionEntryKind Flags) {
13218 assert(EntryInfo.Count == 0 && "expected default EntryInfo");
13219
13220 // Update the EntryInfo with the next available count for this location.
13221 EntryInfo.Count = getTargetRegionEntryInfoCount(EntryInfo);
13222
13223 // If we are emitting code for a target, the entry is already initialized,
13224 // only has to be registered.
13225 if (OMPBuilder->Config.isTargetDevice()) {
13226 // This could happen if the device compilation is invoked standalone.
13227 if (!hasTargetRegionEntryInfo(EntryInfo)) {
13228 return;
13229 }
13230 auto &Entry = OffloadEntriesTargetRegion[EntryInfo];
13231 Entry.setAddress(Addr);
13232 Entry.setID(ID);
13233 Entry.setFlags(Flags);
13234 } else {
13235 if (Flags == OffloadEntriesInfoManager::OMPTargetRegionEntryTargetRegion &&
13236 hasTargetRegionEntryInfo(EntryInfo, /*IgnoreAddressId*/ true))
13237 return;
13238 assert(!hasTargetRegionEntryInfo(EntryInfo) &&
13239 "Target region entry already registered!");
13240 OffloadEntryInfoTargetRegion Entry(OffloadingEntriesNum, Addr, ID, Flags);
13241 OffloadEntriesTargetRegion[EntryInfo] = Entry;
13242 ++OffloadingEntriesNum;
13243 }
13244 incrementTargetRegionEntryInfoCount(EntryInfo);
13245}
13246
13247bool OffloadEntriesInfoManager::hasTargetRegionEntryInfo(
13248 TargetRegionEntryInfo EntryInfo, bool IgnoreAddressId) const {
13249
13250 // Update the EntryInfo with the next available count for this location.
13251 EntryInfo.Count = getTargetRegionEntryInfoCount(EntryInfo);
13252
13253 auto It = OffloadEntriesTargetRegion.find(x: EntryInfo);
13254 if (It == OffloadEntriesTargetRegion.end()) {
13255 return false;
13256 }
13257 // Fail if this entry is already registered.
13258 if (!IgnoreAddressId && (It->second.getAddress() || It->second.getID()))
13259 return false;
13260 return true;
13261}
13262
13263void OffloadEntriesInfoManager::actOnTargetRegionEntriesInfo(
13264 const OffloadTargetRegionEntryInfoActTy &Action) {
13265 // Scan all target region entries and perform the provided action.
13266 for (const auto &It : OffloadEntriesTargetRegion) {
13267 Action(It.first, It.second);
13268 }
13269}
13270
13271void OffloadEntriesInfoManager::initializeDeviceGlobalVarEntryInfo(
13272 StringRef Name, OMPTargetGlobalVarEntryKind Flags, unsigned Order) {
13273 OffloadEntriesDeviceGlobalVar.try_emplace(Key: Name, Args&: Order, Args&: Flags);
13274 ++OffloadingEntriesNum;
13275}
13276
13277void OffloadEntriesInfoManager::registerDeviceGlobalVarEntryInfo(
13278 StringRef VarName, Constant *Addr, int64_t VarSize,
13279 OMPTargetGlobalVarEntryKind Flags, GlobalValue::LinkageTypes Linkage) {
13280 if (OMPBuilder->Config.isTargetDevice()) {
13281 // This could happen if the device compilation is invoked standalone.
13282 if (!hasDeviceGlobalVarEntryInfo(VarName))
13283 return;
13284 auto &Entry = OffloadEntriesDeviceGlobalVar[VarName];
13285 if (Entry.getAddress() && hasDeviceGlobalVarEntryInfo(VarName)) {
13286 if (Entry.getVarSize() == 0) {
13287 Entry.setVarSize(VarSize);
13288 Entry.setLinkage(Linkage);
13289 }
13290 return;
13291 }
13292 Entry.setVarSize(VarSize);
13293 Entry.setLinkage(Linkage);
13294 Entry.setAddress(Addr);
13295 } else {
13296 if (hasDeviceGlobalVarEntryInfo(VarName)) {
13297 auto &Entry = OffloadEntriesDeviceGlobalVar[VarName];
13298 assert(Entry.isValid() && Entry.getFlags() == Flags &&
13299 "Entry not initialized!");
13300 if (Entry.getVarSize() == 0) {
13301 Entry.setVarSize(VarSize);
13302 Entry.setLinkage(Linkage);
13303 }
13304 return;
13305 }
13306 if (Flags == OffloadEntriesInfoManager::OMPTargetGlobalVarEntryIndirect ||
13307 Flags ==
13308 OffloadEntriesInfoManager::OMPTargetGlobalVarEntryIndirectVTable)
13309 OffloadEntriesDeviceGlobalVar.try_emplace(Key: VarName, Args&: OffloadingEntriesNum,
13310 Args&: Addr, Args&: VarSize, Args&: Flags, Args&: Linkage,
13311 Args: VarName.str());
13312 else
13313 OffloadEntriesDeviceGlobalVar.try_emplace(
13314 Key: VarName, Args&: OffloadingEntriesNum, Args&: Addr, Args&: VarSize, Args&: Flags, Args&: Linkage, Args: "");
13315 ++OffloadingEntriesNum;
13316 }
13317}
13318
13319void OffloadEntriesInfoManager::actOnDeviceGlobalVarEntriesInfo(
13320 const OffloadDeviceGlobalVarEntryInfoActTy &Action) {
13321 // Scan all target region entries and perform the provided action.
13322 for (const auto &E : OffloadEntriesDeviceGlobalVar)
13323 Action(E.getKey(), E.getValue());
13324}
13325
13326//===----------------------------------------------------------------------===//
13327// CanonicalLoopInfo
13328//===----------------------------------------------------------------------===//
13329
13330void CanonicalLoopInfo::collectControlBlocks(
13331 SmallVectorImpl<BasicBlock *> &BBs) {
13332 // We only count those BBs as control block for which we do not need to
13333 // reverse the CFG, i.e. not the loop body which can contain arbitrary control
13334 // flow. For consistency, this also means we do not add the Body block, which
13335 // is just the entry to the body code.
13336 BBs.reserve(N: BBs.size() + 6);
13337 BBs.append(IL: {getPreheader(), Header, Cond, Latch, Exit, getAfter()});
13338}
13339
13340BasicBlock *CanonicalLoopInfo::getPreheader() const {
13341 assert(isValid() && "Requires a valid canonical loop");
13342 for (BasicBlock *Pred : predecessors(BB: Header)) {
13343 if (Pred != Latch)
13344 return Pred;
13345 }
13346 llvm_unreachable("Missing preheader");
13347}
13348
13349void CanonicalLoopInfo::setTripCount(Value *TripCount) {
13350 assert(isValid() && "Requires a valid canonical loop");
13351
13352 Instruction *CmpI = &getCond()->front();
13353 assert(isa<CmpInst>(CmpI) && "First inst must compare IV with TripCount");
13354 CmpI->setOperand(i: 1, Val: TripCount);
13355
13356#ifndef NDEBUG
13357 assertOK();
13358#endif
13359}
13360
13361void CanonicalLoopInfo::mapIndVar(
13362 llvm::function_ref<Value *(Instruction *)> Updater) {
13363 assert(isValid() && "Requires a valid canonical loop");
13364
13365 Instruction *OldIV = getIndVar();
13366
13367 // Record all uses excluding those introduced by the updater. Uses by the
13368 // CanonicalLoopInfo itself to keep track of the number of iterations are
13369 // excluded.
13370 SmallVector<Use *> ReplacableUses;
13371 for (Use &U : OldIV->uses()) {
13372 auto *User = dyn_cast<Instruction>(Val: U.getUser());
13373 if (!User)
13374 continue;
13375 if (User->getParent() == getCond())
13376 continue;
13377 if (User->getParent() == getLatch())
13378 continue;
13379 ReplacableUses.push_back(Elt: &U);
13380 }
13381
13382 // Run the updater that may introduce new uses
13383 Value *NewIV = Updater(OldIV);
13384
13385 // Replace the old uses with the value returned by the updater.
13386 for (Use *U : ReplacableUses)
13387 U->set(NewIV);
13388
13389#ifndef NDEBUG
13390 assertOK();
13391#endif
13392}
13393
13394void CanonicalLoopInfo::assertOK() const {
13395#ifndef NDEBUG
13396 // No constraints if this object currently does not describe a loop.
13397 if (!isValid())
13398 return;
13399
13400 BasicBlock *Preheader = getPreheader();
13401 BasicBlock *Body = getBody();
13402 BasicBlock *After = getAfter();
13403
13404 // Verify standard control-flow we use for OpenMP loops.
13405 assert(Preheader);
13406 assert(isa<UncondBrInst>(Preheader->getTerminator()) &&
13407 "Preheader must terminate with unconditional branch");
13408 assert(Preheader->getSingleSuccessor() == Header &&
13409 "Preheader must jump to header");
13410
13411 assert(Header);
13412 assert(isa<UncondBrInst>(Header->getTerminator()) &&
13413 "Header must terminate with unconditional branch");
13414 assert(Header->getSingleSuccessor() == Cond &&
13415 "Header must jump to exiting block");
13416
13417 assert(Cond);
13418 assert(Cond->getSinglePredecessor() == Header &&
13419 "Exiting block only reachable from header");
13420
13421 assert(isa<CondBrInst>(Cond->getTerminator()) &&
13422 "Exiting block must terminate with conditional branch");
13423 assert(cast<CondBrInst>(Cond->getTerminator())->getSuccessor(0) == Body &&
13424 "Exiting block's first successor jump to the body");
13425 assert(cast<CondBrInst>(Cond->getTerminator())->getSuccessor(1) == Exit &&
13426 "Exiting block's second successor must exit the loop");
13427
13428 assert(Body);
13429 assert(Body->getSinglePredecessor() == Cond &&
13430 "Body only reachable from exiting block");
13431 assert(!isa<PHINode>(Body->front()));
13432
13433 assert(Latch);
13434 assert(isa<UncondBrInst>(Latch->getTerminator()) &&
13435 "Latch must terminate with unconditional branch");
13436 assert(Latch->getSingleSuccessor() == Header && "Latch must jump to header");
13437 // TODO: To support simple redirecting of the end of the body code that has
13438 // multiple; introduce another auxiliary basic block like preheader and after.
13439 assert(Latch->getSinglePredecessor() != nullptr);
13440 assert(!isa<PHINode>(Latch->front()));
13441
13442 assert(Exit);
13443 assert(isa<UncondBrInst>(Exit->getTerminator()) &&
13444 "Exit block must terminate with unconditional branch");
13445 assert(Exit->getSingleSuccessor() == After &&
13446 "Exit block must jump to after block");
13447
13448 assert(After);
13449 assert(After->getSinglePredecessor() == Exit &&
13450 "After block only reachable from exit block");
13451 assert(After->empty() || !isa<PHINode>(After->front()));
13452
13453 Instruction *IndVar = getIndVar();
13454 assert(IndVar && "Canonical induction variable not found?");
13455 assert(isa<IntegerType>(IndVar->getType()) &&
13456 "Induction variable must be an integer");
13457 assert(cast<PHINode>(IndVar)->getParent() == Header &&
13458 "Induction variable must be a PHI in the loop header");
13459 assert(cast<PHINode>(IndVar)->getIncomingBlock(0) == Preheader);
13460 assert(
13461 cast<ConstantInt>(cast<PHINode>(IndVar)->getIncomingValue(0))->isZero());
13462 assert(cast<PHINode>(IndVar)->getIncomingBlock(1) == Latch);
13463
13464 auto *NextIndVar = cast<PHINode>(IndVar)->getIncomingValue(1);
13465 assert(cast<Instruction>(NextIndVar)->getParent() == Latch);
13466 assert(cast<BinaryOperator>(NextIndVar)->getOpcode() == BinaryOperator::Add);
13467 assert(cast<BinaryOperator>(NextIndVar)->getOperand(0) == IndVar);
13468 assert(cast<ConstantInt>(cast<BinaryOperator>(NextIndVar)->getOperand(1))
13469 ->isOne());
13470
13471 Value *TripCount = getTripCount();
13472 assert(TripCount && "Loop trip count not found?");
13473 assert(IndVar->getType() == TripCount->getType() &&
13474 "Trip count and induction variable must have the same type");
13475
13476 auto *CmpI = cast<CmpInst>(&Cond->front());
13477 assert(CmpI->getPredicate() == CmpInst::ICMP_ULT &&
13478 "Exit condition must be a signed less-than comparison");
13479 assert(CmpI->getOperand(0) == IndVar &&
13480 "Exit condition must compare the induction variable");
13481 assert(CmpI->getOperand(1) == TripCount &&
13482 "Exit condition must compare with the trip count");
13483#endif
13484}
13485
13486void CanonicalLoopInfo::invalidate() {
13487 Header = nullptr;
13488 Cond = nullptr;
13489 Latch = nullptr;
13490 Exit = nullptr;
13491}
13492