1//===-- IPO/OpenMPOpt.cpp - Collection of OpenMP specific optimizations ---===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// OpenMP specific optimizations:
10//
11// - Deduplication of runtime calls, e.g., omp_get_thread_num.
12// - Replacing globalized device memory with stack memory.
13// - Replacing globalized device memory with shared memory.
14// - Parallel region merging.
15// - Transforming generic-mode device kernels to SPMD mode.
16// - Specializing the state machine for generic-mode device kernels.
17//
18//===----------------------------------------------------------------------===//
19
20#include "llvm/Transforms/IPO/OpenMPOpt.h"
21
22#include "llvm/ADT/DenseSet.h"
23#include "llvm/ADT/EnumeratedArray.h"
24#include "llvm/ADT/PostOrderIterator.h"
25#include "llvm/ADT/SetVector.h"
26#include "llvm/ADT/SmallPtrSet.h"
27#include "llvm/ADT/SmallVector.h"
28#include "llvm/ADT/Statistic.h"
29#include "llvm/ADT/StringExtras.h"
30#include "llvm/ADT/StringRef.h"
31#include "llvm/Analysis/CallGraph.h"
32#include "llvm/Analysis/MemoryBuiltins.h"
33#include "llvm/Analysis/MemoryLocation.h"
34#include "llvm/Analysis/OptimizationRemarkEmitter.h"
35#include "llvm/Analysis/ValueTracking.h"
36#include "llvm/Frontend/OpenMP/OMPConstants.h"
37#include "llvm/Frontend/OpenMP/OMPDeviceConstants.h"
38#include "llvm/Frontend/OpenMP/OMPIRBuilder.h"
39#include "llvm/IR/Assumptions.h"
40#include "llvm/IR/BasicBlock.h"
41#include "llvm/IR/Constants.h"
42#include "llvm/IR/DiagnosticInfo.h"
43#include "llvm/IR/Dominators.h"
44#include "llvm/IR/Function.h"
45#include "llvm/IR/GlobalValue.h"
46#include "llvm/IR/GlobalVariable.h"
47#include "llvm/IR/InstrTypes.h"
48#include "llvm/IR/Instruction.h"
49#include "llvm/IR/Instructions.h"
50#include "llvm/IR/IntrinsicInst.h"
51#include "llvm/IR/IntrinsicsAMDGPU.h"
52#include "llvm/IR/IntrinsicsNVPTX.h"
53#include "llvm/IR/LLVMContext.h"
54#include "llvm/IR/MDBuilder.h"
55#include "llvm/IR/ProfDataUtils.h"
56#include "llvm/IR/ProfileSummary.h"
57#include "llvm/Support/Casting.h"
58#include "llvm/Support/CommandLine.h"
59#include "llvm/Support/Debug.h"
60#include "llvm/Transforms/IPO/Attributor.h"
61#include "llvm/Transforms/Utils/BasicBlockUtils.h"
62#include "llvm/Transforms/Utils/CallGraphUpdater.h"
63
64#include <algorithm>
65#include <memory>
66#include <optional>
67#include <string>
68
69using namespace llvm;
70using namespace omp;
71
72#define DEBUG_TYPE "openmp-opt"
73
74static cl::opt<bool> DisableOpenMPOptimizations(
75 "openmp-opt-disable", cl::desc("Disable OpenMP specific optimizations."),
76 cl::Hidden, cl::init(Val: false));
77
78static cl::opt<bool> EnableParallelRegionMerging(
79 "openmp-opt-enable-merging",
80 cl::desc("Enable the OpenMP region merging optimization."), cl::Hidden,
81 cl::init(Val: false));
82
83static cl::opt<bool>
84 DisableInternalization("openmp-opt-disable-internalization",
85 cl::desc("Disable function internalization."),
86 cl::Hidden, cl::init(Val: false));
87
88static cl::opt<bool> DeduceICVValues("openmp-deduce-icv-values",
89 cl::init(Val: false), cl::Hidden);
90static cl::opt<bool> PrintICVValues("openmp-print-icv-values", cl::init(Val: false),
91 cl::Hidden);
92static cl::opt<bool> PrintOpenMPKernels("openmp-print-gpu-kernels",
93 cl::init(Val: false), cl::Hidden);
94
95static cl::opt<bool> HideMemoryTransferLatency(
96 "openmp-hide-memory-transfer-latency",
97 cl::desc("[WIP] Tries to hide the latency of host to device memory"
98 " transfers"),
99 cl::Hidden, cl::init(Val: false));
100
101static cl::opt<bool> DisableOpenMPOptDeglobalization(
102 "openmp-opt-disable-deglobalization",
103 cl::desc("Disable OpenMP optimizations involving deglobalization."),
104 cl::Hidden, cl::init(Val: false));
105
106static cl::opt<bool> DisableOpenMPOptSPMDization(
107 "openmp-opt-disable-spmdization",
108 cl::desc("Disable OpenMP optimizations involving SPMD-ization."),
109 cl::Hidden, cl::init(Val: false));
110
111static cl::opt<bool> DisableOpenMPOptFolding(
112 "openmp-opt-disable-folding",
113 cl::desc("Disable OpenMP optimizations involving folding."), cl::Hidden,
114 cl::init(Val: false));
115
116static cl::opt<bool> DisableOpenMPOptStateMachineRewrite(
117 "openmp-opt-disable-state-machine-rewrite",
118 cl::desc("Disable OpenMP optimizations that replace the state machine."),
119 cl::Hidden, cl::init(Val: false));
120
121static cl::opt<bool> DisableOpenMPOptBarrierElimination(
122 "openmp-opt-disable-barrier-elimination",
123 cl::desc("Disable OpenMP optimizations that eliminate barriers."),
124 cl::Hidden, cl::init(Val: false));
125
126static cl::opt<bool> PrintModuleAfterOptimizations(
127 "openmp-opt-print-module-after",
128 cl::desc("Print the current module after OpenMP optimizations."),
129 cl::Hidden, cl::init(Val: false));
130
131static cl::opt<bool> PrintModuleBeforeOptimizations(
132 "openmp-opt-print-module-before",
133 cl::desc("Print the current module before OpenMP optimizations."),
134 cl::Hidden, cl::init(Val: false));
135
136static cl::opt<bool> AlwaysInlineDeviceFunctions(
137 "openmp-opt-inline-device",
138 cl::desc("Inline all applicable functions on the device."), cl::Hidden,
139 cl::init(Val: false));
140
141static cl::opt<bool>
142 EnableVerboseRemarks("openmp-opt-verbose-remarks",
143 cl::desc("Enables more verbose remarks."), cl::Hidden,
144 cl::init(Val: false));
145
146static cl::opt<unsigned>
147 SetFixpointIterations("openmp-opt-max-iterations", cl::Hidden,
148 cl::desc("Maximal number of attributor iterations."),
149 cl::init(Val: 256));
150
151static cl::opt<unsigned>
152 SharedMemoryLimit("openmp-opt-shared-limit", cl::Hidden,
153 cl::desc("Maximum amount of shared memory to use."),
154 cl::init(Val: std::numeric_limits<unsigned>::max()));
155
156static cl::opt<unsigned> MaxCalleesForSpecialization(
157 "openmp-opt-max-callees-for-specialization", cl::Hidden,
158 cl::desc("Number of possible callees above which an indirect call site is "
159 "left alone rather than specialized into an if-cascade."),
160 cl::init(Val: 3));
161
162STATISTIC(NumOpenMPRuntimeCallsDeduplicated,
163 "Number of OpenMP runtime calls deduplicated");
164STATISTIC(NumOpenMPParallelRegionsDeleted,
165 "Number of OpenMP parallel regions deleted");
166STATISTIC(NumOpenMPRuntimeFunctionsIdentified,
167 "Number of OpenMP runtime functions identified");
168STATISTIC(NumOpenMPRuntimeFunctionUsesIdentified,
169 "Number of OpenMP runtime function uses identified");
170STATISTIC(NumOpenMPTargetRegionKernels,
171 "Number of OpenMP target region entry points (=kernels) identified");
172STATISTIC(NumNonOpenMPTargetRegionKernels,
173 "Number of non-OpenMP target region kernels identified");
174STATISTIC(NumOpenMPTargetRegionKernelsSPMD,
175 "Number of OpenMP target region entry points (=kernels) executed in "
176 "SPMD-mode instead of generic-mode");
177STATISTIC(NumOpenMPTargetRegionKernelsWithoutStateMachine,
178 "Number of OpenMP target region entry points (=kernels) executed in "
179 "generic-mode without a state machines");
180STATISTIC(NumOpenMPTargetRegionKernelsCustomStateMachineWithFallback,
181 "Number of OpenMP target region entry points (=kernels) executed in "
182 "generic-mode with customized state machines with fallback");
183STATISTIC(NumOpenMPTargetRegionKernelsCustomStateMachineWithoutFallback,
184 "Number of OpenMP target region entry points (=kernels) executed in "
185 "generic-mode with customized state machines without fallback");
186STATISTIC(
187 NumOpenMPParallelRegionsReplacedInGPUStateMachine,
188 "Number of OpenMP parallel regions replaced with ID in GPU state machines");
189STATISTIC(NumOpenMPParallelRegionsMerged,
190 "Number of OpenMP parallel regions merged");
191STATISTIC(NumBytesMovedToSharedMemory,
192 "Amount of memory pushed to shared memory");
193STATISTIC(NumBarriersEliminated, "Number of redundant barriers eliminated");
194
195#if !defined(NDEBUG)
196static constexpr auto TAG = "[" DEBUG_TYPE "]";
197#endif
198
199namespace KernelInfo {
200
201// struct ConfigurationEnvironmentTy {
202// uint8_t UseGenericStateMachine;
203// uint8_t MayUseNestedParallelism;
204// llvm::omp::OMPTgtExecModeFlags ExecMode;
205// int32_t MinThreads;
206// int32_t MaxThreads;
207// int32_t MinTeams;
208// int32_t MaxTeams;
209// };
210
211// struct DynamicEnvironmentTy {
212// uint16_t DebugIndentionLevel;
213// };
214
215// struct KernelEnvironmentTy {
216// ConfigurationEnvironmentTy Configuration;
217// IdentTy *Ident;
218// DynamicEnvironmentTy *DynamicEnv;
219// };
220
221#define KERNEL_ENVIRONMENT_IDX(MEMBER, IDX) \
222 constexpr unsigned MEMBER##Idx = IDX;
223
224KERNEL_ENVIRONMENT_IDX(Configuration, 0)
225KERNEL_ENVIRONMENT_IDX(Ident, 1)
226
227#undef KERNEL_ENVIRONMENT_IDX
228
229#define KERNEL_ENVIRONMENT_CONFIGURATION_IDX(MEMBER, IDX) \
230 constexpr unsigned MEMBER##Idx = IDX;
231
232KERNEL_ENVIRONMENT_CONFIGURATION_IDX(UseGenericStateMachine, 0)
233KERNEL_ENVIRONMENT_CONFIGURATION_IDX(MayUseNestedParallelism, 1)
234KERNEL_ENVIRONMENT_CONFIGURATION_IDX(ExecMode, 2)
235KERNEL_ENVIRONMENT_CONFIGURATION_IDX(MinThreads, 3)
236KERNEL_ENVIRONMENT_CONFIGURATION_IDX(MaxThreads, 4)
237KERNEL_ENVIRONMENT_CONFIGURATION_IDX(MinTeams, 5)
238KERNEL_ENVIRONMENT_CONFIGURATION_IDX(MaxTeams, 6)
239
240#undef KERNEL_ENVIRONMENT_CONFIGURATION_IDX
241
242#define KERNEL_ENVIRONMENT_GETTER(MEMBER, RETURNTYPE) \
243 RETURNTYPE *get##MEMBER##FromKernelEnvironment(ConstantStruct *KernelEnvC) { \
244 return cast<RETURNTYPE>(KernelEnvC->getAggregateElement(MEMBER##Idx)); \
245 }
246
247KERNEL_ENVIRONMENT_GETTER(Ident, Constant)
248KERNEL_ENVIRONMENT_GETTER(Configuration, ConstantStruct)
249
250#undef KERNEL_ENVIRONMENT_GETTER
251
252#define KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(MEMBER) \
253 ConstantInt *get##MEMBER##FromKernelEnvironment( \
254 ConstantStruct *KernelEnvC) { \
255 ConstantStruct *ConfigC = \
256 getConfigurationFromKernelEnvironment(KernelEnvC); \
257 return dyn_cast<ConstantInt>(ConfigC->getAggregateElement(MEMBER##Idx)); \
258 }
259
260KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(UseGenericStateMachine)
261KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(MayUseNestedParallelism)
262KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(ExecMode)
263KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(MinThreads)
264KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(MaxThreads)
265KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(MinTeams)
266KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(MaxTeams)
267
268#undef KERNEL_ENVIRONMENT_CONFIGURATION_GETTER
269
270GlobalVariable *
271getKernelEnvironementGVFromKernelInitCB(CallBase *KernelInitCB) {
272 constexpr int InitKernelEnvironmentArgNo = 0;
273 return cast<GlobalVariable>(
274 Val: KernelInitCB->getArgOperand(i: InitKernelEnvironmentArgNo)
275 ->stripPointerCasts());
276}
277
278ConstantStruct *getKernelEnvironementFromKernelInitCB(CallBase *KernelInitCB) {
279 GlobalVariable *KernelEnvGV =
280 getKernelEnvironementGVFromKernelInitCB(KernelInitCB);
281 return cast<ConstantStruct>(Val: KernelEnvGV->getInitializer());
282}
283} // namespace KernelInfo
284
285namespace {
286
287struct AAHeapToShared;
288
289struct AAICVTracker;
290
291/// OpenMP specific information. For now, stores RFIs and ICVs also needed for
292/// Attributor runs.
293struct OMPInformationCache : public InformationCache {
294 OMPInformationCache(Module &M, AnalysisGetter &AG,
295 BumpPtrAllocator &Allocator, SetVector<Function *> *CGSCC,
296 bool OpenMPPostLink)
297 : InformationCache(M, AG, Allocator, CGSCC), OMPBuilder(M),
298 OpenMPPostLink(OpenMPPostLink) {
299
300 OMPBuilder.Config.IsTargetDevice = isOpenMPDevice(M&: OMPBuilder.M);
301 const Triple T(OMPBuilder.M.getTargetTriple());
302 switch (T.getArch()) {
303 case llvm::Triple::nvptx:
304 case llvm::Triple::nvptx64:
305 case llvm::Triple::amdgpu:
306 assert(OMPBuilder.Config.IsTargetDevice &&
307 "OpenMP AMDGPU/NVPTX is only prepared to deal with device code.");
308 OMPBuilder.Config.IsGPU = true;
309 break;
310 default:
311 OMPBuilder.Config.IsGPU = false;
312 break;
313 }
314 OMPBuilder.initialize();
315 initializeRuntimeFunctions(M);
316 initializeInternalControlVars();
317 }
318
319 /// Generic information that describes an internal control variable.
320 struct InternalControlVarInfo {
321 /// The kind, as described by InternalControlVar enum.
322 InternalControlVar Kind;
323
324 /// The name of the ICV.
325 StringRef Name;
326
327 /// Environment variable associated with this ICV.
328 StringRef EnvVarName;
329
330 /// Initial value kind.
331 ICVInitValue InitKind;
332
333 /// Initial value.
334 ConstantInt *InitValue;
335
336 /// Setter RTL function associated with this ICV.
337 RuntimeFunction Setter;
338
339 /// Getter RTL function associated with this ICV.
340 RuntimeFunction Getter;
341
342 /// RTL Function corresponding to the override clause of this ICV
343 RuntimeFunction Clause;
344 };
345
346 /// Generic information that describes a runtime function
347 struct RuntimeFunctionInfo {
348
349 /// The kind, as described by the RuntimeFunction enum.
350 RuntimeFunction Kind;
351
352 /// The name of the function.
353 StringRef Name;
354
355 /// Flag to indicate a variadic function.
356 bool IsVarArg;
357
358 /// The return type of the function.
359 Type *ReturnType;
360
361 /// The argument types of the function.
362 SmallVector<Type *, 8> ArgumentTypes;
363
364 /// The declaration if available.
365 Function *Declaration = nullptr;
366
367 /// Uses of this runtime function per function containing the use.
368 using UseVector = SmallVector<Use *, 16>;
369
370 /// Clear UsesMap for runtime function.
371 void clearUsesMap() { UsesMap.clear(); }
372
373 /// Boolean conversion that is true if the runtime function was found.
374 operator bool() const { return Declaration; }
375
376 /// Return the vector of uses in function \p F.
377 UseVector &getOrCreateUseVector(Function *F) {
378 std::shared_ptr<UseVector> &UV = UsesMap[F];
379 if (!UV)
380 UV = std::make_shared<UseVector>();
381 return *UV;
382 }
383
384 /// Return the vector of uses in function \p F or `nullptr` if there are
385 /// none.
386 const UseVector *getUseVector(Function &F) const {
387 auto I = UsesMap.find(Val: &F);
388 if (I != UsesMap.end())
389 return I->second.get();
390 return nullptr;
391 }
392
393 /// Return how many functions contain uses of this runtime function.
394 size_t getNumFunctionsWithUses() const { return UsesMap.size(); }
395
396 /// Return the number of arguments (or the minimal number for variadic
397 /// functions).
398 size_t getNumArgs() const { return ArgumentTypes.size(); }
399
400 /// Run the callback \p CB on each use and forget the use if the result is
401 /// true. The callback will be fed the function in which the use was
402 /// encountered as second argument.
403 void foreachUse(SmallVectorImpl<Function *> &SCC,
404 function_ref<bool(Use &, Function &)> CB) {
405 for (Function *F : SCC)
406 foreachUse(CB, F);
407 }
408
409 /// Run the callback \p CB on each use within the function \p F and forget
410 /// the use if the result is true.
411 void foreachUse(function_ref<bool(Use &, Function &)> CB, Function *F) {
412 SmallVector<unsigned, 8> ToBeDeleted;
413 ToBeDeleted.clear();
414
415 unsigned Idx = 0;
416 UseVector &UV = getOrCreateUseVector(F);
417
418 for (Use *U : UV) {
419 if (CB(*U, *F))
420 ToBeDeleted.push_back(Elt: Idx);
421 ++Idx;
422 }
423
424 // Remove the to-be-deleted indices in reverse order as prior
425 // modifications will not modify the smaller indices.
426 while (!ToBeDeleted.empty()) {
427 unsigned Idx = ToBeDeleted.pop_back_val();
428 UV[Idx] = UV.back();
429 UV.pop_back();
430 }
431 }
432
433 private:
434 /// Map from functions to all uses of this runtime function contained in
435 /// them.
436 DenseMap<Function *, std::shared_ptr<UseVector>> UsesMap;
437
438 public:
439 /// Iterators for the uses of this runtime function.
440 decltype(UsesMap)::iterator begin() { return UsesMap.begin(); }
441 decltype(UsesMap)::iterator end() { return UsesMap.end(); }
442 };
443
444 /// An OpenMP-IR-Builder instance
445 OpenMPIRBuilder OMPBuilder;
446
447 /// Map from runtime function kind to the runtime function description.
448 EnumeratedArray<RuntimeFunctionInfo, RuntimeFunction,
449 RuntimeFunction::OMPRTL___last>
450 RFIs;
451
452 /// Map from function declarations/definitions to their runtime enum type.
453 DenseMap<Function *, RuntimeFunction> RuntimeFunctionIDMap;
454
455 /// Map from ICV kind to the ICV description.
456 EnumeratedArray<InternalControlVarInfo, InternalControlVar,
457 InternalControlVar::ICV___last>
458 ICVs;
459
460 /// Helper to initialize all internal control variable information for those
461 /// defined in OMPKinds.def.
462 void initializeInternalControlVars() {
463#define ICV_RT_SET(_Name, RTL) \
464 { \
465 auto &ICV = ICVs[_Name]; \
466 ICV.Setter = RTL; \
467 }
468#define ICV_RT_GET(Name, RTL) \
469 { \
470 auto &ICV = ICVs[Name]; \
471 ICV.Getter = RTL; \
472 }
473#define ICV_DATA_ENV(Enum, _Name, _EnvVarName, Init) \
474 { \
475 auto &ICV = ICVs[Enum]; \
476 ICV.Name = _Name; \
477 ICV.Kind = Enum; \
478 ICV.InitKind = Init; \
479 ICV.EnvVarName = _EnvVarName; \
480 switch (ICV.InitKind) { \
481 case ICV_IMPLEMENTATION_DEFINED: \
482 ICV.InitValue = nullptr; \
483 break; \
484 case ICV_ZERO: \
485 ICV.InitValue = ConstantInt::get( \
486 Type::getInt32Ty(OMPBuilder.Int32->getContext()), 0); \
487 break; \
488 case ICV_FALSE: \
489 ICV.InitValue = ConstantInt::getFalse(OMPBuilder.Int1->getContext()); \
490 break; \
491 case ICV_LAST: \
492 break; \
493 } \
494 }
495#include "llvm/Frontend/OpenMP/OMPKinds.def"
496 }
497
498 /// Returns true if the function declaration \p F matches the runtime
499 /// function types, that is, return type \p RTFRetType, and argument types
500 /// \p RTFArgTypes.
501 static bool declMatchesRTFTypes(Function *F, Type *RTFRetType,
502 SmallVector<Type *, 8> &RTFArgTypes) {
503 // TODO: We should output information to the user (under debug output
504 // and via remarks).
505
506 if (!F)
507 return false;
508 if (F->getReturnType() != RTFRetType)
509 return false;
510 if (F->arg_size() != RTFArgTypes.size())
511 return false;
512
513 auto *RTFTyIt = RTFArgTypes.begin();
514 for (Argument &Arg : F->args()) {
515 if (Arg.getType() != *RTFTyIt)
516 return false;
517
518 ++RTFTyIt;
519 }
520
521 return true;
522 }
523
524 // Helper to collect all uses of the declaration in the UsesMap.
525 unsigned collectUses(RuntimeFunctionInfo &RFI, bool CollectStats = true) {
526 unsigned NumUses = 0;
527 if (!RFI.Declaration)
528 return NumUses;
529 OMPBuilder.addAttributes(FnID: RFI.Kind, Fn&: *RFI.Declaration);
530
531 if (CollectStats) {
532 NumOpenMPRuntimeFunctionsIdentified += 1;
533 NumOpenMPRuntimeFunctionUsesIdentified += RFI.Declaration->getNumUses();
534 }
535
536 // TODO: We directly convert uses into proper calls and unknown uses.
537 for (Use &U : RFI.Declaration->uses()) {
538 if (Instruction *UserI = dyn_cast<Instruction>(Val: U.getUser())) {
539 if (!CGSCC || CGSCC->empty() || CGSCC->contains(key: UserI->getFunction())) {
540 RFI.getOrCreateUseVector(F: UserI->getFunction()).push_back(Elt: &U);
541 ++NumUses;
542 }
543 } else {
544 RFI.getOrCreateUseVector(F: nullptr).push_back(Elt: &U);
545 ++NumUses;
546 }
547 }
548 return NumUses;
549 }
550
551 // Helper function to recollect uses of a runtime function.
552 void recollectUsesForFunction(RuntimeFunction RTF) {
553 auto &RFI = RFIs[RTF];
554 RFI.clearUsesMap();
555 collectUses(RFI, /*CollectStats*/ false);
556 }
557
558 /// Attach !callback metadata to a runtime function that takes one, so that
559 /// the Attributor sees the edge from the runtime call to the callback and
560 /// AAKernelInfo can look inside it. The runtime declares these functions
561 /// without the metadata, so OpenMPOpt supplies it from the table in
562 /// OMPKinds.def.
563 void setCallbackMetadata(Function *F, unsigned ArgNo, ArrayRef<int> Indices,
564 bool IsVarArg) {
565 if (!F || F->hasMetadata(KindID: LLVMContext::MD_callback))
566 return;
567
568 LLVMContext &Ctx = F->getContext();
569 MDBuilder MDB(Ctx);
570 F->addMetadata(KindID: LLVMContext::MD_callback,
571 MD&: *MDNode::get(Context&: Ctx, MDs: {MDB.createCallbackEncoding(CalleeArgNo: ArgNo, Arguments: Indices,
572 VarArgsArePassed: IsVarArg)}));
573 }
574
575 /// The callback a runtime function was handed, if it is one we can analyze.
576 /// Returns null when the call takes no callback, or when the callback is not
577 /// a definition this module can see, in which case its contents are unknown
578 /// and callers have to stay conservative.
579 static Function *getAnalyzableCallback(const CallBase &CB) {
580 Function *Callee = CB.getCalledFunction();
581 if (!Callee)
582 return nullptr;
583 MDNode *CallbackMD = Callee->getMetadata(KindID: LLVMContext::MD_callback);
584 if (!CallbackMD || CallbackMD->getNumOperands() == 0)
585 return nullptr;
586 // TODO: A runtime function with more than one callback would need each of
587 // them checked; none of the ones in the table have more than one.
588 auto *Encoding = dyn_cast<MDNode>(Val: CallbackMD->getOperand(I: 0));
589 if (!Encoding || Encoding->getNumOperands() == 0)
590 return nullptr;
591 auto *ArgNoMD = dyn_cast<ConstantAsMetadata>(Val: Encoding->getOperand(I: 0));
592 if (!ArgNoMD)
593 return nullptr;
594 uint64_t ArgNo =
595 cast<ConstantInt>(Val: ArgNoMD->getValue())->getLimitedValue(UINT64_MAX);
596 if (ArgNo >= CB.arg_size())
597 return nullptr;
598 auto *Callback =
599 dyn_cast<Function>(Val: CB.getArgOperand(i: ArgNo)->stripPointerCasts());
600 if (!Callback || Callback->isDeclaration())
601 return nullptr;
602 return Callback;
603 }
604
605 // Helper function to recollect uses of all runtime functions.
606 void recollectUses() {
607 for (int Idx = 0; Idx < RFIs.size(); ++Idx)
608 recollectUsesForFunction(RTF: static_cast<RuntimeFunction>(Idx));
609 }
610
611 // Helper function to inherit the calling convention of the function callee.
612 void setCallingConvention(FunctionCallee Callee, CallInst *CI) {
613 if (Function *Fn = dyn_cast<Function>(Val: Callee.getCallee()))
614 CI->setCallingConv(Fn->getCallingConv());
615 }
616
617 // Helper function to determine if it's legal to create a call to the runtime
618 // functions.
619 bool runtimeFnsAvailable(ArrayRef<RuntimeFunction> Fns) {
620 // We can always emit calls if we haven't yet linked in the runtime.
621 if (!OpenMPPostLink)
622 return true;
623
624 // Once the runtime has been already been linked in we cannot emit calls to
625 // any undefined functions.
626 for (RuntimeFunction Fn : Fns) {
627 RuntimeFunctionInfo &RFI = RFIs[Fn];
628
629 if (!RFI.Declaration || RFI.Declaration->isDeclaration())
630 return false;
631 }
632 return true;
633 }
634
635 /// Helper to initialize all runtime function information for those defined
636 /// in OpenMPKinds.def.
637 void initializeRuntimeFunctions(Module &M) {
638
639 // Helper macros for handling __VA_ARGS__ in OMP_RTL
640#define OMP_TYPE(VarName, ...) \
641 Type *VarName = OMPBuilder.VarName; \
642 (void)VarName;
643
644#define OMP_ARRAY_TYPE(VarName, ...) \
645 ArrayType *VarName##Ty = OMPBuilder.VarName##Ty; \
646 (void)VarName##Ty; \
647 PointerType *VarName##PtrTy = OMPBuilder.VarName##PtrTy; \
648 (void)VarName##PtrTy;
649
650#define OMP_FUNCTION_TYPE(VarName, ...) \
651 FunctionType *VarName = OMPBuilder.VarName; \
652 (void)VarName; \
653 PointerType *VarName##Ptr = OMPBuilder.VarName##Ptr; \
654 (void)VarName##Ptr;
655
656#define OMP_STRUCT_TYPE(VarName, ...) \
657 StructType *VarName = OMPBuilder.VarName; \
658 (void)VarName; \
659 PointerType *VarName##Ptr = OMPBuilder.VarName##Ptr; \
660 (void)VarName##Ptr;
661
662#define OMP_RTL(_Enum, _Name, _IsVarArg, _ReturnType, ...) \
663 { \
664 SmallVector<Type *, 8> ArgsTypes({__VA_ARGS__}); \
665 Function *F = M.getFunction(_Name); \
666 RTLFunctions.insert(F); \
667 if (declMatchesRTFTypes(F, OMPBuilder._ReturnType, ArgsTypes)) { \
668 RuntimeFunctionIDMap[F] = _Enum; \
669 auto &RFI = RFIs[_Enum]; \
670 RFI.Kind = _Enum; \
671 RFI.Name = _Name; \
672 RFI.IsVarArg = _IsVarArg; \
673 RFI.ReturnType = OMPBuilder._ReturnType; \
674 RFI.ArgumentTypes = std::move(ArgsTypes); \
675 RFI.Declaration = F; \
676 unsigned NumUses = collectUses(RFI); \
677 (void)NumUses; \
678 LLVM_DEBUG({ \
679 dbgs() << TAG << RFI.Name << (RFI.Declaration ? "" : " not") \
680 << " found\n"; \
681 if (RFI.Declaration) \
682 dbgs() << TAG << "-> got " << NumUses << " uses in " \
683 << RFI.getNumFunctionsWithUses() \
684 << " different functions.\n"; \
685 }); \
686 } \
687 }
688
689#define OMP_RTL_CB_INFO(_Enum, _Name, _ArgNo, _ArgIndices, _IsVarArg) \
690 setCallbackMetadata(M.getFunction(_Name), _ArgNo, _ArgIndices, _IsVarArg);
691
692#include "llvm/Frontend/OpenMP/OMPKinds.def"
693
694 // Remove the `noinline` attribute from `__kmpc`, `ompx::` and `omp_`
695 // functions, except if `optnone` is present.
696 if (isOpenMPDevice(M)) {
697 for (Function &F : M) {
698 for (StringRef Prefix : {"__kmpc", "_ZN4ompx", "omp_"})
699 if (F.hasFnAttribute(Kind: Attribute::NoInline) &&
700 F.getName().starts_with(Prefix) &&
701 !F.hasFnAttribute(Kind: Attribute::OptimizeNone))
702 F.removeFnAttr(Kind: Attribute::NoInline);
703 }
704 }
705
706 // TODO: We should attach the attributes defined in OMPKinds.def.
707 }
708
709 /// Collection of known OpenMP runtime functions..
710 DenseSet<const Function *> RTLFunctions;
711
712 /// Indicates if we have already linked in the OpenMP device library.
713 bool OpenMPPostLink = false;
714
715 /// Kernels that OpenMPOpt transformed from generic to SPMD mode. Recorded at
716 /// the transform (changeToSPMDMode) so later cleanup does not have to
717 /// re-derive the mode. Such kernels no longer run a generic-mode state
718 /// machine, so the parallel data-sharing wrapper passed to __kmpc_parallel_60
719 /// is dead in them.
720 SmallPtrSet<Function *, 8> SPMDizedKernels;
721};
722
723template <typename Ty, bool InsertInvalidates = true>
724struct BooleanStateWithSetVector : public BooleanState {
725 bool contains(const Ty &Elem) const { return Set.contains(Elem); }
726 bool insert(const Ty &Elem) {
727 if (InsertInvalidates)
728 BooleanState::indicatePessimisticFixpoint();
729 return Set.insert(Elem);
730 }
731
732 const Ty &operator[](int Idx) const { return Set[Idx]; }
733 bool operator==(const BooleanStateWithSetVector &RHS) const {
734 return BooleanState::operator==(R: RHS) && Set == RHS.Set;
735 }
736 bool operator!=(const BooleanStateWithSetVector &RHS) const {
737 return !(*this == RHS);
738 }
739
740 bool empty() const { return Set.empty(); }
741 size_t size() const { return Set.size(); }
742
743 /// "Clamp" this state with \p RHS.
744 BooleanStateWithSetVector &operator^=(const BooleanStateWithSetVector &RHS) {
745 BooleanState::operator^=(R: RHS);
746 Set.insert_range(RHS.Set);
747 return *this;
748 }
749
750private:
751 /// A set to keep track of elements.
752 SetVector<Ty> Set;
753
754public:
755 typename decltype(Set)::iterator begin() { return Set.begin(); }
756 typename decltype(Set)::iterator end() { return Set.end(); }
757 typename decltype(Set)::const_iterator begin() const { return Set.begin(); }
758 typename decltype(Set)::const_iterator end() const { return Set.end(); }
759};
760
761template <typename Ty, bool InsertInvalidates = true>
762using BooleanStateWithPtrSetVector =
763 BooleanStateWithSetVector<Ty *, InsertInvalidates>;
764
765struct KernelInfoState : AbstractState {
766 /// Flag to track if we reached a fixpoint.
767 bool IsAtFixpoint = false;
768
769 /// The parallel regions (identified by the outlined parallel functions) that
770 /// can be reached from the associated function.
771 BooleanStateWithPtrSetVector<CallBase, /* InsertInvalidates */ false>
772 ReachedKnownParallelRegions;
773
774 /// State to track what parallel region we might reach.
775 BooleanStateWithPtrSetVector<CallBase> ReachedUnknownParallelRegions;
776
777 /// State to track if we are in SPMD-mode, assumed or know, and why we decided
778 /// we cannot be. If it is assumed, then RequiresFullRuntime should also be
779 /// false.
780 BooleanStateWithPtrSetVector<Instruction, false> SPMDCompatibilityTracker;
781
782 /// The __kmpc_target_init call in this kernel, if any. If we find more than
783 /// one we abort as the kernel is malformed.
784 CallBase *KernelInitCB = nullptr;
785
786 /// The constant kernel environement as taken from and passed to
787 /// __kmpc_target_init.
788 ConstantStruct *KernelEnvC = nullptr;
789
790 /// The __kmpc_target_deinit call in this kernel, if any. If we find more than
791 /// one we abort as the kernel is malformed.
792 CallBase *KernelDeinitCB = nullptr;
793
794 /// Flag to indicate if the associated function is a kernel entry.
795 bool IsKernelEntry = false;
796
797 /// State to track what kernel entries can reach the associated function.
798 BooleanStateWithPtrSetVector<Function, false> ReachingKernelEntries;
799
800 /// State to indicate if we can track parallel level of the associated
801 /// function. We will give up tracking if we encounter unknown caller or the
802 /// caller is __kmpc_parallel_60.
803 BooleanStateWithSetVector<uint8_t> ParallelLevels;
804
805 /// Flag that indicates if the kernel has nested Parallelism
806 bool NestedParallelism = false;
807
808 /// Abstract State interface
809 ///{
810
811 KernelInfoState() = default;
812 KernelInfoState(bool BestState) {
813 if (!BestState)
814 indicatePessimisticFixpoint();
815 }
816
817 /// See AbstractState::isValidState(...)
818 bool isValidState() const override { return true; }
819
820 /// See AbstractState::isAtFixpoint(...)
821 bool isAtFixpoint() const override { return IsAtFixpoint; }
822
823 /// See AbstractState::indicatePessimisticFixpoint(...)
824 ChangeStatus indicatePessimisticFixpoint() override {
825 IsAtFixpoint = true;
826 ParallelLevels.indicatePessimisticFixpoint();
827 ReachingKernelEntries.indicatePessimisticFixpoint();
828 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
829 ReachedKnownParallelRegions.indicatePessimisticFixpoint();
830 ReachedUnknownParallelRegions.indicatePessimisticFixpoint();
831 NestedParallelism = true;
832 return ChangeStatus::CHANGED;
833 }
834
835 /// See AbstractState::indicateOptimisticFixpoint(...)
836 ChangeStatus indicateOptimisticFixpoint() override {
837 IsAtFixpoint = true;
838 ParallelLevels.indicateOptimisticFixpoint();
839 ReachingKernelEntries.indicateOptimisticFixpoint();
840 SPMDCompatibilityTracker.indicateOptimisticFixpoint();
841 ReachedKnownParallelRegions.indicateOptimisticFixpoint();
842 ReachedUnknownParallelRegions.indicateOptimisticFixpoint();
843 return ChangeStatus::UNCHANGED;
844 }
845
846 /// Return the assumed state
847 KernelInfoState &getAssumed() { return *this; }
848 const KernelInfoState &getAssumed() const { return *this; }
849
850 bool operator==(const KernelInfoState &RHS) const {
851 if (SPMDCompatibilityTracker != RHS.SPMDCompatibilityTracker)
852 return false;
853 if (ReachedKnownParallelRegions != RHS.ReachedKnownParallelRegions)
854 return false;
855 if (ReachedUnknownParallelRegions != RHS.ReachedUnknownParallelRegions)
856 return false;
857 if (ReachingKernelEntries != RHS.ReachingKernelEntries)
858 return false;
859 if (ParallelLevels != RHS.ParallelLevels)
860 return false;
861 if (NestedParallelism != RHS.NestedParallelism)
862 return false;
863 return true;
864 }
865
866 /// Returns true if this kernel contains any OpenMP parallel regions.
867 bool mayContainParallelRegion() {
868 return !ReachedKnownParallelRegions.empty() ||
869 !ReachedUnknownParallelRegions.empty();
870 }
871
872 /// Return empty set as the best state of potential values.
873 static KernelInfoState getBestState() { return KernelInfoState(true); }
874
875 static KernelInfoState getBestState(KernelInfoState &KIS) {
876 return getBestState();
877 }
878
879 /// Return full set as the worst state of potential values.
880 static KernelInfoState getWorstState() { return KernelInfoState(false); }
881
882 /// "Clamp" this state with \p KIS.
883 KernelInfoState operator^=(const KernelInfoState &KIS) {
884 // Do not merge two different _init and _deinit call sites.
885 if (KIS.KernelInitCB) {
886 if (KernelInitCB && KernelInitCB != KIS.KernelInitCB)
887 llvm_unreachable("Kernel that calls another kernel violates OpenMP-Opt "
888 "assumptions.");
889 KernelInitCB = KIS.KernelInitCB;
890 }
891 if (KIS.KernelDeinitCB) {
892 if (KernelDeinitCB && KernelDeinitCB != KIS.KernelDeinitCB)
893 llvm_unreachable("Kernel that calls another kernel violates OpenMP-Opt "
894 "assumptions.");
895 KernelDeinitCB = KIS.KernelDeinitCB;
896 }
897 if (KIS.KernelEnvC) {
898 if (KernelEnvC && KernelEnvC != KIS.KernelEnvC)
899 llvm_unreachable("Kernel that calls another kernel violates OpenMP-Opt "
900 "assumptions.");
901 KernelEnvC = KIS.KernelEnvC;
902 }
903 SPMDCompatibilityTracker ^= KIS.SPMDCompatibilityTracker;
904 ReachedKnownParallelRegions ^= KIS.ReachedKnownParallelRegions;
905 ReachedUnknownParallelRegions ^= KIS.ReachedUnknownParallelRegions;
906 NestedParallelism |= KIS.NestedParallelism;
907 return *this;
908 }
909
910 KernelInfoState operator&=(const KernelInfoState &KIS) {
911 return (*this ^= KIS);
912 }
913
914 ///}
915};
916
917/// Used to map the values physically (in the IR) stored in an offload
918/// array, to a vector in memory.
919struct OffloadArray {
920 /// Physical array (in the IR).
921 AllocaInst *Array = nullptr;
922 /// Mapped values.
923 SmallVector<Value *, 8> StoredValues;
924 /// Last stores made in the offload array.
925 SmallVector<StoreInst *, 8> LastAccesses;
926
927 OffloadArray() = default;
928
929 /// Initializes the OffloadArray with the values stored in \p Array before
930 /// instruction \p Before is reached. Returns false if the initialization
931 /// fails.
932 /// This MUST be used immediately after the construction of the object.
933 bool initialize(AllocaInst &Array, Instruction &Before) {
934 if (!getValues(Array, Before))
935 return false;
936
937 this->Array = &Array;
938 return true;
939 }
940
941 static const unsigned DeviceIDArgNum = 1;
942 static const unsigned BasePtrsArgNum = 3;
943 static const unsigned PtrsArgNum = 4;
944 static const unsigned SizesArgNum = 5;
945
946private:
947 /// Traverses the BasicBlock where \p Array is, collecting the stores made to
948 /// \p Array, leaving StoredValues with the values stored before the
949 /// instruction \p Before is reached.
950 bool getValues(AllocaInst &Array, Instruction &Before) {
951 // Initialize containers.
952 const DataLayout &DL = Array.getDataLayout();
953 std::optional<TypeSize> ArraySize = Array.getAllocationSize(DL);
954 if (!ArraySize || !ArraySize->isFixed())
955 return false;
956 const unsigned int PointerSize = DL.getPointerSize();
957 const uint64_t NumValues = ArraySize->getFixedValue() / PointerSize;
958 StoredValues.assign(NumElts: NumValues, Elt: nullptr);
959 LastAccesses.assign(NumElts: NumValues, Elt: nullptr);
960
961 // TODO: This assumes the instruction \p Before is in the same
962 // BasicBlock as Array. Make it general, for any control flow graph.
963 BasicBlock *BB = Array.getParent();
964 if (BB != Before.getParent())
965 return false;
966
967 for (Instruction &I : *BB) {
968 if (&I == &Before)
969 break;
970
971 if (!isa<StoreInst>(Val: &I))
972 continue;
973
974 auto *S = cast<StoreInst>(Val: &I);
975 int64_t Offset = -1;
976 auto *Dst =
977 GetPointerBaseWithConstantOffset(Ptr: S->getPointerOperand(), Offset, DL);
978 if (Dst == &Array) {
979 int64_t Idx = Offset / PointerSize;
980 // Ignore updates that must be UB (probably in dead code at runtime)
981 if ((uint64_t)Idx < NumValues) {
982 StoredValues[Idx] = getUnderlyingObject(V: S->getValueOperand());
983 LastAccesses[Idx] = S;
984 }
985 }
986 }
987
988 return isFilled();
989 }
990
991 /// Returns true if all values in StoredValues and
992 /// LastAccesses are not nullptrs.
993 bool isFilled() {
994 const unsigned NumValues = StoredValues.size();
995 for (unsigned I = 0; I < NumValues; ++I) {
996 if (!StoredValues[I] || !LastAccesses[I])
997 return false;
998 }
999
1000 return true;
1001 }
1002};
1003
1004// Use the max outlined entry count. Instrumentation counts should already
1005// match across merged callbacks, but sample profiles can differ. Returns
1006// nullopt when no callback entry count is available.
1007static std::optional<uint64_t>
1008getMergedWrapperEntryCount(ArrayRef<CallInst *> ForkCalls,
1009 unsigned CallbackOpNo) {
1010 std::optional<uint64_t> EntryCount;
1011 for (CallInst *CI : ForkCalls) {
1012 auto *Callback = dyn_cast<Function>(
1013 Val: CI->getArgOperand(i: CallbackOpNo)->stripPointerCasts());
1014 if (!Callback)
1015 continue;
1016 if (std::optional<uint64_t> EC = Callback->getEntryCount())
1017 // Each callback runs once per wrapper entry, so the counts should match.
1018 // The Sample profiles can disagree slightly, so take the largest.
1019 EntryCount = EntryCount ? std::max(a: *EntryCount, b: *EC) : *EC;
1020 }
1021 return EntryCount;
1022}
1023
1024static bool moduleHasSampleProfile(const Module &M) {
1025 std::unique_ptr<ProfileSummary> Summary(
1026 ProfileSummary::getFromMD(MD: M.getProfileSummary(/*IsCS=*/false)));
1027 return Summary && Summary->getKind() == ProfileSummary::PSK_Sample;
1028}
1029
1030struct OpenMPOpt {
1031
1032 using OptimizationRemarkGetter =
1033 function_ref<OptimizationRemarkEmitter &(Function *)>;
1034
1035 OpenMPOpt(SmallVectorImpl<Function *> &SCC, CallGraphUpdater &CGUpdater,
1036 OptimizationRemarkGetter OREGetter,
1037 OMPInformationCache &OMPInfoCache, Attributor &A)
1038 : M(*(*SCC.begin())->getParent()), SCC(SCC), CGUpdater(CGUpdater),
1039 OREGetter(OREGetter), OMPInfoCache(OMPInfoCache), A(A) {}
1040
1041 /// Check if any remarks are enabled for openmp-opt
1042 bool remarksEnabled() {
1043 auto &Ctx = M.getContext();
1044 return Ctx.getDiagHandlerPtr()->isAnyRemarkEnabled(DEBUG_TYPE);
1045 }
1046
1047 /// Run all OpenMP optimizations on the underlying SCC.
1048 bool run(bool IsModulePass) {
1049 if (SCC.empty())
1050 return false;
1051
1052 bool Changed = false;
1053
1054 LLVM_DEBUG(dbgs() << TAG << "Run on SCC with " << SCC.size()
1055 << " functions\n");
1056
1057 if (IsModulePass) {
1058 Changed |= runAttributor(IsModulePass);
1059
1060 // Recollect uses, in case Attributor deleted any.
1061 OMPInfoCache.recollectUses();
1062
1063 // TODO: This should be folded into buildCustomStateMachine.
1064 Changed |= rewriteDeviceCodeStateMachine();
1065
1066 // Drop the parallel data-sharing wrapper from __kmpc_parallel_60 calls in
1067 // SPMD kernels, where the runtime never uses it, so the (otherwise dead)
1068 // wrapper can be eliminated instead of lingering as a non-kernel LDS
1069 // user.
1070 Changed |= removeSPMDParallelWrappers();
1071
1072 if (remarksEnabled())
1073 analysisGlobalization();
1074 } else {
1075 if (PrintICVValues)
1076 printICVs();
1077 if (PrintOpenMPKernels)
1078 printKernels();
1079
1080 Changed |= runAttributor(IsModulePass);
1081
1082 // Recollect uses, in case Attributor deleted any.
1083 OMPInfoCache.recollectUses();
1084
1085 Changed |= deleteParallelRegions();
1086
1087 if (HideMemoryTransferLatency)
1088 Changed |= hideMemTransfersLatency();
1089 Changed |= deduplicateRuntimeCalls();
1090 if (EnableParallelRegionMerging) {
1091 if (mergeParallelRegions()) {
1092 deduplicateRuntimeCalls();
1093 Changed = true;
1094 }
1095 }
1096 }
1097
1098 if (OMPInfoCache.OpenMPPostLink)
1099 Changed |= removeRuntimeSymbols();
1100
1101 return Changed;
1102 }
1103
1104 /// Print initial ICV values for testing.
1105 /// FIXME: This should be done from the Attributor once it is added.
1106 void printICVs() const {
1107 InternalControlVar ICVs[] = {ICV_nthreads, ICV_active_levels, ICV_cancel,
1108 ICV_proc_bind};
1109
1110 for (Function *F : SCC) {
1111 for (auto ICV : ICVs) {
1112 auto ICVInfo = OMPInfoCache.ICVs[ICV];
1113 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
1114 return ORA << "OpenMP ICV " << ore::NV("OpenMPICV", ICVInfo.Name)
1115 << " Value: "
1116 << (ICVInfo.InitValue
1117 ? toString(I: ICVInfo.InitValue->getValue(), Radix: 10, Signed: true)
1118 : "IMPLEMENTATION_DEFINED");
1119 };
1120
1121 emitRemark<OptimizationRemarkAnalysis>(F, RemarkName: "OpenMPICVTracker", RemarkCB&: Remark);
1122 }
1123 }
1124 }
1125
1126 /// Print OpenMP GPU kernels for testing.
1127 void printKernels() const {
1128 for (Function *F : SCC) {
1129 if (!omp::isOpenMPKernel(Fn&: *F))
1130 continue;
1131
1132 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
1133 return ORA << "OpenMP GPU kernel "
1134 << ore::NV("OpenMPGPUKernel", F->getName()) << "\n";
1135 };
1136
1137 emitRemark<OptimizationRemarkAnalysis>(F, RemarkName: "OpenMPGPU", RemarkCB&: Remark);
1138 }
1139 }
1140
1141 /// Return the call if \p U is a callee use in a regular call. If \p RFI is
1142 /// given it has to be the callee or a nullptr is returned.
1143 static CallInst *getCallIfRegularCall(
1144 Use &U, OMPInformationCache::RuntimeFunctionInfo *RFI = nullptr) {
1145 CallInst *CI = dyn_cast<CallInst>(Val: U.getUser());
1146 if (CI && CI->isCallee(U: &U) && !CI->hasOperandBundles() &&
1147 (!RFI ||
1148 (RFI->Declaration && CI->getCalledFunction() == RFI->Declaration)))
1149 return CI;
1150 return nullptr;
1151 }
1152
1153 /// Return the call if \p V is a regular call. If \p RFI is given it has to be
1154 /// the callee or a nullptr is returned.
1155 static CallInst *getCallIfRegularCall(
1156 Value &V, OMPInformationCache::RuntimeFunctionInfo *RFI = nullptr) {
1157 CallInst *CI = dyn_cast<CallInst>(Val: &V);
1158 if (CI && !CI->hasOperandBundles() &&
1159 (!RFI ||
1160 (RFI->Declaration && CI->getCalledFunction() == RFI->Declaration)))
1161 return CI;
1162 return nullptr;
1163 }
1164
1165private:
1166 /// Merge parallel regions when it is safe.
1167 bool mergeParallelRegions() {
1168 const unsigned CallbackCalleeOperand = 2;
1169 const unsigned CallbackFirstArgOperand = 3;
1170 using InsertPointTy = OpenMPIRBuilder::InsertPointTy;
1171
1172 // Check if there are any __kmpc_fork_call calls to merge.
1173 OMPInformationCache::RuntimeFunctionInfo &RFI =
1174 OMPInfoCache.RFIs[OMPRTL___kmpc_fork_call];
1175
1176 if (!RFI.Declaration)
1177 return false;
1178
1179 // Unmergable calls that prevent merging a parallel region.
1180 OMPInformationCache::RuntimeFunctionInfo UnmergableCallsInfo[] = {
1181 OMPInfoCache.RFIs[OMPRTL___kmpc_push_proc_bind],
1182 OMPInfoCache.RFIs[OMPRTL___kmpc_push_num_threads],
1183 };
1184
1185 bool Changed = false;
1186 LoopInfo *LI = nullptr;
1187 DominatorTree *DT = nullptr;
1188
1189 SmallDenseMap<BasicBlock *, SmallPtrSet<Instruction *, 4>> BB2PRMap;
1190
1191 BasicBlock *StartBB = nullptr, *EndBB = nullptr;
1192 auto BodyGenCB = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
1193 ArrayRef<BasicBlock *> DeallocBlocks) {
1194 BasicBlock *CGStartBB = CodeGenIP.getNodeParent();
1195 BasicBlock *CGEndBB = SplitBlock(Old: CGStartBB, SplitPt: &*CodeGenIP, DT, LI);
1196 assert(StartBB != nullptr && "StartBB should not be null");
1197 CGStartBB->getTerminator()->setSuccessor(Idx: 0, BB: StartBB);
1198 assert(EndBB != nullptr && "EndBB should not be null");
1199 EndBB->getTerminator()->setSuccessor(Idx: 0, BB: CGEndBB);
1200 return Error::success();
1201 };
1202
1203 auto PrivCB = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP, Value &,
1204 Value &Inner, Value *&ReplacementValue) -> InsertPointTy {
1205 ReplacementValue = &Inner;
1206 return CodeGenIP;
1207 };
1208
1209 auto FiniCB = [&](InsertPointTy CodeGenIP) { return Error::success(); };
1210
1211 /// Create a sequential execution region within a merged parallel region,
1212 /// encapsulated in a master construct with a barrier for synchronization.
1213 auto CreateSequentialRegion = [&](Function *OuterFn,
1214 BasicBlock *OuterPredBB,
1215 Instruction *SeqStartI,
1216 Instruction *SeqEndI) {
1217 // Isolate the instructions of the sequential region to a separate
1218 // block.
1219 BasicBlock *ParentBB = SeqStartI->getParent();
1220 BasicBlock *SeqEndBB =
1221 SplitBlock(Old: ParentBB, SplitPt: SeqEndI->getNextNode(), DT, LI);
1222 BasicBlock *SeqAfterBB =
1223 SplitBlock(Old: SeqEndBB, SplitPt: &*SeqEndBB->getFirstInsertionPt(), DT, LI);
1224 BasicBlock *SeqStartBB =
1225 SplitBlock(Old: ParentBB, SplitPt: SeqStartI, DT, LI, MSSAU: nullptr, BBName: "seq.par.merged");
1226
1227 assert(ParentBB->getUniqueSuccessor() == SeqStartBB &&
1228 "Expected a different CFG");
1229 const DebugLoc DL = ParentBB->getTerminator()->getDebugLoc();
1230 ParentBB->getTerminator()->eraseFromParent();
1231
1232 auto BodyGenCB = [&](InsertPointTy AllocaIP, InsertPointTy CodeGenIP,
1233 ArrayRef<BasicBlock *> DeallocBlocks) {
1234 BasicBlock *CGStartBB = CodeGenIP.getNodeParent();
1235 BasicBlock *CGEndBB = SplitBlock(Old: CGStartBB, SplitPt: &*CodeGenIP, DT, LI);
1236 assert(SeqStartBB != nullptr && "SeqStartBB should not be null");
1237 CGStartBB->getTerminator()->setSuccessor(Idx: 0, BB: SeqStartBB);
1238 assert(SeqEndBB != nullptr && "SeqEndBB should not be null");
1239 SeqEndBB->getTerminator()->setSuccessor(Idx: 0, BB: CGEndBB);
1240 return Error::success();
1241 };
1242 auto FiniCB = [&](InsertPointTy CodeGenIP) { return Error::success(); };
1243
1244 // Find outputs from the sequential region to outside users and
1245 // broadcast their values to them.
1246 for (Instruction &I : *SeqStartBB) {
1247 SmallPtrSet<Instruction *, 4> OutsideUsers;
1248 for (User *Usr : I.users()) {
1249 Instruction &UsrI = *cast<Instruction>(Val: Usr);
1250 // Ignore outputs to LT intrinsics, code extraction for the merged
1251 // parallel region will fix them.
1252 if (UsrI.isLifetimeStartOrEnd())
1253 continue;
1254
1255 if (UsrI.getParent() != SeqStartBB)
1256 OutsideUsers.insert(Ptr: &UsrI);
1257 }
1258
1259 if (OutsideUsers.empty())
1260 continue;
1261
1262 // Emit an alloca in the outer region to store the broadcasted
1263 // value.
1264 const DataLayout &DL = M.getDataLayout();
1265 AllocaInst *AllocaI = new AllocaInst(
1266 I.getType(), DL.getAllocaAddrSpace(), nullptr,
1267 I.getName() + ".seq.output.alloc", OuterFn->front().begin());
1268
1269 // Emit a store instruction in the sequential BB to update the
1270 // value.
1271 new StoreInst(&I, AllocaI, SeqStartBB->getTerminator()->getIterator());
1272
1273 // Emit a load instruction and replace the use of the output value
1274 // with it.
1275 for (Instruction *UsrI : OutsideUsers) {
1276 LoadInst *LoadI = new LoadInst(I.getType(), AllocaI,
1277 I.getName() + ".seq.output.load",
1278 UsrI->getIterator());
1279 UsrI->replaceUsesOfWith(From: &I, To: LoadI);
1280 }
1281 }
1282
1283 OpenMPIRBuilder::LocationDescription Loc(ParentBB->end(), DL);
1284 OpenMPIRBuilder::InsertPointTy SeqAfterIP = cantFail(
1285 ValOrErr: OMPInfoCache.OMPBuilder.createMaster(Loc, BodyGenCB, FiniCB));
1286 cantFail(ValOrErr: OMPInfoCache.OMPBuilder.createBarrier(Loc: {SeqAfterIP, DL},
1287 Kind: OMPD_parallel));
1288
1289 UncondBrInst::Create(Target: SeqAfterBB, InsertBefore: SeqAfterIP.getNodeParent());
1290
1291 LLVM_DEBUG(dbgs() << TAG << "After sequential inlining " << *OuterFn
1292 << "\n");
1293 };
1294
1295 // Helper to merge the __kmpc_fork_call calls in MergableCIs. They are all
1296 // contained in BB and only separated by instructions that can be
1297 // redundantly executed in parallel. The block BB is split before the first
1298 // call (in MergableCIs) and after the last so the entire region we merge
1299 // into a single parallel region is contained in a single basic block
1300 // without any other instructions. We use the OpenMPIRBuilder to outline
1301 // that block and call the resulting function via __kmpc_fork_call.
1302 auto Merge = [&](const SmallVectorImpl<CallInst *> &MergableCIs,
1303 BasicBlock *BB) {
1304 // TODO: Change the interface to allow single CIs expanded, e.g, to
1305 // include an outer loop.
1306 assert(MergableCIs.size() > 1 && "Assumed multiple mergable CIs");
1307
1308 auto Remark = [&](OptimizationRemark OR) {
1309 OR << "Parallel region merged with parallel region"
1310 << (MergableCIs.size() > 2 ? "s" : "") << " at ";
1311 for (auto *CI : llvm::drop_begin(RangeOrContainer: MergableCIs)) {
1312 OR << ore::NV("OpenMPParallelMerge", CI->getDebugLoc());
1313 if (CI != MergableCIs.back())
1314 OR << ", ";
1315 }
1316 return OR << ".";
1317 };
1318
1319 emitRemark<OptimizationRemark>(I: MergableCIs.front(), RemarkName: "OMP150", RemarkCB&: Remark);
1320
1321 Function *OriginalFn = BB->getParent();
1322 LLVM_DEBUG(dbgs() << TAG << "Merge " << MergableCIs.size()
1323 << " parallel regions in " << OriginalFn->getName()
1324 << "\n");
1325
1326 // Isolate the calls to merge in a separate block.
1327 EndBB = SplitBlock(Old: BB, SplitPt: MergableCIs.back()->getNextNode(), DT, LI);
1328 BasicBlock *AfterBB =
1329 SplitBlock(Old: EndBB, SplitPt: &*EndBB->getFirstInsertionPt(), DT, LI);
1330 StartBB = SplitBlock(Old: BB, SplitPt: MergableCIs.front(), DT, LI, MSSAU: nullptr,
1331 BBName: "omp.par.merged");
1332
1333 assert(BB->getUniqueSuccessor() == StartBB && "Expected a different CFG");
1334 const DebugLoc DL = BB->getTerminator()->getDebugLoc();
1335 BB->getTerminator()->eraseFromParent();
1336
1337 // Create sequential regions for sequential instructions that are
1338 // in-between mergable parallel regions.
1339 for (auto *It = MergableCIs.begin(), *End = MergableCIs.end() - 1;
1340 It != End; ++It) {
1341 Instruction *ForkCI = *It;
1342 Instruction *NextForkCI = *(It + 1);
1343
1344 // Continue if there are not in-between instructions.
1345 if (ForkCI->getNextNode() == NextForkCI)
1346 continue;
1347
1348 CreateSequentialRegion(OriginalFn, BB, ForkCI->getNextNode(),
1349 NextForkCI->getPrevNode());
1350 }
1351
1352 OpenMPIRBuilder::LocationDescription Loc(BB->end(), DL);
1353 IRBuilder<>::InsertPoint AllocaIP(
1354 OriginalFn->getEntryBlock().getFirstInsertionPt());
1355 // Create the merged parallel region with default proc binding, to
1356 // avoid overriding binding settings, and without explicit cancellation.
1357 OpenMPIRBuilder::InsertPointTy AfterIP =
1358 cantFail(ValOrErr: OMPInfoCache.OMPBuilder.createParallel(
1359 Loc, AllocaIP, /* DeallocBlocks */ {}, BodyGenCB, PrivCB, FiniCB,
1360 IfCondition: nullptr, NumThreads: nullptr, ProcBind: OMP_PROC_BIND_default,
1361 /* IsCancellable */ false));
1362 UncondBrInst::Create(Target: AfterBB, InsertBefore: AfterIP.getNodeParent());
1363
1364 // Perform the actual outlining.
1365 OMPInfoCache.OMPBuilder.finalize(Fn: OriginalFn);
1366
1367 Function *OutlinedFn = MergableCIs.front()->getCaller();
1368 std::optional<uint64_t> WrapperCount =
1369 getMergedWrapperEntryCount(ForkCalls: MergableCIs, CallbackOpNo: CallbackCalleeOperand);
1370 // Leave the wrapper unprofiled when no callback has an entry count.
1371 if (WrapperCount)
1372 OutlinedFn->setEntryCount(Count: *WrapperCount);
1373 // Only sample PGO treats a profiled caller with no callsite weight as
1374 // cold. Instrumentation profiles derive that count from the entry count.
1375 const bool SampleProfile =
1376 moduleHasSampleProfile(M: *OriginalFn->getParent());
1377
1378 // Replace the __kmpc_fork_call calls with direct calls to the outlined
1379 // callbacks.
1380 SmallVector<Value *, 8> Args;
1381 for (auto *CI : MergableCIs) {
1382 Value *Callee = CI->getArgOperand(i: CallbackCalleeOperand);
1383 FunctionType *FT = OMPInfoCache.OMPBuilder.ParallelTask;
1384 Args.clear();
1385 Args.push_back(Elt: OutlinedFn->getArg(i: 0));
1386 Args.push_back(Elt: OutlinedFn->getArg(i: 1));
1387 for (unsigned U = CallbackFirstArgOperand, E = CI->arg_size(); U < E;
1388 ++U)
1389 Args.push_back(Elt: CI->getArgOperand(i: U));
1390
1391 CallInst *NewCI =
1392 CallInst::Create(Ty: FT, Func: Callee, Args, NameStr: "", InsertBefore: CI->getIterator());
1393 if (CI->getDebugLoc())
1394 NewCI->setDebugLoc(CI->getDebugLoc());
1395 // Each body runs once per wrapper entry. Without a callsite weight,
1396 // sample PGO treats these calls as cold.
1397 if (WrapperCount && SampleProfile) {
1398 setFittedBranchWeights(I&: *NewCI, Weights: {*WrapperCount},
1399 /*IsExpected=*/false);
1400 }
1401
1402 // Forward parameter attributes from the callback to the callee.
1403 for (unsigned U = CallbackFirstArgOperand, E = CI->arg_size(); U < E;
1404 ++U)
1405 for (const Attribute &A : CI->getAttributes().getParamAttrs(ArgNo: U))
1406 NewCI->addParamAttr(
1407 ArgNo: U - (CallbackFirstArgOperand - CallbackCalleeOperand), Attr: A);
1408
1409 // Emit an explicit barrier to replace the implicit fork-join barrier.
1410 if (CI != MergableCIs.back()) {
1411 // TODO: Remove barrier if the merged parallel region includes the
1412 // 'nowait' clause.
1413 cantFail(ValOrErr: OMPInfoCache.OMPBuilder.createBarrier(
1414 Loc: {NewCI->getNextNode()->getIterator(), NewCI->getDebugLoc()},
1415 Kind: OMPD_parallel));
1416 }
1417
1418 CI->eraseFromParent();
1419 }
1420
1421 assert(OutlinedFn != OriginalFn && "Outlining failed");
1422 CGUpdater.registerOutlinedFunction(OriginalFn&: *OriginalFn, NewFn&: *OutlinedFn);
1423 CGUpdater.reanalyzeFunction(Fn&: *OriginalFn);
1424
1425 NumOpenMPParallelRegionsMerged += MergableCIs.size();
1426
1427 return true;
1428 };
1429
1430 // Helper function that identifes sequences of
1431 // __kmpc_fork_call uses in a basic block.
1432 auto DetectPRsCB = [&](Use &U, Function &F) {
1433 CallInst *CI = getCallIfRegularCall(U, RFI: &RFI);
1434 BB2PRMap[CI->getParent()].insert(Ptr: CI);
1435
1436 return false;
1437 };
1438
1439 BB2PRMap.clear();
1440 RFI.foreachUse(SCC, CB: DetectPRsCB);
1441 SmallVector<SmallVector<CallInst *, 4>, 4> MergableCIsVector;
1442 // Find mergable parallel regions within a basic block that are
1443 // safe to merge, that is any in-between instructions can safely
1444 // execute in parallel after merging.
1445 // TODO: support merging across basic-blocks.
1446 for (auto &It : BB2PRMap) {
1447 auto &CIs = It.getSecond();
1448 if (CIs.size() < 2)
1449 continue;
1450
1451 BasicBlock *BB = It.getFirst();
1452 SmallVector<CallInst *, 4> MergableCIs;
1453
1454 /// Returns true if the instruction is mergable, false otherwise.
1455 /// A terminator instruction is unmergable by definition since merging
1456 /// works within a BB. Instructions before the mergable region are
1457 /// mergable if they are not calls to OpenMP runtime functions that may
1458 /// set different execution parameters for subsequent parallel regions.
1459 /// Instructions in-between parallel regions are mergable if they are not
1460 /// calls to any non-intrinsic function since that may call a non-mergable
1461 /// OpenMP runtime function.
1462 auto IsMergable = [&](Instruction &I, bool IsBeforeMergableRegion) {
1463 // We do not merge across BBs, hence return false (unmergable) if the
1464 // instruction is a terminator.
1465 if (I.isTerminator())
1466 return false;
1467
1468 if (!isa<CallInst>(Val: &I))
1469 return true;
1470
1471 CallInst *CI = cast<CallInst>(Val: &I);
1472 if (IsBeforeMergableRegion) {
1473 Function *CalledFunction = CI->getCalledFunction();
1474 if (!CalledFunction)
1475 return false;
1476 // Return false (unmergable) if the call before the parallel
1477 // region calls an explicit affinity (proc_bind) or number of
1478 // threads (num_threads) compiler-generated function. Those settings
1479 // may be incompatible with following parallel regions.
1480 // TODO: ICV tracking to detect compatibility.
1481 for (const auto &RFI : UnmergableCallsInfo) {
1482 if (CalledFunction == RFI.Declaration)
1483 return false;
1484 }
1485 } else {
1486 // Return false (unmergable) if there is a call instruction
1487 // in-between parallel regions when it is not an intrinsic. It
1488 // may call an unmergable OpenMP runtime function in its callpath.
1489 // TODO: Keep track of possible OpenMP calls in the callpath.
1490 if (!isa<IntrinsicInst>(Val: CI))
1491 return false;
1492 }
1493
1494 return true;
1495 };
1496 // Find maximal number of parallel region CIs that are safe to merge.
1497 for (auto It = BB->begin(), End = BB->end(); It != End;) {
1498 Instruction &I = *It;
1499 ++It;
1500
1501 if (CIs.count(Ptr: &I)) {
1502 MergableCIs.push_back(Elt: cast<CallInst>(Val: &I));
1503 continue;
1504 }
1505
1506 // Continue expanding if the instruction is mergable.
1507 if (IsMergable(I, MergableCIs.empty()))
1508 continue;
1509
1510 // Forward the instruction iterator to skip the next parallel region
1511 // since there is an unmergable instruction which can affect it.
1512 for (; It != End; ++It) {
1513 Instruction &SkipI = *It;
1514 if (CIs.count(Ptr: &SkipI)) {
1515 LLVM_DEBUG(dbgs() << TAG << "Skip parallel region " << SkipI
1516 << " due to " << I << "\n");
1517 ++It;
1518 break;
1519 }
1520 }
1521
1522 // Store mergable regions found.
1523 if (MergableCIs.size() > 1) {
1524 MergableCIsVector.push_back(Elt: MergableCIs);
1525 LLVM_DEBUG(dbgs() << TAG << "Found " << MergableCIs.size()
1526 << " parallel regions in block " << BB->getName()
1527 << " of function " << BB->getParent()->getName()
1528 << "\n";);
1529 }
1530
1531 MergableCIs.clear();
1532 }
1533
1534 if (!MergableCIsVector.empty()) {
1535 Changed = true;
1536
1537 for (auto &MergableCIs : MergableCIsVector)
1538 Merge(MergableCIs, BB);
1539 MergableCIsVector.clear();
1540 }
1541 }
1542
1543 if (Changed) {
1544 /// Re-collect use for fork calls, emitted barrier calls, and
1545 /// any emitted master/end_master calls.
1546 OMPInfoCache.recollectUsesForFunction(RTF: OMPRTL___kmpc_fork_call);
1547 OMPInfoCache.recollectUsesForFunction(RTF: OMPRTL___kmpc_barrier);
1548 OMPInfoCache.recollectUsesForFunction(RTF: OMPRTL___kmpc_master);
1549 OMPInfoCache.recollectUsesForFunction(RTF: OMPRTL___kmpc_end_master);
1550 }
1551
1552 return Changed;
1553 }
1554
1555 /// Try to delete parallel regions if possible.
1556 bool deleteParallelRegions() {
1557 const unsigned CallbackCalleeOperand = 2;
1558
1559 OMPInformationCache::RuntimeFunctionInfo &RFI =
1560 OMPInfoCache.RFIs[OMPRTL___kmpc_fork_call];
1561
1562 if (!RFI.Declaration)
1563 return false;
1564
1565 bool Changed = false;
1566 auto DeleteCallCB = [&](Use &U, Function &) {
1567 CallInst *CI = getCallIfRegularCall(U);
1568 if (!CI)
1569 return false;
1570 auto *Fn = dyn_cast<Function>(
1571 Val: CI->getArgOperand(i: CallbackCalleeOperand)->stripPointerCasts());
1572 if (!Fn)
1573 return false;
1574 if (!Fn->onlyReadsMemory())
1575 return false;
1576 if (!Fn->hasFnAttribute(Kind: Attribute::WillReturn))
1577 return false;
1578
1579 LLVM_DEBUG(dbgs() << TAG << "Delete read-only parallel region in "
1580 << CI->getCaller()->getName() << "\n");
1581
1582 auto Remark = [&](OptimizationRemark OR) {
1583 return OR << "Removing parallel region with no side-effects.";
1584 };
1585 emitRemark<OptimizationRemark>(I: CI, RemarkName: "OMP160", RemarkCB&: Remark);
1586
1587 CI->eraseFromParent();
1588 Changed = true;
1589 ++NumOpenMPParallelRegionsDeleted;
1590 return true;
1591 };
1592
1593 RFI.foreachUse(SCC, CB: DeleteCallCB);
1594
1595 return Changed;
1596 }
1597
1598 /// Try to eliminate runtime calls by reusing existing ones.
1599 bool deduplicateRuntimeCalls() {
1600 bool Changed = false;
1601
1602 RuntimeFunction DeduplicableRuntimeCallIDs[] = {
1603 OMPRTL_omp_get_num_threads,
1604 OMPRTL_omp_in_parallel,
1605 OMPRTL_omp_get_cancellation,
1606 OMPRTL_omp_get_supported_active_levels,
1607 OMPRTL_omp_get_level,
1608 OMPRTL_omp_get_ancestor_thread_num,
1609 OMPRTL_omp_get_team_size,
1610 OMPRTL_omp_get_active_level,
1611 OMPRTL_omp_in_final,
1612 OMPRTL_omp_get_proc_bind,
1613 OMPRTL_omp_get_num_places,
1614 OMPRTL_omp_get_num_procs,
1615 OMPRTL_omp_get_place_num,
1616 OMPRTL_omp_get_partition_num_places,
1617 OMPRTL_omp_get_partition_place_nums};
1618
1619 // Global-tid is handled separately.
1620 SmallSetVector<Value *, 16> GTIdArgs;
1621 collectGlobalThreadIdArguments(GTIdArgs);
1622 LLVM_DEBUG(dbgs() << TAG << "Found " << GTIdArgs.size()
1623 << " global thread ID arguments\n");
1624
1625 for (Function *F : SCC) {
1626 for (auto DeduplicableRuntimeCallID : DeduplicableRuntimeCallIDs)
1627 Changed |= deduplicateRuntimeCalls(
1628 F&: *F, RFI&: OMPInfoCache.RFIs[DeduplicableRuntimeCallID]);
1629
1630 // __kmpc_global_thread_num is special as we can replace it with an
1631 // argument in enough cases to make it worth trying.
1632 Value *GTIdArg = nullptr;
1633 for (Argument &Arg : F->args())
1634 if (GTIdArgs.count(key: &Arg)) {
1635 GTIdArg = &Arg;
1636 break;
1637 }
1638 Changed |= deduplicateRuntimeCalls(
1639 F&: *F, RFI&: OMPInfoCache.RFIs[OMPRTL___kmpc_global_thread_num], ReplVal: GTIdArg);
1640 }
1641
1642 return Changed;
1643 }
1644
1645 /// Tries to remove known runtime symbols that are optional from the module.
1646 bool removeRuntimeSymbols() {
1647 // The RPC client symbol is defined in `libc` and indicates that something
1648 // required an RPC server. If its users were all optimized out then we can
1649 // safely remove it.
1650 // TODO: This should be somewhere more common in the future.
1651 if (GlobalVariable *GV = M.getNamedGlobal(Name: "__llvm_rpc_client")) {
1652 if (GV->hasNUsesOrMore(N: 1))
1653 return false;
1654
1655 GV->replaceAllUsesWith(V: PoisonValue::get(T: GV->getType()));
1656 GV->eraseFromParent();
1657 return true;
1658 }
1659 return false;
1660 }
1661
1662 /// Tries to hide the latency of runtime calls that involve host to
1663 /// device memory transfers by splitting them into their "issue" and "wait"
1664 /// versions. The "issue" is moved upwards as much as possible. The "wait" is
1665 /// moved downards as much as possible. The "issue" issues the memory transfer
1666 /// asynchronously, returning a handle. The "wait" waits in the returned
1667 /// handle for the memory transfer to finish.
1668 bool hideMemTransfersLatency() {
1669 auto &RFI = OMPInfoCache.RFIs[OMPRTL___tgt_target_data_begin_mapper];
1670 bool Changed = false;
1671 auto SplitMemTransfers = [&](Use &U, Function &Decl) {
1672 auto *RTCall = getCallIfRegularCall(U, RFI: &RFI);
1673 if (!RTCall)
1674 return false;
1675
1676 OffloadArray OffloadArrays[3];
1677 if (!getValuesInOffloadArrays(RuntimeCall&: *RTCall, OAs: OffloadArrays))
1678 return false;
1679
1680 LLVM_DEBUG(dumpValuesInOffloadArrays(OffloadArrays));
1681
1682 // TODO: Check if can be moved upwards.
1683 bool WasSplit = false;
1684 Instruction *WaitMovementPoint = canBeMovedDownwards(RuntimeCall&: *RTCall);
1685 if (WaitMovementPoint)
1686 WasSplit = splitTargetDataBeginRTC(RuntimeCall&: *RTCall, WaitMovementPoint&: *WaitMovementPoint);
1687
1688 Changed |= WasSplit;
1689 return WasSplit;
1690 };
1691 if (OMPInfoCache.runtimeFnsAvailable(
1692 Fns: {OMPRTL___tgt_target_data_begin_mapper_issue,
1693 OMPRTL___tgt_target_data_begin_mapper_wait}))
1694 RFI.foreachUse(SCC, CB: SplitMemTransfers);
1695
1696 return Changed;
1697 }
1698
1699 void analysisGlobalization() {
1700 auto &RFI = OMPInfoCache.RFIs[OMPRTL___kmpc_alloc_shared];
1701
1702 auto CheckGlobalization = [&](Use &U, Function &Decl) {
1703 if (CallInst *CI = getCallIfRegularCall(U, RFI: &RFI)) {
1704 auto Remark = [&](OptimizationRemarkMissed ORM) {
1705 return ORM
1706 << "Found thread data sharing on the GPU. "
1707 << "Expect degraded performance due to data globalization.";
1708 };
1709 emitRemark<OptimizationRemarkMissed>(I: CI, RemarkName: "OMP112", RemarkCB&: Remark);
1710 }
1711
1712 return false;
1713 };
1714
1715 RFI.foreachUse(SCC, CB: CheckGlobalization);
1716 }
1717
1718 /// Maps the values stored in the offload arrays passed as arguments to
1719 /// \p RuntimeCall into the offload arrays in \p OAs.
1720 bool getValuesInOffloadArrays(CallInst &RuntimeCall,
1721 MutableArrayRef<OffloadArray> OAs) {
1722 assert(OAs.size() == 3 && "Need space for three offload arrays!");
1723
1724 // A runtime call that involves memory offloading looks something like:
1725 // call void @__tgt_target_data_begin_mapper(arg0, arg1,
1726 // i8** %offload_baseptrs, i8** %offload_ptrs, i64* %offload_sizes,
1727 // ...)
1728 // So, the idea is to access the allocas that allocate space for these
1729 // offload arrays, offload_baseptrs, offload_ptrs, offload_sizes.
1730 // Therefore:
1731 // i8** %offload_baseptrs.
1732 Value *BasePtrsArg =
1733 RuntimeCall.getArgOperand(i: OffloadArray::BasePtrsArgNum);
1734 // i8** %offload_ptrs.
1735 Value *PtrsArg = RuntimeCall.getArgOperand(i: OffloadArray::PtrsArgNum);
1736 // i8** %offload_sizes.
1737 Value *SizesArg = RuntimeCall.getArgOperand(i: OffloadArray::SizesArgNum);
1738
1739 // Get values stored in **offload_baseptrs.
1740 auto *V = getUnderlyingObject(V: BasePtrsArg);
1741 if (!isa<AllocaInst>(Val: V))
1742 return false;
1743 auto *BasePtrsArray = cast<AllocaInst>(Val: V);
1744 if (!OAs[0].initialize(Array&: *BasePtrsArray, Before&: RuntimeCall))
1745 return false;
1746
1747 // Get values stored in **offload_baseptrs.
1748 V = getUnderlyingObject(V: PtrsArg);
1749 if (!isa<AllocaInst>(Val: V))
1750 return false;
1751 auto *PtrsArray = cast<AllocaInst>(Val: V);
1752 if (!OAs[1].initialize(Array&: *PtrsArray, Before&: RuntimeCall))
1753 return false;
1754
1755 // Get values stored in **offload_sizes.
1756 V = getUnderlyingObject(V: SizesArg);
1757 // If it's a [constant] global array don't analyze it.
1758 if (isa<GlobalValue>(Val: V))
1759 return isa<Constant>(Val: V);
1760 if (!isa<AllocaInst>(Val: V))
1761 return false;
1762
1763 auto *SizesArray = cast<AllocaInst>(Val: V);
1764 if (!OAs[2].initialize(Array&: *SizesArray, Before&: RuntimeCall))
1765 return false;
1766
1767 return true;
1768 }
1769
1770 /// Prints the values in the OffloadArrays \p OAs using LLVM_DEBUG.
1771 /// For now this is a way to test that the function getValuesInOffloadArrays
1772 /// is working properly.
1773 /// TODO: Move this to a unittest when unittests are available for OpenMPOpt.
1774 void dumpValuesInOffloadArrays(ArrayRef<OffloadArray> OAs) {
1775 assert(OAs.size() == 3 && "There are three offload arrays to debug!");
1776
1777 LLVM_DEBUG(dbgs() << TAG << " Successfully got offload values:\n");
1778 std::string ValuesStr;
1779 raw_string_ostream Printer(ValuesStr);
1780 std::string Separator = " --- ";
1781
1782 for (auto *BP : OAs[0].StoredValues) {
1783 BP->print(O&: Printer);
1784 Printer << Separator;
1785 }
1786 LLVM_DEBUG(dbgs() << "\t\toffload_baseptrs: " << ValuesStr << "\n");
1787 ValuesStr.clear();
1788
1789 for (auto *P : OAs[1].StoredValues) {
1790 P->print(O&: Printer);
1791 Printer << Separator;
1792 }
1793 LLVM_DEBUG(dbgs() << "\t\toffload_ptrs: " << ValuesStr << "\n");
1794 ValuesStr.clear();
1795
1796 for (auto *S : OAs[2].StoredValues) {
1797 S->print(O&: Printer);
1798 Printer << Separator;
1799 }
1800 LLVM_DEBUG(dbgs() << "\t\toffload_sizes: " << ValuesStr << "\n");
1801 }
1802
1803 /// Returns the instruction where the "wait" counterpart \p RuntimeCall can be
1804 /// moved. Returns nullptr if the movement is not possible, or not worth it.
1805 Instruction *canBeMovedDownwards(CallInst &RuntimeCall) {
1806 // FIXME: This traverses only the BasicBlock where RuntimeCall is.
1807 // Make it traverse the CFG.
1808
1809 Instruction *CurrentI = &RuntimeCall;
1810 bool IsWorthIt = false;
1811 while ((CurrentI = CurrentI->getNextNode())) {
1812
1813 // TODO: Once we detect the regions to be offloaded we should use the
1814 // alias analysis manager to check if CurrentI may modify one of
1815 // the offloaded regions.
1816 if (CurrentI->mayHaveSideEffects() || CurrentI->mayReadFromMemory()) {
1817 if (IsWorthIt)
1818 return CurrentI;
1819
1820 return nullptr;
1821 }
1822
1823 // FIXME: For now if we move it over anything without side effect
1824 // is worth it.
1825 IsWorthIt = true;
1826 }
1827
1828 // Return end of BasicBlock.
1829 return RuntimeCall.getParent()->getTerminator();
1830 }
1831
1832 /// Splits \p RuntimeCall into its "issue" and "wait" counterparts.
1833 bool splitTargetDataBeginRTC(CallInst &RuntimeCall,
1834 Instruction &WaitMovementPoint) {
1835 // Create stack allocated handle (__tgt_async_info) at the beginning of the
1836 // function. Used for storing information of the async transfer, allowing to
1837 // wait on it later.
1838 auto &IRBuilder = OMPInfoCache.OMPBuilder;
1839 Function *F = RuntimeCall.getCaller();
1840 BasicBlock &Entry = F->getEntryBlock();
1841 IRBuilder.Builder.SetInsertPoint(Entry.getFirstNonPHIOrDbgOrAlloca());
1842 Value *Handle = IRBuilder.Builder.CreateAlloca(
1843 Ty: IRBuilder.AsyncInfo, /*ArraySize=*/nullptr, Name: "handle");
1844 Handle =
1845 IRBuilder.Builder.CreateAddrSpaceCast(V: Handle, DestTy: IRBuilder.AsyncInfoPtr);
1846
1847 // Add "issue" runtime call declaration:
1848 // declare %struct.tgt_async_info @__tgt_target_data_begin_issue(i64, i32,
1849 // i8**, i8**, i64*, i64*)
1850 FunctionCallee IssueDecl = IRBuilder.getOrCreateRuntimeFunction(
1851 M, FnID: OMPRTL___tgt_target_data_begin_mapper_issue);
1852
1853 // Change RuntimeCall call site for its asynchronous version.
1854 SmallVector<Value *, 16> Args;
1855 for (auto &Arg : RuntimeCall.args())
1856 Args.push_back(Elt: Arg.get());
1857 Args.push_back(Elt: Handle);
1858
1859 CallInst *IssueCallsite = CallInst::Create(Func: IssueDecl, Args, /*NameStr=*/"",
1860 InsertBefore: RuntimeCall.getIterator());
1861 OMPInfoCache.setCallingConvention(Callee: IssueDecl, CI: IssueCallsite);
1862 RuntimeCall.eraseFromParent();
1863
1864 // Add "wait" runtime call declaration:
1865 // declare void @__tgt_target_data_begin_wait(i64, %struct.__tgt_async_info)
1866 FunctionCallee WaitDecl = IRBuilder.getOrCreateRuntimeFunction(
1867 M, FnID: OMPRTL___tgt_target_data_begin_mapper_wait);
1868
1869 Value *WaitParams[2] = {
1870 IssueCallsite->getArgOperand(
1871 i: OffloadArray::DeviceIDArgNum), // device_id.
1872 Handle // handle to wait on.
1873 };
1874 CallInst *WaitCallsite = CallInst::Create(
1875 Func: WaitDecl, Args: WaitParams, /*NameStr=*/"", InsertBefore: WaitMovementPoint.getIterator());
1876 OMPInfoCache.setCallingConvention(Callee: WaitDecl, CI: WaitCallsite);
1877
1878 return true;
1879 }
1880
1881 static Value *combinedIdentStruct(Value *CurrentIdent, Value *NextIdent,
1882 bool GlobalOnly, bool &SingleChoice) {
1883 if (CurrentIdent == NextIdent)
1884 return CurrentIdent;
1885
1886 // TODO: Figure out how to actually combine multiple debug locations. For
1887 // now we just keep an existing one if there is a single choice.
1888 if (!GlobalOnly || isa<GlobalValue>(Val: NextIdent)) {
1889 SingleChoice = !CurrentIdent;
1890 return NextIdent;
1891 }
1892 return nullptr;
1893 }
1894
1895 /// Return an `struct ident_t*` value that represents the ones used in the
1896 /// calls of \p RFI inside of \p F. If \p GlobalOnly is true, we will not
1897 /// return a local `struct ident_t*`. For now, if we cannot find a suitable
1898 /// return value we create one from scratch. We also do not yet combine
1899 /// information, e.g., the source locations, see combinedIdentStruct.
1900 Value *
1901 getCombinedIdentFromCallUsesIn(OMPInformationCache::RuntimeFunctionInfo &RFI,
1902 Function &F, bool GlobalOnly) {
1903 bool SingleChoice = true;
1904 Value *Ident = nullptr;
1905 auto CombineIdentStruct = [&](Use &U, Function &Caller) {
1906 CallInst *CI = getCallIfRegularCall(U, RFI: &RFI);
1907 if (!CI || &F != &Caller)
1908 return false;
1909 Ident = combinedIdentStruct(CurrentIdent: Ident, NextIdent: CI->getArgOperand(i: 0),
1910 /* GlobalOnly */ true, SingleChoice);
1911 return false;
1912 };
1913 RFI.foreachUse(SCC, CB: CombineIdentStruct);
1914
1915 if (!Ident || !SingleChoice) {
1916 // The IRBuilder uses the insertion block to get to the module, this is
1917 // unfortunate but we work around it for now. No instruction is emitted
1918 // here, so there is no debug location to preserve.
1919 if (!OMPInfoCache.OMPBuilder.getInsertionPoint().isValid())
1920 OMPInfoCache.OMPBuilder.updateToLocation(
1921 Loc: {F.getEntryBlock().begin(), DebugLoc()});
1922 // Create a fallback location if non was found.
1923 // TODO: Use the debug locations of the calls instead.
1924 uint32_t SrcLocStrSize;
1925 Constant *Loc =
1926 OMPInfoCache.OMPBuilder.getOrCreateDefaultSrcLocStr(SrcLocStrSize);
1927 Ident = OMPInfoCache.OMPBuilder.getOrCreateIdent(SrcLocStr: Loc, SrcLocStrSize);
1928 }
1929 return Ident;
1930 }
1931
1932 /// Try to eliminate calls of \p RFI in \p F by reusing an existing one or
1933 /// \p ReplVal if given.
1934 bool deduplicateRuntimeCalls(Function &F,
1935 OMPInformationCache::RuntimeFunctionInfo &RFI,
1936 Value *ReplVal = nullptr) {
1937 auto *UV = RFI.getUseVector(F);
1938 if (!UV || UV->size() + (ReplVal != nullptr) < 2)
1939 return false;
1940
1941 LLVM_DEBUG(
1942 dbgs() << TAG << "Deduplicate " << UV->size() << " uses of " << RFI.Name
1943 << (ReplVal ? " with an existing value\n" : "\n") << "\n");
1944
1945 assert((!ReplVal || (isa<Argument>(ReplVal) &&
1946 cast<Argument>(ReplVal)->getParent() == &F)) &&
1947 "Unexpected replacement value!");
1948
1949 // TODO: Use dominance to find a good position instead.
1950 auto CanBeMoved = [this](CallBase &CB) {
1951 unsigned NumArgs = CB.arg_size();
1952 if (NumArgs == 0)
1953 return true;
1954 if (CB.getArgOperand(i: 0)->getType() != OMPInfoCache.OMPBuilder.IdentPtr)
1955 return false;
1956 for (unsigned U = 1; U < NumArgs; ++U)
1957 if (isa<Instruction>(Val: CB.getArgOperand(i: U)))
1958 return false;
1959 return true;
1960 };
1961
1962 if (!ReplVal) {
1963 auto *DT =
1964 OMPInfoCache.getAnalysisResultForFunction<DominatorTreeAnalysis>(F);
1965 if (!DT)
1966 return false;
1967 Instruction *IP = nullptr;
1968 for (Use *U : *UV) {
1969 if (CallInst *CI = getCallIfRegularCall(U&: *U, RFI: &RFI)) {
1970 if (IP)
1971 IP = DT->findNearestCommonDominator(I1: IP, I2: CI);
1972 else
1973 IP = CI;
1974 if (!CanBeMoved(*CI))
1975 continue;
1976 if (!ReplVal)
1977 ReplVal = CI;
1978 }
1979 }
1980 if (!ReplVal)
1981 return false;
1982 assert(IP && "Expected insertion point!");
1983 cast<Instruction>(Val: ReplVal)->moveBefore(InsertPos: IP->getIterator());
1984 }
1985
1986 // If we use a call as a replacement value we need to make sure the ident is
1987 // valid at the new location. For now we just pick a global one, either
1988 // existing and used by one of the calls, or created from scratch.
1989 if (CallBase *CI = dyn_cast<CallBase>(Val: ReplVal)) {
1990 if (!CI->arg_empty() &&
1991 CI->getArgOperand(i: 0)->getType() == OMPInfoCache.OMPBuilder.IdentPtr) {
1992 Value *Ident = getCombinedIdentFromCallUsesIn(RFI, F,
1993 /* GlobalOnly */ true);
1994 CI->setArgOperand(i: 0, v: Ident);
1995 }
1996 }
1997
1998 bool Changed = false;
1999 auto ReplaceAndDeleteCB = [&](Use &U, Function &Caller) {
2000 CallInst *CI = getCallIfRegularCall(U, RFI: &RFI);
2001 if (!CI || CI == ReplVal || &F != &Caller)
2002 return false;
2003 assert(CI->getCaller() == &F && "Unexpected call!");
2004
2005 auto Remark = [&](OptimizationRemark OR) {
2006 return OR << "OpenMP runtime call "
2007 << ore::NV("OpenMPOptRuntime", RFI.Name) << " deduplicated.";
2008 };
2009 if (CI->getDebugLoc())
2010 emitRemark<OptimizationRemark>(I: CI, RemarkName: "OMP170", RemarkCB&: Remark);
2011 else
2012 emitRemark<OptimizationRemark>(F: &F, RemarkName: "OMP170", RemarkCB&: Remark);
2013
2014 CI->replaceAllUsesWith(V: ReplVal);
2015 CI->eraseFromParent();
2016 ++NumOpenMPRuntimeCallsDeduplicated;
2017 Changed = true;
2018 return true;
2019 };
2020 RFI.foreachUse(SCC, CB: ReplaceAndDeleteCB);
2021
2022 return Changed;
2023 }
2024
2025 /// Collect arguments that represent the global thread id in \p GTIdArgs.
2026 void collectGlobalThreadIdArguments(SmallSetVector<Value *, 16> &GTIdArgs) {
2027 // TODO: Below we basically perform a fixpoint iteration with a pessimistic
2028 // initialization. We could define an AbstractAttribute instead and
2029 // run the Attributor here once it can be run as an SCC pass.
2030
2031 // Helper to check the argument \p ArgNo at all call sites of \p F for
2032 // a GTId.
2033 auto CallArgOpIsGTId = [&](Function &F, unsigned ArgNo, CallInst &RefCI) {
2034 if (!F.hasLocalLinkage())
2035 return false;
2036 for (Use &U : F.uses()) {
2037 if (CallInst *CI = getCallIfRegularCall(U)) {
2038 Value *ArgOp = CI->getArgOperand(i: ArgNo);
2039 if (CI == &RefCI || GTIdArgs.count(key: ArgOp) ||
2040 getCallIfRegularCall(
2041 V&: *ArgOp, RFI: &OMPInfoCache.RFIs[OMPRTL___kmpc_global_thread_num]))
2042 continue;
2043 }
2044 return false;
2045 }
2046 return true;
2047 };
2048
2049 // Helper to identify uses of a GTId as GTId arguments.
2050 auto AddUserArgs = [&](Value &GTId) {
2051 for (Use &U : GTId.uses())
2052 if (CallInst *CI = dyn_cast<CallInst>(Val: U.getUser()))
2053 if (CI->isArgOperand(U: &U))
2054 if (Function *Callee = CI->getCalledFunction())
2055 if (CallArgOpIsGTId(*Callee, U.getOperandNo(), *CI))
2056 GTIdArgs.insert(X: Callee->getArg(i: U.getOperandNo()));
2057 };
2058
2059 // The argument users of __kmpc_global_thread_num calls are GTIds.
2060 OMPInformationCache::RuntimeFunctionInfo &GlobThreadNumRFI =
2061 OMPInfoCache.RFIs[OMPRTL___kmpc_global_thread_num];
2062
2063 GlobThreadNumRFI.foreachUse(SCC, CB: [&](Use &U, Function &F) {
2064 if (CallInst *CI = getCallIfRegularCall(U, RFI: &GlobThreadNumRFI))
2065 AddUserArgs(*CI);
2066 return false;
2067 });
2068
2069 // Transitively search for more arguments by looking at the users of the
2070 // ones we know already. During the search the GTIdArgs vector is extended
2071 // so we cannot cache the size nor can we use a range based for.
2072 for (unsigned U = 0; U < GTIdArgs.size(); ++U)
2073 AddUserArgs(*GTIdArgs[U]);
2074 }
2075
2076 /// Kernel (=GPU) optimizations and utility functions
2077 ///
2078 ///{{
2079
2080 /// Cache to remember the unique kernel for a function.
2081 DenseMap<Function *, std::optional<Kernel>> UniqueKernelMap;
2082
2083 /// Find the unique kernel that will execute \p F, if any.
2084 Kernel getUniqueKernelFor(Function &F);
2085
2086 /// Find the unique kernel that will execute \p I, if any.
2087 Kernel getUniqueKernelFor(Instruction &I) {
2088 return getUniqueKernelFor(F&: *I.getFunction());
2089 }
2090
2091 /// Rewrite the device (=GPU) code state machine create in non-SPMD mode in
2092 /// the cases we can avoid taking the address of a function.
2093 bool rewriteDeviceCodeStateMachine();
2094
2095 /// In SPMD kernels the parallel data-sharing wrapper passed to
2096 /// __kmpc_parallel_60 is never used by the runtime; null it out so the dead
2097 /// wrapper (and any LDS it references) can be removed.
2098 bool removeSPMDParallelWrappers();
2099
2100 ///
2101 ///}}
2102
2103 /// Emit a remark generically
2104 ///
2105 /// This template function can be used to generically emit a remark. The
2106 /// RemarkKind should be one of the following:
2107 /// - OptimizationRemark to indicate a successful optimization attempt
2108 /// - OptimizationRemarkMissed to report a failed optimization attempt
2109 /// - OptimizationRemarkAnalysis to provide additional information about an
2110 /// optimization attempt
2111 ///
2112 /// The remark is built using a callback function provided by the caller that
2113 /// takes a RemarkKind as input and returns a RemarkKind.
2114 template <typename RemarkKind, typename RemarkCallBack>
2115 void emitRemark(Instruction *I, StringRef RemarkName,
2116 RemarkCallBack &&RemarkCB) const {
2117 Function *F = I->getParent()->getParent();
2118 auto &ORE = OREGetter(F);
2119
2120 if (RemarkName.starts_with(Prefix: "OMP"))
2121 ORE.emit([&]() {
2122 return RemarkCB(RemarkKind(DEBUG_TYPE, RemarkName, I))
2123 << " [" << RemarkName << "]";
2124 });
2125 else
2126 ORE.emit(
2127 [&]() { return RemarkCB(RemarkKind(DEBUG_TYPE, RemarkName, I)); });
2128 }
2129
2130 /// Emit a remark on a function.
2131 template <typename RemarkKind, typename RemarkCallBack>
2132 void emitRemark(Function *F, StringRef RemarkName,
2133 RemarkCallBack &&RemarkCB) const {
2134 auto &ORE = OREGetter(F);
2135
2136 if (RemarkName.starts_with(Prefix: "OMP"))
2137 ORE.emit([&]() {
2138 return RemarkCB(RemarkKind(DEBUG_TYPE, RemarkName, F))
2139 << " [" << RemarkName << "]";
2140 });
2141 else
2142 ORE.emit(
2143 [&]() { return RemarkCB(RemarkKind(DEBUG_TYPE, RemarkName, F)); });
2144 }
2145
2146 /// The underlying module.
2147 Module &M;
2148
2149 /// The SCC we are operating on.
2150 SmallVectorImpl<Function *> &SCC;
2151
2152 /// Callback to update the call graph, the first argument is a removed call,
2153 /// the second an optional replacement call.
2154 CallGraphUpdater &CGUpdater;
2155
2156 /// Callback to get an OptimizationRemarkEmitter from a Function *
2157 OptimizationRemarkGetter OREGetter;
2158
2159 /// OpenMP-specific information cache. Also Used for Attributor runs.
2160 OMPInformationCache &OMPInfoCache;
2161
2162 /// Attributor instance.
2163 Attributor &A;
2164
2165 /// Helper function to run Attributor on SCC.
2166 bool runAttributor(bool IsModulePass) {
2167 if (SCC.empty())
2168 return false;
2169
2170 registerAAs(IsModulePass);
2171
2172 ChangeStatus Changed = A.run();
2173
2174 LLVM_DEBUG(dbgs() << "[Attributor] Done with " << SCC.size()
2175 << " functions, result: " << Changed << ".\n");
2176
2177 if (Changed == ChangeStatus::CHANGED)
2178 OMPInfoCache.invalidateAnalyses();
2179
2180 return Changed == ChangeStatus::CHANGED;
2181 }
2182
2183 void registerFoldRuntimeCall(RuntimeFunction RF);
2184
2185 /// Populate the Attributor with abstract attribute opportunities in the
2186 /// functions.
2187 void registerAAs(bool IsModulePass);
2188
2189public:
2190 /// Callback to register AAs for live functions, including internal functions
2191 /// marked live during the traversal.
2192 static void registerAAsForFunction(Attributor &A, const Function &F);
2193};
2194
2195Kernel OpenMPOpt::getUniqueKernelFor(Function &F) {
2196 if (OMPInfoCache.CGSCC && !OMPInfoCache.CGSCC->empty() &&
2197 !OMPInfoCache.CGSCC->contains(key: &F))
2198 return nullptr;
2199
2200 // Use a scope to keep the lifetime of the CachedKernel short.
2201 {
2202 std::optional<Kernel> &CachedKernel = UniqueKernelMap[&F];
2203 if (CachedKernel)
2204 return *CachedKernel;
2205
2206 // TODO: We should use an AA to create an (optimistic and callback
2207 // call-aware) call graph. For now we stick to simple patterns that
2208 // are less powerful, basically the worst fixpoint.
2209 if (isOpenMPKernel(Fn&: F)) {
2210 CachedKernel = Kernel(&F);
2211 return *CachedKernel;
2212 }
2213
2214 CachedKernel = nullptr;
2215 if (!F.hasLocalLinkage()) {
2216
2217 // See https://openmp.llvm.org/remarks/OptimizationRemarks.html
2218 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
2219 return ORA << "Potentially unknown OpenMP target region caller.";
2220 };
2221 emitRemark<OptimizationRemarkAnalysis>(F: &F, RemarkName: "OMP100", RemarkCB&: Remark);
2222
2223 return nullptr;
2224 }
2225 }
2226
2227 auto GetUniqueKernelForUse = [&](const Use &U) -> Kernel {
2228 if (auto *Cmp = dyn_cast<ICmpInst>(Val: U.getUser())) {
2229 // Allow use in equality comparisons.
2230 if (Cmp->isEquality())
2231 return getUniqueKernelFor(I&: *Cmp);
2232 return nullptr;
2233 }
2234 if (auto *CB = dyn_cast<CallBase>(Val: U.getUser())) {
2235 // Allow direct calls.
2236 if (CB->isCallee(U: &U))
2237 return getUniqueKernelFor(I&: *CB);
2238
2239 OMPInformationCache::RuntimeFunctionInfo &KernelParallelRFI =
2240 OMPInfoCache.RFIs[OMPRTL___kmpc_parallel_60];
2241 // Allow the use in __kmpc_parallel_60 calls.
2242 if (OpenMPOpt::getCallIfRegularCall(V&: *U.getUser(), RFI: &KernelParallelRFI))
2243 return getUniqueKernelFor(I&: *CB);
2244 return nullptr;
2245 }
2246 // Disallow every other use.
2247 return nullptr;
2248 };
2249
2250 // TODO: In the future we want to track more than just a unique kernel.
2251 SmallPtrSet<Kernel, 2> PotentialKernels;
2252 OMPInformationCache::foreachUse(F, CB: [&](const Use &U) {
2253 PotentialKernels.insert(Ptr: GetUniqueKernelForUse(U));
2254 });
2255
2256 Kernel K = nullptr;
2257 if (PotentialKernels.size() == 1)
2258 K = *PotentialKernels.begin();
2259
2260 // Cache the result.
2261 UniqueKernelMap[&F] = K;
2262
2263 return K;
2264}
2265
2266bool OpenMPOpt::rewriteDeviceCodeStateMachine() {
2267 OMPInformationCache::RuntimeFunctionInfo &KernelParallelRFI =
2268 OMPInfoCache.RFIs[OMPRTL___kmpc_parallel_60];
2269
2270 bool Changed = false;
2271 if (!KernelParallelRFI)
2272 return Changed;
2273
2274 // If we have disabled state machine changes, exit
2275 if (DisableOpenMPOptStateMachineRewrite)
2276 return Changed;
2277
2278 for (Function *F : SCC) {
2279
2280 // Check if the function is a use in a __kmpc_parallel_60 call at
2281 // all.
2282 bool UnknownUse = false;
2283 bool KernelParallelUse = false;
2284 unsigned NumDirectCalls = 0;
2285
2286 SmallVector<Use *, 2> ToBeReplacedStateMachineUses;
2287 OMPInformationCache::foreachUse(F&: *F, CB: [&](Use &U) {
2288 if (auto *CB = dyn_cast<CallBase>(Val: U.getUser()))
2289 if (CB->isCallee(U: &U)) {
2290 ++NumDirectCalls;
2291 return;
2292 }
2293
2294 if (isa<ICmpInst>(Val: U.getUser())) {
2295 ToBeReplacedStateMachineUses.push_back(Elt: &U);
2296 return;
2297 }
2298
2299 // Find wrapper functions that represent parallel kernels.
2300 CallInst *CI =
2301 OpenMPOpt::getCallIfRegularCall(V&: *U.getUser(), RFI: &KernelParallelRFI);
2302 const unsigned int WrapperFunctionArgNo = 6;
2303 if (!KernelParallelUse && CI &&
2304 CI->getArgOperandNo(U: &U) == WrapperFunctionArgNo) {
2305 KernelParallelUse = true;
2306 ToBeReplacedStateMachineUses.push_back(Elt: &U);
2307 return;
2308 }
2309 UnknownUse = true;
2310 });
2311
2312 // Do not emit a remark if we haven't seen a __kmpc_parallel_60
2313 // use.
2314 if (!KernelParallelUse)
2315 continue;
2316
2317 // If this ever hits, we should investigate.
2318 // TODO: Checking the number of uses is not a necessary restriction and
2319 // should be lifted.
2320 if (UnknownUse || NumDirectCalls != 1 ||
2321 ToBeReplacedStateMachineUses.size() > 2) {
2322 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
2323 return ORA << "Parallel region is used in "
2324 << (UnknownUse ? "unknown" : "unexpected")
2325 << " ways. Will not attempt to rewrite the state machine.";
2326 };
2327 emitRemark<OptimizationRemarkAnalysis>(F, RemarkName: "OMP101", RemarkCB&: Remark);
2328 continue;
2329 }
2330
2331 // Even if we have __kmpc_parallel_60 calls, we (for now) give
2332 // up if the function is not called from a unique kernel.
2333 Kernel K = getUniqueKernelFor(F&: *F);
2334 if (!K) {
2335 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
2336 return ORA << "Parallel region is not called from a unique kernel. "
2337 "Will not attempt to rewrite the state machine.";
2338 };
2339 emitRemark<OptimizationRemarkAnalysis>(F, RemarkName: "OMP102", RemarkCB&: Remark);
2340 continue;
2341 }
2342
2343 // We now know F is a parallel body function called only from the kernel K.
2344 // We also identified the state machine uses in which we replace the
2345 // function pointer by a new global symbol for identification purposes. This
2346 // ensures only direct calls to the function are left.
2347
2348 Module &M = *F->getParent();
2349 Type *Int8Ty = Type::getInt8Ty(C&: M.getContext());
2350
2351 auto *ID = new GlobalVariable(
2352 M, Int8Ty, /* isConstant */ true, GlobalValue::PrivateLinkage,
2353 UndefValue::get(T: Int8Ty), F->getName() + ".ID");
2354
2355 for (Use *U : ToBeReplacedStateMachineUses)
2356 U->set(ConstantExpr::getPointerBitCastOrAddrSpaceCast(
2357 C: ID, Ty: U->get()->getType()));
2358
2359 ++NumOpenMPParallelRegionsReplacedInGPUStateMachine;
2360
2361 Changed = true;
2362 }
2363
2364 return Changed;
2365}
2366
2367bool OpenMPOpt::removeSPMDParallelWrappers() {
2368 // Nothing to clean up unless we SPMD-ized at least one kernel.
2369 if (OMPInfoCache.SPMDizedKernels.empty())
2370 return false;
2371
2372 OMPInformationCache::RuntimeFunctionInfo &KernelParallelRFI =
2373 OMPInfoCache.RFIs[OMPRTL___kmpc_parallel_60];
2374 if (!KernelParallelRFI || !KernelParallelRFI.Declaration)
2375 return false;
2376
2377 constexpr unsigned WrapperFunctionArgNo = 6;
2378 bool Changed = false;
2379 for (User *U : KernelParallelRFI.Declaration->users()) {
2380 auto *CI = dyn_cast<CallInst>(Val: U);
2381 if (!CI || CI->getCalledOperand() != KernelParallelRFI.Declaration ||
2382 CI->arg_size() <= WrapperFunctionArgNo)
2383 continue;
2384
2385 Value *Wrapper = CI->getArgOperand(i: WrapperFunctionArgNo);
2386 if (isa<ConstantPointerNull>(Val: Wrapper))
2387 continue;
2388
2389 // Only drop the wrapper for a parallel region reached from a single kernel
2390 // that we transformed to SPMD mode. A region also reachable from a
2391 // generic-mode kernel still needs its wrapper for that kernel's state
2392 // machine, and getUniqueKernelFor conservatively bails on such shared
2393 // regions. (Mirrors the unique-kernel requirement in
2394 // rewriteDeviceCodeStateMachine.)
2395 Kernel K = getUniqueKernelFor(F&: *CI->getFunction());
2396 if (!K || !OMPInfoCache.SPMDizedKernels.contains(Ptr: K))
2397 continue;
2398
2399 CI->setArgOperand(
2400 i: WrapperFunctionArgNo,
2401 v: ConstantPointerNull::get(T: cast<PointerType>(Val: Wrapper->getType())));
2402 Changed = true;
2403 }
2404
2405 return Changed;
2406}
2407
2408/// Abstract Attribute for tracking ICV values.
2409struct AAICVTracker : public StateWrapper<BooleanState, AbstractAttribute> {
2410 using Base = StateWrapper<BooleanState, AbstractAttribute>;
2411 AAICVTracker(const IRPosition &IRP, Attributor &A) : Base(IRP) {}
2412
2413 /// Returns true if value is assumed to be tracked.
2414 bool isAssumedTracked() const { return getAssumed(); }
2415
2416 /// Returns true if value is known to be tracked.
2417 bool isKnownTracked() const { return getAssumed(); }
2418
2419 /// Create an abstract attribute biew for the position \p IRP.
2420 static AAICVTracker &createForPosition(const IRPosition &IRP, Attributor &A);
2421
2422 /// Return the value with which \p I can be replaced for specific \p ICV.
2423 virtual std::optional<Value *> getReplacementValue(InternalControlVar ICV,
2424 const Instruction *I,
2425 Attributor &A) const {
2426 return std::nullopt;
2427 }
2428
2429 /// Return an assumed unique ICV value if a single candidate is found. If
2430 /// there cannot be one, return a nullptr. If it is not clear yet, return
2431 /// std::nullopt.
2432 virtual std::optional<Value *>
2433 getUniqueReplacementValue(InternalControlVar ICV) const = 0;
2434
2435 // Currently only nthreads is being tracked.
2436 // this array will only grow with time.
2437 InternalControlVar TrackableICVs[1] = {ICV_nthreads};
2438
2439 /// See AbstractAttribute::getName()
2440 StringRef getName() const override { return "AAICVTracker"; }
2441
2442 /// See AbstractAttribute::getIdAddr()
2443 const char *getIdAddr() const override { return &ID; }
2444
2445 /// This function should return true if the type of the \p AA is AAICVTracker
2446 static bool classof(const AbstractAttribute *AA) {
2447 return (AA->getIdAddr() == &ID);
2448 }
2449
2450 static const char ID;
2451};
2452
2453struct AAICVTrackerFunction : public AAICVTracker {
2454 AAICVTrackerFunction(const IRPosition &IRP, Attributor &A)
2455 : AAICVTracker(IRP, A) {}
2456
2457 // FIXME: come up with better string.
2458 const std::string getAsStr(Attributor *) const override {
2459 return "ICVTrackerFunction";
2460 }
2461
2462 // FIXME: come up with some stats.
2463 void trackStatistics() const override {}
2464
2465 /// We don't manifest anything for this AA.
2466 ChangeStatus manifest(Attributor &A) override {
2467 return ChangeStatus::UNCHANGED;
2468 }
2469
2470 // Map of ICV to their values at specific program point.
2471 EnumeratedArray<DenseMap<Instruction *, Value *>, InternalControlVar,
2472 InternalControlVar::ICV___last>
2473 ICVReplacementValuesMap;
2474
2475 ChangeStatus updateImpl(Attributor &A) override {
2476 ChangeStatus HasChanged = ChangeStatus::UNCHANGED;
2477
2478 Function *F = getAnchorScope();
2479
2480 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
2481
2482 for (InternalControlVar ICV : TrackableICVs) {
2483 auto &SetterRFI = OMPInfoCache.RFIs[OMPInfoCache.ICVs[ICV].Setter];
2484
2485 auto &ValuesMap = ICVReplacementValuesMap[ICV];
2486 auto TrackValues = [&](Use &U, Function &) {
2487 CallInst *CI = OpenMPOpt::getCallIfRegularCall(U);
2488 if (!CI)
2489 return false;
2490
2491 // FIXME: handle setters with more that 1 arguments.
2492 /// Track new value.
2493 if (ValuesMap.insert(KV: std::make_pair(x&: CI, y: CI->getArgOperand(i: 0))).second)
2494 HasChanged = ChangeStatus::CHANGED;
2495
2496 return false;
2497 };
2498
2499 auto CallCheck = [&](Instruction &I) {
2500 std::optional<Value *> ReplVal = getValueForCall(A, I, ICV);
2501 if (ReplVal && ValuesMap.insert(KV: std::make_pair(x: &I, y&: *ReplVal)).second)
2502 HasChanged = ChangeStatus::CHANGED;
2503
2504 return true;
2505 };
2506
2507 // Track all changes of an ICV.
2508 SetterRFI.foreachUse(CB: TrackValues, F);
2509
2510 bool UsedAssumedInformation = false;
2511 A.checkForAllInstructions(Pred: CallCheck, QueryingAA: *this, Opcodes: {Instruction::Call},
2512 UsedAssumedInformation,
2513 /* CheckBBLivenessOnly */ true);
2514
2515 /// TODO: Figure out a way to avoid adding entry in
2516 /// ICVReplacementValuesMap
2517 Instruction *Entry = &F->getEntryBlock().front();
2518 if (HasChanged == ChangeStatus::CHANGED)
2519 ValuesMap.try_emplace(Key: Entry);
2520 }
2521
2522 return HasChanged;
2523 }
2524
2525 /// Helper to check if \p I is a call and get the value for it if it is
2526 /// unique.
2527 std::optional<Value *> getValueForCall(Attributor &A, const Instruction &I,
2528 InternalControlVar &ICV) const {
2529
2530 const auto *CB = dyn_cast<CallBase>(Val: &I);
2531 if (!CB || CB->hasFnAttr(Kind: "no_openmp") ||
2532 CB->hasFnAttr(Kind: "no_openmp_routines") ||
2533 CB->hasFnAttr(Kind: "no_openmp_constructs"))
2534 return std::nullopt;
2535
2536 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
2537 auto &GetterRFI = OMPInfoCache.RFIs[OMPInfoCache.ICVs[ICV].Getter];
2538 auto &SetterRFI = OMPInfoCache.RFIs[OMPInfoCache.ICVs[ICV].Setter];
2539 Function *CalledFunction = CB->getCalledFunction();
2540
2541 // Indirect call, assume ICV changes.
2542 if (CalledFunction == nullptr)
2543 return nullptr;
2544 if (CalledFunction == GetterRFI.Declaration)
2545 return std::nullopt;
2546 if (CalledFunction == SetterRFI.Declaration) {
2547 if (ICVReplacementValuesMap[ICV].count(Val: &I))
2548 return ICVReplacementValuesMap[ICV].lookup(Val: &I);
2549
2550 return nullptr;
2551 }
2552
2553 // Since we don't know, assume it changes the ICV.
2554 if (CalledFunction->isDeclaration())
2555 return nullptr;
2556
2557 const auto *ICVTrackingAA = A.getAAFor<AAICVTracker>(
2558 QueryingAA: *this, IRP: IRPosition::callsite_returned(CB: *CB), DepClass: DepClassTy::REQUIRED);
2559
2560 if (ICVTrackingAA->isAssumedTracked()) {
2561 std::optional<Value *> URV =
2562 ICVTrackingAA->getUniqueReplacementValue(ICV);
2563 if (!URV || (*URV && AA::isValidAtPosition(VAC: AA::ValueAndContext(**URV, I),
2564 InfoCache&: OMPInfoCache)))
2565 return URV;
2566 }
2567
2568 // If we don't know, assume it changes.
2569 return nullptr;
2570 }
2571
2572 // We don't check unique value for a function, so return std::nullopt.
2573 std::optional<Value *>
2574 getUniqueReplacementValue(InternalControlVar ICV) const override {
2575 return std::nullopt;
2576 }
2577
2578 /// Return the value with which \p I can be replaced for specific \p ICV.
2579 std::optional<Value *> getReplacementValue(InternalControlVar ICV,
2580 const Instruction *I,
2581 Attributor &A) const override {
2582 const auto &ValuesMap = ICVReplacementValuesMap[ICV];
2583 if (ValuesMap.count(Val: I))
2584 return ValuesMap.lookup(Val: I);
2585
2586 SmallVector<const Instruction *, 16> Worklist;
2587 SmallPtrSet<const Instruction *, 16> Visited;
2588 Worklist.push_back(Elt: I);
2589
2590 std::optional<Value *> ReplVal;
2591
2592 while (!Worklist.empty()) {
2593 const Instruction *CurrInst = Worklist.pop_back_val();
2594 if (!Visited.insert(Ptr: CurrInst).second)
2595 continue;
2596
2597 const BasicBlock *CurrBB = CurrInst->getParent();
2598
2599 // Go up and look for all potential setters/calls that might change the
2600 // ICV.
2601 while ((CurrInst = CurrInst->getPrevNode())) {
2602 if (ValuesMap.count(Val: CurrInst)) {
2603 std::optional<Value *> NewReplVal = ValuesMap.lookup(Val: CurrInst);
2604 // Unknown value, track new.
2605 if (!ReplVal) {
2606 ReplVal = NewReplVal;
2607 break;
2608 }
2609
2610 // If we found a new value, we can't know the icv value anymore.
2611 if (NewReplVal)
2612 if (ReplVal != NewReplVal)
2613 return nullptr;
2614
2615 break;
2616 }
2617
2618 std::optional<Value *> NewReplVal = getValueForCall(A, I: *CurrInst, ICV);
2619 if (!NewReplVal)
2620 continue;
2621
2622 // Unknown value, track new.
2623 if (!ReplVal) {
2624 ReplVal = NewReplVal;
2625 break;
2626 }
2627
2628 // if (NewReplVal.hasValue())
2629 // We found a new value, we can't know the icv value anymore.
2630 if (ReplVal != NewReplVal)
2631 return nullptr;
2632 }
2633
2634 // If we are in the same BB and we have a value, we are done.
2635 if (CurrBB == I->getParent() && ReplVal)
2636 return ReplVal;
2637
2638 // Go through all predecessors and add terminators for analysis.
2639 for (const BasicBlock *Pred : predecessors(BB: CurrBB))
2640 if (const Instruction *Terminator = Pred->getTerminator())
2641 Worklist.push_back(Elt: Terminator);
2642 }
2643
2644 return ReplVal;
2645 }
2646};
2647
2648struct AAICVTrackerFunctionReturned : AAICVTracker {
2649 AAICVTrackerFunctionReturned(const IRPosition &IRP, Attributor &A)
2650 : AAICVTracker(IRP, A) {}
2651
2652 // FIXME: come up with better string.
2653 const std::string getAsStr(Attributor *) const override {
2654 return "ICVTrackerFunctionReturned";
2655 }
2656
2657 // FIXME: come up with some stats.
2658 void trackStatistics() const override {}
2659
2660 /// We don't manifest anything for this AA.
2661 ChangeStatus manifest(Attributor &A) override {
2662 return ChangeStatus::UNCHANGED;
2663 }
2664
2665 // Map of ICV to their values at specific program point.
2666 EnumeratedArray<std::optional<Value *>, InternalControlVar,
2667 InternalControlVar::ICV___last>
2668 ICVReplacementValuesMap;
2669
2670 /// Return the value with which \p I can be replaced for specific \p ICV.
2671 std::optional<Value *>
2672 getUniqueReplacementValue(InternalControlVar ICV) const override {
2673 return ICVReplacementValuesMap[ICV];
2674 }
2675
2676 ChangeStatus updateImpl(Attributor &A) override {
2677 ChangeStatus Changed = ChangeStatus::UNCHANGED;
2678 const auto *ICVTrackingAA = A.getAAFor<AAICVTracker>(
2679 QueryingAA: *this, IRP: IRPosition::function(F: *getAnchorScope()), DepClass: DepClassTy::REQUIRED);
2680
2681 if (!ICVTrackingAA->isAssumedTracked())
2682 return indicatePessimisticFixpoint();
2683
2684 for (InternalControlVar ICV : TrackableICVs) {
2685 std::optional<Value *> &ReplVal = ICVReplacementValuesMap[ICV];
2686 std::optional<Value *> UniqueICVValue;
2687
2688 auto CheckReturnInst = [&](Instruction &I) {
2689 std::optional<Value *> NewReplVal =
2690 ICVTrackingAA->getReplacementValue(ICV, I: &I, A);
2691
2692 // If we found a second ICV value there is no unique returned value.
2693 if (UniqueICVValue && UniqueICVValue != NewReplVal)
2694 return false;
2695
2696 UniqueICVValue = NewReplVal;
2697
2698 return true;
2699 };
2700
2701 bool UsedAssumedInformation = false;
2702 if (!A.checkForAllInstructions(Pred: CheckReturnInst, QueryingAA: *this, Opcodes: {Instruction::Ret},
2703 UsedAssumedInformation,
2704 /* CheckBBLivenessOnly */ true))
2705 UniqueICVValue = nullptr;
2706
2707 if (UniqueICVValue == ReplVal)
2708 continue;
2709
2710 ReplVal = UniqueICVValue;
2711 Changed = ChangeStatus::CHANGED;
2712 }
2713
2714 return Changed;
2715 }
2716};
2717
2718struct AAICVTrackerCallSite : AAICVTracker {
2719 AAICVTrackerCallSite(const IRPosition &IRP, Attributor &A)
2720 : AAICVTracker(IRP, A) {}
2721
2722 void initialize(Attributor &A) override {
2723 assert(getAnchorScope() && "Expected anchor function");
2724
2725 // We only initialize this AA for getters, so we need to know which ICV it
2726 // gets.
2727 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
2728 for (InternalControlVar ICV : TrackableICVs) {
2729 auto ICVInfo = OMPInfoCache.ICVs[ICV];
2730 auto &Getter = OMPInfoCache.RFIs[ICVInfo.Getter];
2731 if (Getter.Declaration == getAssociatedFunction()) {
2732 AssociatedICV = ICVInfo.Kind;
2733 return;
2734 }
2735 }
2736
2737 /// Unknown ICV.
2738 indicatePessimisticFixpoint();
2739 }
2740
2741 ChangeStatus manifest(Attributor &A) override {
2742 if (!ReplVal || !*ReplVal)
2743 return ChangeStatus::UNCHANGED;
2744
2745 A.changeAfterManifest(IRP: IRPosition::inst(I: *getCtxI()), NV&: **ReplVal);
2746 A.deleteAfterManifest(I&: *getCtxI());
2747
2748 return ChangeStatus::CHANGED;
2749 }
2750
2751 // FIXME: come up with better string.
2752 const std::string getAsStr(Attributor *) const override {
2753 return "ICVTrackerCallSite";
2754 }
2755
2756 // FIXME: come up with some stats.
2757 void trackStatistics() const override {}
2758
2759 InternalControlVar AssociatedICV;
2760 std::optional<Value *> ReplVal;
2761
2762 ChangeStatus updateImpl(Attributor &A) override {
2763 const auto *ICVTrackingAA = A.getAAFor<AAICVTracker>(
2764 QueryingAA: *this, IRP: IRPosition::function(F: *getAnchorScope()), DepClass: DepClassTy::REQUIRED);
2765
2766 // We don't have any information, so we assume it changes the ICV.
2767 if (!ICVTrackingAA->isAssumedTracked())
2768 return indicatePessimisticFixpoint();
2769
2770 std::optional<Value *> NewReplVal =
2771 ICVTrackingAA->getReplacementValue(ICV: AssociatedICV, I: getCtxI(), A);
2772
2773 if (ReplVal == NewReplVal)
2774 return ChangeStatus::UNCHANGED;
2775
2776 ReplVal = NewReplVal;
2777 return ChangeStatus::CHANGED;
2778 }
2779
2780 // Return the value with which associated value can be replaced for specific
2781 // \p ICV.
2782 std::optional<Value *>
2783 getUniqueReplacementValue(InternalControlVar ICV) const override {
2784 return ReplVal;
2785 }
2786};
2787
2788struct AAICVTrackerCallSiteReturned : AAICVTracker {
2789 AAICVTrackerCallSiteReturned(const IRPosition &IRP, Attributor &A)
2790 : AAICVTracker(IRP, A) {}
2791
2792 // FIXME: come up with better string.
2793 const std::string getAsStr(Attributor *) const override {
2794 return "ICVTrackerCallSiteReturned";
2795 }
2796
2797 // FIXME: come up with some stats.
2798 void trackStatistics() const override {}
2799
2800 /// We don't manifest anything for this AA.
2801 ChangeStatus manifest(Attributor &A) override {
2802 return ChangeStatus::UNCHANGED;
2803 }
2804
2805 // Map of ICV to their values at specific program point.
2806 EnumeratedArray<std::optional<Value *>, InternalControlVar,
2807 InternalControlVar::ICV___last>
2808 ICVReplacementValuesMap;
2809
2810 /// Return the value with which associated value can be replaced for specific
2811 /// \p ICV.
2812 std::optional<Value *>
2813 getUniqueReplacementValue(InternalControlVar ICV) const override {
2814 return ICVReplacementValuesMap[ICV];
2815 }
2816
2817 ChangeStatus updateImpl(Attributor &A) override {
2818 ChangeStatus Changed = ChangeStatus::UNCHANGED;
2819 const auto *ICVTrackingAA = A.getAAFor<AAICVTracker>(
2820 QueryingAA: *this, IRP: IRPosition::returned(F: *getAssociatedFunction()),
2821 DepClass: DepClassTy::REQUIRED);
2822
2823 // We don't have any information, so we assume it changes the ICV.
2824 if (!ICVTrackingAA->isAssumedTracked())
2825 return indicatePessimisticFixpoint();
2826
2827 for (InternalControlVar ICV : TrackableICVs) {
2828 std::optional<Value *> &ReplVal = ICVReplacementValuesMap[ICV];
2829 std::optional<Value *> NewReplVal =
2830 ICVTrackingAA->getUniqueReplacementValue(ICV);
2831
2832 if (ReplVal == NewReplVal)
2833 continue;
2834
2835 ReplVal = NewReplVal;
2836 Changed = ChangeStatus::CHANGED;
2837 }
2838 return Changed;
2839 }
2840};
2841
2842/// Determines if \p BB exits the function unconditionally itself or reaches a
2843/// block that does through only unique successors.
2844static bool hasFunctionEndAsUniqueSuccessor(const BasicBlock *BB) {
2845 if (succ_empty(BB))
2846 return true;
2847 const BasicBlock *const Successor = BB->getUniqueSuccessor();
2848 if (!Successor)
2849 return false;
2850 return hasFunctionEndAsUniqueSuccessor(BB: Successor);
2851}
2852
2853struct AAExecutionDomainFunction : public AAExecutionDomain {
2854 AAExecutionDomainFunction(const IRPosition &IRP, Attributor &A)
2855 : AAExecutionDomain(IRP, A) {}
2856
2857 ~AAExecutionDomainFunction() override { delete RPOT; }
2858
2859 void initialize(Attributor &A) override {
2860 Function *F = getAnchorScope();
2861 assert(F && "Expected anchor function");
2862 RPOT = new ReversePostOrderTraversal<Function *>(F);
2863 }
2864
2865 const std::string getAsStr(Attributor *) const override {
2866 unsigned TotalBlocks = 0, InitialThreadBlocks = 0, AlignedBlocks = 0;
2867 for (auto &It : BEDMap) {
2868 if (!It.getFirst())
2869 continue;
2870 TotalBlocks++;
2871 InitialThreadBlocks += It.getSecond().IsExecutedByInitialThreadOnly;
2872 AlignedBlocks += It.getSecond().IsReachedFromAlignedBarrierOnly &&
2873 It.getSecond().IsReachingAlignedBarrierOnly;
2874 }
2875 return "[AAExecutionDomain] " + std::to_string(val: InitialThreadBlocks) + "/" +
2876 std::to_string(val: AlignedBlocks) + " of " +
2877 std::to_string(val: TotalBlocks) +
2878 " executed by initial thread / aligned";
2879 }
2880
2881 /// See AbstractAttribute::trackStatistics().
2882 void trackStatistics() const override {}
2883
2884 ChangeStatus manifest(Attributor &A) override {
2885 LLVM_DEBUG({
2886 for (const BasicBlock &BB : *getAnchorScope()) {
2887 if (!isExecutedByInitialThreadOnly(BB))
2888 continue;
2889 dbgs() << TAG << " Basic block @" << getAnchorScope()->getName() << " "
2890 << BB.getName() << " is executed by a single thread.\n";
2891 }
2892 });
2893
2894 ChangeStatus Changed = ChangeStatus::UNCHANGED;
2895
2896 if (DisableOpenMPOptBarrierElimination)
2897 return Changed;
2898
2899 SmallPtrSet<CallBase *, 16> DeletedBarriers;
2900 auto HandleAlignedBarrier = [&](CallBase *CB) {
2901 const ExecutionDomainTy &ED = CB ? CEDMap[{CB, PRE}] : BEDMap[nullptr];
2902 if (!ED.IsReachedFromAlignedBarrierOnly ||
2903 ED.EncounteredNonLocalSideEffect)
2904 return;
2905 if (!ED.EncounteredAssumes.empty() && !A.isModulePass())
2906 return;
2907
2908 // We can remove this barrier, if it is one, or aligned barriers reaching
2909 // the kernel end (if CB is nullptr). Aligned barriers reaching the kernel
2910 // end should only be removed if the kernel end is their unique successor;
2911 // otherwise, they may have side-effects that aren't accounted for in the
2912 // kernel end in their other successors. If those barriers have other
2913 // barriers reaching them, those can be transitively removed as well as
2914 // long as the kernel end is also their unique successor.
2915 if (CB) {
2916 DeletedBarriers.insert(Ptr: CB);
2917 A.deleteAfterManifest(I&: *CB);
2918 ++NumBarriersEliminated;
2919 Changed = ChangeStatus::CHANGED;
2920 } else if (!ED.AlignedBarriers.empty()) {
2921 Changed = ChangeStatus::CHANGED;
2922 SmallVector<CallBase *> Worklist(ED.AlignedBarriers.begin(),
2923 ED.AlignedBarriers.end());
2924 SmallSetVector<CallBase *, 16> Visited;
2925 while (!Worklist.empty()) {
2926 CallBase *LastCB = Worklist.pop_back_val();
2927 if (!Visited.insert(X: LastCB))
2928 continue;
2929 if (LastCB->getFunction() != getAnchorScope())
2930 continue;
2931 if (!hasFunctionEndAsUniqueSuccessor(BB: LastCB->getParent()))
2932 continue;
2933 if (!DeletedBarriers.count(Ptr: LastCB)) {
2934 ++NumBarriersEliminated;
2935 A.deleteAfterManifest(I&: *LastCB);
2936 continue;
2937 }
2938 // The final aligned barrier (LastCB) reaching the kernel end was
2939 // removed already. This means we can go one step further and remove
2940 // the barriers encoutered last before (LastCB).
2941 const ExecutionDomainTy &LastED = CEDMap[{LastCB, PRE}];
2942 Worklist.append(in_start: LastED.AlignedBarriers.begin(),
2943 in_end: LastED.AlignedBarriers.end());
2944 }
2945 }
2946
2947 // If we actually eliminated a barrier we need to eliminate the associated
2948 // llvm.assumes as well to avoid creating UB.
2949 if (!ED.EncounteredAssumes.empty() && (CB || !ED.AlignedBarriers.empty()))
2950 for (auto *AssumeCB : ED.EncounteredAssumes)
2951 A.deleteAfterManifest(I&: *AssumeCB);
2952 };
2953
2954 for (auto *CB : AlignedBarriers)
2955 HandleAlignedBarrier(CB);
2956
2957 // Handle the "kernel end barrier" for kernels too.
2958 if (omp::isOpenMPKernel(Fn&: *getAnchorScope()))
2959 HandleAlignedBarrier(nullptr);
2960
2961 return Changed;
2962 }
2963
2964 bool isNoOpFence(const FenceInst &FI) const override {
2965 return getState().isValidState() && !NonNoOpFences.count(Ptr: &FI);
2966 }
2967
2968 /// Merge barrier and assumption information from \p PredED into the successor
2969 /// \p ED.
2970 void
2971 mergeInPredecessorBarriersAndAssumptions(Attributor &A, ExecutionDomainTy &ED,
2972 const ExecutionDomainTy &PredED);
2973
2974 /// Merge all information from \p PredED into the successor \p ED. If
2975 /// \p InitialEdgeOnly is set, only the initial edge will enter the block
2976 /// represented by \p ED from this predecessor.
2977 bool mergeInPredecessor(Attributor &A, ExecutionDomainTy &ED,
2978 const ExecutionDomainTy &PredED,
2979 bool InitialEdgeOnly = false);
2980
2981 /// Accumulate information for the entry block in \p EntryBBED.
2982 bool handleCallees(Attributor &A, ExecutionDomainTy &EntryBBED);
2983
2984 /// See AbstractAttribute::updateImpl.
2985 ChangeStatus updateImpl(Attributor &A) override;
2986
2987 /// Query interface, see AAExecutionDomain
2988 ///{
2989 bool isExecutedByInitialThreadOnly(const BasicBlock &BB) const override {
2990 if (!isValidState())
2991 return false;
2992 assert(BB.getParent() == getAnchorScope() && "Block is out of scope!");
2993 return BEDMap.lookup(Val: &BB).IsExecutedByInitialThreadOnly;
2994 }
2995
2996 bool isExecutedInAlignedRegion(Attributor &A,
2997 const Instruction &I) const override {
2998 assert(I.getFunction() == getAnchorScope() &&
2999 "Instruction is out of scope!");
3000 if (!isValidState())
3001 return false;
3002
3003 bool ForwardIsOk = true;
3004 const Instruction *CurI;
3005
3006 // Check forward until a call or the block end is reached.
3007 CurI = &I;
3008 do {
3009 auto *CB = dyn_cast<CallBase>(Val: CurI);
3010 if (!CB)
3011 continue;
3012 if (CB != &I && AlignedBarriers.contains(key: const_cast<CallBase *>(CB)))
3013 return true;
3014 const auto &It = CEDMap.find(Val: {CB, PRE});
3015 if (It == CEDMap.end())
3016 continue;
3017 if (!It->getSecond().IsReachingAlignedBarrierOnly)
3018 ForwardIsOk = false;
3019 break;
3020 } while ((CurI = CurI->getNextNode()));
3021
3022 if (!CurI && !BEDMap.lookup(Val: I.getParent()).IsReachingAlignedBarrierOnly)
3023 ForwardIsOk = false;
3024
3025 // Check backward until a call or the block beginning is reached.
3026 CurI = &I;
3027 do {
3028 auto *CB = dyn_cast<CallBase>(Val: CurI);
3029 if (!CB)
3030 continue;
3031 if (CB != &I && AlignedBarriers.contains(key: const_cast<CallBase *>(CB)))
3032 return true;
3033 const auto &It = CEDMap.find(Val: {CB, POST});
3034 if (It == CEDMap.end())
3035 continue;
3036 if (It->getSecond().IsReachedFromAlignedBarrierOnly)
3037 break;
3038 return false;
3039 } while ((CurI = CurI->getPrevNode()));
3040
3041 // Delayed decision on the forward pass to allow aligned barrier detection
3042 // in the backwards traversal.
3043 if (!ForwardIsOk)
3044 return false;
3045
3046 if (!CurI) {
3047 const BasicBlock *BB = I.getParent();
3048 if (BB == &BB->getParent()->getEntryBlock())
3049 return BEDMap.lookup(Val: nullptr).IsReachedFromAlignedBarrierOnly;
3050 if (!llvm::all_of(Range: predecessors(BB), P: [&](const BasicBlock *PredBB) {
3051 return BEDMap.lookup(Val: PredBB).IsReachedFromAlignedBarrierOnly;
3052 })) {
3053 return false;
3054 }
3055 }
3056
3057 // On neither traversal we found a anything but aligned barriers.
3058 return true;
3059 }
3060
3061 ExecutionDomainTy getExecutionDomain(const BasicBlock &BB) const override {
3062 assert(isValidState() &&
3063 "No request should be made against an invalid state!");
3064 return BEDMap.lookup(Val: &BB);
3065 }
3066 std::pair<ExecutionDomainTy, ExecutionDomainTy>
3067 getExecutionDomain(const CallBase &CB) const override {
3068 assert(isValidState() &&
3069 "No request should be made against an invalid state!");
3070 return {CEDMap.lookup(Val: {&CB, PRE}), CEDMap.lookup(Val: {&CB, POST})};
3071 }
3072 ExecutionDomainTy getFunctionExecutionDomain() const override {
3073 assert(isValidState() &&
3074 "No request should be made against an invalid state!");
3075 return InterProceduralED;
3076 }
3077 ///}
3078
3079 // Check if the edge into the successor block contains a condition that only
3080 // lets the main thread execute it.
3081 static bool isInitialThreadOnlyEdge(Attributor &A, CondBrInst *Edge,
3082 BasicBlock &SuccessorBB) {
3083 if (!Edge)
3084 return false;
3085 if (Edge->getSuccessor(i: 0) != &SuccessorBB)
3086 return false;
3087
3088 auto *Cmp = dyn_cast<CmpInst>(Val: Edge->getCondition());
3089 if (!Cmp || !Cmp->isTrueWhenEqual() || !Cmp->isEquality())
3090 return false;
3091
3092 ConstantInt *C = dyn_cast<ConstantInt>(Val: Cmp->getOperand(i_nocapture: 1));
3093 if (!C)
3094 return false;
3095
3096 // Match: -1 == __kmpc_target_init (for non-SPMD kernels only!)
3097 if (C->isAllOnesValue()) {
3098 auto *CB = dyn_cast<CallBase>(Val: Cmp->getOperand(i_nocapture: 0));
3099 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
3100 auto &RFI = OMPInfoCache.RFIs[OMPRTL___kmpc_target_init];
3101 CB = CB ? OpenMPOpt::getCallIfRegularCall(V&: *CB, RFI: &RFI) : nullptr;
3102 if (!CB)
3103 return false;
3104 ConstantStruct *KernelEnvC =
3105 KernelInfo::getKernelEnvironementFromKernelInitCB(KernelInitCB: CB);
3106 ConstantInt *ExecModeC =
3107 KernelInfo::getExecModeFromKernelEnvironment(KernelEnvC);
3108 return ExecModeC->getSExtValue() & OMP_TGT_EXEC_MODE_GENERIC;
3109 }
3110
3111 if (C->isZero()) {
3112 // Match: 0 == llvm.nvvm.read.ptx.sreg.tid.x()
3113 if (auto *II = dyn_cast<IntrinsicInst>(Val: Cmp->getOperand(i_nocapture: 0)))
3114 if (II->getIntrinsicID() == Intrinsic::nvvm_read_ptx_sreg_tid_x)
3115 return true;
3116
3117 // Match: 0 == llvm.amdgcn.workitem.id.x()
3118 if (auto *II = dyn_cast<IntrinsicInst>(Val: Cmp->getOperand(i_nocapture: 0)))
3119 if (II->getIntrinsicID() == Intrinsic::amdgcn_workitem_id_x)
3120 return true;
3121 }
3122
3123 return false;
3124 };
3125
3126 /// Mapping containing information about the function for other AAs.
3127 ExecutionDomainTy InterProceduralED;
3128
3129 enum Direction { PRE = 0, POST = 1 };
3130 /// Mapping containing information per block.
3131 DenseMap<const BasicBlock *, ExecutionDomainTy> BEDMap;
3132 DenseMap<PointerIntPair<const CallBase *, 1, Direction>, ExecutionDomainTy>
3133 CEDMap;
3134 SmallSetVector<CallBase *, 16> AlignedBarriers;
3135
3136 ReversePostOrderTraversal<Function *> *RPOT = nullptr;
3137
3138 /// Set \p R to \V and report true if that changed \p R.
3139 static bool setAndRecord(bool &R, bool V) {
3140 bool Eq = (R == V);
3141 R = V;
3142 return !Eq;
3143 }
3144
3145 /// Collection of fences known to be non-no-opt. All fences not in this set
3146 /// can be assumed no-opt.
3147 SmallPtrSet<const FenceInst *, 8> NonNoOpFences;
3148};
3149
3150void AAExecutionDomainFunction::mergeInPredecessorBarriersAndAssumptions(
3151 Attributor &A, ExecutionDomainTy &ED, const ExecutionDomainTy &PredED) {
3152 for (auto *EA : PredED.EncounteredAssumes)
3153 ED.addAssumeInst(A, AI&: *EA);
3154
3155 for (auto *AB : PredED.AlignedBarriers)
3156 ED.addAlignedBarrier(A, CB&: *AB);
3157}
3158
3159bool AAExecutionDomainFunction::mergeInPredecessor(
3160 Attributor &A, ExecutionDomainTy &ED, const ExecutionDomainTy &PredED,
3161 bool InitialEdgeOnly) {
3162
3163 bool Changed = false;
3164 Changed |=
3165 setAndRecord(R&: ED.IsExecutedByInitialThreadOnly,
3166 V: InitialEdgeOnly || (PredED.IsExecutedByInitialThreadOnly &&
3167 ED.IsExecutedByInitialThreadOnly));
3168
3169 Changed |= setAndRecord(R&: ED.IsReachedFromAlignedBarrierOnly,
3170 V: ED.IsReachedFromAlignedBarrierOnly &&
3171 PredED.IsReachedFromAlignedBarrierOnly);
3172 Changed |= setAndRecord(R&: ED.EncounteredNonLocalSideEffect,
3173 V: ED.EncounteredNonLocalSideEffect |
3174 PredED.EncounteredNonLocalSideEffect);
3175 // Do not track assumptions and barriers as part of Changed.
3176 if (ED.IsReachedFromAlignedBarrierOnly)
3177 mergeInPredecessorBarriersAndAssumptions(A, ED, PredED);
3178 else
3179 ED.clearAssumeInstAndAlignedBarriers();
3180 return Changed;
3181}
3182
3183bool AAExecutionDomainFunction::handleCallees(Attributor &A,
3184 ExecutionDomainTy &EntryBBED) {
3185 SmallVector<std::pair<ExecutionDomainTy, ExecutionDomainTy>, 4> CallSiteEDs;
3186 auto PredForCallSite = [&](AbstractCallSite ACS) {
3187 const auto *EDAA = A.getAAFor<AAExecutionDomain>(
3188 QueryingAA: *this, IRP: IRPosition::function(F: *ACS.getInstruction()->getFunction()),
3189 DepClass: DepClassTy::OPTIONAL);
3190 if (!EDAA || !EDAA->getState().isValidState())
3191 return false;
3192 CallSiteEDs.emplace_back(
3193 Args: EDAA->getExecutionDomain(CB: *cast<CallBase>(Val: ACS.getInstruction())));
3194 return true;
3195 };
3196
3197 ExecutionDomainTy ExitED;
3198 bool AllCallSitesKnown;
3199 if (A.checkForAllCallSites(Pred: PredForCallSite, QueryingAA: *this,
3200 /* RequiresAllCallSites */ RequireAllCallSites: true,
3201 UsedAssumedInformation&: AllCallSitesKnown)) {
3202 for (const auto &[CSInED, CSOutED] : CallSiteEDs) {
3203 mergeInPredecessor(A, ED&: EntryBBED, PredED: CSInED);
3204 ExitED.IsReachingAlignedBarrierOnly &=
3205 CSOutED.IsReachingAlignedBarrierOnly;
3206 }
3207
3208 } else {
3209 // We could not find all predecessors, so this is either a kernel or a
3210 // function with external linkage (or with some other weird uses).
3211 if (omp::isOpenMPKernel(Fn&: *getAnchorScope())) {
3212 EntryBBED.IsExecutedByInitialThreadOnly = false;
3213 EntryBBED.IsReachedFromAlignedBarrierOnly = true;
3214 EntryBBED.EncounteredNonLocalSideEffect = false;
3215 ExitED.IsReachingAlignedBarrierOnly = false;
3216 } else {
3217 EntryBBED.IsExecutedByInitialThreadOnly = false;
3218 EntryBBED.IsReachedFromAlignedBarrierOnly = false;
3219 EntryBBED.EncounteredNonLocalSideEffect = true;
3220 ExitED.IsReachingAlignedBarrierOnly = false;
3221 }
3222 }
3223
3224 bool Changed = false;
3225 auto &FnED = BEDMap[nullptr];
3226 Changed |= setAndRecord(R&: FnED.IsReachedFromAlignedBarrierOnly,
3227 V: FnED.IsReachedFromAlignedBarrierOnly &
3228 EntryBBED.IsReachedFromAlignedBarrierOnly);
3229 Changed |= setAndRecord(R&: FnED.IsReachingAlignedBarrierOnly,
3230 V: FnED.IsReachingAlignedBarrierOnly &
3231 ExitED.IsReachingAlignedBarrierOnly);
3232 Changed |= setAndRecord(R&: FnED.IsExecutedByInitialThreadOnly,
3233 V: EntryBBED.IsExecutedByInitialThreadOnly);
3234 return Changed;
3235}
3236
3237ChangeStatus AAExecutionDomainFunction::updateImpl(Attributor &A) {
3238
3239 bool Changed = false;
3240
3241 // Helper to deal with an aligned barrier encountered during the forward
3242 // traversal. \p CB is the aligned barrier, \p ED is the execution domain when
3243 // it was encountered.
3244 auto HandleAlignedBarrier = [&](CallBase &CB, ExecutionDomainTy &ED) {
3245 Changed |= AlignedBarriers.insert(X: &CB);
3246 // First, update the barrier ED kept in the separate CEDMap.
3247 auto &CallInED = CEDMap[{&CB, PRE}];
3248 Changed |= mergeInPredecessor(A, ED&: CallInED, PredED: ED);
3249 CallInED.IsReachingAlignedBarrierOnly = true;
3250 // Next adjust the ED we use for the traversal.
3251 ED.EncounteredNonLocalSideEffect = false;
3252 ED.IsReachedFromAlignedBarrierOnly = true;
3253 // Aligned barrier collection has to come last.
3254 ED.clearAssumeInstAndAlignedBarriers();
3255 ED.addAlignedBarrier(A, CB);
3256 auto &CallOutED = CEDMap[{&CB, POST}];
3257 Changed |= mergeInPredecessor(A, ED&: CallOutED, PredED: ED);
3258 };
3259
3260 auto *LivenessAA =
3261 A.getAAFor<AAIsDead>(QueryingAA: *this, IRP: getIRPosition(), DepClass: DepClassTy::OPTIONAL);
3262
3263 Function *F = getAnchorScope();
3264 BasicBlock &EntryBB = F->getEntryBlock();
3265 bool IsKernel = omp::isOpenMPKernel(Fn&: *F);
3266
3267 SmallVector<Instruction *> SyncInstWorklist;
3268 for (auto &RIt : *RPOT) {
3269 BasicBlock &BB = *RIt;
3270
3271 bool IsEntryBB = &BB == &EntryBB;
3272 // TODO: We use local reasoning since we don't have a divergence analysis
3273 // running as well. We could basically allow uniform branches here.
3274 bool AlignedBarrierLastInBlock = IsEntryBB && IsKernel;
3275 bool IsExplicitlyAligned = IsEntryBB && IsKernel;
3276 ExecutionDomainTy ED;
3277 // Propagate "incoming edges" into information about this block.
3278 if (IsEntryBB) {
3279 Changed |= handleCallees(A, EntryBBED&: ED);
3280 } else {
3281 // For live non-entry blocks we only propagate
3282 // information via live edges.
3283 if (LivenessAA && LivenessAA->isAssumedDead(BB: &BB))
3284 continue;
3285
3286 for (auto *PredBB : predecessors(BB: &BB)) {
3287 if (LivenessAA && LivenessAA->isEdgeDead(From: PredBB, To: &BB))
3288 continue;
3289 bool InitialEdgeOnly = isInitialThreadOnlyEdge(
3290 A, Edge: dyn_cast<CondBrInst>(Val: PredBB->getTerminator()), SuccessorBB&: BB);
3291 mergeInPredecessor(A, ED, PredED: BEDMap[PredBB], InitialEdgeOnly);
3292 }
3293 }
3294
3295 // Now we traverse the block, accumulate effects in ED and attach
3296 // information to calls.
3297 for (Instruction &I : BB) {
3298 bool UsedAssumedInformation;
3299 if (A.isAssumedDead(I, QueryingAA: *this, LivenessAA, UsedAssumedInformation,
3300 /* CheckBBLivenessOnly */ false, DepClass: DepClassTy::OPTIONAL,
3301 /* CheckForDeadStore */ true))
3302 continue;
3303
3304 // Asummes and "assume-like" (dbg, lifetime, ...) are handled first, the
3305 // former is collected the latter is ignored.
3306 if (auto *II = dyn_cast<IntrinsicInst>(Val: &I)) {
3307 if (auto *AI = dyn_cast_or_null<AssumeInst>(Val: II)) {
3308 ED.addAssumeInst(A, AI&: *AI);
3309 continue;
3310 }
3311 // TODO: Should we also collect and delete lifetime markers?
3312 if (II->isAssumeLikeIntrinsic())
3313 continue;
3314 }
3315
3316 if (auto *FI = dyn_cast<FenceInst>(Val: &I)) {
3317 if (!ED.EncounteredNonLocalSideEffect) {
3318 // An aligned fence without non-local side-effects is a no-op.
3319 if (ED.IsReachedFromAlignedBarrierOnly)
3320 continue;
3321 // A non-aligned fence without non-local side-effects is a no-op
3322 // if the ordering only publishes non-local side-effects (or less).
3323 switch (FI->getOrdering()) {
3324 case AtomicOrdering::NotAtomic:
3325 continue;
3326 case AtomicOrdering::Unordered:
3327 continue;
3328 case AtomicOrdering::Monotonic:
3329 continue;
3330 case AtomicOrdering::Acquire:
3331 break;
3332 case AtomicOrdering::Release:
3333 continue;
3334 case AtomicOrdering::AcquireRelease:
3335 break;
3336 case AtomicOrdering::SequentiallyConsistent:
3337 break;
3338 };
3339 }
3340 NonNoOpFences.insert(Ptr: FI);
3341 }
3342
3343 auto *CB = dyn_cast<CallBase>(Val: &I);
3344 bool IsNoSync = AA::isNoSyncInst(A, I, QueryingAA: *this);
3345 bool IsAlignedBarrier =
3346 !IsNoSync && CB &&
3347 AANoSync::isAlignedBarrier(CB: *CB, ExecutedAligned: AlignedBarrierLastInBlock);
3348
3349 AlignedBarrierLastInBlock &= IsNoSync;
3350 IsExplicitlyAligned &= IsNoSync;
3351
3352 // Next we check for calls. Aligned barriers are handled
3353 // explicitly, everything else is kept for the backward traversal and will
3354 // also affect our state.
3355 if (CB) {
3356 if (IsAlignedBarrier) {
3357 HandleAlignedBarrier(*CB, ED);
3358 AlignedBarrierLastInBlock = true;
3359 IsExplicitlyAligned = true;
3360 continue;
3361 }
3362
3363 // Check the pointer(s) of a memory intrinsic explicitly.
3364 if (isa<MemIntrinsic>(Val: &I)) {
3365 if (!ED.EncounteredNonLocalSideEffect &&
3366 AA::isPotentiallyAffectedByBarrier(A, I, QueryingAA: *this))
3367 ED.EncounteredNonLocalSideEffect = true;
3368 if (!IsNoSync) {
3369 ED.IsReachedFromAlignedBarrierOnly = false;
3370 SyncInstWorklist.push_back(Elt: &I);
3371 }
3372 continue;
3373 }
3374
3375 // Record how we entered the call, then accumulate the effect of the
3376 // call in ED for potential use by the callee.
3377 auto &CallInED = CEDMap[{CB, PRE}];
3378 Changed |= mergeInPredecessor(A, ED&: CallInED, PredED: ED);
3379
3380 // If we have a sync-definition we can check if it starts/ends in an
3381 // aligned barrier. If we are unsure we assume any sync breaks
3382 // alignment.
3383 Function *Callee = CB->getCalledFunction();
3384 if (!IsNoSync && Callee && !Callee->isDeclaration()) {
3385 const auto *EDAA = A.getAAFor<AAExecutionDomain>(
3386 QueryingAA: *this, IRP: IRPosition::function(F: *Callee), DepClass: DepClassTy::OPTIONAL);
3387 if (EDAA && EDAA->getState().isValidState()) {
3388 const auto &CalleeED = EDAA->getFunctionExecutionDomain();
3389 ED.IsReachedFromAlignedBarrierOnly =
3390 CalleeED.IsReachedFromAlignedBarrierOnly;
3391 AlignedBarrierLastInBlock = ED.IsReachedFromAlignedBarrierOnly;
3392 if (IsNoSync || !CalleeED.IsReachedFromAlignedBarrierOnly)
3393 ED.EncounteredNonLocalSideEffect |=
3394 CalleeED.EncounteredNonLocalSideEffect;
3395 else
3396 ED.EncounteredNonLocalSideEffect =
3397 CalleeED.EncounteredNonLocalSideEffect;
3398 if (!CalleeED.IsReachingAlignedBarrierOnly) {
3399 Changed |=
3400 setAndRecord(R&: CallInED.IsReachingAlignedBarrierOnly, V: false);
3401 SyncInstWorklist.push_back(Elt: &I);
3402 }
3403 if (CalleeED.IsReachedFromAlignedBarrierOnly)
3404 mergeInPredecessorBarriersAndAssumptions(A, ED, PredED: CalleeED);
3405 auto &CallOutED = CEDMap[{CB, POST}];
3406 Changed |= mergeInPredecessor(A, ED&: CallOutED, PredED: ED);
3407 continue;
3408 }
3409 }
3410 if (!IsNoSync) {
3411 ED.IsReachedFromAlignedBarrierOnly = false;
3412 Changed |= setAndRecord(R&: CallInED.IsReachingAlignedBarrierOnly, V: false);
3413 SyncInstWorklist.push_back(Elt: &I);
3414 }
3415 AlignedBarrierLastInBlock &= ED.IsReachedFromAlignedBarrierOnly;
3416 ED.EncounteredNonLocalSideEffect |= !CB->doesNotAccessMemory();
3417 auto &CallOutED = CEDMap[{CB, POST}];
3418 Changed |= mergeInPredecessor(A, ED&: CallOutED, PredED: ED);
3419 }
3420
3421 if (!I.mayHaveSideEffects() && !I.mayReadFromMemory())
3422 continue;
3423
3424 // If we have a callee we try to use fine-grained information to
3425 // determine local side-effects.
3426 if (CB) {
3427 const auto *MemAA = A.getAAFor<AAMemoryLocation>(
3428 QueryingAA: *this, IRP: IRPosition::callsite_function(CB: *CB), DepClass: DepClassTy::OPTIONAL);
3429
3430 auto AccessPred = [&](const Instruction *I, const Value *Ptr,
3431 AAMemoryLocation::AccessKind,
3432 AAMemoryLocation::MemoryLocationsKind) {
3433 return !AA::isPotentiallyAffectedByBarrier(A, Ptrs: {Ptr}, QueryingAA: *this, CtxI: I);
3434 };
3435 if (MemAA && MemAA->getState().isValidState() &&
3436 MemAA->checkForAllAccessesToMemoryKind(
3437 Pred: AccessPred, MLK: AAMemoryLocation::ALL_LOCATIONS))
3438 continue;
3439 }
3440
3441 auto &InfoCache = A.getInfoCache();
3442 if (!I.mayHaveSideEffects() && InfoCache.isOnlyUsedByAssume(I))
3443 continue;
3444
3445 if (auto *LI = dyn_cast<LoadInst>(Val: &I))
3446 if (LI->hasMetadata(KindID: LLVMContext::MD_invariant_load))
3447 continue;
3448
3449 if (!ED.EncounteredNonLocalSideEffect &&
3450 AA::isPotentiallyAffectedByBarrier(A, I, QueryingAA: *this))
3451 ED.EncounteredNonLocalSideEffect = true;
3452 }
3453
3454 bool IsEndAndNotReachingAlignedBarriersOnly = false;
3455 if (!isa<UnreachableInst>(Val: BB.getTerminator()) &&
3456 !BB.getTerminator()->getNumSuccessors()) {
3457
3458 Changed |= mergeInPredecessor(A, ED&: InterProceduralED, PredED: ED);
3459
3460 auto &FnED = BEDMap[nullptr];
3461 if (IsKernel && !IsExplicitlyAligned)
3462 FnED.IsReachingAlignedBarrierOnly = false;
3463 Changed |= mergeInPredecessor(A, ED&: FnED, PredED: ED);
3464
3465 if (!FnED.IsReachingAlignedBarrierOnly) {
3466 IsEndAndNotReachingAlignedBarriersOnly = true;
3467 SyncInstWorklist.push_back(Elt: BB.getTerminator());
3468 auto &BBED = BEDMap[&BB];
3469 Changed |= setAndRecord(R&: BBED.IsReachingAlignedBarrierOnly, V: false);
3470 }
3471 }
3472
3473 ExecutionDomainTy &StoredED = BEDMap[&BB];
3474 ED.IsReachingAlignedBarrierOnly = StoredED.IsReachingAlignedBarrierOnly &&
3475 !IsEndAndNotReachingAlignedBarriersOnly;
3476
3477 // Check if we computed anything different as part of the forward
3478 // traversal. We do not take assumptions and aligned barriers into account
3479 // as they do not influence the state we iterate. Backward traversal values
3480 // are handled later on.
3481 if (ED.IsExecutedByInitialThreadOnly !=
3482 StoredED.IsExecutedByInitialThreadOnly ||
3483 ED.IsReachedFromAlignedBarrierOnly !=
3484 StoredED.IsReachedFromAlignedBarrierOnly ||
3485 ED.EncounteredNonLocalSideEffect !=
3486 StoredED.EncounteredNonLocalSideEffect)
3487 Changed = true;
3488
3489 // Update the state with the new value.
3490 StoredED = std::move(ED);
3491 }
3492
3493 // Propagate (non-aligned) sync instruction effects backwards until the
3494 // entry is hit or an aligned barrier.
3495 SmallSetVector<BasicBlock *, 16> Visited;
3496 while (!SyncInstWorklist.empty()) {
3497 Instruction *SyncInst = SyncInstWorklist.pop_back_val();
3498 Instruction *CurInst = SyncInst;
3499 bool HitAlignedBarrierOrKnownEnd = false;
3500 while ((CurInst = CurInst->getPrevNode())) {
3501 auto *CB = dyn_cast<CallBase>(Val: CurInst);
3502 if (!CB)
3503 continue;
3504 auto &CallOutED = CEDMap[{CB, POST}];
3505 Changed |= setAndRecord(R&: CallOutED.IsReachingAlignedBarrierOnly, V: false);
3506 auto &CallInED = CEDMap[{CB, PRE}];
3507 HitAlignedBarrierOrKnownEnd =
3508 AlignedBarriers.count(key: CB) || !CallInED.IsReachingAlignedBarrierOnly;
3509 if (HitAlignedBarrierOrKnownEnd)
3510 break;
3511 Changed |= setAndRecord(R&: CallInED.IsReachingAlignedBarrierOnly, V: false);
3512 }
3513 if (HitAlignedBarrierOrKnownEnd)
3514 continue;
3515 BasicBlock *SyncBB = SyncInst->getParent();
3516 for (auto *PredBB : predecessors(BB: SyncBB)) {
3517 if (LivenessAA && LivenessAA->isEdgeDead(From: PredBB, To: SyncBB))
3518 continue;
3519 if (!Visited.insert(X: PredBB))
3520 continue;
3521 auto &PredED = BEDMap[PredBB];
3522 if (setAndRecord(R&: PredED.IsReachingAlignedBarrierOnly, V: false)) {
3523 Changed = true;
3524 SyncInstWorklist.push_back(Elt: PredBB->getTerminator());
3525 }
3526 }
3527 if (SyncBB != &EntryBB)
3528 continue;
3529 Changed |=
3530 setAndRecord(R&: InterProceduralED.IsReachingAlignedBarrierOnly, V: false);
3531 }
3532
3533 return Changed ? ChangeStatus::CHANGED : ChangeStatus::UNCHANGED;
3534}
3535
3536/// Try to replace memory allocation calls called by a single thread with a
3537/// static buffer of shared memory.
3538struct AAHeapToShared : public StateWrapper<BooleanState, AbstractAttribute> {
3539 using Base = StateWrapper<BooleanState, AbstractAttribute>;
3540 AAHeapToShared(const IRPosition &IRP, Attributor &A) : Base(IRP) {}
3541
3542 /// Create an abstract attribute view for the position \p IRP.
3543 static AAHeapToShared &createForPosition(const IRPosition &IRP,
3544 Attributor &A);
3545
3546 /// Returns true if HeapToShared conversion is assumed to be possible.
3547 virtual bool isAssumedHeapToShared(CallBase &CB) const = 0;
3548
3549 /// Returns true if HeapToShared conversion is assumed and the CB is a
3550 /// callsite to a free operation to be removed.
3551 virtual bool isAssumedHeapToSharedRemovedFree(CallBase &CB) const = 0;
3552
3553 /// See AbstractAttribute::getName().
3554 StringRef getName() const override { return "AAHeapToShared"; }
3555
3556 /// See AbstractAttribute::getIdAddr().
3557 const char *getIdAddr() const override { return &ID; }
3558
3559 /// This function should return true if the type of the \p AA is
3560 /// AAHeapToShared.
3561 static bool classof(const AbstractAttribute *AA) {
3562 return (AA->getIdAddr() == &ID);
3563 }
3564
3565 /// Unique ID (due to the unique address)
3566 static const char ID;
3567};
3568
3569struct AAHeapToSharedFunction : public AAHeapToShared {
3570 AAHeapToSharedFunction(const IRPosition &IRP, Attributor &A)
3571 : AAHeapToShared(IRP, A) {}
3572
3573 const std::string getAsStr(Attributor *) const override {
3574 return "[AAHeapToShared] " + std::to_string(val: MallocCalls.size()) +
3575 " malloc calls eligible.";
3576 }
3577
3578 /// See AbstractAttribute::trackStatistics().
3579 void trackStatistics() const override {}
3580
3581 /// This functions finds free calls that will be removed by the
3582 /// HeapToShared transformation.
3583 void findPotentialRemovedFreeCalls(Attributor &A) {
3584 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
3585 auto &FreeRFI = OMPInfoCache.RFIs[OMPRTL___kmpc_free_shared];
3586
3587 PotentialRemovedFreeCalls.clear();
3588 // Update free call users of found malloc calls.
3589 for (CallBase *CB : MallocCalls) {
3590 SmallVector<CallBase *, 4> FreeCalls;
3591 for (auto *U : CB->users()) {
3592 CallBase *C = dyn_cast<CallBase>(Val: U);
3593 if (C && C->getCalledFunction() == FreeRFI.Declaration)
3594 FreeCalls.push_back(Elt: C);
3595 }
3596
3597 if (FreeCalls.size() != 1)
3598 continue;
3599
3600 PotentialRemovedFreeCalls.insert(Ptr: FreeCalls.front());
3601 }
3602 }
3603
3604 void initialize(Attributor &A) override {
3605 if (DisableOpenMPOptDeglobalization) {
3606 indicatePessimisticFixpoint();
3607 return;
3608 }
3609
3610 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
3611 auto &RFI = OMPInfoCache.RFIs[OMPRTL___kmpc_alloc_shared];
3612 if (!RFI.Declaration)
3613 return;
3614
3615 Attributor::SimplifictionCallbackTy SCB =
3616 [](const IRPosition &, const AbstractAttribute *,
3617 bool &) -> std::optional<Value *> { return nullptr; };
3618
3619 Function *F = getAnchorScope();
3620 const OMPInformationCache::RuntimeFunctionInfo::UseVector *Uses =
3621 RFI.getUseVector(F&: *F);
3622 if (!Uses)
3623 return;
3624
3625 for (Use *U : *Uses)
3626 if (CallBase *CB = dyn_cast<CallBase>(Val: U->getUser())) {
3627 MallocCalls.insert(X: CB);
3628 A.registerSimplificationCallback(IRP: IRPosition::callsite_returned(CB: *CB),
3629 CB: SCB);
3630 }
3631
3632 findPotentialRemovedFreeCalls(A);
3633 }
3634
3635 bool isAssumedHeapToShared(CallBase &CB) const override {
3636 return isValidState() && MallocCalls.count(key: &CB);
3637 }
3638
3639 bool isAssumedHeapToSharedRemovedFree(CallBase &CB) const override {
3640 return isValidState() && PotentialRemovedFreeCalls.count(Ptr: &CB);
3641 }
3642
3643 ChangeStatus manifest(Attributor &A) override {
3644 if (MallocCalls.empty())
3645 return ChangeStatus::UNCHANGED;
3646
3647 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
3648 auto &FreeCall = OMPInfoCache.RFIs[OMPRTL___kmpc_free_shared];
3649
3650 Function *F = getAnchorScope();
3651 auto *HS = A.lookupAAFor<AAHeapToStack>(IRP: IRPosition::function(F: *F), QueryingAA: this,
3652 DepClass: DepClassTy::OPTIONAL);
3653
3654 ChangeStatus Changed = ChangeStatus::UNCHANGED;
3655 for (CallBase *CB : MallocCalls) {
3656 // Skip replacing this if HeapToStack has already claimed it.
3657 if (HS && HS->isAssumedHeapToStack(CB: *CB))
3658 continue;
3659
3660 // Find the unique free call to remove it.
3661 SmallVector<CallBase *, 4> FreeCalls;
3662 for (auto *U : CB->users()) {
3663 CallBase *C = dyn_cast<CallBase>(Val: U);
3664 if (C && C->getCalledFunction() == FreeCall.Declaration)
3665 FreeCalls.push_back(Elt: C);
3666 }
3667 if (FreeCalls.size() != 1)
3668 continue;
3669
3670 auto *AllocSize = cast<ConstantInt>(Val: CB->getArgOperand(i: 0));
3671
3672 if (AllocSize->getZExtValue() + SharedMemoryUsed > SharedMemoryLimit) {
3673 LLVM_DEBUG(dbgs() << TAG << "Cannot replace call " << *CB
3674 << " with shared memory."
3675 << " Shared memory usage is limited to "
3676 << SharedMemoryLimit << " bytes\n");
3677 continue;
3678 }
3679
3680 LLVM_DEBUG(dbgs() << TAG << "Replace globalization call " << *CB
3681 << " with " << AllocSize->getZExtValue()
3682 << " bytes of shared memory\n");
3683
3684 // Create a new shared memory buffer of the same size as the allocation
3685 // and replace all the uses of the original allocation with it.
3686 Module *M = CB->getModule();
3687 Type *Int8Ty = Type::getInt8Ty(C&: M->getContext());
3688 Type *Int8ArrTy = ArrayType::get(ElementType: Int8Ty, NumElements: AllocSize->getZExtValue());
3689 auto *SharedMem = new GlobalVariable(
3690 *M, Int8ArrTy, /* IsConstant */ false, GlobalValue::InternalLinkage,
3691 PoisonValue::get(T: Int8ArrTy), CB->getName() + "_shared", nullptr,
3692 GlobalValue::NotThreadLocal,
3693 static_cast<unsigned>(AddressSpace::Shared));
3694 auto *NewBuffer = ConstantExpr::getPointerCast(
3695 C: SharedMem, Ty: PointerType::getUnqual(C&: M->getContext()));
3696
3697 auto Remark = [&](OptimizationRemark OR) {
3698 return OR << "Replaced globalized variable with "
3699 << ore::NV("SharedMemory", AllocSize->getZExtValue())
3700 << (AllocSize->isOne() ? " byte " : " bytes ")
3701 << "of shared memory.";
3702 };
3703 A.emitRemark<OptimizationRemark>(I: CB, RemarkName: "OMP111", RemarkCB&: Remark);
3704
3705 MaybeAlign Alignment = CB->getRetAlign();
3706 assert(Alignment &&
3707 "HeapToShared on allocation without alignment attribute");
3708 SharedMem->setAlignment(*Alignment);
3709
3710 A.changeAfterManifest(IRP: IRPosition::callsite_returned(CB: *CB), NV&: *NewBuffer);
3711 A.deleteAfterManifest(I&: *CB);
3712 A.deleteAfterManifest(I&: *FreeCalls.front());
3713
3714 SharedMemoryUsed += AllocSize->getZExtValue();
3715 NumBytesMovedToSharedMemory = SharedMemoryUsed;
3716 Changed = ChangeStatus::CHANGED;
3717 }
3718
3719 return Changed;
3720 }
3721
3722 ChangeStatus updateImpl(Attributor &A) override {
3723 if (MallocCalls.empty())
3724 return indicatePessimisticFixpoint();
3725 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
3726 auto &RFI = OMPInfoCache.RFIs[OMPRTL___kmpc_alloc_shared];
3727 if (!RFI.Declaration)
3728 return ChangeStatus::UNCHANGED;
3729
3730 Function *F = getAnchorScope();
3731
3732 auto NumMallocCalls = MallocCalls.size();
3733
3734 // Only consider malloc calls executed by a single thread with a constant.
3735 for (User *U : RFI.Declaration->users()) {
3736 if (CallBase *CB = dyn_cast<CallBase>(Val: U)) {
3737 if (CB->getCaller() != F)
3738 continue;
3739 if (!MallocCalls.count(key: CB))
3740 continue;
3741 if (!isa<ConstantInt>(Val: CB->getArgOperand(i: 0))) {
3742 MallocCalls.remove(X: CB);
3743 continue;
3744 }
3745 const auto *ED = A.getAAFor<AAExecutionDomain>(
3746 QueryingAA: *this, IRP: IRPosition::function(F: *F), DepClass: DepClassTy::REQUIRED);
3747 if (!ED || !ED->isExecutedByInitialThreadOnly(I: *CB))
3748 MallocCalls.remove(X: CB);
3749 }
3750 }
3751
3752 findPotentialRemovedFreeCalls(A);
3753
3754 if (NumMallocCalls != MallocCalls.size())
3755 return ChangeStatus::CHANGED;
3756
3757 return ChangeStatus::UNCHANGED;
3758 }
3759
3760 /// Collection of all malloc calls in a function.
3761 SmallSetVector<CallBase *, 4> MallocCalls;
3762 /// Collection of potentially removed free calls in a function.
3763 SmallPtrSet<CallBase *, 4> PotentialRemovedFreeCalls;
3764 /// The total amount of shared memory that has been used for HeapToShared.
3765 unsigned SharedMemoryUsed = 0;
3766};
3767
3768struct AAKernelInfo : public StateWrapper<KernelInfoState, AbstractAttribute> {
3769 using Base = StateWrapper<KernelInfoState, AbstractAttribute>;
3770 AAKernelInfo(const IRPosition &IRP, Attributor &A) : Base(IRP) {}
3771
3772 /// The callee value is tracked beyond a simple stripPointerCasts, so we allow
3773 /// unknown callees.
3774 static bool requiresCalleeForCallBase() { return false; }
3775
3776 /// Statistics are tracked as part of manifest for now.
3777 void trackStatistics() const override {}
3778
3779 /// See AbstractAttribute::getAsStr()
3780 const std::string getAsStr(Attributor *) const override {
3781 if (!isValidState())
3782 return "<invalid>";
3783 return std::string(SPMDCompatibilityTracker.isAssumed() ? "SPMD"
3784 : "generic") +
3785 std::string(SPMDCompatibilityTracker.isAtFixpoint() ? " [FIX]"
3786 : "") +
3787 std::string(" #PRs: ") +
3788 (ReachedKnownParallelRegions.isValidState()
3789 ? std::to_string(val: ReachedKnownParallelRegions.size())
3790 : "<invalid>") +
3791 ", #Unknown PRs: " +
3792 (ReachedUnknownParallelRegions.isValidState()
3793 ? std::to_string(val: ReachedUnknownParallelRegions.size())
3794 : "<invalid>") +
3795 ", #Reaching Kernels: " +
3796 (ReachingKernelEntries.isValidState()
3797 ? std::to_string(val: ReachingKernelEntries.size())
3798 : "<invalid>") +
3799 ", #ParLevels: " +
3800 (ParallelLevels.isValidState()
3801 ? std::to_string(val: ParallelLevels.size())
3802 : "<invalid>") +
3803 ", NestedPar: " + (NestedParallelism ? "yes" : "no");
3804 }
3805
3806 /// Create an abstract attribute biew for the position \p IRP.
3807 static AAKernelInfo &createForPosition(const IRPosition &IRP, Attributor &A);
3808
3809 /// See AbstractAttribute::getName()
3810 StringRef getName() const override { return "AAKernelInfo"; }
3811
3812 /// See AbstractAttribute::getIdAddr()
3813 const char *getIdAddr() const override { return &ID; }
3814
3815 /// This function should return true if the type of the \p AA is AAKernelInfo
3816 static bool classof(const AbstractAttribute *AA) {
3817 return (AA->getIdAddr() == &ID);
3818 }
3819
3820 static const char ID;
3821};
3822
3823/// The function kernel info abstract attribute, basically, what can we say
3824/// about a function with regards to the KernelInfoState.
3825struct AAKernelInfoFunction : AAKernelInfo {
3826 AAKernelInfoFunction(const IRPosition &IRP, Attributor &A)
3827 : AAKernelInfo(IRP, A) {}
3828
3829 SmallPtrSet<Instruction *, 4> GuardedInstructions;
3830
3831 SmallPtrSetImpl<Instruction *> &getGuardedInstructions() {
3832 return GuardedInstructions;
3833 }
3834
3835 void setConfigurationOfKernelEnvironment(ConstantStruct *ConfigC) {
3836 Constant *NewKernelEnvC = ConstantFoldInsertValueInstruction(
3837 Agg: KernelEnvC, Val: ConfigC, Idxs: {KernelInfo::ConfigurationIdx});
3838 assert(NewKernelEnvC && "Failed to create new kernel environment");
3839 KernelEnvC = cast<ConstantStruct>(Val: NewKernelEnvC);
3840 }
3841
3842#define KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(MEMBER) \
3843 void set##MEMBER##OfKernelEnvironment(ConstantInt *NewVal) { \
3844 ConstantStruct *ConfigC = \
3845 KernelInfo::getConfigurationFromKernelEnvironment(KernelEnvC); \
3846 Constant *NewConfigC = ConstantFoldInsertValueInstruction( \
3847 ConfigC, NewVal, {KernelInfo::MEMBER##Idx}); \
3848 assert(NewConfigC && "Failed to create new configuration environment"); \
3849 setConfigurationOfKernelEnvironment(cast<ConstantStruct>(NewConfigC)); \
3850 }
3851
3852 KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(UseGenericStateMachine)
3853 KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(MayUseNestedParallelism)
3854 KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(ExecMode)
3855 KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(MinThreads)
3856 KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(MaxThreads)
3857 KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(MinTeams)
3858 KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(MaxTeams)
3859
3860#undef KERNEL_ENVIRONMENT_CONFIGURATION_SETTER
3861
3862 /// See AbstractAttribute::initialize(...).
3863 void initialize(Attributor &A) override {
3864 // This is a high-level transform that might change the constant arguments
3865 // of the init and dinit calls. We need to tell the Attributor about this
3866 // to avoid other parts using the current constant value for simpliication.
3867 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
3868
3869 Function *Fn = getAnchorScope();
3870
3871 OMPInformationCache::RuntimeFunctionInfo &InitRFI =
3872 OMPInfoCache.RFIs[OMPRTL___kmpc_target_init];
3873 OMPInformationCache::RuntimeFunctionInfo &DeinitRFI =
3874 OMPInfoCache.RFIs[OMPRTL___kmpc_target_deinit];
3875
3876 // For kernels we perform more initialization work, first we find the init
3877 // and deinit calls.
3878 auto StoreCallBase = [](Use &U,
3879 OMPInformationCache::RuntimeFunctionInfo &RFI,
3880 CallBase *&Storage) {
3881 CallBase *CB = OpenMPOpt::getCallIfRegularCall(U, RFI: &RFI);
3882 assert(CB &&
3883 "Unexpected use of __kmpc_target_init or __kmpc_target_deinit!");
3884 assert(!Storage &&
3885 "Multiple uses of __kmpc_target_init or __kmpc_target_deinit!");
3886 Storage = CB;
3887 return false;
3888 };
3889 InitRFI.foreachUse(
3890 CB: [&](Use &U, Function &) {
3891 StoreCallBase(U, InitRFI, KernelInitCB);
3892 return false;
3893 },
3894 F: Fn);
3895 DeinitRFI.foreachUse(
3896 CB: [&](Use &U, Function &) {
3897 StoreCallBase(U, DeinitRFI, KernelDeinitCB);
3898 return false;
3899 },
3900 F: Fn);
3901
3902 // Ignore kernels without initializers such as global constructors.
3903 if (!KernelInitCB || !KernelDeinitCB)
3904 return;
3905
3906 // Add itself to the reaching kernel and set IsKernelEntry.
3907 ReachingKernelEntries.insert(Elem: Fn);
3908 IsKernelEntry = true;
3909
3910 KernelEnvC =
3911 KernelInfo::getKernelEnvironementFromKernelInitCB(KernelInitCB);
3912 GlobalVariable *KernelEnvGV =
3913 KernelInfo::getKernelEnvironementGVFromKernelInitCB(KernelInitCB);
3914
3915 Attributor::GlobalVariableSimplifictionCallbackTy
3916 KernelConfigurationSimplifyCB =
3917 [&](const GlobalVariable &GV, const AbstractAttribute *AA,
3918 bool &UsedAssumedInformation) -> std::optional<Constant *> {
3919 if (!isAtFixpoint()) {
3920 if (!AA)
3921 return nullptr;
3922 UsedAssumedInformation = true;
3923 A.recordDependence(FromAA: *this, ToAA: *AA, DepClass: DepClassTy::OPTIONAL);
3924 }
3925 return KernelEnvC;
3926 };
3927
3928 A.registerGlobalVariableSimplificationCallback(
3929 GV: *KernelEnvGV, CB: KernelConfigurationSimplifyCB);
3930
3931 // We cannot change to SPMD mode if the runtime functions aren't availible.
3932 bool CanChangeToSPMD = OMPInfoCache.runtimeFnsAvailable(
3933 Fns: {OMPRTL___kmpc_get_hardware_thread_id_in_block,
3934 OMPRTL___kmpc_barrier_simple_spmd});
3935
3936 // Check if we know we are in SPMD-mode already.
3937 ConstantInt *ExecModeC =
3938 KernelInfo::getExecModeFromKernelEnvironment(KernelEnvC);
3939 ConstantInt *AssumedExecModeC = ConstantInt::get(
3940 Ty: ExecModeC->getIntegerType(),
3941 V: ExecModeC->getSExtValue() | OMP_TGT_EXEC_MODE_GENERIC_SPMD);
3942 if (ExecModeC->getSExtValue() & OMP_TGT_EXEC_MODE_SPMD)
3943 SPMDCompatibilityTracker.indicateOptimisticFixpoint();
3944 else if (DisableOpenMPOptSPMDization || !CanChangeToSPMD)
3945 // This is a generic region but SPMDization is disabled so stop
3946 // tracking.
3947 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
3948 else
3949 setExecModeOfKernelEnvironment(AssumedExecModeC);
3950
3951 const Triple T(Fn->getParent()->getTargetTriple());
3952 auto *Int32Ty = Type::getInt32Ty(C&: Fn->getContext());
3953 auto [MinThreads, MaxThreads] =
3954 OpenMPIRBuilder::readThreadBoundsForKernel(T, Kernel&: *Fn);
3955 if (MinThreads)
3956 setMinThreadsOfKernelEnvironment(ConstantInt::get(Ty: Int32Ty, V: MinThreads));
3957 if (MaxThreads)
3958 setMaxThreadsOfKernelEnvironment(ConstantInt::get(Ty: Int32Ty, V: MaxThreads));
3959 auto [MinTeams, MaxTeams] =
3960 OpenMPIRBuilder::readTeamBoundsForKernel(T, Kernel&: *Fn);
3961 if (MinTeams)
3962 setMinTeamsOfKernelEnvironment(ConstantInt::get(Ty: Int32Ty, V: MinTeams));
3963 if (MaxTeams)
3964 setMaxTeamsOfKernelEnvironment(ConstantInt::get(Ty: Int32Ty, V: MaxTeams));
3965
3966 ConstantInt *MayUseNestedParallelismC =
3967 KernelInfo::getMayUseNestedParallelismFromKernelEnvironment(KernelEnvC);
3968 ConstantInt *AssumedMayUseNestedParallelismC = ConstantInt::get(
3969 Ty: MayUseNestedParallelismC->getIntegerType(), V: NestedParallelism);
3970 setMayUseNestedParallelismOfKernelEnvironment(
3971 AssumedMayUseNestedParallelismC);
3972
3973 if (!DisableOpenMPOptStateMachineRewrite) {
3974 ConstantInt *UseGenericStateMachineC =
3975 KernelInfo::getUseGenericStateMachineFromKernelEnvironment(
3976 KernelEnvC);
3977 ConstantInt *AssumedUseGenericStateMachineC =
3978 ConstantInt::get(Ty: UseGenericStateMachineC->getIntegerType(), V: false);
3979 setUseGenericStateMachineOfKernelEnvironment(
3980 AssumedUseGenericStateMachineC);
3981 }
3982
3983 // Register virtual uses of functions we might need to preserve.
3984 auto RegisterVirtualUse = [&](RuntimeFunction RFKind,
3985 Attributor::VirtualUseCallbackTy &CB) {
3986 if (!OMPInfoCache.RFIs[RFKind].Declaration)
3987 return;
3988 A.registerVirtualUseCallback(V: *OMPInfoCache.RFIs[RFKind].Declaration, CB);
3989 };
3990
3991 // Add a dependence to ensure updates if the state changes.
3992 auto AddDependence = [](Attributor &A, const AAKernelInfo *KI,
3993 const AbstractAttribute *QueryingAA) {
3994 if (QueryingAA) {
3995 A.recordDependence(FromAA: *KI, ToAA: *QueryingAA, DepClass: DepClassTy::OPTIONAL);
3996 }
3997 return true;
3998 };
3999
4000 Attributor::VirtualUseCallbackTy CustomStateMachineUseCB =
4001 [&](Attributor &A, const AbstractAttribute *QueryingAA) {
4002 // Whenever we create a custom state machine we will insert calls to
4003 // __kmpc_get_max_team_threads,
4004 // __kmpc_barrier_simple_generic,
4005 // __kmpc_kernel_parallel, and
4006 // __kmpc_kernel_end_parallel.
4007 // Not needed if we are on track for SPMDzation.
4008 if (SPMDCompatibilityTracker.isValidState())
4009 return AddDependence(A, this, QueryingAA);
4010 // Not needed if we can't rewrite due to an invalid state.
4011 if (!ReachedKnownParallelRegions.isValidState())
4012 return AddDependence(A, this, QueryingAA);
4013 return false;
4014 };
4015
4016 // Not needed if we are pre-runtime merge.
4017 if (!KernelInitCB->getCalledFunction()->isDeclaration()) {
4018 RegisterVirtualUse(OMPRTL___kmpc_get_max_team_threads,
4019 CustomStateMachineUseCB);
4020 RegisterVirtualUse(OMPRTL___kmpc_barrier_simple_generic,
4021 CustomStateMachineUseCB);
4022 RegisterVirtualUse(OMPRTL___kmpc_kernel_parallel,
4023 CustomStateMachineUseCB);
4024 RegisterVirtualUse(OMPRTL___kmpc_kernel_end_parallel,
4025 CustomStateMachineUseCB);
4026 }
4027
4028 // If we do not perform SPMDzation we do not need the virtual uses below.
4029 if (SPMDCompatibilityTracker.isAtFixpoint())
4030 return;
4031
4032 Attributor::VirtualUseCallbackTy HWThreadIdUseCB =
4033 [&](Attributor &A, const AbstractAttribute *QueryingAA) {
4034 // Whenever we perform SPMDzation we will insert
4035 // __kmpc_get_hardware_thread_id_in_block calls.
4036 if (!SPMDCompatibilityTracker.isValidState())
4037 return AddDependence(A, this, QueryingAA);
4038 return false;
4039 };
4040 RegisterVirtualUse(OMPRTL___kmpc_get_hardware_thread_id_in_block,
4041 HWThreadIdUseCB);
4042
4043 Attributor::VirtualUseCallbackTy SPMDBarrierUseCB =
4044 [&](Attributor &A, const AbstractAttribute *QueryingAA) {
4045 // Whenever we perform SPMDzation with guarding we will insert
4046 // __kmpc_simple_barrier_spmd calls. If SPMDzation failed, there is
4047 // nothing to guard, or there are no parallel regions, we don't need
4048 // the calls.
4049 if (!SPMDCompatibilityTracker.isValidState())
4050 return AddDependence(A, this, QueryingAA);
4051 if (SPMDCompatibilityTracker.empty())
4052 return AddDependence(A, this, QueryingAA);
4053 if (!mayContainParallelRegion())
4054 return AddDependence(A, this, QueryingAA);
4055 return false;
4056 };
4057 RegisterVirtualUse(OMPRTL___kmpc_barrier_simple_spmd, SPMDBarrierUseCB);
4058 }
4059
4060 /// Sanitize the string \p S such that it is a suitable global symbol name.
4061 static std::string sanitizeForGlobalName(std::string S) {
4062 std::replace_if(
4063 first: S.begin(), last: S.end(),
4064 pred: [](const char C) {
4065 return !((C >= 'a' && C <= 'z') || (C >= 'A' && C <= 'Z') ||
4066 (C >= '0' && C <= '9') || C == '_');
4067 },
4068 new_value: '.');
4069 return S;
4070 }
4071
4072 /// Modify the IR based on the KernelInfoState as the fixpoint iteration is
4073 /// finished now.
4074 ChangeStatus manifest(Attributor &A) override {
4075 // If we are not looking at a kernel with __kmpc_target_init and
4076 // __kmpc_target_deinit call we cannot actually manifest the information.
4077 if (!KernelInitCB || !KernelDeinitCB)
4078 return ChangeStatus::UNCHANGED;
4079
4080 ChangeStatus Changed = ChangeStatus::UNCHANGED;
4081
4082 bool HasBuiltStateMachine = true;
4083 if (!changeToSPMDMode(A, Changed)) {
4084 if (!KernelInitCB->getCalledFunction()->isDeclaration())
4085 HasBuiltStateMachine = buildCustomStateMachine(A, Changed);
4086 else
4087 HasBuiltStateMachine = false;
4088 }
4089
4090 // We need to reset KernelEnvC if specific rewriting is not done.
4091 ConstantStruct *ExistingKernelEnvC =
4092 KernelInfo::getKernelEnvironementFromKernelInitCB(KernelInitCB);
4093 ConstantInt *OldUseGenericStateMachineVal =
4094 KernelInfo::getUseGenericStateMachineFromKernelEnvironment(
4095 KernelEnvC: ExistingKernelEnvC);
4096 if (!HasBuiltStateMachine)
4097 setUseGenericStateMachineOfKernelEnvironment(
4098 OldUseGenericStateMachineVal);
4099
4100 // At last, update the KernelEnvc
4101 GlobalVariable *KernelEnvGV =
4102 KernelInfo::getKernelEnvironementGVFromKernelInitCB(KernelInitCB);
4103 if (KernelEnvGV->getInitializer() != KernelEnvC) {
4104 KernelEnvGV->setInitializer(KernelEnvC);
4105 Changed = ChangeStatus::CHANGED;
4106 }
4107
4108 return Changed;
4109 }
4110
4111 void insertInstructionGuardsHelper(Attributor &A) {
4112 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
4113
4114 auto CreateGuardedRegion = [&](Instruction *RegionStartI,
4115 Instruction *RegionEndI) {
4116 LoopInfo *LI = nullptr;
4117 DominatorTree *DT = nullptr;
4118 MemorySSAUpdater *MSU = nullptr;
4119
4120 BasicBlock *ParentBB = RegionStartI->getParent();
4121 Function *Fn = ParentBB->getParent();
4122 Module &M = *Fn->getParent();
4123
4124 // Create all the blocks and logic.
4125 // ParentBB:
4126 // goto RegionCheckTidBB
4127 // RegionCheckTidBB:
4128 // Tid = __kmpc_hardware_thread_id()
4129 // if (Tid != 0)
4130 // goto RegionBarrierBB
4131 // RegionStartBB:
4132 // <execute instructions guarded>
4133 // goto RegionEndBB
4134 // RegionEndBB:
4135 // <store escaping values to shared mem>
4136 // goto RegionBarrierBB
4137 // RegionBarrierBB:
4138 // __kmpc_simple_barrier_spmd()
4139 // // second barrier is omitted if lacking escaping values.
4140 // <load escaping values from shared mem>
4141 // __kmpc_simple_barrier_spmd()
4142 // goto RegionExitBB
4143 // RegionExitBB:
4144 // <execute rest of instructions>
4145
4146 BasicBlock *RegionEndBB = SplitBlock(Old: ParentBB, SplitPt: RegionEndI->getNextNode(),
4147 DT, LI, MSSAU: MSU, BBName: "region.guarded.end");
4148 BasicBlock *RegionBarrierBB =
4149 SplitBlock(Old: RegionEndBB, SplitPt: &*RegionEndBB->getFirstInsertionPt(), DT, LI,
4150 MSSAU: MSU, BBName: "region.barrier");
4151 BasicBlock *RegionExitBB =
4152 SplitBlock(Old: RegionBarrierBB, SplitPt: &*RegionBarrierBB->getFirstInsertionPt(),
4153 DT, LI, MSSAU: MSU, BBName: "region.exit");
4154 BasicBlock *RegionStartBB =
4155 SplitBlock(Old: ParentBB, SplitPt: RegionStartI, DT, LI, MSSAU: MSU, BBName: "region.guarded");
4156
4157 assert(ParentBB->getUniqueSuccessor() == RegionStartBB &&
4158 "Expected a different CFG");
4159
4160 BasicBlock *RegionCheckTidBB = SplitBlock(
4161 Old: ParentBB, SplitPt: ParentBB->getTerminator(), DT, LI, MSSAU: MSU, BBName: "region.check.tid");
4162
4163 // Register basic blocks with the Attributor.
4164 A.registerManifestAddedBasicBlock(BB&: *RegionEndBB);
4165 A.registerManifestAddedBasicBlock(BB&: *RegionBarrierBB);
4166 A.registerManifestAddedBasicBlock(BB&: *RegionExitBB);
4167 A.registerManifestAddedBasicBlock(BB&: *RegionStartBB);
4168 A.registerManifestAddedBasicBlock(BB&: *RegionCheckTidBB);
4169
4170 bool HasBroadcastValues = false;
4171 // Find escaping outputs from the guarded region to outside users and
4172 // broadcast their values to them.
4173 for (Instruction &I : *RegionStartBB) {
4174 SmallVector<Use *, 4> OutsideUses;
4175 for (Use &U : I.uses()) {
4176 Instruction &UsrI = *cast<Instruction>(Val: U.getUser());
4177 if (UsrI.getParent() != RegionStartBB)
4178 OutsideUses.push_back(Elt: &U);
4179 }
4180
4181 if (OutsideUses.empty())
4182 continue;
4183
4184 HasBroadcastValues = true;
4185
4186 // Emit a global variable in shared memory to store the broadcasted
4187 // value.
4188 auto *SharedMem = new GlobalVariable(
4189 M, I.getType(), /* IsConstant */ false,
4190 GlobalValue::InternalLinkage, UndefValue::get(T: I.getType()),
4191 sanitizeForGlobalName(
4192 S: (I.getName() + ".guarded.output.alloc").str()),
4193 nullptr, GlobalValue::NotThreadLocal,
4194 static_cast<unsigned>(AddressSpace::Shared));
4195
4196 // Emit a store instruction to update the value.
4197 new StoreInst(&I, SharedMem,
4198 RegionEndBB->getTerminator()->getIterator());
4199
4200 LoadInst *LoadI = new LoadInst(
4201 I.getType(), SharedMem, I.getName() + ".guarded.output.load",
4202 RegionBarrierBB->getTerminator()->getIterator());
4203
4204 // Emit a load instruction and replace uses of the output value.
4205 for (Use *U : OutsideUses)
4206 A.changeUseAfterManifest(U&: *U, NV&: *LoadI);
4207 }
4208
4209 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
4210
4211 // Go to tid check BB in ParentBB.
4212 const DebugLoc DL = ParentBB->getTerminator()->getDebugLoc();
4213 ParentBB->getTerminator()->eraseFromParent();
4214 OpenMPIRBuilder::LocationDescription Loc(ParentBB->end(), DL);
4215 OMPInfoCache.OMPBuilder.updateToLocation(Loc);
4216 uint32_t SrcLocStrSize;
4217 auto *SrcLocStr =
4218 OMPInfoCache.OMPBuilder.getOrCreateSrcLocStr(Loc, SrcLocStrSize);
4219 Value *Ident =
4220 OMPInfoCache.OMPBuilder.getOrCreateIdent(SrcLocStr, SrcLocStrSize);
4221 UncondBrInst::Create(Target: RegionCheckTidBB, InsertBefore: ParentBB)->setDebugLoc(DL);
4222
4223 // Add check for Tid in RegionCheckTidBB
4224 RegionCheckTidBB->getTerminator()->eraseFromParent();
4225 OpenMPIRBuilder::LocationDescription LocRegionCheckTid(
4226 RegionCheckTidBB->end(), DL);
4227 OMPInfoCache.OMPBuilder.updateToLocation(Loc: LocRegionCheckTid);
4228 FunctionCallee HardwareTidFn =
4229 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4230 M, FnID: OMPRTL___kmpc_get_hardware_thread_id_in_block);
4231 CallInst *Tid =
4232 OMPInfoCache.OMPBuilder.Builder.CreateCall(Callee: HardwareTidFn, Args: {});
4233 Tid->setDebugLoc(DL);
4234 OMPInfoCache.setCallingConvention(Callee: HardwareTidFn, CI: Tid);
4235 Value *TidCheck = OMPInfoCache.OMPBuilder.Builder.CreateIsNull(Arg: Tid);
4236 OMPInfoCache.OMPBuilder.Builder
4237 .CreateCondBr(Cond: TidCheck, True: RegionStartBB, False: RegionBarrierBB)
4238 ->setDebugLoc(DL);
4239
4240 // First barrier for synchronization, ensures main thread has updated
4241 // values.
4242 FunctionCallee BarrierFn =
4243 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4244 M, FnID: OMPRTL___kmpc_barrier_simple_spmd);
4245 OMPInfoCache.OMPBuilder.updateToLocation(
4246 Loc: {RegionBarrierBB->getFirstInsertionPt(), DL});
4247 CallInst *Barrier =
4248 OMPInfoCache.OMPBuilder.Builder.CreateCall(Callee: BarrierFn, Args: {Ident, Tid});
4249 OMPInfoCache.setCallingConvention(Callee: BarrierFn, CI: Barrier);
4250
4251 // Second barrier ensures workers have read broadcast values.
4252 if (HasBroadcastValues) {
4253 CallInst *Barrier =
4254 CallInst::Create(Func: BarrierFn, Args: {Ident, Tid}, NameStr: "",
4255 InsertBefore: RegionBarrierBB->getTerminator()->getIterator());
4256 Barrier->setDebugLoc(DL);
4257 OMPInfoCache.setCallingConvention(Callee: BarrierFn, CI: Barrier);
4258 }
4259 };
4260
4261 auto &AllocSharedRFI = OMPInfoCache.RFIs[OMPRTL___kmpc_alloc_shared];
4262 SmallPtrSet<BasicBlock *, 8> Visited;
4263 for (Instruction *GuardedI : SPMDCompatibilityTracker) {
4264 BasicBlock *BB = GuardedI->getParent();
4265 if (!Visited.insert(Ptr: BB).second)
4266 continue;
4267
4268 SmallVector<std::pair<Instruction *, Instruction *>> Reorders;
4269 Instruction *LastEffect = nullptr;
4270 BasicBlock::reverse_iterator IP = BB->rbegin(), IPEnd = BB->rend();
4271 while (++IP != IPEnd) {
4272 if (!IP->mayHaveSideEffects() && !IP->mayReadFromMemory())
4273 continue;
4274 Instruction *I = &*IP;
4275 if (OpenMPOpt::getCallIfRegularCall(V&: *I, RFI: &AllocSharedRFI))
4276 continue;
4277 if (!I->user_empty() || !SPMDCompatibilityTracker.contains(Elem: I)) {
4278 LastEffect = nullptr;
4279 continue;
4280 }
4281 if (LastEffect)
4282 Reorders.push_back(Elt: {I, LastEffect});
4283 LastEffect = &*IP;
4284 }
4285 for (auto &Reorder : Reorders)
4286 Reorder.first->moveBefore(InsertPos: Reorder.second->getIterator());
4287 }
4288
4289 SmallVector<std::pair<Instruction *, Instruction *>, 4> GuardedRegions;
4290
4291 for (Instruction *GuardedI : SPMDCompatibilityTracker) {
4292 BasicBlock *BB = GuardedI->getParent();
4293 auto *CalleeAA = A.lookupAAFor<AAKernelInfo>(
4294 IRP: IRPosition::function(F: *GuardedI->getFunction()), QueryingAA: nullptr,
4295 DepClass: DepClassTy::NONE);
4296 assert(CalleeAA != nullptr && "Expected Callee AAKernelInfo");
4297 auto &CalleeAAFunction = *cast<AAKernelInfoFunction>(Val: CalleeAA);
4298 // Continue if instruction is already guarded.
4299 if (CalleeAAFunction.getGuardedInstructions().contains(Ptr: GuardedI))
4300 continue;
4301
4302 Instruction *GuardedRegionStart = nullptr, *GuardedRegionEnd = nullptr;
4303 for (Instruction &I : *BB) {
4304 // If instruction I needs to be guarded update the guarded region
4305 // bounds.
4306 if (SPMDCompatibilityTracker.contains(Elem: &I)) {
4307 CalleeAAFunction.getGuardedInstructions().insert(Ptr: &I);
4308 if (GuardedRegionStart)
4309 GuardedRegionEnd = &I;
4310 else
4311 GuardedRegionStart = GuardedRegionEnd = &I;
4312
4313 continue;
4314 }
4315
4316 // Instruction I does not need guarding, store
4317 // any region found and reset bounds.
4318 if (GuardedRegionStart) {
4319 GuardedRegions.push_back(
4320 Elt: std::make_pair(x&: GuardedRegionStart, y&: GuardedRegionEnd));
4321 GuardedRegionStart = nullptr;
4322 GuardedRegionEnd = nullptr;
4323 }
4324 }
4325 }
4326
4327 for (auto &GR : GuardedRegions)
4328 CreateGuardedRegion(GR.first, GR.second);
4329 }
4330
4331 void forceSingleThreadPerWorkgroupHelper(Attributor &A) {
4332 // Only allow 1 thread per workgroup to continue executing the user code.
4333 //
4334 // InitCB = __kmpc_target_init(...)
4335 // ThreadIdInBlock = __kmpc_get_hardware_thread_id_in_block();
4336 // if (ThreadIdInBlock != 0) return;
4337 // UserCode:
4338 // // user code
4339 //
4340 auto &Ctx = getAnchorValue().getContext();
4341 Function *Kernel = getAssociatedFunction();
4342 assert(Kernel && "Expected an associated function!");
4343
4344 // Create block for user code to branch to from initial block.
4345 BasicBlock *InitBB = KernelInitCB->getParent();
4346 BasicBlock *UserCodeBB = InitBB->splitBasicBlock(
4347 I: KernelInitCB->getNextNode(), BBName: "main.thread.user_code");
4348 BasicBlock *ReturnBB =
4349 BasicBlock::Create(Context&: Ctx, Name: "exit.threads", Parent: Kernel, InsertBefore: UserCodeBB);
4350
4351 // Register blocks with attributor:
4352 A.registerManifestAddedBasicBlock(BB&: *InitBB);
4353 A.registerManifestAddedBasicBlock(BB&: *UserCodeBB);
4354 A.registerManifestAddedBasicBlock(BB&: *ReturnBB);
4355
4356 // Debug location:
4357 const DebugLoc &DLoc = KernelInitCB->getDebugLoc();
4358 ReturnInst::Create(C&: Ctx, InsertAtEnd: ReturnBB)->setDebugLoc(DLoc);
4359 InitBB->getTerminator()->eraseFromParent();
4360
4361 // Prepare call to OMPRTL___kmpc_get_hardware_thread_id_in_block.
4362 Module &M = *Kernel->getParent();
4363 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
4364 FunctionCallee ThreadIdInBlockFn =
4365 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4366 M, FnID: OMPRTL___kmpc_get_hardware_thread_id_in_block);
4367
4368 // Get thread ID in block.
4369 CallInst *ThreadIdInBlock =
4370 CallInst::Create(Func: ThreadIdInBlockFn, NameStr: "thread_id.in.block", InsertBefore: InitBB);
4371 OMPInfoCache.setCallingConvention(Callee: ThreadIdInBlockFn, CI: ThreadIdInBlock);
4372 ThreadIdInBlock->setDebugLoc(DLoc);
4373
4374 // Eliminate all threads in the block with ID not equal to 0:
4375 Instruction *IsMainThread =
4376 ICmpInst::Create(Op: ICmpInst::ICmp, Pred: CmpInst::ICMP_NE, S1: ThreadIdInBlock,
4377 S2: ConstantInt::get(Ty: ThreadIdInBlock->getType(), V: 0),
4378 Name: "thread.is_main", InsertBefore: InitBB);
4379 IsMainThread->setDebugLoc(DLoc);
4380 CondBrInst::Create(Cond: IsMainThread, IfTrue: ReturnBB, IfFalse: UserCodeBB, InsertBefore: InitBB);
4381 }
4382
4383 bool changeToSPMDMode(Attributor &A, ChangeStatus &Changed) {
4384 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
4385
4386 if (!SPMDCompatibilityTracker.isAssumed()) {
4387 for (Instruction *NonCompatibleI : SPMDCompatibilityTracker) {
4388 if (!NonCompatibleI)
4389 continue;
4390
4391 // Skip diagnostics on calls to known OpenMP runtime functions for now.
4392 if (auto *CB = dyn_cast<CallBase>(Val: NonCompatibleI))
4393 if (OMPInfoCache.RTLFunctions.contains(V: CB->getCalledFunction()))
4394 continue;
4395
4396 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
4397 ORA << "Value has potential side effects preventing SPMD-mode "
4398 "execution";
4399 if (isa<CallBase>(Val: NonCompatibleI)) {
4400 ORA << ". Add `[[omp::assume(\"ompx_spmd_amenable\")]]` to "
4401 "the called function to override";
4402 }
4403 return ORA << ".";
4404 };
4405 A.emitRemark<OptimizationRemarkAnalysis>(I: NonCompatibleI, RemarkName: "OMP121",
4406 RemarkCB&: Remark);
4407
4408 LLVM_DEBUG(dbgs() << TAG << "SPMD-incompatible side-effect: "
4409 << *NonCompatibleI << "\n");
4410 }
4411
4412 return false;
4413 }
4414
4415 // Get the actual kernel, could be the caller of the anchor scope if we have
4416 // a debug wrapper.
4417 Function *Kernel = getAnchorScope();
4418 if (Kernel->hasLocalLinkage()) {
4419 assert(Kernel->hasOneUse() && "Unexpected use of debug kernel wrapper.");
4420 auto *CB = cast<CallBase>(Val: Kernel->user_back());
4421 Kernel = CB->getCaller();
4422 }
4423 assert(omp::isOpenMPKernel(*Kernel) && "Expected kernel function!");
4424
4425 // Check if the kernel is already in SPMD mode, if so, return success.
4426 ConstantStruct *ExistingKernelEnvC =
4427 KernelInfo::getKernelEnvironementFromKernelInitCB(KernelInitCB);
4428 auto *ExecModeC =
4429 KernelInfo::getExecModeFromKernelEnvironment(KernelEnvC: ExistingKernelEnvC);
4430 const int8_t ExecModeVal = ExecModeC->getSExtValue();
4431 if (ExecModeVal != OMP_TGT_EXEC_MODE_GENERIC)
4432 return true;
4433
4434 // We will now unconditionally modify the IR, indicate a change.
4435 Changed = ChangeStatus::CHANGED;
4436
4437 // Do not use instruction guards when no parallel is present inside
4438 // the target region.
4439 if (mayContainParallelRegion())
4440 insertInstructionGuardsHelper(A);
4441 else
4442 forceSingleThreadPerWorkgroupHelper(A);
4443
4444 // Adjust the global exec mode flag that tells the runtime what mode this
4445 // kernel is executed in.
4446 assert(ExecModeVal == OMP_TGT_EXEC_MODE_GENERIC &&
4447 "Initially non-SPMD kernel has SPMD exec mode!");
4448 setExecModeOfKernelEnvironment(
4449 ConstantInt::get(Ty: ExecModeC->getIntegerType(),
4450 V: ExecModeVal | OMP_TGT_EXEC_MODE_GENERIC_SPMD));
4451
4452 ++NumOpenMPTargetRegionKernelsSPMD;
4453
4454 // Record that this kernel now runs SPMD so post-Attributor cleanup can drop
4455 // the now-dead parallel data-sharing wrapper without re-deriving the mode.
4456 OMPInfoCache.SPMDizedKernels.insert(Ptr: Kernel);
4457
4458 auto Remark = [&](OptimizationRemark OR) {
4459 return OR << "Transformed generic-mode kernel to SPMD-mode.";
4460 };
4461 A.emitRemark<OptimizationRemark>(I: KernelInitCB, RemarkName: "OMP120", RemarkCB&: Remark);
4462 return true;
4463 };
4464
4465 bool buildCustomStateMachine(Attributor &A, ChangeStatus &Changed) {
4466 // If we have disabled state machine rewrites, don't make a custom one
4467 if (DisableOpenMPOptStateMachineRewrite)
4468 return false;
4469
4470 // Don't rewrite the state machine if we are not in a valid state.
4471 if (!ReachedKnownParallelRegions.isValidState())
4472 return false;
4473
4474 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
4475 if (!OMPInfoCache.runtimeFnsAvailable(Fns: {OMPRTL___kmpc_get_max_team_threads,
4476 OMPRTL___kmpc_barrier_simple_generic,
4477 OMPRTL___kmpc_kernel_parallel,
4478 OMPRTL___kmpc_kernel_end_parallel}))
4479 return false;
4480
4481 ConstantStruct *ExistingKernelEnvC =
4482 KernelInfo::getKernelEnvironementFromKernelInitCB(KernelInitCB);
4483
4484 // Check if the current configuration is non-SPMD and generic state machine.
4485 // If we already have SPMD mode or a custom state machine we do not need to
4486 // go any further. If it is anything but a constant something is weird and
4487 // we give up.
4488 ConstantInt *UseStateMachineC =
4489 KernelInfo::getUseGenericStateMachineFromKernelEnvironment(
4490 KernelEnvC: ExistingKernelEnvC);
4491 ConstantInt *ModeC =
4492 KernelInfo::getExecModeFromKernelEnvironment(KernelEnvC: ExistingKernelEnvC);
4493
4494 // If we are stuck with generic mode, try to create a custom device (=GPU)
4495 // state machine which is specialized for the parallel regions that are
4496 // reachable by the kernel.
4497 if (UseStateMachineC->isZero() ||
4498 (ModeC->getSExtValue() & OMP_TGT_EXEC_MODE_SPMD))
4499 return false;
4500
4501 Changed = ChangeStatus::CHANGED;
4502
4503 // If not SPMD mode, indicate we use a custom state machine now.
4504 setUseGenericStateMachineOfKernelEnvironment(
4505 ConstantInt::get(Ty: UseStateMachineC->getIntegerType(), V: false));
4506
4507 // If we don't actually need a state machine we are done here. This can
4508 // happen if there simply are no parallel regions. In the resulting kernel
4509 // all worker threads will simply exit right away, leaving the main thread
4510 // to do the work alone.
4511 if (!mayContainParallelRegion()) {
4512 ++NumOpenMPTargetRegionKernelsWithoutStateMachine;
4513
4514 auto Remark = [&](OptimizationRemark OR) {
4515 return OR << "Removing unused state machine from generic-mode kernel.";
4516 };
4517 A.emitRemark<OptimizationRemark>(I: KernelInitCB, RemarkName: "OMP130", RemarkCB&: Remark);
4518
4519 return true;
4520 }
4521
4522 // Keep track in the statistics of our new shiny custom state machine.
4523 if (ReachedUnknownParallelRegions.empty()) {
4524 ++NumOpenMPTargetRegionKernelsCustomStateMachineWithoutFallback;
4525
4526 auto Remark = [&](OptimizationRemark OR) {
4527 return OR << "Rewriting generic-mode kernel with a customized state "
4528 "machine.";
4529 };
4530 A.emitRemark<OptimizationRemark>(I: KernelInitCB, RemarkName: "OMP131", RemarkCB&: Remark);
4531 } else {
4532 ++NumOpenMPTargetRegionKernelsCustomStateMachineWithFallback;
4533
4534 auto Remark = [&](OptimizationRemarkAnalysis OR) {
4535 return OR << "Generic-mode kernel is executed with a customized state "
4536 "machine that requires a fallback.";
4537 };
4538 A.emitRemark<OptimizationRemarkAnalysis>(I: KernelInitCB, RemarkName: "OMP132", RemarkCB&: Remark);
4539
4540 // Tell the user why we ended up with a fallback.
4541 for (CallBase *UnknownParallelRegionCB : ReachedUnknownParallelRegions) {
4542 if (!UnknownParallelRegionCB)
4543 continue;
4544 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
4545 return ORA << "Call may contain unknown parallel regions. Use "
4546 << "`[[omp::assume(\"omp_no_parallelism\")]]` to "
4547 "override.";
4548 };
4549 A.emitRemark<OptimizationRemarkAnalysis>(I: UnknownParallelRegionCB,
4550 RemarkName: "OMP133", RemarkCB&: Remark);
4551 }
4552 }
4553
4554 // Create all the blocks:
4555 //
4556 // InitCB = __kmpc_target_init(...)
4557 // MaxTeamThreads =
4558 // __kmpc_get_max_team_threads(/*IsSPMD=*/false);
4559 // IsWorkerCheckBB: bool IsWorker = InitCB != -1;
4560 // if (IsWorker) {
4561 // if (InitCB >= MaxTeamThreads) return;
4562 // SMBeginBB: __kmpc_barrier_simple_generic(...);
4563 // void *WorkFn;
4564 // bool Active = __kmpc_kernel_parallel(&WorkFn);
4565 // if (!WorkFn) return;
4566 // SMIsActiveCheckBB: if (Active) {
4567 // SMIfCascadeCurrentBB: if (WorkFn == <ParFn0>)
4568 // ParFn0(...);
4569 // SMIfCascadeCurrentBB: else if (WorkFn == <ParFn1>)
4570 // ParFn1(...);
4571 // ...
4572 // SMIfCascadeCurrentBB: else
4573 // ((WorkFnTy*)WorkFn)(...);
4574 // SMEndParallelBB: __kmpc_kernel_end_parallel(...);
4575 // }
4576 // SMDoneBB: __kmpc_barrier_simple_generic(...);
4577 // goto SMBeginBB;
4578 // }
4579 // UserCodeEntryBB: // user code
4580 // __kmpc_target_deinit(...)
4581 //
4582 auto &Ctx = getAnchorValue().getContext();
4583 Function *Kernel = getAssociatedFunction();
4584 assert(Kernel && "Expected an associated function!");
4585
4586 BasicBlock *InitBB = KernelInitCB->getParent();
4587 BasicBlock *UserCodeEntryBB = InitBB->splitBasicBlock(
4588 I: KernelInitCB->getNextNode(), BBName: "thread.user_code.check");
4589 BasicBlock *IsWorkerCheckBB =
4590 BasicBlock::Create(Context&: Ctx, Name: "is_worker_check", Parent: Kernel, InsertBefore: UserCodeEntryBB);
4591 BasicBlock *StateMachineBeginBB = BasicBlock::Create(
4592 Context&: Ctx, Name: "worker_state_machine.begin", Parent: Kernel, InsertBefore: UserCodeEntryBB);
4593 BasicBlock *StateMachineFinishedBB = BasicBlock::Create(
4594 Context&: Ctx, Name: "worker_state_machine.finished", Parent: Kernel, InsertBefore: UserCodeEntryBB);
4595 BasicBlock *StateMachineIsActiveCheckBB = BasicBlock::Create(
4596 Context&: Ctx, Name: "worker_state_machine.is_active.check", Parent: Kernel, InsertBefore: UserCodeEntryBB);
4597 BasicBlock *StateMachineIfCascadeCurrentBB =
4598 BasicBlock::Create(Context&: Ctx, Name: "worker_state_machine.parallel_region.check",
4599 Parent: Kernel, InsertBefore: UserCodeEntryBB);
4600 BasicBlock *StateMachineEndParallelBB =
4601 BasicBlock::Create(Context&: Ctx, Name: "worker_state_machine.parallel_region.end",
4602 Parent: Kernel, InsertBefore: UserCodeEntryBB);
4603 BasicBlock *StateMachineDoneBarrierBB = BasicBlock::Create(
4604 Context&: Ctx, Name: "worker_state_machine.done.barrier", Parent: Kernel, InsertBefore: UserCodeEntryBB);
4605 A.registerManifestAddedBasicBlock(BB&: *InitBB);
4606 A.registerManifestAddedBasicBlock(BB&: *UserCodeEntryBB);
4607 A.registerManifestAddedBasicBlock(BB&: *IsWorkerCheckBB);
4608 A.registerManifestAddedBasicBlock(BB&: *StateMachineBeginBB);
4609 A.registerManifestAddedBasicBlock(BB&: *StateMachineFinishedBB);
4610 A.registerManifestAddedBasicBlock(BB&: *StateMachineIsActiveCheckBB);
4611 A.registerManifestAddedBasicBlock(BB&: *StateMachineIfCascadeCurrentBB);
4612 A.registerManifestAddedBasicBlock(BB&: *StateMachineEndParallelBB);
4613 A.registerManifestAddedBasicBlock(BB&: *StateMachineDoneBarrierBB);
4614
4615 const DebugLoc &DLoc = KernelInitCB->getDebugLoc();
4616 ReturnInst::Create(C&: Ctx, InsertAtEnd: StateMachineFinishedBB)->setDebugLoc(DLoc);
4617 InitBB->getTerminator()->eraseFromParent();
4618
4619 Instruction *IsWorker =
4620 ICmpInst::Create(Op: ICmpInst::ICmp, Pred: llvm::CmpInst::ICMP_NE, S1: KernelInitCB,
4621 S2: ConstantInt::getAllOnesValue(Ty: KernelInitCB->getType()),
4622 Name: "thread.is_worker", InsertBefore: InitBB);
4623 IsWorker->setDebugLoc(DLoc);
4624 CondBrInst::Create(Cond: IsWorker, IfTrue: IsWorkerCheckBB, IfFalse: UserCodeEntryBB, InsertBefore: InitBB);
4625
4626 // How much of the block the main thread takes is the runtime's to know, so
4627 // ask it rather than subtracting a warp here. The mode is passed in because
4628 // this runs before the barrier that would make the shared one visible; it
4629 // is a constant, a custom state machine being built only for generic mode.
4630 Module &M = *Kernel->getParent();
4631 FunctionCallee MaxTeamThreadsFn =
4632 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4633 M, FnID: OMPRTL___kmpc_get_max_team_threads);
4634 Constant *IsSPMDArg = ConstantInt::get(Ty: OMPInfoCache.OMPBuilder.Int32, V: 0);
4635 CallInst *MaxTeamThreads = CallInst::Create(
4636 Func: MaxTeamThreadsFn, Args: {IsSPMDArg}, NameStr: "max_team_threads", InsertBefore: IsWorkerCheckBB);
4637 OMPInfoCache.setCallingConvention(Callee: MaxTeamThreadsFn, CI: MaxTeamThreads);
4638 MaxTeamThreads->setDebugLoc(DLoc);
4639 Instruction *IsMainOrWorker = ICmpInst::Create(
4640 Op: ICmpInst::ICmp, Pred: llvm::CmpInst::ICMP_SLT, S1: KernelInitCB, S2: MaxTeamThreads,
4641 Name: "thread.is_main_or_worker", InsertBefore: IsWorkerCheckBB);
4642 IsMainOrWorker->setDebugLoc(DLoc);
4643 CondBrInst::Create(Cond: IsMainOrWorker, IfTrue: StateMachineBeginBB,
4644 IfFalse: StateMachineFinishedBB, InsertBefore: IsWorkerCheckBB);
4645
4646 // Create local storage for the work function pointer.
4647 const DataLayout &DL = M.getDataLayout();
4648 Type *VoidPtrTy = PointerType::getUnqual(C&: Ctx);
4649 Instruction *WorkFnAI =
4650 new AllocaInst(VoidPtrTy, DL.getAllocaAddrSpace(), nullptr,
4651 "worker.work_fn.addr", Kernel->getEntryBlock().begin());
4652 WorkFnAI->setDebugLoc(DLoc);
4653
4654 OMPInfoCache.OMPBuilder.updateToLocation(
4655 Loc: OpenMPIRBuilder::LocationDescription(StateMachineBeginBB->end(), DLoc));
4656
4657 Value *Ident = KernelInfo::getIdentFromKernelEnvironment(KernelEnvC);
4658 Value *GTid = KernelInitCB;
4659
4660 FunctionCallee BarrierFn =
4661 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4662 M, FnID: OMPRTL___kmpc_barrier_simple_generic);
4663 CallInst *Barrier =
4664 CallInst::Create(Func: BarrierFn, Args: {Ident, GTid}, NameStr: "", InsertBefore: StateMachineBeginBB);
4665 OMPInfoCache.setCallingConvention(Callee: BarrierFn, CI: Barrier);
4666 Barrier->setDebugLoc(DLoc);
4667
4668 if (WorkFnAI->getType()->getPointerAddressSpace() !=
4669 (unsigned int)AddressSpace::Generic) {
4670 WorkFnAI = new AddrSpaceCastInst(
4671 WorkFnAI, PointerType::get(C&: Ctx, AddressSpace: (unsigned int)AddressSpace::Generic),
4672 WorkFnAI->getName() + ".generic", StateMachineBeginBB);
4673 WorkFnAI->setDebugLoc(DLoc);
4674 }
4675
4676 FunctionCallee KernelParallelFn =
4677 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4678 M, FnID: OMPRTL___kmpc_kernel_parallel);
4679 CallInst *IsActiveWorker = CallInst::Create(
4680 Func: KernelParallelFn, Args: {WorkFnAI}, NameStr: "worker.is_active", InsertBefore: StateMachineBeginBB);
4681 OMPInfoCache.setCallingConvention(Callee: KernelParallelFn, CI: IsActiveWorker);
4682 IsActiveWorker->setDebugLoc(DLoc);
4683 Instruction *WorkFn = new LoadInst(VoidPtrTy, WorkFnAI, "worker.work_fn",
4684 StateMachineBeginBB);
4685 WorkFn->setDebugLoc(DLoc);
4686
4687 FunctionType *ParallelRegionFnTy = FunctionType::get(
4688 Result: Type::getVoidTy(C&: Ctx), Params: {Type::getInt16Ty(C&: Ctx), Type::getInt32Ty(C&: Ctx)},
4689 isVarArg: false);
4690
4691 Instruction *IsDone =
4692 ICmpInst::Create(Op: ICmpInst::ICmp, Pred: llvm::CmpInst::ICMP_EQ, S1: WorkFn,
4693 S2: Constant::getNullValue(Ty: VoidPtrTy), Name: "worker.is_done",
4694 InsertBefore: StateMachineBeginBB);
4695 IsDone->setDebugLoc(DLoc);
4696 CondBrInst::Create(Cond: IsDone, IfTrue: StateMachineFinishedBB,
4697 IfFalse: StateMachineIsActiveCheckBB, InsertBefore: StateMachineBeginBB)
4698 ->setDebugLoc(DLoc);
4699
4700 CondBrInst::Create(Cond: IsActiveWorker, IfTrue: StateMachineIfCascadeCurrentBB,
4701 IfFalse: StateMachineDoneBarrierBB, InsertBefore: StateMachineIsActiveCheckBB)
4702 ->setDebugLoc(DLoc);
4703
4704 Value *ZeroArg =
4705 Constant::getNullValue(Ty: ParallelRegionFnTy->getParamType(i: 0));
4706
4707 const unsigned int WrapperFunctionArgNo = 6;
4708
4709 // Now that we have most of the CFG skeleton it is time for the if-cascade
4710 // that checks the function pointer we got from the runtime against the
4711 // parallel regions we expect, if there are any.
4712 for (int I = 0, E = ReachedKnownParallelRegions.size(); I < E; ++I) {
4713 auto *CB = ReachedKnownParallelRegions[I];
4714 auto *ParallelRegion = dyn_cast<Function>(
4715 Val: CB->getArgOperand(i: WrapperFunctionArgNo)->stripPointerCasts());
4716 BasicBlock *PRExecuteBB = BasicBlock::Create(
4717 Context&: Ctx, Name: "worker_state_machine.parallel_region.execute", Parent: Kernel,
4718 InsertBefore: StateMachineEndParallelBB);
4719 CallInst::Create(Func: ParallelRegion, Args: {ZeroArg, GTid}, NameStr: "", InsertBefore: PRExecuteBB)
4720 ->setDebugLoc(DLoc);
4721 UncondBrInst::Create(Target: StateMachineEndParallelBB, InsertBefore: PRExecuteBB)
4722 ->setDebugLoc(DLoc);
4723
4724 BasicBlock *PRNextBB =
4725 BasicBlock::Create(Context&: Ctx, Name: "worker_state_machine.parallel_region.check",
4726 Parent: Kernel, InsertBefore: StateMachineEndParallelBB);
4727 A.registerManifestAddedBasicBlock(BB&: *PRExecuteBB);
4728 A.registerManifestAddedBasicBlock(BB&: *PRNextBB);
4729
4730 // Check if we need to compare the pointer at all or if we can just
4731 // call the parallel region function.
4732 Value *IsPR;
4733 if (I + 1 < E || !ReachedUnknownParallelRegions.empty()) {
4734 Instruction *CmpI = ICmpInst::Create(
4735 Op: ICmpInst::ICmp, Pred: llvm::CmpInst::ICMP_EQ, S1: WorkFn, S2: ParallelRegion,
4736 Name: "worker.check_parallel_region", InsertBefore: StateMachineIfCascadeCurrentBB);
4737 CmpI->setDebugLoc(DLoc);
4738 IsPR = CmpI;
4739 } else {
4740 IsPR = ConstantInt::getTrue(Context&: Ctx);
4741 }
4742
4743 CondBrInst::Create(Cond: IsPR, IfTrue: PRExecuteBB, IfFalse: PRNextBB,
4744 InsertBefore: StateMachineIfCascadeCurrentBB)
4745 ->setDebugLoc(DLoc);
4746 StateMachineIfCascadeCurrentBB = PRNextBB;
4747 }
4748
4749 // At the end of the if-cascade we place the indirect function pointer call
4750 // in case we might need it, that is if there can be parallel regions we
4751 // have not handled in the if-cascade above.
4752 if (!ReachedUnknownParallelRegions.empty()) {
4753 StateMachineIfCascadeCurrentBB->setName(
4754 "worker_state_machine.parallel_region.fallback.execute");
4755 CallInst::Create(Ty: ParallelRegionFnTy, Func: WorkFn, Args: {ZeroArg, GTid}, NameStr: "",
4756 InsertBefore: StateMachineIfCascadeCurrentBB)
4757 ->setDebugLoc(DLoc);
4758 }
4759 UncondBrInst::Create(Target: StateMachineEndParallelBB,
4760 InsertBefore: StateMachineIfCascadeCurrentBB)
4761 ->setDebugLoc(DLoc);
4762
4763 FunctionCallee EndParallelFn =
4764 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4765 M, FnID: OMPRTL___kmpc_kernel_end_parallel);
4766 CallInst *EndParallel =
4767 CallInst::Create(Func: EndParallelFn, Args: {}, NameStr: "", InsertBefore: StateMachineEndParallelBB);
4768 OMPInfoCache.setCallingConvention(Callee: EndParallelFn, CI: EndParallel);
4769 EndParallel->setDebugLoc(DLoc);
4770 UncondBrInst::Create(Target: StateMachineDoneBarrierBB, InsertBefore: StateMachineEndParallelBB)
4771 ->setDebugLoc(DLoc);
4772
4773 CallInst::Create(Func: BarrierFn, Args: {Ident, GTid}, NameStr: "", InsertBefore: StateMachineDoneBarrierBB)
4774 ->setDebugLoc(DLoc);
4775 UncondBrInst::Create(Target: StateMachineBeginBB, InsertBefore: StateMachineDoneBarrierBB)
4776 ->setDebugLoc(DLoc);
4777
4778 return true;
4779 }
4780
4781 /// Fixpoint iteration update function. Will be called every time a dependence
4782 /// changed its state (and in the beginning).
4783 ChangeStatus updateImpl(Attributor &A) override {
4784 KernelInfoState StateBefore = getState();
4785
4786 // When we leave this function this RAII will make sure the member
4787 // KernelEnvC is updated properly depending on the state. That member is
4788 // used for simplification of values and needs to be up to date at all
4789 // times.
4790 struct UpdateKernelEnvCRAII {
4791 AAKernelInfoFunction &AA;
4792
4793 UpdateKernelEnvCRAII(AAKernelInfoFunction &AA) : AA(AA) {}
4794
4795 ~UpdateKernelEnvCRAII() {
4796 if (!AA.KernelEnvC)
4797 return;
4798
4799 ConstantStruct *ExistingKernelEnvC =
4800 KernelInfo::getKernelEnvironementFromKernelInitCB(KernelInitCB: AA.KernelInitCB);
4801
4802 if (!AA.isValidState()) {
4803 AA.KernelEnvC = ExistingKernelEnvC;
4804 return;
4805 }
4806
4807 if (!AA.ReachedKnownParallelRegions.isValidState())
4808 AA.setUseGenericStateMachineOfKernelEnvironment(
4809 KernelInfo::getUseGenericStateMachineFromKernelEnvironment(
4810 KernelEnvC: ExistingKernelEnvC));
4811
4812 if (!AA.SPMDCompatibilityTracker.isValidState())
4813 AA.setExecModeOfKernelEnvironment(
4814 KernelInfo::getExecModeFromKernelEnvironment(KernelEnvC: ExistingKernelEnvC));
4815
4816 ConstantInt *MayUseNestedParallelismC =
4817 KernelInfo::getMayUseNestedParallelismFromKernelEnvironment(
4818 KernelEnvC: AA.KernelEnvC);
4819 ConstantInt *NewMayUseNestedParallelismC = ConstantInt::get(
4820 Ty: MayUseNestedParallelismC->getIntegerType(), V: AA.NestedParallelism);
4821 AA.setMayUseNestedParallelismOfKernelEnvironment(
4822 NewMayUseNestedParallelismC);
4823 }
4824 } RAII(*this);
4825
4826 // Callback to check a read/write instruction.
4827 auto CheckRWInst = [&](Instruction &I) {
4828 // We handle calls later.
4829 if (isa<CallBase>(Val: I))
4830 return true;
4831 // We only care about write effects.
4832 if (!I.mayWriteToMemory())
4833 return true;
4834 if (auto *SI = dyn_cast<StoreInst>(Val: &I)) {
4835 const auto *UnderlyingObjsAA = A.getAAFor<AAUnderlyingObjects>(
4836 QueryingAA: *this, IRP: IRPosition::value(V: *SI->getPointerOperand()),
4837 DepClass: DepClassTy::OPTIONAL);
4838 auto *HS = A.getAAFor<AAHeapToStack>(
4839 QueryingAA: *this, IRP: IRPosition::function(F: *I.getFunction()),
4840 DepClass: DepClassTy::OPTIONAL);
4841 if (UnderlyingObjsAA &&
4842 UnderlyingObjsAA->forallUnderlyingObjects(Pred: [&](Value &Obj) {
4843 if (AA::isAssumedThreadLocalObject(A, Obj, QueryingAA: *this))
4844 return true;
4845 // Check for AAHeapToStack moved objects which must not be
4846 // guarded.
4847 auto *CB = dyn_cast<CallBase>(Val: &Obj);
4848 return CB && HS && HS->isAssumedHeapToStack(CB: *CB);
4849 }))
4850 return true;
4851 }
4852
4853 // Insert instruction that needs guarding.
4854 SPMDCompatibilityTracker.insert(Elem: &I);
4855 return true;
4856 };
4857
4858 bool UsedAssumedInformationInCheckRWInst = false;
4859 if (!SPMDCompatibilityTracker.isAtFixpoint())
4860 if (!A.checkForAllReadWriteInstructions(
4861 Pred: CheckRWInst, QueryingAA&: *this, UsedAssumedInformation&: UsedAssumedInformationInCheckRWInst))
4862 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
4863
4864 bool UsedAssumedInformationFromReachingKernels = false;
4865 if (!IsKernelEntry) {
4866 updateParallelLevels(A);
4867
4868 bool AllReachingKernelsKnown = true;
4869 updateReachingKernelEntries(A, AllReachingKernelsKnown);
4870 UsedAssumedInformationFromReachingKernels = !AllReachingKernelsKnown;
4871
4872 if (!SPMDCompatibilityTracker.empty()) {
4873 if (!ParallelLevels.isValidState())
4874 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
4875 else if (!ReachingKernelEntries.isValidState())
4876 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
4877 else {
4878 // Check if all reaching kernels agree on the mode as we can otherwise
4879 // not guard instructions. We might not be sure about the mode so we
4880 // we cannot fix the internal spmd-zation state either.
4881 int SPMD = 0, Generic = 0;
4882 for (auto *Kernel : ReachingKernelEntries) {
4883 auto *CBAA = A.getAAFor<AAKernelInfo>(
4884 QueryingAA: *this, IRP: IRPosition::function(F: *Kernel), DepClass: DepClassTy::OPTIONAL);
4885 if (CBAA && CBAA->SPMDCompatibilityTracker.isValidState() &&
4886 CBAA->SPMDCompatibilityTracker.isAssumed())
4887 ++SPMD;
4888 else
4889 ++Generic;
4890 if (!CBAA || !CBAA->SPMDCompatibilityTracker.isAtFixpoint())
4891 UsedAssumedInformationFromReachingKernels = true;
4892 }
4893 if (SPMD != 0 && Generic != 0)
4894 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
4895 }
4896 }
4897 }
4898
4899 // Callback to check a call instruction.
4900 bool AllParallelRegionStatesWereFixed = true;
4901 bool AllSPMDStatesWereFixed = true;
4902 auto CheckCallInst = [&](Instruction &I) {
4903 auto &CB = cast<CallBase>(Val&: I);
4904 // A runtime function that takes a callback runs the user's code inside
4905 // it, so whatever the callback reaches this kernel reaches too. Fold the
4906 // callback's state in; without this the call tells us nothing about the
4907 // parallel regions on the other side of it.
4908 if (Function *Callback = OMPInformationCache::getAnalyzableCallback(CB)) {
4909 LLVM_DEBUG(dbgs() << TAG << "folding in callback "
4910 << Callback->getName() << " of " << CB << "\n");
4911 if (auto *CallbackAA = A.getAAFor<AAKernelInfo>(
4912 QueryingAA: *this, IRP: IRPosition::function(F: *Callback), DepClass: DepClassTy::OPTIONAL)) {
4913 getState() ^= CallbackAA->getState();
4914 AllSPMDStatesWereFixed &=
4915 CallbackAA->SPMDCompatibilityTracker.isAtFixpoint();
4916 AllParallelRegionStatesWereFixed &=
4917 CallbackAA->ReachedKnownParallelRegions.isAtFixpoint();
4918 AllParallelRegionStatesWereFixed &=
4919 CallbackAA->ReachedUnknownParallelRegions.isAtFixpoint();
4920 }
4921 }
4922 auto *CBAA = A.getAAFor<AAKernelInfo>(
4923 QueryingAA: *this, IRP: IRPosition::callsite_function(CB), DepClass: DepClassTy::OPTIONAL);
4924 if (!CBAA)
4925 return false;
4926 getState() ^= CBAA->getState();
4927 AllSPMDStatesWereFixed &= CBAA->SPMDCompatibilityTracker.isAtFixpoint();
4928 AllParallelRegionStatesWereFixed &=
4929 CBAA->ReachedKnownParallelRegions.isAtFixpoint();
4930 AllParallelRegionStatesWereFixed &=
4931 CBAA->ReachedUnknownParallelRegions.isAtFixpoint();
4932 return true;
4933 };
4934
4935 bool UsedAssumedInformationInCheckCallInst = false;
4936 if (!A.checkForAllCallLikeInstructions(
4937 Pred: CheckCallInst, QueryingAA: *this, UsedAssumedInformation&: UsedAssumedInformationInCheckCallInst)) {
4938 LLVM_DEBUG(dbgs() << TAG
4939 << "Failed to visit all call-like instructions!\n";);
4940 return indicatePessimisticFixpoint();
4941 }
4942
4943 // If we haven't used any assumed information for the reached parallel
4944 // region states we can fix it.
4945 if (!UsedAssumedInformationInCheckCallInst &&
4946 AllParallelRegionStatesWereFixed) {
4947 ReachedKnownParallelRegions.indicateOptimisticFixpoint();
4948 ReachedUnknownParallelRegions.indicateOptimisticFixpoint();
4949 }
4950
4951 // If we haven't used any assumed information for the SPMD state we can fix
4952 // it.
4953 if (!UsedAssumedInformationInCheckRWInst &&
4954 !UsedAssumedInformationInCheckCallInst &&
4955 !UsedAssumedInformationFromReachingKernels && AllSPMDStatesWereFixed)
4956 SPMDCompatibilityTracker.indicateOptimisticFixpoint();
4957
4958 return StateBefore == getState() ? ChangeStatus::UNCHANGED
4959 : ChangeStatus::CHANGED;
4960 }
4961
4962private:
4963 /// Update info regarding reaching kernels.
4964 void updateReachingKernelEntries(Attributor &A,
4965 bool &AllReachingKernelsKnown) {
4966 auto PredCallSite = [&](AbstractCallSite ACS) {
4967 Function *Caller = ACS.getInstruction()->getFunction();
4968
4969 assert(Caller && "Caller is nullptr");
4970
4971 auto *CAA = A.getOrCreateAAFor<AAKernelInfo>(
4972 IRP: IRPosition::function(F: *Caller), QueryingAA: this, DepClass: DepClassTy::REQUIRED);
4973 if (CAA && CAA->ReachingKernelEntries.isValidState()) {
4974 ReachingKernelEntries ^= CAA->ReachingKernelEntries;
4975 return true;
4976 }
4977
4978 // We lost track of the caller of the associated function, any kernel
4979 // could reach now.
4980 ReachingKernelEntries.indicatePessimisticFixpoint();
4981
4982 return true;
4983 };
4984
4985 if (!A.checkForAllCallSites(Pred: PredCallSite, QueryingAA: *this,
4986 RequireAllCallSites: true /* RequireAllCallSites */,
4987 UsedAssumedInformation&: AllReachingKernelsKnown))
4988 ReachingKernelEntries.indicatePessimisticFixpoint();
4989 }
4990
4991 /// Update info regarding parallel levels.
4992 void updateParallelLevels(Attributor &A) {
4993 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
4994 OMPInformationCache::RuntimeFunctionInfo &Parallel60RFI =
4995 OMPInfoCache.RFIs[OMPRTL___kmpc_parallel_60];
4996
4997 auto PredCallSite = [&](AbstractCallSite ACS) {
4998 Function *Caller = ACS.getInstruction()->getFunction();
4999
5000 assert(Caller && "Caller is nullptr");
5001
5002 auto *CAA =
5003 A.getOrCreateAAFor<AAKernelInfo>(IRP: IRPosition::function(F: *Caller));
5004 if (CAA && CAA->ParallelLevels.isValidState()) {
5005 // Any function that is called by `__kmpc_parallel_60` will not be
5006 // folded as the parallel level in the function is updated. In order to
5007 // get it right, all the analysis would depend on the implentation. That
5008 // said, if in the future any change to the implementation, the analysis
5009 // could be wrong. As a consequence, we are just conservative here.
5010 if (Caller == Parallel60RFI.Declaration) {
5011 ParallelLevels.indicatePessimisticFixpoint();
5012 return true;
5013 }
5014
5015 ParallelLevels ^= CAA->ParallelLevels;
5016
5017 return true;
5018 }
5019
5020 // We lost track of the caller of the associated function, any kernel
5021 // could reach now.
5022 ParallelLevels.indicatePessimisticFixpoint();
5023
5024 return true;
5025 };
5026
5027 bool AllCallSitesKnown = true;
5028 if (!A.checkForAllCallSites(Pred: PredCallSite, QueryingAA: *this,
5029 RequireAllCallSites: true /* RequireAllCallSites */,
5030 UsedAssumedInformation&: AllCallSitesKnown))
5031 ParallelLevels.indicatePessimisticFixpoint();
5032 }
5033};
5034
5035/// The call site kernel info abstract attribute, basically, what can we say
5036/// about a call site with regards to the KernelInfoState. For now this simply
5037/// forwards the information from the callee.
5038struct AAKernelInfoCallSite : AAKernelInfo {
5039 AAKernelInfoCallSite(const IRPosition &IRP, Attributor &A)
5040 : AAKernelInfo(IRP, A) {}
5041
5042 /// See AbstractAttribute::initialize(...).
5043 void initialize(Attributor &A) override {
5044 AAKernelInfo::initialize(A);
5045
5046 CallBase &CB = cast<CallBase>(Val&: getAssociatedValue());
5047 auto *AssumptionAA = A.getAAFor<AAAssumptionInfo>(
5048 QueryingAA: *this, IRP: IRPosition::callsite_function(CB), DepClass: DepClassTy::OPTIONAL);
5049
5050 // Check for SPMD-mode assumptions.
5051 if (AssumptionAA && AssumptionAA->hasAssumption(Assumption: "ompx_spmd_amenable")) {
5052 indicateOptimisticFixpoint();
5053 return;
5054 }
5055
5056 // First weed out calls we do not care about, that is readonly/readnone
5057 // calls, intrinsics, and "no_openmp" calls. Neither of these can reach a
5058 // parallel region or anything else we are looking for.
5059 if (!CB.mayWriteToMemory() || isa<IntrinsicInst>(Val: CB)) {
5060 indicateOptimisticFixpoint();
5061 return;
5062 }
5063
5064 // Next we check if we know the callee. If it is a known OpenMP function
5065 // we will handle them explicitly in the switch below. If it is not, we
5066 // will use an AAKernelInfo object on the callee to gather information and
5067 // merge that into the current state. The latter happens in the updateImpl.
5068 auto CheckCallee = [&](Function *Callee, unsigned NumCallees) {
5069 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
5070 const auto &It = OMPInfoCache.RuntimeFunctionIDMap.find(Val: Callee);
5071 if (It == OMPInfoCache.RuntimeFunctionIDMap.end()) {
5072 // Unknown caller or declarations are not analyzable, we give up.
5073 if (!Callee || !A.isFunctionIPOAmendable(F: *Callee)) {
5074
5075 // Unknown callees might contain parallel regions, except if they have
5076 // an appropriate assumption attached.
5077 if (!AssumptionAA ||
5078 !(AssumptionAA->hasAssumption(Assumption: "omp_no_openmp") ||
5079 AssumptionAA->hasAssumption(Assumption: "omp_no_parallelism")))
5080 ReachedUnknownParallelRegions.insert(Elem: &CB);
5081
5082 // If SPMDCompatibilityTracker is not fixed, we need to give up on the
5083 // idea we can run something unknown in SPMD-mode.
5084 if (!SPMDCompatibilityTracker.isAtFixpoint()) {
5085 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5086 SPMDCompatibilityTracker.insert(Elem: &CB);
5087 }
5088
5089 // We have updated the state for this unknown call properly, there
5090 // won't be any change so we indicate a fixpoint.
5091 indicateOptimisticFixpoint();
5092 }
5093 // If the callee is known and can be used in IPO, we will update the
5094 // state based on the callee state in updateImpl.
5095 return;
5096 }
5097 // More than one callee normally means an indirect call we cannot resolve.
5098 // A runtime function carrying !callback is the exception: the extra edge
5099 // is the callback, which we analyze rather than give up on.
5100 if (NumCallees > 1 && !Callee->hasMetadata(KindID: LLVMContext::MD_callback)) {
5101 indicatePessimisticFixpoint();
5102 return;
5103 }
5104
5105 RuntimeFunction RF = It->getSecond();
5106 switch (RF) {
5107 // All the functions we know are compatible with SPMD mode.
5108 case OMPRTL___kmpc_is_spmd_exec_mode:
5109 case OMPRTL___kmpc_distribute_static_fini:
5110 case OMPRTL___kmpc_for_static_fini:
5111 case OMPRTL___kmpc_global_thread_num:
5112 case OMPRTL___kmpc_get_hardware_num_threads_in_block:
5113 case OMPRTL___kmpc_get_hardware_num_blocks:
5114 case OMPRTL___kmpc_single:
5115 case OMPRTL___kmpc_end_single:
5116 case OMPRTL___kmpc_master:
5117 case OMPRTL___kmpc_end_master:
5118 case OMPRTL___kmpc_barrier:
5119 case OMPRTL___kmpc_nvptx_parallel_reduce_nowait_v2:
5120 case OMPRTL___kmpc_gpu_xteam_reduce_nowait:
5121 case OMPRTL___kmpc_error:
5122 case OMPRTL___kmpc_flush:
5123 case OMPRTL___kmpc_get_hardware_thread_id_in_block:
5124 case OMPRTL___kmpc_get_warp_size:
5125 case OMPRTL_omp_get_thread_num:
5126 case OMPRTL_omp_get_num_threads:
5127 case OMPRTL_omp_get_max_threads:
5128 case OMPRTL_omp_in_parallel:
5129 case OMPRTL_omp_get_dynamic:
5130 case OMPRTL_omp_get_cancellation:
5131 case OMPRTL_omp_get_nested:
5132 case OMPRTL_omp_get_schedule:
5133 case OMPRTL_omp_get_thread_limit:
5134 case OMPRTL_omp_get_supported_active_levels:
5135 case OMPRTL_omp_get_max_active_levels:
5136 case OMPRTL_omp_get_level:
5137 case OMPRTL_omp_get_ancestor_thread_num:
5138 case OMPRTL_omp_get_team_size:
5139 case OMPRTL_omp_get_active_level:
5140 case OMPRTL_omp_in_final:
5141 case OMPRTL_omp_get_proc_bind:
5142 case OMPRTL_omp_get_num_places:
5143 case OMPRTL_omp_get_num_procs:
5144 case OMPRTL_omp_get_place_proc_ids:
5145 case OMPRTL_omp_get_place_num:
5146 case OMPRTL_omp_get_partition_num_places:
5147 case OMPRTL_omp_get_partition_place_nums:
5148 case OMPRTL_omp_get_wtime:
5149 break;
5150 case OMPRTL___kmpc_distribute_static_init_4:
5151 case OMPRTL___kmpc_distribute_static_init_4u:
5152 case OMPRTL___kmpc_distribute_static_init_8:
5153 case OMPRTL___kmpc_distribute_static_init_8u:
5154 case OMPRTL___kmpc_for_static_init_4:
5155 case OMPRTL___kmpc_for_static_init_4u:
5156 case OMPRTL___kmpc_for_static_init_8:
5157 case OMPRTL___kmpc_for_static_init_8u: {
5158 // Check the schedule and allow static schedule in SPMD mode.
5159 unsigned ScheduleArgOpNo = 2;
5160 auto *ScheduleTypeCI =
5161 dyn_cast<ConstantInt>(Val: CB.getArgOperand(i: ScheduleArgOpNo));
5162 unsigned ScheduleTypeVal =
5163 ScheduleTypeCI ? ScheduleTypeCI->getZExtValue() : 0;
5164 switch (OMPScheduleType(ScheduleTypeVal)) {
5165 case OMPScheduleType::UnorderedStatic:
5166 case OMPScheduleType::UnorderedStaticChunked:
5167 case OMPScheduleType::OrderedDistribute:
5168 case OMPScheduleType::OrderedDistributeChunked:
5169 break;
5170 default:
5171 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5172 SPMDCompatibilityTracker.insert(Elem: &CB);
5173 break;
5174 };
5175 } break;
5176 case OMPRTL___kmpc_target_init:
5177 KernelInitCB = &CB;
5178 break;
5179 case OMPRTL___kmpc_target_deinit:
5180 KernelDeinitCB = &CB;
5181 break;
5182 case OMPRTL___kmpc_parallel_60:
5183 if (!handleParallel60(A, CB))
5184 indicatePessimisticFixpoint();
5185 return;
5186 case OMPRTL___kmpc_omp_task:
5187 // We do not look into tasks right now, just give up.
5188 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5189 SPMDCompatibilityTracker.insert(Elem: &CB);
5190 ReachedUnknownParallelRegions.insert(Elem: &CB);
5191 break;
5192 case OMPRTL___kmpc_alloc_shared:
5193 case OMPRTL___kmpc_free_shared:
5194 // Return without setting a fixpoint, to be resolved in updateImpl.
5195 return;
5196 // The twelve static-loop entry points split into the two groups below.
5197 // Both come out SPMD-incompatible, but for different reasons: the first
5198 // because the call is single-threaded by construction, the second only
5199 // because SPMD-ization cannot yet guard per iteration. They are kept
5200 // apart so the second can be relaxed on its own once it can.
5201 case OMPRTL___kmpc_distribute_static_loop_4:
5202 case OMPRTL___kmpc_distribute_static_loop_4u:
5203 case OMPRTL___kmpc_distribute_static_loop_8:
5204 case OMPRTL___kmpc_distribute_static_loop_8u:
5205 // A plain `distribute` spreads its iterations over the teams, not over
5206 // the threads of a team: the runtime runs it with TId 0 and a team size
5207 // of one, and asserts the kernel is at parallel level 0. One thread per
5208 // block calls it, which is what generic mode gives it. In SPMD mode
5209 // every thread would call it, each running the whole of its block's
5210 // share of the loop body, so the kernel cannot be SPMD-ized however
5211 // analyzable the body is.
5212 if (!OMPInformationCache::getAnalyzableCallback(CB))
5213 ReachedUnknownParallelRegions.insert(Elem: &CB);
5214 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5215 SPMDCompatibilityTracker.insert(Elem: &CB);
5216 break;
5217 case OMPRTL___kmpc_distribute_for_static_loop_4:
5218 case OMPRTL___kmpc_distribute_for_static_loop_4u:
5219 case OMPRTL___kmpc_distribute_for_static_loop_8:
5220 case OMPRTL___kmpc_distribute_for_static_loop_8u:
5221 case OMPRTL___kmpc_for_static_loop_4:
5222 case OMPRTL___kmpc_for_static_loop_4u:
5223 case OMPRTL___kmpc_for_static_loop_8:
5224 case OMPRTL___kmpc_for_static_loop_8u:
5225 // These index by the thread's own id, so unlike a plain distribute they
5226 // are meant to be called by every thread of the block, and a kernel
5227 // reaching one is not SPMD-incompatible for that reason alone. What
5228 // stops us is the transform rather than the analysis: SPMD-ization
5229 // guards whatever has to stay single-threaded with a block-wide
5230 // barrier, and a barrier placed inside a loop body only some threads
5231 // run is divergent. Until guarding can express "the thread that owns
5232 // this iteration", stay conservative here too.
5233 if (!OMPInformationCache::getAnalyzableCallback(CB))
5234 ReachedUnknownParallelRegions.insert(Elem: &CB);
5235 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5236 SPMDCompatibilityTracker.insert(Elem: &CB);
5237 break;
5238 default:
5239 // Unknown OpenMP runtime calls cannot be executed in SPMD-mode,
5240 // generally. However, they do not hide parallel regions.
5241 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5242 SPMDCompatibilityTracker.insert(Elem: &CB);
5243 break;
5244 }
5245 // All other OpenMP runtime calls will not reach parallel regions so they
5246 // can be safely ignored for now. Since it is a known OpenMP runtime call
5247 // we have now modeled all effects and there is no need for any update.
5248 indicateOptimisticFixpoint();
5249 };
5250
5251 const auto *AACE =
5252 A.getAAFor<AACallEdges>(QueryingAA: *this, IRP: getIRPosition(), DepClass: DepClassTy::OPTIONAL);
5253 if (!AACE || !AACE->getState().isValidState() || AACE->hasUnknownCallee()) {
5254 CheckCallee(getAssociatedFunction(), 1);
5255 return;
5256 }
5257 const auto &OptimisticEdges = AACE->getOptimisticEdges();
5258 for (auto *Callee : OptimisticEdges) {
5259 CheckCallee(Callee, OptimisticEdges.size());
5260 if (isAtFixpoint())
5261 break;
5262 }
5263 }
5264
5265 ChangeStatus updateImpl(Attributor &A) override {
5266 // TODO: Once we have call site specific value information we can provide
5267 // call site specific liveness information and then it makes
5268 // sense to specialize attributes for call sites arguments instead of
5269 // redirecting requests to the callee argument.
5270 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
5271 KernelInfoState StateBefore = getState();
5272
5273 auto CheckCallee = [&](Function *F, int NumCallees) {
5274 const auto &It = OMPInfoCache.RuntimeFunctionIDMap.find(Val: F);
5275
5276 // If F is not a runtime function, propagate the AAKernelInfo of the
5277 // callee.
5278 if (It == OMPInfoCache.RuntimeFunctionIDMap.end()) {
5279 const IRPosition &FnPos = IRPosition::function(F: *F);
5280 auto *FnAA =
5281 A.getAAFor<AAKernelInfo>(QueryingAA: *this, IRP: FnPos, DepClass: DepClassTy::REQUIRED);
5282 if (!FnAA)
5283 return indicatePessimisticFixpoint();
5284 if (getState() == FnAA->getState())
5285 return ChangeStatus::UNCHANGED;
5286 getState() = FnAA->getState();
5287 return ChangeStatus::CHANGED;
5288 }
5289 // See the matching check in initialize: a !callback runtime function has
5290 // a second call edge by construction, and it is one we can analyze.
5291 if (NumCallees > 1 && !F->hasMetadata(KindID: LLVMContext::MD_callback))
5292 return indicatePessimisticFixpoint();
5293
5294 CallBase &CB = cast<CallBase>(Val&: getAssociatedValue());
5295 if (It->getSecond() == OMPRTL___kmpc_parallel_60) {
5296 if (!handleParallel60(A, CB))
5297 return indicatePessimisticFixpoint();
5298 return StateBefore == getState() ? ChangeStatus::UNCHANGED
5299 : ChangeStatus::CHANGED;
5300 }
5301
5302 // F is a runtime function that allocates or frees memory, check
5303 // AAHeapToStack and AAHeapToShared.
5304 assert(
5305 (It->getSecond() == OMPRTL___kmpc_alloc_shared ||
5306 It->getSecond() == OMPRTL___kmpc_free_shared) &&
5307 "Expected a __kmpc_alloc_shared or __kmpc_free_shared runtime call");
5308
5309 auto *HeapToStackAA = A.getAAFor<AAHeapToStack>(
5310 QueryingAA: *this, IRP: IRPosition::function(F: *CB.getCaller()), DepClass: DepClassTy::OPTIONAL);
5311 auto *HeapToSharedAA = A.getAAFor<AAHeapToShared>(
5312 QueryingAA: *this, IRP: IRPosition::function(F: *CB.getCaller()), DepClass: DepClassTy::OPTIONAL);
5313
5314 RuntimeFunction RF = It->getSecond();
5315
5316 switch (RF) {
5317 // If neither HeapToStack nor HeapToShared assume the call is removed,
5318 // assume SPMD incompatibility.
5319 case OMPRTL___kmpc_alloc_shared:
5320 if ((!HeapToStackAA || !HeapToStackAA->isAssumedHeapToStack(CB)) &&
5321 (!HeapToSharedAA || !HeapToSharedAA->isAssumedHeapToShared(CB)))
5322 SPMDCompatibilityTracker.insert(Elem: &CB);
5323 break;
5324 case OMPRTL___kmpc_free_shared:
5325 if ((!HeapToStackAA ||
5326 !HeapToStackAA->isAssumedHeapToStackRemovedFree(CB)) &&
5327 (!HeapToSharedAA ||
5328 !HeapToSharedAA->isAssumedHeapToSharedRemovedFree(CB)))
5329 SPMDCompatibilityTracker.insert(Elem: &CB);
5330 break;
5331 default:
5332 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5333 SPMDCompatibilityTracker.insert(Elem: &CB);
5334 }
5335 return ChangeStatus::CHANGED;
5336 };
5337
5338 const auto *AACE =
5339 A.getAAFor<AACallEdges>(QueryingAA: *this, IRP: getIRPosition(), DepClass: DepClassTy::OPTIONAL);
5340 if (!AACE || !AACE->getState().isValidState() || AACE->hasUnknownCallee()) {
5341 if (Function *F = getAssociatedFunction())
5342 CheckCallee(F, /*NumCallees=*/1);
5343 } else {
5344 const auto &OptimisticEdges = AACE->getOptimisticEdges();
5345 for (auto *Callee : OptimisticEdges) {
5346 CheckCallee(Callee, OptimisticEdges.size());
5347 if (isAtFixpoint())
5348 break;
5349 }
5350 }
5351
5352 return StateBefore == getState() ? ChangeStatus::UNCHANGED
5353 : ChangeStatus::CHANGED;
5354 }
5355
5356 /// Deal with a __kmpc_parallel_60 call (\p CB). Returns true if the call was
5357 /// handled, if a problem occurred, false is returned.
5358 bool handleParallel60(Attributor &A, CallBase &CB) {
5359 const unsigned int NonWrapperFunctionArgNo = 5;
5360 const unsigned int WrapperFunctionArgNo = 6;
5361 auto ParallelRegionOpArgNo = SPMDCompatibilityTracker.isAssumed()
5362 ? NonWrapperFunctionArgNo
5363 : WrapperFunctionArgNo;
5364
5365 auto *ParallelRegion = dyn_cast<Function>(
5366 Val: CB.getArgOperand(i: ParallelRegionOpArgNo)->stripPointerCasts());
5367 if (!ParallelRegion)
5368 return false;
5369
5370 ReachedKnownParallelRegions.insert(Elem: &CB);
5371 /// Check nested parallelism
5372 auto *FnAA = A.getAAFor<AAKernelInfo>(
5373 QueryingAA: *this, IRP: IRPosition::function(F: *ParallelRegion), DepClass: DepClassTy::OPTIONAL);
5374 NestedParallelism |= !FnAA || !FnAA->getState().isValidState() ||
5375 !FnAA->ReachedKnownParallelRegions.empty() ||
5376 !FnAA->ReachedKnownParallelRegions.isValidState() ||
5377 !FnAA->ReachedUnknownParallelRegions.isValidState() ||
5378 !FnAA->ReachedUnknownParallelRegions.empty();
5379 return true;
5380 }
5381};
5382
5383struct AAFoldRuntimeCall
5384 : public StateWrapper<BooleanState, AbstractAttribute> {
5385 using Base = StateWrapper<BooleanState, AbstractAttribute>;
5386
5387 AAFoldRuntimeCall(const IRPosition &IRP, Attributor &A) : Base(IRP) {}
5388
5389 /// Statistics are tracked as part of manifest for now.
5390 void trackStatistics() const override {}
5391
5392 /// Create an abstract attribute biew for the position \p IRP.
5393 static AAFoldRuntimeCall &createForPosition(const IRPosition &IRP,
5394 Attributor &A);
5395
5396 /// See AbstractAttribute::getName()
5397 StringRef getName() const override { return "AAFoldRuntimeCall"; }
5398
5399 /// See AbstractAttribute::getIdAddr()
5400 const char *getIdAddr() const override { return &ID; }
5401
5402 /// This function should return true if the type of the \p AA is
5403 /// AAFoldRuntimeCall
5404 static bool classof(const AbstractAttribute *AA) {
5405 return (AA->getIdAddr() == &ID);
5406 }
5407
5408 static const char ID;
5409};
5410
5411struct AAFoldRuntimeCallCallSiteReturned : AAFoldRuntimeCall {
5412 AAFoldRuntimeCallCallSiteReturned(const IRPosition &IRP, Attributor &A)
5413 : AAFoldRuntimeCall(IRP, A) {}
5414
5415 /// See AbstractAttribute::getAsStr()
5416 const std::string getAsStr(Attributor *) const override {
5417 if (!isValidState())
5418 return "<invalid>";
5419
5420 std::string Str("simplified value: ");
5421
5422 if (!SimplifiedValue)
5423 return Str + std::string("none");
5424
5425 if (!*SimplifiedValue)
5426 return Str + std::string("nullptr");
5427
5428 if (ConstantInt *CI = dyn_cast<ConstantInt>(Val: *SimplifiedValue))
5429 return Str + std::to_string(val: CI->getSExtValue());
5430
5431 return Str + std::string("unknown");
5432 }
5433
5434 void initialize(Attributor &A) override {
5435 if (DisableOpenMPOptFolding)
5436 indicatePessimisticFixpoint();
5437
5438 Function *Callee = getAssociatedFunction();
5439
5440 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
5441 const auto &It = OMPInfoCache.RuntimeFunctionIDMap.find(Val: Callee);
5442 assert(It != OMPInfoCache.RuntimeFunctionIDMap.end() &&
5443 "Expected a known OpenMP runtime function");
5444
5445 RFKind = It->getSecond();
5446
5447 CallBase &CB = cast<CallBase>(Val&: getAssociatedValue());
5448 A.registerSimplificationCallback(
5449 IRP: IRPosition::callsite_returned(CB),
5450 CB: [&](const IRPosition &IRP, const AbstractAttribute *AA,
5451 bool &UsedAssumedInformation) -> std::optional<Value *> {
5452 assert((isValidState() || SimplifiedValue == nullptr) &&
5453 "Unexpected invalid state!");
5454
5455 if (!isAtFixpoint()) {
5456 UsedAssumedInformation = true;
5457 if (AA)
5458 A.recordDependence(FromAA: *this, ToAA: *AA, DepClass: DepClassTy::OPTIONAL);
5459 }
5460 return SimplifiedValue;
5461 });
5462 }
5463
5464 ChangeStatus updateImpl(Attributor &A) override {
5465 ChangeStatus Changed = ChangeStatus::UNCHANGED;
5466 switch (RFKind) {
5467 case OMPRTL___kmpc_is_spmd_exec_mode:
5468 Changed |= foldIsSPMDExecMode(A);
5469 break;
5470 case OMPRTL___kmpc_parallel_level:
5471 Changed |= foldParallelLevel(A);
5472 break;
5473 case OMPRTL___kmpc_get_hardware_num_threads_in_block:
5474 Changed = Changed | foldKernelFnAttribute(A, Attr: "omp_target_thread_limit");
5475 break;
5476 case OMPRTL___kmpc_get_hardware_num_blocks:
5477 Changed = Changed | foldKernelFnAttribute(A, Attr: "omp_target_num_teams");
5478 break;
5479 default:
5480 llvm_unreachable("Unhandled OpenMP runtime function!");
5481 }
5482
5483 return Changed;
5484 }
5485
5486 ChangeStatus manifest(Attributor &A) override {
5487 ChangeStatus Changed = ChangeStatus::UNCHANGED;
5488
5489 if (SimplifiedValue && *SimplifiedValue) {
5490 Instruction &I = *getCtxI();
5491 A.changeAfterManifest(IRP: IRPosition::inst(I), NV&: **SimplifiedValue);
5492 A.deleteAfterManifest(I);
5493
5494 CallBase *CB = dyn_cast<CallBase>(Val: &I);
5495 auto Remark = [&](OptimizationRemark OR) {
5496 if (auto *C = dyn_cast<ConstantInt>(Val: *SimplifiedValue))
5497 return OR << "Replacing OpenMP runtime call "
5498 << CB->getCalledFunction()->getName() << " with "
5499 << ore::NV("FoldedValue", C->getZExtValue()) << ".";
5500 return OR << "Replacing OpenMP runtime call "
5501 << CB->getCalledFunction()->getName() << ".";
5502 };
5503
5504 if (CB && EnableVerboseRemarks)
5505 A.emitRemark<OptimizationRemark>(I: CB, RemarkName: "OMP180", RemarkCB&: Remark);
5506
5507 LLVM_DEBUG(dbgs() << TAG << "Replacing runtime call: " << I << " with "
5508 << **SimplifiedValue << "\n");
5509
5510 Changed = ChangeStatus::CHANGED;
5511 }
5512
5513 return Changed;
5514 }
5515
5516 ChangeStatus indicatePessimisticFixpoint() override {
5517 SimplifiedValue = nullptr;
5518 return AAFoldRuntimeCall::indicatePessimisticFixpoint();
5519 }
5520
5521private:
5522 /// Fold __kmpc_is_spmd_exec_mode into a constant if possible.
5523 ChangeStatus foldIsSPMDExecMode(Attributor &A) {
5524 std::optional<Value *> SimplifiedValueBefore = SimplifiedValue;
5525
5526 unsigned AssumedSPMDCount = 0, KnownSPMDCount = 0;
5527 unsigned AssumedNonSPMDCount = 0, KnownNonSPMDCount = 0;
5528 auto *CallerKernelInfoAA = A.getAAFor<AAKernelInfo>(
5529 QueryingAA: *this, IRP: IRPosition::function(F: *getAnchorScope()), DepClass: DepClassTy::REQUIRED);
5530
5531 if (!CallerKernelInfoAA ||
5532 !CallerKernelInfoAA->ReachingKernelEntries.isValidState())
5533 return indicatePessimisticFixpoint();
5534
5535 for (Kernel K : CallerKernelInfoAA->ReachingKernelEntries) {
5536 auto *AA = A.getAAFor<AAKernelInfo>(QueryingAA: *this, IRP: IRPosition::function(F: *K),
5537 DepClass: DepClassTy::REQUIRED);
5538
5539 if (!AA || !AA->isValidState()) {
5540 SimplifiedValue = nullptr;
5541 return indicatePessimisticFixpoint();
5542 }
5543
5544 if (AA->SPMDCompatibilityTracker.isAssumed()) {
5545 if (AA->SPMDCompatibilityTracker.isAtFixpoint())
5546 ++KnownSPMDCount;
5547 else
5548 ++AssumedSPMDCount;
5549 } else {
5550 if (AA->SPMDCompatibilityTracker.isAtFixpoint())
5551 ++KnownNonSPMDCount;
5552 else
5553 ++AssumedNonSPMDCount;
5554 }
5555 }
5556
5557 if ((AssumedSPMDCount + KnownSPMDCount) &&
5558 (AssumedNonSPMDCount + KnownNonSPMDCount))
5559 return indicatePessimisticFixpoint();
5560
5561 auto &Ctx = getAnchorValue().getContext();
5562 if (KnownSPMDCount || AssumedSPMDCount) {
5563 assert(KnownNonSPMDCount == 0 && AssumedNonSPMDCount == 0 &&
5564 "Expected only SPMD kernels!");
5565 // All reaching kernels are in SPMD mode. Update all function calls to
5566 // __kmpc_is_spmd_exec_mode to 1.
5567 SimplifiedValue = ConstantInt::get(Ty: Type::getInt8Ty(C&: Ctx), V: true);
5568 } else if (KnownNonSPMDCount || AssumedNonSPMDCount) {
5569 assert(KnownSPMDCount == 0 && AssumedSPMDCount == 0 &&
5570 "Expected only non-SPMD kernels!");
5571 // All reaching kernels are in non-SPMD mode. Update all function
5572 // calls to __kmpc_is_spmd_exec_mode to 0.
5573 SimplifiedValue = ConstantInt::get(Ty: Type::getInt8Ty(C&: Ctx), V: false);
5574 } else {
5575 // We have empty reaching kernels, therefore we cannot tell if the
5576 // associated call site can be folded. At this moment, SimplifiedValue
5577 // must be none.
5578 assert(!SimplifiedValue && "SimplifiedValue should be none");
5579 }
5580
5581 return SimplifiedValue == SimplifiedValueBefore ? ChangeStatus::UNCHANGED
5582 : ChangeStatus::CHANGED;
5583 }
5584
5585 /// Fold __kmpc_parallel_level into a constant if possible.
5586 ChangeStatus foldParallelLevel(Attributor &A) {
5587 std::optional<Value *> SimplifiedValueBefore = SimplifiedValue;
5588
5589 auto *CallerKernelInfoAA = A.getAAFor<AAKernelInfo>(
5590 QueryingAA: *this, IRP: IRPosition::function(F: *getAnchorScope()), DepClass: DepClassTy::REQUIRED);
5591
5592 if (!CallerKernelInfoAA ||
5593 !CallerKernelInfoAA->ParallelLevels.isValidState())
5594 return indicatePessimisticFixpoint();
5595
5596 if (!CallerKernelInfoAA->ReachingKernelEntries.isValidState())
5597 return indicatePessimisticFixpoint();
5598
5599 if (CallerKernelInfoAA->ReachingKernelEntries.empty()) {
5600 assert(!SimplifiedValue &&
5601 "SimplifiedValue should keep none at this point");
5602 return ChangeStatus::UNCHANGED;
5603 }
5604
5605 unsigned AssumedSPMDCount = 0, KnownSPMDCount = 0;
5606 unsigned AssumedNonSPMDCount = 0, KnownNonSPMDCount = 0;
5607 for (Kernel K : CallerKernelInfoAA->ReachingKernelEntries) {
5608 auto *AA = A.getAAFor<AAKernelInfo>(QueryingAA: *this, IRP: IRPosition::function(F: *K),
5609 DepClass: DepClassTy::REQUIRED);
5610 if (!AA || !AA->SPMDCompatibilityTracker.isValidState())
5611 return indicatePessimisticFixpoint();
5612
5613 if (AA->SPMDCompatibilityTracker.isAssumed()) {
5614 if (AA->SPMDCompatibilityTracker.isAtFixpoint())
5615 ++KnownSPMDCount;
5616 else
5617 ++AssumedSPMDCount;
5618 } else {
5619 if (AA->SPMDCompatibilityTracker.isAtFixpoint())
5620 ++KnownNonSPMDCount;
5621 else
5622 ++AssumedNonSPMDCount;
5623 }
5624 }
5625
5626 if ((AssumedSPMDCount + KnownSPMDCount) &&
5627 (AssumedNonSPMDCount + KnownNonSPMDCount))
5628 return indicatePessimisticFixpoint();
5629
5630 auto &Ctx = getAnchorValue().getContext();
5631 // If the caller can only be reached by SPMD kernel entries, the parallel
5632 // level is 1. Similarly, if the caller can only be reached by non-SPMD
5633 // kernel entries, it is 0.
5634 if (AssumedSPMDCount || KnownSPMDCount) {
5635 assert(KnownNonSPMDCount == 0 && AssumedNonSPMDCount == 0 &&
5636 "Expected only SPMD kernels!");
5637 SimplifiedValue = ConstantInt::get(Ty: Type::getInt8Ty(C&: Ctx), V: 1);
5638 } else {
5639 assert(KnownSPMDCount == 0 && AssumedSPMDCount == 0 &&
5640 "Expected only non-SPMD kernels!");
5641 SimplifiedValue = ConstantInt::get(Ty: Type::getInt8Ty(C&: Ctx), V: 0);
5642 }
5643 return SimplifiedValue == SimplifiedValueBefore ? ChangeStatus::UNCHANGED
5644 : ChangeStatus::CHANGED;
5645 }
5646
5647 ChangeStatus foldKernelFnAttribute(Attributor &A, llvm::StringRef Attr) {
5648 // Specialize only if all the calls agree with the attribute constant value
5649 int32_t CurrentAttrValue = -1;
5650 std::optional<Value *> SimplifiedValueBefore = SimplifiedValue;
5651
5652 auto *CallerKernelInfoAA = A.getAAFor<AAKernelInfo>(
5653 QueryingAA: *this, IRP: IRPosition::function(F: *getAnchorScope()), DepClass: DepClassTy::REQUIRED);
5654
5655 if (!CallerKernelInfoAA ||
5656 !CallerKernelInfoAA->ReachingKernelEntries.isValidState())
5657 return indicatePessimisticFixpoint();
5658
5659 // Iterate over the kernels that reach this function
5660 for (Kernel K : CallerKernelInfoAA->ReachingKernelEntries) {
5661 int32_t NextAttrVal = K->getFnAttributeAsParsedInteger(Kind: Attr, Default: -1);
5662
5663 if (NextAttrVal == -1 ||
5664 (CurrentAttrValue != -1 && CurrentAttrValue != NextAttrVal))
5665 return indicatePessimisticFixpoint();
5666 CurrentAttrValue = NextAttrVal;
5667 }
5668
5669 if (CurrentAttrValue != -1) {
5670 auto &Ctx = getAnchorValue().getContext();
5671 SimplifiedValue =
5672 ConstantInt::get(Ty: Type::getInt32Ty(C&: Ctx), V: CurrentAttrValue);
5673 }
5674 return SimplifiedValue == SimplifiedValueBefore ? ChangeStatus::UNCHANGED
5675 : ChangeStatus::CHANGED;
5676 }
5677
5678 /// An optional value the associated value is assumed to fold to. That is, we
5679 /// assume the associated value (which is a call) can be replaced by this
5680 /// simplified value.
5681 std::optional<Value *> SimplifiedValue;
5682
5683 /// The runtime function kind of the callee of the associated call site.
5684 RuntimeFunction RFKind;
5685};
5686
5687} // namespace
5688
5689/// Register folding callsite
5690void OpenMPOpt::registerFoldRuntimeCall(RuntimeFunction RF) {
5691 auto &RFI = OMPInfoCache.RFIs[RF];
5692 RFI.foreachUse(SCC, CB: [&](Use &U, Function &F) {
5693 CallInst *CI = OpenMPOpt::getCallIfRegularCall(U, RFI: &RFI);
5694 if (!CI)
5695 return false;
5696 A.getOrCreateAAFor<AAFoldRuntimeCall>(
5697 IRP: IRPosition::callsite_returned(CB: *CI), /* QueryingAA */ nullptr,
5698 DepClass: DepClassTy::NONE, /* ForceUpdate */ false,
5699 /* UpdateAfterInit */ false);
5700 return false;
5701 });
5702}
5703
5704void OpenMPOpt::registerAAs(bool IsModulePass) {
5705 if (SCC.empty())
5706 return;
5707
5708 if (IsModulePass) {
5709 // Ensure we create the AAKernelInfo AAs first and without triggering an
5710 // update. This will make sure we register all value simplification
5711 // callbacks before any other AA has the chance to create an AAValueSimplify
5712 // or similar.
5713 auto CreateKernelInfoCB = [&](Use &, Function &Kernel) {
5714 A.getOrCreateAAFor<AAKernelInfo>(
5715 IRP: IRPosition::function(F: Kernel), /* QueryingAA */ nullptr,
5716 DepClass: DepClassTy::NONE, /* ForceUpdate */ false,
5717 /* UpdateAfterInit */ false);
5718 return false;
5719 };
5720 OMPInformationCache::RuntimeFunctionInfo &InitRFI =
5721 OMPInfoCache.RFIs[OMPRTL___kmpc_target_init];
5722 InitRFI.foreachUse(SCC, CB: CreateKernelInfoCB);
5723
5724 registerFoldRuntimeCall(RF: OMPRTL___kmpc_is_spmd_exec_mode);
5725 registerFoldRuntimeCall(RF: OMPRTL___kmpc_parallel_level);
5726 registerFoldRuntimeCall(RF: OMPRTL___kmpc_get_hardware_num_threads_in_block);
5727 registerFoldRuntimeCall(RF: OMPRTL___kmpc_get_hardware_num_blocks);
5728 }
5729
5730 // Create CallSite AA for all Getters.
5731 if (DeduceICVValues) {
5732 for (int Idx = 0; Idx < OMPInfoCache.ICVs.size() - 1; ++Idx) {
5733 auto ICVInfo = OMPInfoCache.ICVs[static_cast<InternalControlVar>(Idx)];
5734
5735 auto &GetterRFI = OMPInfoCache.RFIs[ICVInfo.Getter];
5736
5737 auto CreateAA = [&](Use &U, Function &Caller) {
5738 CallInst *CI = OpenMPOpt::getCallIfRegularCall(U, RFI: &GetterRFI);
5739 if (!CI)
5740 return false;
5741
5742 auto &CB = cast<CallBase>(Val&: *CI);
5743
5744 IRPosition CBPos = IRPosition::callsite_function(CB);
5745 A.getOrCreateAAFor<AAICVTracker>(IRP: CBPos);
5746 return false;
5747 };
5748
5749 GetterRFI.foreachUse(SCC, CB: CreateAA);
5750 }
5751 }
5752
5753 // Create an ExecutionDomain AA for every function and a HeapToStack AA for
5754 // every function if there is a device kernel.
5755 if (!isOpenMPDevice(M))
5756 return;
5757
5758 for (auto *F : SCC) {
5759 if (F->isDeclaration())
5760 continue;
5761
5762 // We look at internal functions only on-demand but if any use is not a
5763 // direct call or outside the current set of analyzed functions, we have
5764 // to do it eagerly.
5765 if (F->hasLocalLinkage()) {
5766 if (llvm::all_of(Range: F->uses(), P: [this](const Use &U) {
5767 const auto *CB = dyn_cast<CallBase>(Val: U.getUser());
5768 return CB && CB->isCallee(U: &U) &&
5769 A.isRunOn(Fn: const_cast<Function *>(CB->getCaller()));
5770 }))
5771 continue;
5772 }
5773 registerAAsForFunction(A, F: *F);
5774 }
5775}
5776
5777void OpenMPOpt::registerAAsForFunction(Attributor &A, const Function &F) {
5778 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
5779
5780 IRPosition FPos = IRPosition::function(F);
5781 A.getOrCreateAAFor<AAExecutionDomain>(IRP: FPos);
5782 if (F.hasFnAttribute(Kind: Attribute::Convergent))
5783 A.getOrCreateAAFor<AANonConvergent>(IRP: FPos);
5784
5785 bool FunctionUsesSharedAlloc = false;
5786 if (!DisableOpenMPOptDeglobalization) {
5787 const OMPInformationCache::RuntimeFunctionInfo::UseVector *SharedAllocUses =
5788 OMPInfoCache.RFIs[OMPRTL___kmpc_alloc_shared].getUseVector(
5789 F&: const_cast<Function &>(F));
5790 FunctionUsesSharedAlloc = SharedAllocUses && !SharedAllocUses->empty();
5791 }
5792 bool HasHeapToStackCandidate = false;
5793 const TargetLibraryInfo *TLI = nullptr;
5794
5795 for (auto &I : instructions(F)) {
5796 if (auto *LI = dyn_cast<LoadInst>(Val: &I)) {
5797 bool UsedAssumedInformation = false;
5798 A.getAssumedSimplified(V: IRPosition::value(V: *LI), /* AA */ nullptr,
5799 UsedAssumedInformation, S: AA::Interprocedural);
5800 A.getOrCreateAAFor<AAAddressSpace>(
5801 IRP: IRPosition::value(V: *LI->getPointerOperand()));
5802 continue;
5803 }
5804 if (auto *CI = dyn_cast<CallBase>(Val: &I)) {
5805 if (!DisableOpenMPOptDeglobalization && !HasHeapToStackCandidate) {
5806 if (!TLI)
5807 TLI = A.getInfoCache().getTargetLibraryInfoForFunction(F);
5808 HasHeapToStackCandidate =
5809 isRemovableAlloc(V: CI, TLI) || getFreedOperand(CB: CI, TLI);
5810 }
5811 if (CI->isIndirectCall())
5812 A.getOrCreateAAFor<AAIndirectCallInfo>(
5813 IRP: IRPosition::callsite_function(CB: *CI));
5814 }
5815 if (auto *SI = dyn_cast<StoreInst>(Val: &I)) {
5816 A.getOrCreateAAFor<AAIsDead>(IRP: IRPosition::value(V: *SI));
5817 A.getOrCreateAAFor<AAAddressSpace>(
5818 IRP: IRPosition::value(V: *SI->getPointerOperand()));
5819 continue;
5820 }
5821 if (auto *FI = dyn_cast<FenceInst>(Val: &I)) {
5822 A.getOrCreateAAFor<AAIsDead>(IRP: IRPosition::value(V: *FI));
5823 continue;
5824 }
5825 if (auto *II = dyn_cast<IntrinsicInst>(Val: &I)) {
5826 if (II->getIntrinsicID() == Intrinsic::assume) {
5827 A.getOrCreateAAFor<AAPotentialValues>(
5828 IRP: IRPosition::value(V: *II->getArgOperand(i: 0)));
5829 continue;
5830 }
5831 }
5832 }
5833
5834 if (FunctionUsesSharedAlloc)
5835 A.getOrCreateAAFor<AAHeapToShared>(IRP: FPos);
5836 if (HasHeapToStackCandidate)
5837 A.getOrCreateAAFor<AAHeapToStack>(IRP: FPos);
5838}
5839
5840const char AAICVTracker::ID = 0;
5841const char AAKernelInfo::ID = 0;
5842const char AAExecutionDomain::ID = 0;
5843const char AAHeapToShared::ID = 0;
5844const char AAFoldRuntimeCall::ID = 0;
5845
5846AAICVTracker &AAICVTracker::createForPosition(const IRPosition &IRP,
5847 Attributor &A) {
5848 AAICVTracker *AA = nullptr;
5849 switch (IRP.getPositionKind()) {
5850 case IRPosition::IRP_INVALID:
5851 case IRPosition::IRP_FLOAT:
5852 case IRPosition::IRP_ARGUMENT:
5853 case IRPosition::IRP_CALL_SITE_ARGUMENT:
5854 llvm_unreachable("ICVTracker can only be created for function position!");
5855 case IRPosition::IRP_RETURNED:
5856 AA = new (A.Allocator) AAICVTrackerFunctionReturned(IRP, A);
5857 break;
5858 case IRPosition::IRP_CALL_SITE_RETURNED:
5859 AA = new (A.Allocator) AAICVTrackerCallSiteReturned(IRP, A);
5860 break;
5861 case IRPosition::IRP_CALL_SITE:
5862 AA = new (A.Allocator) AAICVTrackerCallSite(IRP, A);
5863 break;
5864 case IRPosition::IRP_FUNCTION:
5865 AA = new (A.Allocator) AAICVTrackerFunction(IRP, A);
5866 break;
5867 }
5868
5869 return *AA;
5870}
5871
5872AAExecutionDomain &AAExecutionDomain::createForPosition(const IRPosition &IRP,
5873 Attributor &A) {
5874 AAExecutionDomainFunction *AA = nullptr;
5875 switch (IRP.getPositionKind()) {
5876 case IRPosition::IRP_INVALID:
5877 case IRPosition::IRP_FLOAT:
5878 case IRPosition::IRP_ARGUMENT:
5879 case IRPosition::IRP_CALL_SITE_ARGUMENT:
5880 case IRPosition::IRP_RETURNED:
5881 case IRPosition::IRP_CALL_SITE_RETURNED:
5882 case IRPosition::IRP_CALL_SITE:
5883 llvm_unreachable(
5884 "AAExecutionDomain can only be created for function position!");
5885 case IRPosition::IRP_FUNCTION:
5886 AA = new (A.Allocator) AAExecutionDomainFunction(IRP, A);
5887 break;
5888 }
5889
5890 return *AA;
5891}
5892
5893AAHeapToShared &AAHeapToShared::createForPosition(const IRPosition &IRP,
5894 Attributor &A) {
5895 AAHeapToSharedFunction *AA = nullptr;
5896 switch (IRP.getPositionKind()) {
5897 case IRPosition::IRP_INVALID:
5898 case IRPosition::IRP_FLOAT:
5899 case IRPosition::IRP_ARGUMENT:
5900 case IRPosition::IRP_CALL_SITE_ARGUMENT:
5901 case IRPosition::IRP_RETURNED:
5902 case IRPosition::IRP_CALL_SITE_RETURNED:
5903 case IRPosition::IRP_CALL_SITE:
5904 llvm_unreachable(
5905 "AAHeapToShared can only be created for function position!");
5906 case IRPosition::IRP_FUNCTION:
5907 AA = new (A.Allocator) AAHeapToSharedFunction(IRP, A);
5908 break;
5909 }
5910
5911 return *AA;
5912}
5913
5914AAKernelInfo &AAKernelInfo::createForPosition(const IRPosition &IRP,
5915 Attributor &A) {
5916 AAKernelInfo *AA = nullptr;
5917 switch (IRP.getPositionKind()) {
5918 case IRPosition::IRP_INVALID:
5919 case IRPosition::IRP_FLOAT:
5920 case IRPosition::IRP_ARGUMENT:
5921 case IRPosition::IRP_RETURNED:
5922 case IRPosition::IRP_CALL_SITE_RETURNED:
5923 case IRPosition::IRP_CALL_SITE_ARGUMENT:
5924 llvm_unreachable("KernelInfo can only be created for function position!");
5925 case IRPosition::IRP_CALL_SITE:
5926 AA = new (A.Allocator) AAKernelInfoCallSite(IRP, A);
5927 break;
5928 case IRPosition::IRP_FUNCTION:
5929 AA = new (A.Allocator) AAKernelInfoFunction(IRP, A);
5930 break;
5931 }
5932
5933 return *AA;
5934}
5935
5936AAFoldRuntimeCall &AAFoldRuntimeCall::createForPosition(const IRPosition &IRP,
5937 Attributor &A) {
5938 AAFoldRuntimeCall *AA = nullptr;
5939 switch (IRP.getPositionKind()) {
5940 case IRPosition::IRP_INVALID:
5941 case IRPosition::IRP_FLOAT:
5942 case IRPosition::IRP_ARGUMENT:
5943 case IRPosition::IRP_RETURNED:
5944 case IRPosition::IRP_FUNCTION:
5945 case IRPosition::IRP_CALL_SITE:
5946 case IRPosition::IRP_CALL_SITE_ARGUMENT:
5947 llvm_unreachable("KernelInfo can only be created for call site position!");
5948 case IRPosition::IRP_CALL_SITE_RETURNED:
5949 AA = new (A.Allocator) AAFoldRuntimeCallCallSiteReturned(IRP, A);
5950 break;
5951 }
5952
5953 return *AA;
5954}
5955
5956/// Bound the if-cascade AAIndirectCallInfo builds for an indirect call. Device
5957/// code reaches its callees through function-pointer tables and virtual
5958/// dispatch, so a call site can see every address-taken candidate in the
5959/// module; specializing all of them costs more in code size and compile time
5960/// than the direct calls are worth.
5961///
5962/// This is a threshold on the call site rather than a limit on how many callees
5963/// get specialized: the Attributor asks about each callee with the same total,
5964/// so a site above the threshold keeps its indirect call instead of getting
5965/// this many direct ones plus a fallback.
5966static bool shouldSpecializeIndirectCallee(Attributor &,
5967 const AbstractAttribute &,
5968 CallBase &, Function &,
5969 unsigned NumAssumedCallees) {
5970 return NumAssumedCallees <= MaxCalleesForSpecialization;
5971}
5972
5973PreservedAnalyses OpenMPOptPass::run(Module &M, ModuleAnalysisManager &AM) {
5974 if (!containsOpenMP(M))
5975 return PreservedAnalyses::all();
5976 if (DisableOpenMPOptimizations)
5977 return PreservedAnalyses::all();
5978
5979 FunctionAnalysisManager &FAM =
5980 AM.getResult<FunctionAnalysisManagerModuleProxy>(IR&: M).getManager();
5981 KernelSet Kernels = getDeviceKernels(M);
5982
5983 if (PrintModuleBeforeOptimizations)
5984 LLVM_DEBUG(dbgs() << TAG << "Module before OpenMPOpt Module Pass:\n" << M);
5985
5986 auto IsCalled = [&](Function &F) {
5987 if (Kernels.contains(key: &F))
5988 return true;
5989 return !F.use_empty();
5990 };
5991
5992 auto EmitRemark = [&](Function &F) {
5993 auto &ORE = FAM.getResult<OptimizationRemarkEmitterAnalysis>(IR&: F);
5994 ORE.emit(RemarkBuilder: [&]() {
5995 OptimizationRemarkAnalysis ORA(DEBUG_TYPE, "OMP140", &F);
5996 return ORA << "Could not internalize function. "
5997 << "Some optimizations may not be possible. [OMP140]";
5998 });
5999 };
6000
6001 bool Changed = false;
6002
6003 // Create internal copies of each function if this is a kernel Module. This
6004 // allows iterprocedural passes to see every call edge.
6005 DenseMap<Function *, Function *> InternalizedMap;
6006 if (isOpenMPDevice(M)) {
6007 SmallPtrSet<Function *, 16> InternalizeFns;
6008 for (Function &F : M)
6009 if (!F.isDeclaration() && !Kernels.contains(key: &F) && IsCalled(F) &&
6010 !DisableInternalization) {
6011 if (Attributor::isInternalizable(F)) {
6012 InternalizeFns.insert(Ptr: &F);
6013 } else if (!F.hasLocalLinkage() && !F.hasFnAttribute(Kind: Attribute::Cold)) {
6014 EmitRemark(F);
6015 }
6016 }
6017
6018 Changed |=
6019 Attributor::internalizeFunctions(FnSet&: InternalizeFns, FnMap&: InternalizedMap);
6020 }
6021
6022 // Look at every function in the Module unless it was internalized.
6023 SetVector<Function *> Functions;
6024 SmallVector<Function *, 16> SCC;
6025 for (Function &F : M)
6026 if (!F.isDeclaration() && !InternalizedMap.lookup(Val: &F)) {
6027 SCC.push_back(Elt: &F);
6028 Functions.insert(X: &F);
6029 }
6030
6031 if (SCC.empty())
6032 return Changed ? PreservedAnalyses::none() : PreservedAnalyses::all();
6033
6034 AnalysisGetter AG(FAM);
6035
6036 auto OREGetter = [&FAM](Function *F) -> OptimizationRemarkEmitter & {
6037 return FAM.getResult<OptimizationRemarkEmitterAnalysis>(IR&: *F);
6038 };
6039
6040 BumpPtrAllocator Allocator;
6041 CallGraphUpdater CGUpdater;
6042
6043 bool PostLink = LTOPhase == ThinOrFullLTOPhase::FullLTOPostLink ||
6044 LTOPhase == ThinOrFullLTOPhase::ThinLTOPostLink ||
6045 LTOPhase == ThinOrFullLTOPhase::ThinLTOPreLink;
6046 OMPInformationCache InfoCache(M, AG, Allocator, /*CGSCC*/ nullptr, PostLink);
6047
6048 unsigned MaxFixpointIterations =
6049 (isOpenMPDevice(M)) ? SetFixpointIterations : 32;
6050
6051 AttributorConfig AC(CGUpdater);
6052 AC.DefaultInitializeLiveInternals = false;
6053 AC.IsModulePass = true;
6054 AC.RewriteSignatures = false;
6055 AC.MaxFixpointIterations = MaxFixpointIterations;
6056 AC.OREGetter = OREGetter;
6057 AC.PassName = DEBUG_TYPE;
6058 AC.InitializationCallback = OpenMPOpt::registerAAsForFunction;
6059 AC.IndirectCalleeSpecializationCallback = shouldSpecializeIndirectCallee;
6060 AC.IPOAmendableCB = [](const Function &F) {
6061 return F.hasFnAttribute(Kind: "kernel");
6062 };
6063
6064 Attributor A(Functions, InfoCache, AC);
6065
6066 OpenMPOpt OMPOpt(SCC, CGUpdater, OREGetter, InfoCache, A);
6067 Changed |= OMPOpt.run(IsModulePass: true);
6068
6069 // Optionally inline device functions for potentially better performance.
6070 if (AlwaysInlineDeviceFunctions && isOpenMPDevice(M))
6071 for (Function &F : M)
6072 if (!F.isDeclaration() && !Kernels.contains(key: &F) &&
6073 !F.hasFnAttribute(Kind: Attribute::NoInline))
6074 F.addFnAttr(Kind: Attribute::AlwaysInline);
6075
6076 if (PrintModuleAfterOptimizations)
6077 LLVM_DEBUG(dbgs() << TAG << "Module after OpenMPOpt Module Pass:\n" << M);
6078
6079 if (Changed)
6080 return PreservedAnalyses::none();
6081
6082 return PreservedAnalyses::all();
6083}
6084
6085PreservedAnalyses OpenMPOptCGSCCPass::run(LazyCallGraph::SCC &C,
6086 CGSCCAnalysisManager &AM,
6087 LazyCallGraph &CG,
6088 CGSCCUpdateResult &UR) {
6089 if (!containsOpenMP(M&: *C.begin()->getFunction().getParent()))
6090 return PreservedAnalyses::all();
6091 if (DisableOpenMPOptimizations)
6092 return PreservedAnalyses::all();
6093
6094 SmallVector<Function *, 16> SCC;
6095 // If there are kernels in the module, we have to run on all SCC's.
6096 for (LazyCallGraph::Node &N : C) {
6097 Function *Fn = &N.getFunction();
6098 SCC.push_back(Elt: Fn);
6099 }
6100
6101 if (SCC.empty())
6102 return PreservedAnalyses::all();
6103
6104 Module &M = *C.begin()->getFunction().getParent();
6105
6106 if (PrintModuleBeforeOptimizations)
6107 LLVM_DEBUG(dbgs() << TAG << "Module before OpenMPOpt CGSCC Pass:\n" << M);
6108
6109 FunctionAnalysisManager &FAM =
6110 AM.getResult<FunctionAnalysisManagerCGSCCProxy>(IR&: C, ExtraArgs&: CG).getManager();
6111
6112 AnalysisGetter AG(FAM);
6113
6114 auto OREGetter = [&FAM](Function *F) -> OptimizationRemarkEmitter & {
6115 return FAM.getResult<OptimizationRemarkEmitterAnalysis>(IR&: *F);
6116 };
6117
6118 BumpPtrAllocator Allocator;
6119 CallGraphUpdater CGUpdater;
6120 CGUpdater.initialize(LCG&: CG, SCC&: C, AM, UR);
6121
6122 bool PostLink = LTOPhase == ThinOrFullLTOPhase::FullLTOPostLink ||
6123 LTOPhase == ThinOrFullLTOPhase::ThinLTOPostLink ||
6124 LTOPhase == ThinOrFullLTOPhase::ThinLTOPreLink;
6125 SetVector<Function *> Functions(llvm::from_range, SCC);
6126 OMPInformationCache InfoCache(*(Functions.back()->getParent()), AG, Allocator,
6127 /*CGSCC*/ &Functions, PostLink);
6128
6129 unsigned MaxFixpointIterations =
6130 (isOpenMPDevice(M)) ? SetFixpointIterations : 32;
6131
6132 AttributorConfig AC(CGUpdater);
6133 AC.DefaultInitializeLiveInternals = false;
6134 AC.IsModulePass = false;
6135 AC.RewriteSignatures = false;
6136 AC.MaxFixpointIterations = MaxFixpointIterations;
6137 AC.OREGetter = OREGetter;
6138 AC.PassName = DEBUG_TYPE;
6139 AC.InitializationCallback = OpenMPOpt::registerAAsForFunction;
6140 AC.IndirectCalleeSpecializationCallback = shouldSpecializeIndirectCallee;
6141
6142 Attributor A(Functions, InfoCache, AC);
6143
6144 OpenMPOpt OMPOpt(SCC, CGUpdater, OREGetter, InfoCache, A);
6145 bool Changed = OMPOpt.run(IsModulePass: false);
6146
6147 if (PrintModuleAfterOptimizations)
6148 LLVM_DEBUG(dbgs() << TAG << "Module after OpenMPOpt CGSCC Pass:\n" << M);
6149
6150 if (Changed)
6151 return PreservedAnalyses::none();
6152
6153 return PreservedAnalyses::all();
6154}
6155
6156bool llvm::omp::isOpenMPKernel(Function &Fn) {
6157 return Fn.hasFnAttribute(Kind: "kernel");
6158}
6159
6160KernelSet llvm::omp::getDeviceKernels(Module &M) {
6161 KernelSet Kernels;
6162
6163 for (Function &F : M)
6164 if (F.hasKernelCallingConv()) {
6165 // We are only interested in OpenMP target regions. Others, such as
6166 // kernels generated by CUDA but linked together, are not interesting to
6167 // this pass.
6168 if (isOpenMPKernel(Fn&: F)) {
6169 ++NumOpenMPTargetRegionKernels;
6170 Kernels.insert(X: &F);
6171 } else
6172 ++NumNonOpenMPTargetRegionKernels;
6173 }
6174
6175 return Kernels;
6176}
6177
6178bool llvm::omp::containsOpenMP(Module &M) {
6179 Metadata *MD = M.getModuleFlag(Key: "openmp");
6180 if (!MD)
6181 return false;
6182
6183 return true;
6184}
6185
6186bool llvm::omp::isOpenMPDevice(Module &M) {
6187 Metadata *MD = M.getModuleFlag(Key: "openmp-device");
6188 if (!MD)
6189 return false;
6190
6191 return true;
6192}
6193