LLVM 24.0.0git
OpenMPOpt.cpp
Go to the documentation of this file.
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
21
22#include "llvm/ADT/DenseSet.h"
25#include "llvm/ADT/SetVector.h"
28#include "llvm/ADT/Statistic.h"
30#include "llvm/ADT/StringRef.h"
39#include "llvm/IR/Assumptions.h"
40#include "llvm/IR/BasicBlock.h"
41#include "llvm/IR/Constants.h"
43#include "llvm/IR/Dominators.h"
44#include "llvm/IR/Function.h"
45#include "llvm/IR/GlobalValue.h"
47#include "llvm/IR/InstrTypes.h"
48#include "llvm/IR/Instruction.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"
59#include "llvm/Support/Debug.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
75 "openmp-opt-disable", cl::desc("Disable OpenMP specific optimizations."),
76 cl::Hidden, cl::init(false));
77
79 "openmp-opt-enable-merging",
80 cl::desc("Enable the OpenMP region merging optimization."), cl::Hidden,
81 cl::init(false));
82
83static cl::opt<bool>
84 DisableInternalization("openmp-opt-disable-internalization",
85 cl::desc("Disable function internalization."),
86 cl::Hidden, cl::init(false));
87
88static cl::opt<bool> DeduceICVValues("openmp-deduce-icv-values",
89 cl::init(false), cl::Hidden);
90static cl::opt<bool> PrintICVValues("openmp-print-icv-values", cl::init(false),
92static cl::opt<bool> PrintOpenMPKernels("openmp-print-gpu-kernels",
93 cl::init(false), cl::Hidden);
94
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(false));
100
102 "openmp-opt-disable-deglobalization",
103 cl::desc("Disable OpenMP optimizations involving deglobalization."),
104 cl::Hidden, cl::init(false));
105
107 "openmp-opt-disable-spmdization",
108 cl::desc("Disable OpenMP optimizations involving SPMD-ization."),
109 cl::Hidden, cl::init(false));
110
112 "openmp-opt-disable-folding",
113 cl::desc("Disable OpenMP optimizations involving folding."), cl::Hidden,
114 cl::init(false));
115
117 "openmp-opt-disable-state-machine-rewrite",
118 cl::desc("Disable OpenMP optimizations that replace the state machine."),
119 cl::Hidden, cl::init(false));
120
122 "openmp-opt-disable-barrier-elimination",
123 cl::desc("Disable OpenMP optimizations that eliminate barriers."),
124 cl::Hidden, cl::init(false));
125
127 "openmp-opt-print-module-after",
128 cl::desc("Print the current module after OpenMP optimizations."),
129 cl::Hidden, cl::init(false));
130
132 "openmp-opt-print-module-before",
133 cl::desc("Print the current module before OpenMP optimizations."),
134 cl::Hidden, cl::init(false));
135
137 "openmp-opt-inline-device",
138 cl::desc("Inline all applicable functions on the device."), cl::Hidden,
139 cl::init(false));
140
141static cl::opt<bool>
142 EnableVerboseRemarks("openmp-opt-verbose-remarks",
143 cl::desc("Enables more verbose remarks."), cl::Hidden,
144 cl::init(false));
145
147 SetFixpointIterations("openmp-opt-max-iterations", cl::Hidden,
148 cl::desc("Maximal number of attributor iterations."),
149 cl::init(256));
150
152 SharedMemoryLimit("openmp-opt-shared-limit", cl::Hidden,
153 cl::desc("Maximum amount of shared memory to use."),
154 cl::init(std::numeric_limits<unsigned>::max()));
155
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(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");
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)
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)
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
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)
267
268#undef KERNEL_ENVIRONMENT_CONFIGURATION_GETTER
269
272 constexpr int InitKernelEnvironmentArgNo = 0;
274 KernelInitCB->getArgOperand(InitKernelEnvironmentArgNo)
276}
277
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(OMPBuilder.M);
301 const Triple T(OMPBuilder.M.getTargetTriple());
302 switch (T.getArch()) {
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.
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(&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(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(RFI.Kind, *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>(U.getUser())) {
539 if (!CGSCC || CGSCC->empty() || CGSCC->contains(UserI->getFunction())) {
540 RFI.getOrCreateUseVector(UserI->getFunction()).push_back(&U);
541 ++NumUses;
542 }
543 } else {
544 RFI.getOrCreateUseVector(nullptr).push_back(&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(LLVMContext::MD_callback))
566 return;
567
568 LLVMContext &Ctx = F->getContext();
569 MDBuilder MDB(Ctx);
570 F->addMetadata(LLVMContext::MD_callback,
571 *MDNode::get(Ctx, {MDB.createCallbackEncoding(ArgNo, Indices,
572 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) {
581 if (!Callee)
582 return nullptr;
583 MDNode *CallbackMD = Callee->getMetadata(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>(CallbackMD->getOperand(0));
589 if (!Encoding || Encoding->getNumOperands() == 0)
590 return nullptr;
591 auto *ArgNoMD = dyn_cast<ConstantAsMetadata>(Encoding->getOperand(0));
592 if (!ArgNoMD)
593 return nullptr;
594 uint64_t ArgNo =
595 cast<ConstantInt>(ArgNoMD->getValue())->getLimitedValue(UINT64_MAX);
596 if (ArgNo >= CB.arg_size())
597 return nullptr;
598 auto *Callback =
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(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>(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(Attribute::NoInline) &&
700 F.getName().starts_with(Prefix) &&
701 !F.hasFnAttribute(Attribute::OptimizeNone))
702 F.removeFnAttr(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==(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^=(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(NumValues, nullptr);
959 LastAccesses.assign(NumValues, 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>(&I))
972 continue;
973
974 auto *S = cast<StoreInst>(&I);
975 int64_t Offset = -1;
976 auto *Dst =
977 GetPointerBaseWithConstantOffset(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(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) {
1013 CI->getArgOperand(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(*EntryCount, *EC) : *EC;
1020 }
1021 return EntryCount;
1022}
1023
1024static bool moduleHasSampleProfile(const Module &M) {
1025 std::unique_ptr<ProfileSummary> Summary(
1026 ProfileSummary::getFromMD(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();
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();
1078 printKernels();
1079
1080 Changed |= runAttributor(IsModulePass);
1081
1082 // Recollect uses, in case Attributor deleted any.
1083 OMPInfoCache.recollectUses();
1084
1085 Changed |= deleteParallelRegions();
1086
1088 Changed |= hideMemTransfersLatency();
1089 Changed |= deduplicateRuntimeCalls();
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(ICVInfo.InitValue->getValue(), 10, true)
1118 : "IMPLEMENTATION_DEFINED");
1119 };
1120
1121 emitRemark<OptimizationRemarkAnalysis>(F, "OpenMPICVTracker", 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(*F))
1130 continue;
1131
1132 auto Remark = [&](OptimizationRemarkAnalysis ORA) {
1133 return ORA << "OpenMP GPU kernel "
1134 << ore::NV("OpenMPGPUKernel", F->getName()) << "\n";
1135 };
1136
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>(U.getUser());
1146 if (CI && CI->isCallee(&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>(&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(CGStartBB, &*CodeGenIP, DT, LI);
1196 assert(StartBB != nullptr && "StartBB should not be null");
1197 CGStartBB->getTerminator()->setSuccessor(0, StartBB);
1198 assert(EndBB != nullptr && "EndBB should not be null");
1199 EndBB->getTerminator()->setSuccessor(0, 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(ParentBB, SeqEndI->getNextNode(), DT, LI);
1222 BasicBlock *SeqAfterBB =
1223 SplitBlock(SeqEndBB, &*SeqEndBB->getFirstInsertionPt(), DT, LI);
1224 BasicBlock *SeqStartBB =
1225 SplitBlock(ParentBB, SeqStartI, DT, LI, nullptr, "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(CGStartBB, &*CodeGenIP, DT, LI);
1236 assert(SeqStartBB != nullptr && "SeqStartBB should not be null");
1237 CGStartBB->getTerminator()->setSuccessor(0, SeqStartBB);
1238 assert(SeqEndBB != nullptr && "SeqEndBB should not be null");
1239 SeqEndBB->getTerminator()->setSuccessor(0, 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>(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(&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(&I, LoadI);
1280 }
1281 }
1282
1283 OpenMPIRBuilder::LocationDescription Loc(ParentBB->end(), DL);
1285 OMPInfoCache.OMPBuilder.createMaster(Loc, BodyGenCB, FiniCB));
1286 cantFail(OMPInfoCache.OMPBuilder.createBarrier({SeqAfterIP, DL},
1287 OMPD_parallel));
1288
1289 UncondBrInst::Create(SeqAfterBB, 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(MergableCIs)) {
1312 OR << ore::NV("OpenMPParallelMerge", CI->getDebugLoc());
1313 if (CI != MergableCIs.back())
1314 OR << ", ";
1315 }
1316 return OR << ".";
1317 };
1318
1319 emitRemark<OptimizationRemark>(MergableCIs.front(), "OMP150", 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(BB, MergableCIs.back()->getNextNode(), DT, LI);
1328 BasicBlock *AfterBB =
1329 SplitBlock(EndBB, &*EndBB->getFirstInsertionPt(), DT, LI);
1330 StartBB = SplitBlock(BB, MergableCIs.front(), DT, LI, nullptr,
1331 "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.
1358 cantFail(OMPInfoCache.OMPBuilder.createParallel(
1359 Loc, AllocaIP, /* DeallocBlocks */ {}, BodyGenCB, PrivCB, FiniCB,
1360 nullptr, nullptr, OMP_PROC_BIND_default,
1361 /* IsCancellable */ false));
1362 UncondBrInst::Create(AfterBB, AfterIP.getNodeParent());
1363
1364 // Perform the actual outlining.
1365 OMPInfoCache.OMPBuilder.finalize(OriginalFn);
1366
1367 Function *OutlinedFn = MergableCIs.front()->getCaller();
1368 std::optional<uint64_t> WrapperCount =
1369 getMergedWrapperEntryCount(MergableCIs, CallbackCalleeOperand);
1370 // Leave the wrapper unprofiled when no callback has an entry count.
1371 if (WrapperCount)
1372 OutlinedFn->setEntryCount(*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(*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(CallbackCalleeOperand);
1383 FunctionType *FT = OMPInfoCache.OMPBuilder.ParallelTask;
1384 Args.clear();
1385 Args.push_back(OutlinedFn->getArg(0));
1386 Args.push_back(OutlinedFn->getArg(1));
1387 for (unsigned U = CallbackFirstArgOperand, E = CI->arg_size(); U < E;
1388 ++U)
1389 Args.push_back(CI->getArgOperand(U));
1390
1391 CallInst *NewCI =
1392 CallInst::Create(FT, Callee, Args, "", 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(*NewCI, {*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(U))
1406 NewCI->addParamAttr(
1407 U - (CallbackFirstArgOperand - CallbackCalleeOperand), 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(OMPInfoCache.OMPBuilder.createBarrier(
1414 {NewCI->getNextNode()->getIterator(), NewCI->getDebugLoc()},
1415 OMPD_parallel));
1416 }
1417
1418 CI->eraseFromParent();
1419 }
1420
1421 assert(OutlinedFn != OriginalFn && "Outlining failed");
1422 CGUpdater.registerOutlinedFunction(*OriginalFn, *OutlinedFn);
1423 CGUpdater.reanalyzeFunction(*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);
1434 BB2PRMap[CI->getParent()].insert(CI);
1435
1436 return false;
1437 };
1438
1439 BB2PRMap.clear();
1440 RFI.foreachUse(SCC, 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>(&I))
1469 return true;
1470
1471 CallInst *CI = cast<CallInst>(&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>(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(&I)) {
1502 MergableCIs.push_back(cast<CallInst>(&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(&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(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(OMPRTL___kmpc_fork_call);
1547 OMPInfoCache.recollectUsesForFunction(OMPRTL___kmpc_barrier);
1548 OMPInfoCache.recollectUsesForFunction(OMPRTL___kmpc_master);
1549 OMPInfoCache.recollectUsesForFunction(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 CI->getArgOperand(CallbackCalleeOperand)->stripPointerCasts());
1572 if (!Fn)
1573 return false;
1574 if (!Fn->onlyReadsMemory())
1575 return false;
1576 if (!Fn->hasFnAttribute(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 };
1586
1587 CI->eraseFromParent();
1588 Changed = true;
1589 ++NumOpenMPParallelRegionsDeleted;
1590 return true;
1591 };
1592
1593 RFI.foreachUse(SCC, 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, 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(&Arg)) {
1635 GTIdArg = &Arg;
1636 break;
1637 }
1638 Changed |= deduplicateRuntimeCalls(
1639 *F, OMPInfoCache.RFIs[OMPRTL___kmpc_global_thread_num], 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("__llvm_rpc_client")) {
1652 if (GV->hasNUsesOrMore(1))
1653 return false;
1654
1655 GV->replaceAllUsesWith(PoisonValue::get(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);
1673 if (!RTCall)
1674 return false;
1675
1676 OffloadArray OffloadArrays[3];
1677 if (!getValuesInOffloadArrays(*RTCall, 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(*RTCall);
1685 if (WaitMovementPoint)
1686 WasSplit = splitTargetDataBeginRTC(*RTCall, *WaitMovementPoint);
1687
1688 Changed |= WasSplit;
1689 return WasSplit;
1690 };
1691 if (OMPInfoCache.runtimeFnsAvailable(
1692 {OMPRTL___tgt_target_data_begin_mapper_issue,
1693 OMPRTL___tgt_target_data_begin_mapper_wait}))
1694 RFI.foreachUse(SCC, 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)) {
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 };
1710 }
1711
1712 return false;
1713 };
1714
1715 RFI.foreachUse(SCC, 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,
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(OffloadArray::BasePtrsArgNum);
1734 // i8** %offload_ptrs.
1735 Value *PtrsArg = RuntimeCall.getArgOperand(OffloadArray::PtrsArgNum);
1736 // i8** %offload_sizes.
1737 Value *SizesArg = RuntimeCall.getArgOperand(OffloadArray::SizesArgNum);
1738
1739 // Get values stored in **offload_baseptrs.
1740 auto *V = getUnderlyingObject(BasePtrsArg);
1741 if (!isa<AllocaInst>(V))
1742 return false;
1743 auto *BasePtrsArray = cast<AllocaInst>(V);
1744 if (!OAs[0].initialize(*BasePtrsArray, RuntimeCall))
1745 return false;
1746
1747 // Get values stored in **offload_baseptrs.
1748 V = getUnderlyingObject(PtrsArg);
1749 if (!isa<AllocaInst>(V))
1750 return false;
1751 auto *PtrsArray = cast<AllocaInst>(V);
1752 if (!OAs[1].initialize(*PtrsArray, RuntimeCall))
1753 return false;
1754
1755 // Get values stored in **offload_sizes.
1756 V = getUnderlyingObject(SizesArg);
1757 // If it's a [constant] global array don't analyze it.
1758 if (isa<GlobalValue>(V))
1759 return isa<Constant>(V);
1760 if (!isa<AllocaInst>(V))
1761 return false;
1762
1763 auto *SizesArray = cast<AllocaInst>(V);
1764 if (!OAs[2].initialize(*SizesArray, 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(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(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(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 IRBuilder.AsyncInfo, /*ArraySize=*/nullptr, "handle");
1844 Handle =
1845 IRBuilder.Builder.CreateAddrSpaceCast(Handle, 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, 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(Arg.get());
1857 Args.push_back(Handle);
1858
1859 CallInst *IssueCallsite = CallInst::Create(IssueDecl, Args, /*NameStr=*/"",
1860 RuntimeCall.getIterator());
1861 OMPInfoCache.setCallingConvention(IssueDecl, 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, OMPRTL___tgt_target_data_begin_mapper_wait);
1868
1869 Value *WaitParams[2] = {
1870 IssueCallsite->getArgOperand(
1871 OffloadArray::DeviceIDArgNum), // device_id.
1872 Handle // handle to wait on.
1873 };
1874 CallInst *WaitCallsite = CallInst::Create(
1875 WaitDecl, WaitParams, /*NameStr=*/"", WaitMovementPoint.getIterator());
1876 OMPInfoCache.setCallingConvention(WaitDecl, 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>(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);
1907 if (!CI || &F != &Caller)
1908 return false;
1909 Ident = combinedIdentStruct(Ident, CI->getArgOperand(0),
1910 /* GlobalOnly */ true, SingleChoice);
1911 return false;
1912 };
1913 RFI.foreachUse(SCC, 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 {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(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(0)->getType() != OMPInfoCache.OMPBuilder.IdentPtr)
1955 return false;
1956 for (unsigned U = 1; U < NumArgs; ++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, &RFI)) {
1970 if (IP)
1971 IP = DT->findNearestCommonDominator(IP, 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>(ReplVal)->moveBefore(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>(ReplVal)) {
1990 if (!CI->arg_empty() &&
1991 CI->getArgOperand(0)->getType() == OMPInfoCache.OMPBuilder.IdentPtr) {
1992 Value *Ident = getCombinedIdentFromCallUsesIn(RFI, F,
1993 /* GlobalOnly */ true);
1994 CI->setArgOperand(0, Ident);
1995 }
1996 }
1997
1998 bool Changed = false;
1999 auto ReplaceAndDeleteCB = [&](Use &U, Function &Caller) {
2000 CallInst *CI = getCallIfRegularCall(U, &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())
2011 else
2013
2014 CI->replaceAllUsesWith(ReplVal);
2015 CI->eraseFromParent();
2016 ++NumOpenMPRuntimeCallsDeduplicated;
2017 Changed = true;
2018 return true;
2019 };
2020 RFI.foreachUse(SCC, 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(ArgNo);
2039 if (CI == &RefCI || GTIdArgs.count(ArgOp) ||
2040 getCallIfRegularCall(
2041 *ArgOp, &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>(U.getUser()))
2053 if (CI->isArgOperand(&U))
2054 if (Function *Callee = CI->getCalledFunction())
2055 if (CallArgOpIsGTId(*Callee, U.getOperandNo(), *CI))
2056 GTIdArgs.insert(Callee->getArg(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, [&](Use &U, Function &F) {
2064 if (CallInst *CI = getCallIfRegularCall(U, &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(*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("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("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(&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(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 };
2222
2223 return nullptr;
2224 }
2225 }
2226
2227 auto GetUniqueKernelForUse = [&](const Use &U) -> Kernel {
2228 if (auto *Cmp = dyn_cast<ICmpInst>(U.getUser())) {
2229 // Allow use in equality comparisons.
2230 if (Cmp->isEquality())
2231 return getUniqueKernelFor(*Cmp);
2232 return nullptr;
2233 }
2234 if (auto *CB = dyn_cast<CallBase>(U.getUser())) {
2235 // Allow direct calls.
2236 if (CB->isCallee(&U))
2237 return getUniqueKernelFor(*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(*U.getUser(), &KernelParallelRFI))
2243 return getUniqueKernelFor(*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, [&](const Use &U) {
2253 PotentialKernels.insert(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
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, [&](Use &U) {
2288 if (auto *CB = dyn_cast<CallBase>(U.getUser()))
2289 if (CB->isCallee(&U)) {
2290 ++NumDirectCalls;
2291 return;
2292 }
2293
2294 if (isa<ICmpInst>(U.getUser())) {
2295 ToBeReplacedStateMachineUses.push_back(&U);
2296 return;
2297 }
2298
2299 // Find wrapper functions that represent parallel kernels.
2300 CallInst *CI =
2301 OpenMPOpt::getCallIfRegularCall(*U.getUser(), &KernelParallelRFI);
2302 const unsigned int WrapperFunctionArgNo = 6;
2303 if (!KernelParallelUse && CI &&
2304 CI->getArgOperandNo(&U) == WrapperFunctionArgNo) {
2305 KernelParallelUse = true;
2306 ToBeReplacedStateMachineUses.push_back(&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 };
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);
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 };
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(M.getContext());
2350
2351 auto *ID = new GlobalVariable(
2352 M, Int8Ty, /* isConstant */ true, GlobalValue::PrivateLinkage,
2353 UndefValue::get(Int8Ty), F->getName() + ".ID");
2354
2355 for (Use *U : ToBeReplacedStateMachineUses)
2357 ID, 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>(U);
2381 if (!CI || CI->getCalledOperand() != KernelParallelRFI.Declaration ||
2382 CI->arg_size() <= WrapperFunctionArgNo)
2383 continue;
2384
2385 Value *Wrapper = CI->getArgOperand(WrapperFunctionArgNo);
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(*CI->getFunction());
2396 if (!K || !OMPInfoCache.SPMDizedKernels.contains(K))
2397 continue;
2398
2399 CI->setArgOperand(
2400 WrapperFunctionArgNo,
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(std::make_pair(CI, CI->getArgOperand(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(std::make_pair(&I, *ReplVal)).second)
2502 HasChanged = ChangeStatus::CHANGED;
2503
2504 return true;
2505 };
2506
2507 // Track all changes of an ICV.
2508 SetterRFI.foreachUse(TrackValues, F);
2509
2510 bool UsedAssumedInformation = false;
2511 A.checkForAllInstructions(CallCheck, *this, {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(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>(&I);
2531 if (!CB || CB->hasFnAttr("no_openmp") ||
2532 CB->hasFnAttr("no_openmp_routines") ||
2533 CB->hasFnAttr("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(&I))
2548 return ICVReplacementValuesMap[ICV].lookup(&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 *this, IRPosition::callsite_returned(*CB), DepClassTy::REQUIRED);
2559
2560 if (ICVTrackingAA->isAssumedTracked()) {
2561 std::optional<Value *> URV =
2562 ICVTrackingAA->getUniqueReplacementValue(ICV);
2563 if (!URV || (*URV && AA::isValidAtPosition(AA::ValueAndContext(**URV, I),
2564 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(I))
2584 return ValuesMap.lookup(I);
2585
2587 SmallPtrSet<const Instruction *, 16> Visited;
2588 Worklist.push_back(I);
2589
2590 std::optional<Value *> ReplVal;
2591
2592 while (!Worklist.empty()) {
2593 const Instruction *CurrInst = Worklist.pop_back_val();
2594 if (!Visited.insert(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(CurrInst)) {
2603 std::optional<Value *> NewReplVal = ValuesMap.lookup(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, *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(CurrBB))
2640 if (const Instruction *Terminator = Pred->getTerminator())
2641 Worklist.push_back(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 *this, IRPosition::function(*getAnchorScope()), 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, 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(CheckReturnInst, *this, {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(IRPosition::inst(*getCtxI()), **ReplVal);
2746 A.deleteAfterManifest(*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 *this, IRPosition::function(*getAnchorScope()), 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(AssociatedICV, 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 *this, IRPosition::returned(*getAssociatedFunction()),
2821 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(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(InitialThreadBlocks) + "/" +
2876 std::to_string(AlignedBlocks) + " of " +
2877 std::to_string(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
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(CB);
2917 A.deleteAfterManifest(*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(LastCB))
2928 continue;
2929 if (LastCB->getFunction() != getAnchorScope())
2930 continue;
2931 if (!hasFunctionEndAsUniqueSuccessor(LastCB->getParent()))
2932 continue;
2933 if (!DeletedBarriers.count(LastCB)) {
2934 ++NumBarriersEliminated;
2935 A.deleteAfterManifest(*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(LastED.AlignedBarriers.begin(),
2943 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(*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(*getAnchorScope()))
2959 HandleAlignedBarrier(nullptr);
2960
2961 return Changed;
2962 }
2963
2964 bool isNoOpFence(const FenceInst &FI) const override {
2965 return getState().isValidState() && !NonNoOpFences.count(&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(&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>(CurI);
3010 if (!CB)
3011 continue;
3012 if (CB != &I && AlignedBarriers.contains(const_cast<CallBase *>(CB)))
3013 return true;
3014 const auto &It = CEDMap.find({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(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>(CurI);
3029 if (!CB)
3030 continue;
3031 if (CB != &I && AlignedBarriers.contains(const_cast<CallBase *>(CB)))
3032 return true;
3033 const auto &It = CEDMap.find({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(nullptr).IsReachedFromAlignedBarrierOnly;
3050 if (!llvm::all_of(predecessors(BB), [&](const BasicBlock *PredBB) {
3051 return BEDMap.lookup(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(&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({&CB, PRE}), CEDMap.lookup({&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(0) != &SuccessorBB)
3086 return false;
3087
3088 auto *Cmp = dyn_cast<CmpInst>(Edge->getCondition());
3089 if (!Cmp || !Cmp->isTrueWhenEqual() || !Cmp->isEquality())
3090 return false;
3091
3092 ConstantInt *C = dyn_cast<ConstantInt>(Cmp->getOperand(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>(Cmp->getOperand(0));
3099 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
3100 auto &RFI = OMPInfoCache.RFIs[OMPRTL___kmpc_target_init];
3101 CB = CB ? OpenMPOpt::getCallIfRegularCall(*CB, &RFI) : nullptr;
3102 if (!CB)
3103 return false;
3104 ConstantStruct *KernelEnvC =
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>(Cmp->getOperand(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>(Cmp->getOperand(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, *EA);
3154
3155 for (auto *AB : PredED.AlignedBarriers)
3156 ED.addAlignedBarrier(A, *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(ED.IsExecutedByInitialThreadOnly,
3166 InitialEdgeOnly || (PredED.IsExecutedByInitialThreadOnly &&
3167 ED.IsExecutedByInitialThreadOnly));
3168
3169 Changed |= setAndRecord(ED.IsReachedFromAlignedBarrierOnly,
3170 ED.IsReachedFromAlignedBarrierOnly &&
3171 PredED.IsReachedFromAlignedBarrierOnly);
3172 Changed |= setAndRecord(ED.EncounteredNonLocalSideEffect,
3173 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) {
3186 auto PredForCallSite = [&](AbstractCallSite ACS) {
3187 const auto *EDAA = A.getAAFor<AAExecutionDomain>(
3188 *this, IRPosition::function(*ACS.getInstruction()->getFunction()),
3189 DepClassTy::OPTIONAL);
3190 if (!EDAA || !EDAA->getState().isValidState())
3191 return false;
3192 CallSiteEDs.emplace_back(
3193 EDAA->getExecutionDomain(*cast<CallBase>(ACS.getInstruction())));
3194 return true;
3195 };
3196
3197 ExecutionDomainTy ExitED;
3198 bool AllCallSitesKnown;
3199 if (A.checkForAllCallSites(PredForCallSite, *this,
3200 /* RequiresAllCallSites */ true,
3201 AllCallSitesKnown)) {
3202 for (const auto &[CSInED, CSOutED] : CallSiteEDs) {
3203 mergeInPredecessor(A, EntryBBED, 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(*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(FnED.IsReachedFromAlignedBarrierOnly,
3227 FnED.IsReachedFromAlignedBarrierOnly &
3228 EntryBBED.IsReachedFromAlignedBarrierOnly);
3229 Changed |= setAndRecord(FnED.IsReachingAlignedBarrierOnly,
3230 FnED.IsReachingAlignedBarrierOnly &
3231 ExitED.IsReachingAlignedBarrierOnly);
3232 Changed |= setAndRecord(FnED.IsExecutedByInitialThreadOnly,
3233 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(&CB);
3246 // First, update the barrier ED kept in the separate CEDMap.
3247 auto &CallInED = CEDMap[{&CB, PRE}];
3248 Changed |= mergeInPredecessor(A, CallInED, 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, CallOutED, ED);
3258 };
3259
3260 auto *LivenessAA =
3261 A.getAAFor<AAIsDead>(*this, getIRPosition(), DepClassTy::OPTIONAL);
3262
3263 Function *F = getAnchorScope();
3264 BasicBlock &EntryBB = F->getEntryBlock();
3265 bool IsKernel = omp::isOpenMPKernel(*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, ED);
3280 } else {
3281 // For live non-entry blocks we only propagate
3282 // information via live edges.
3283 if (LivenessAA && LivenessAA->isAssumedDead(&BB))
3284 continue;
3285
3286 for (auto *PredBB : predecessors(&BB)) {
3287 if (LivenessAA && LivenessAA->isEdgeDead(PredBB, &BB))
3288 continue;
3289 bool InitialEdgeOnly = isInitialThreadOnlyEdge(
3290 A, dyn_cast<CondBrInst>(PredBB->getTerminator()), BB);
3291 mergeInPredecessor(A, ED, 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, *this, LivenessAA, UsedAssumedInformation,
3300 /* CheckBBLivenessOnly */ false, 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>(&I)) {
3307 if (auto *AI = dyn_cast_or_null<AssumeInst>(II)) {
3308 ED.addAssumeInst(A, *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>(&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(FI);
3341 }
3342
3343 auto *CB = dyn_cast<CallBase>(&I);
3344 bool IsNoSync = AA::isNoSyncInst(A, I, *this);
3345 bool IsAlignedBarrier =
3346 !IsNoSync && CB &&
3347 AANoSync::isAlignedBarrier(*CB, 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>(&I)) {
3365 if (!ED.EncounteredNonLocalSideEffect &&
3367 ED.EncounteredNonLocalSideEffect = true;
3368 if (!IsNoSync) {
3369 ED.IsReachedFromAlignedBarrierOnly = false;
3370 SyncInstWorklist.push_back(&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, CallInED, 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.
3384 if (!IsNoSync && Callee && !Callee->isDeclaration()) {
3385 const auto *EDAA = A.getAAFor<AAExecutionDomain>(
3386 *this, IRPosition::function(*Callee), 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(CallInED.IsReachingAlignedBarrierOnly, false);
3401 SyncInstWorklist.push_back(&I);
3402 }
3403 if (CalleeED.IsReachedFromAlignedBarrierOnly)
3404 mergeInPredecessorBarriersAndAssumptions(A, ED, CalleeED);
3405 auto &CallOutED = CEDMap[{CB, POST}];
3406 Changed |= mergeInPredecessor(A, CallOutED, ED);
3407 continue;
3408 }
3409 }
3410 if (!IsNoSync) {
3411 ED.IsReachedFromAlignedBarrierOnly = false;
3412 Changed |= setAndRecord(CallInED.IsReachingAlignedBarrierOnly, false);
3413 SyncInstWorklist.push_back(&I);
3414 }
3415 AlignedBarrierLastInBlock &= ED.IsReachedFromAlignedBarrierOnly;
3416 ED.EncounteredNonLocalSideEffect |= !CB->doesNotAccessMemory();
3417 auto &CallOutED = CEDMap[{CB, POST}];
3418 Changed |= mergeInPredecessor(A, CallOutED, 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 *this, IRPosition::callsite_function(*CB), DepClassTy::OPTIONAL);
3429
3430 auto AccessPred = [&](const Instruction *I, const Value *Ptr,
3433 return !AA::isPotentiallyAffectedByBarrier(A, {Ptr}, *this, I);
3434 };
3435 if (MemAA && MemAA->getState().isValidState() &&
3436 MemAA->checkForAllAccessesToMemoryKind(
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>(&I))
3446 if (LI->hasMetadata(LLVMContext::MD_invariant_load))
3447 continue;
3448
3449 if (!ED.EncounteredNonLocalSideEffect &&
3451 ED.EncounteredNonLocalSideEffect = true;
3452 }
3453
3454 bool IsEndAndNotReachingAlignedBarriersOnly = false;
3455 if (!isa<UnreachableInst>(BB.getTerminator()) &&
3456 !BB.getTerminator()->getNumSuccessors()) {
3457
3458 Changed |= mergeInPredecessor(A, InterProceduralED, ED);
3459
3460 auto &FnED = BEDMap[nullptr];
3461 if (IsKernel && !IsExplicitlyAligned)
3462 FnED.IsReachingAlignedBarrierOnly = false;
3463 Changed |= mergeInPredecessor(A, FnED, ED);
3464
3465 if (!FnED.IsReachingAlignedBarrierOnly) {
3466 IsEndAndNotReachingAlignedBarriersOnly = true;
3467 SyncInstWorklist.push_back(BB.getTerminator());
3468 auto &BBED = BEDMap[&BB];
3469 Changed |= setAndRecord(BBED.IsReachingAlignedBarrierOnly, 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>(CurInst);
3502 if (!CB)
3503 continue;
3504 auto &CallOutED = CEDMap[{CB, POST}];
3505 Changed |= setAndRecord(CallOutED.IsReachingAlignedBarrierOnly, false);
3506 auto &CallInED = CEDMap[{CB, PRE}];
3507 HitAlignedBarrierOrKnownEnd =
3508 AlignedBarriers.count(CB) || !CallInED.IsReachingAlignedBarrierOnly;
3509 if (HitAlignedBarrierOrKnownEnd)
3510 break;
3511 Changed |= setAndRecord(CallInED.IsReachingAlignedBarrierOnly, false);
3512 }
3513 if (HitAlignedBarrierOrKnownEnd)
3514 continue;
3515 BasicBlock *SyncBB = SyncInst->getParent();
3516 for (auto *PredBB : predecessors(SyncBB)) {
3517 if (LivenessAA && LivenessAA->isEdgeDead(PredBB, SyncBB))
3518 continue;
3519 if (!Visited.insert(PredBB))
3520 continue;
3521 auto &PredED = BEDMap[PredBB];
3522 if (setAndRecord(PredED.IsReachingAlignedBarrierOnly, false)) {
3523 Changed = true;
3524 SyncInstWorklist.push_back(PredBB->getTerminator());
3525 }
3526 }
3527 if (SyncBB != &EntryBB)
3528 continue;
3529 Changed |=
3530 setAndRecord(InterProceduralED.IsReachingAlignedBarrierOnly, 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(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) {
3591 for (auto *U : CB->users()) {
3592 CallBase *C = dyn_cast<CallBase>(U);
3593 if (C && C->getCalledFunction() == FreeRFI.Declaration)
3594 FreeCalls.push_back(C);
3595 }
3596
3597 if (FreeCalls.size() != 1)
3598 continue;
3599
3600 PotentialRemovedFreeCalls.insert(FreeCalls.front());
3601 }
3602 }
3603
3604 void initialize(Attributor &A) override {
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
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);
3622 if (!Uses)
3623 return;
3624
3625 for (Use *U : *Uses)
3626 if (CallBase *CB = dyn_cast<CallBase>(U->getUser())) {
3627 MallocCalls.insert(CB);
3628 A.registerSimplificationCallback(IRPosition::callsite_returned(*CB),
3629 SCB);
3630 }
3631
3632 findPotentialRemovedFreeCalls(A);
3633 }
3634
3635 bool isAssumedHeapToShared(CallBase &CB) const override {
3636 return isValidState() && MallocCalls.count(&CB);
3637 }
3638
3639 bool isAssumedHeapToSharedRemovedFree(CallBase &CB) const override {
3640 return isValidState() && PotentialRemovedFreeCalls.count(&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>(IRPosition::function(*F), this,
3652 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))
3658 continue;
3659
3660 // Find the unique free call to remove it.
3662 for (auto *U : CB->users()) {
3663 CallBase *C = dyn_cast<CallBase>(U);
3664 if (C && C->getCalledFunction() == FreeCall.Declaration)
3665 FreeCalls.push_back(C);
3666 }
3667 if (FreeCalls.size() != 1)
3668 continue;
3669
3670 auto *AllocSize = cast<ConstantInt>(CB->getArgOperand(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(M->getContext());
3688 Type *Int8ArrTy = ArrayType::get(Int8Ty, AllocSize->getZExtValue());
3689 auto *SharedMem = new GlobalVariable(
3690 *M, Int8ArrTy, /* IsConstant */ false, GlobalValue::InternalLinkage,
3691 PoisonValue::get(Int8ArrTy), CB->getName() + "_shared", nullptr,
3693 static_cast<unsigned>(AddressSpace::Shared));
3694 auto *NewBuffer = ConstantExpr::getPointerCast(
3695 SharedMem, PointerType::getUnqual(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>(CB, "OMP111", 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(IRPosition::callsite_returned(*CB), *NewBuffer);
3711 A.deleteAfterManifest(*CB);
3712 A.deleteAfterManifest(*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>(U)) {
3737 if (CB->getCaller() != F)
3738 continue;
3739 if (!MallocCalls.count(CB))
3740 continue;
3741 if (!isa<ConstantInt>(CB->getArgOperand(0))) {
3742 MallocCalls.remove(CB);
3743 continue;
3744 }
3745 const auto *ED = A.getAAFor<AAExecutionDomain>(
3746 *this, IRPosition::function(*F), DepClassTy::REQUIRED);
3747 if (!ED || !ED->isExecutedByInitialThreadOnly(*CB))
3748 MallocCalls.remove(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(ReachedKnownParallelRegions.size())
3790 : "<invalid>") +
3791 ", #Unknown PRs: " +
3792 (ReachedUnknownParallelRegions.isValidState()
3793 ? std::to_string(ReachedUnknownParallelRegions.size())
3794 : "<invalid>") +
3795 ", #Reaching Kernels: " +
3796 (ReachingKernelEntries.isValidState()
3797 ? std::to_string(ReachingKernelEntries.size())
3798 : "<invalid>") +
3799 ", #ParLevels: " +
3800 (ParallelLevels.isValidState()
3801 ? std::to_string(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) {
3837 KernelEnvC, ConfigC, {KernelInfo::ConfigurationIdx});
3838 assert(NewKernelEnvC && "Failed to create new kernel environment");
3839 KernelEnvC = cast<ConstantStruct>(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)
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);
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 [&](Use &U, Function &) {
3891 StoreCallBase(U, InitRFI, KernelInitCB);
3892 return false;
3893 },
3894 Fn);
3895 DeinitRFI.foreachUse(
3896 [&](Use &U, Function &) {
3897 StoreCallBase(U, DeinitRFI, KernelDeinitCB);
3898 return false;
3899 },
3900 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(Fn);
3908 IsKernelEntry = true;
3909
3910 KernelEnvC =
3912 GlobalVariable *KernelEnvGV =
3914
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(*this, *AA, DepClassTy::OPTIONAL);
3924 }
3925 return KernelEnvC;
3926 };
3927
3928 A.registerGlobalVariableSimplificationCallback(
3929 *KernelEnvGV, KernelConfigurationSimplifyCB);
3930
3931 // We cannot change to SPMD mode if the runtime functions aren't availible.
3932 bool CanChangeToSPMD = OMPInfoCache.runtimeFnsAvailable(
3933 {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 ExecModeC->getIntegerType(),
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(Fn->getContext());
3953 auto [MinThreads, MaxThreads] =
3955 if (MinThreads)
3956 setMinThreadsOfKernelEnvironment(ConstantInt::get(Int32Ty, MinThreads));
3957 if (MaxThreads)
3958 setMaxThreadsOfKernelEnvironment(ConstantInt::get(Int32Ty, MaxThreads));
3959 auto [MinTeams, MaxTeams] =
3961 if (MinTeams)
3962 setMinTeamsOfKernelEnvironment(ConstantInt::get(Int32Ty, MinTeams));
3963 if (MaxTeams)
3964 setMaxTeamsOfKernelEnvironment(ConstantInt::get(Int32Ty, MaxTeams));
3965
3966 ConstantInt *MayUseNestedParallelismC =
3967 KernelInfo::getMayUseNestedParallelismFromKernelEnvironment(KernelEnvC);
3968 ConstantInt *AssumedMayUseNestedParallelismC = ConstantInt::get(
3969 MayUseNestedParallelismC->getIntegerType(), NestedParallelism);
3970 setMayUseNestedParallelismOfKernelEnvironment(
3971 AssumedMayUseNestedParallelismC);
3972
3974 ConstantInt *UseGenericStateMachineC =
3975 KernelInfo::getUseGenericStateMachineFromKernelEnvironment(
3976 KernelEnvC);
3977 ConstantInt *AssumedUseGenericStateMachineC =
3978 ConstantInt::get(UseGenericStateMachineC->getIntegerType(), false);
3979 setUseGenericStateMachineOfKernelEnvironment(
3980 AssumedUseGenericStateMachineC);
3981 }
3982
3983 // Register virtual uses of functions we might need to preserve.
3984 auto RegisterVirtualUse = [&](RuntimeFunction RFKind,
3986 if (!OMPInfoCache.RFIs[RFKind].Declaration)
3987 return;
3988 A.registerVirtualUseCallback(*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(*KI, *QueryingAA, 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 S.begin(), S.end(),
4064 [](const char C) {
4065 return !((C >= 'a' && C <= 'z') || (C >= 'A' && C <= 'Z') ||
4066 (C >= '0' && C <= '9') || C == '_');
4067 },
4068 '.');
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 =
4093 ConstantInt *OldUseGenericStateMachineVal =
4094 KernelInfo::getUseGenericStateMachineFromKernelEnvironment(
4095 ExistingKernelEnvC);
4096 if (!HasBuiltStateMachine)
4097 setUseGenericStateMachineOfKernelEnvironment(
4098 OldUseGenericStateMachineVal);
4099
4100 // At last, update the KernelEnvc
4101 GlobalVariable *KernelEnvGV =
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(ParentBB, RegionEndI->getNextNode(),
4147 DT, LI, MSU, "region.guarded.end");
4148 BasicBlock *RegionBarrierBB =
4149 SplitBlock(RegionEndBB, &*RegionEndBB->getFirstInsertionPt(), DT, LI,
4150 MSU, "region.barrier");
4151 BasicBlock *RegionExitBB =
4152 SplitBlock(RegionBarrierBB, &*RegionBarrierBB->getFirstInsertionPt(),
4153 DT, LI, MSU, "region.exit");
4154 BasicBlock *RegionStartBB =
4155 SplitBlock(ParentBB, RegionStartI, DT, LI, MSU, "region.guarded");
4156
4157 assert(ParentBB->getUniqueSuccessor() == RegionStartBB &&
4158 "Expected a different CFG");
4159
4160 BasicBlock *RegionCheckTidBB = SplitBlock(
4161 ParentBB, ParentBB->getTerminator(), DT, LI, MSU, "region.check.tid");
4162
4163 // Register basic blocks with the Attributor.
4164 A.registerManifestAddedBasicBlock(*RegionEndBB);
4165 A.registerManifestAddedBasicBlock(*RegionBarrierBB);
4166 A.registerManifestAddedBasicBlock(*RegionExitBB);
4167 A.registerManifestAddedBasicBlock(*RegionStartBB);
4168 A.registerManifestAddedBasicBlock(*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>(U.getUser());
4177 if (UsrI.getParent() != RegionStartBB)
4178 OutsideUses.push_back(&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,
4191 sanitizeForGlobalName(
4192 (I.getName() + ".guarded.output.alloc").str()),
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, *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(RegionCheckTidBB, 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(LocRegionCheckTid);
4228 FunctionCallee HardwareTidFn =
4229 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4230 M, OMPRTL___kmpc_get_hardware_thread_id_in_block);
4231 CallInst *Tid =
4232 OMPInfoCache.OMPBuilder.Builder.CreateCall(HardwareTidFn, {});
4233 Tid->setDebugLoc(DL);
4234 OMPInfoCache.setCallingConvention(HardwareTidFn, Tid);
4235 Value *TidCheck = OMPInfoCache.OMPBuilder.Builder.CreateIsNull(Tid);
4236 OMPInfoCache.OMPBuilder.Builder
4237 .CreateCondBr(TidCheck, RegionStartBB, 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, OMPRTL___kmpc_barrier_simple_spmd);
4245 OMPInfoCache.OMPBuilder.updateToLocation(
4246 {RegionBarrierBB->getFirstInsertionPt(), DL});
4247 CallInst *Barrier =
4248 OMPInfoCache.OMPBuilder.Builder.CreateCall(BarrierFn, {Ident, Tid});
4249 OMPInfoCache.setCallingConvention(BarrierFn, Barrier);
4250
4251 // Second barrier ensures workers have read broadcast values.
4252 if (HasBroadcastValues) {
4253 CallInst *Barrier =
4254 CallInst::Create(BarrierFn, {Ident, Tid}, "",
4255 RegionBarrierBB->getTerminator()->getIterator());
4256 Barrier->setDebugLoc(DL);
4257 OMPInfoCache.setCallingConvention(BarrierFn, 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(BB).second)
4266 continue;
4267
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(*I, &AllocSharedRFI))
4276 continue;
4277 if (!I->user_empty() || !SPMDCompatibilityTracker.contains(I)) {
4278 LastEffect = nullptr;
4279 continue;
4280 }
4281 if (LastEffect)
4282 Reorders.push_back({I, LastEffect});
4283 LastEffect = &*IP;
4284 }
4285 for (auto &Reorder : Reorders)
4286 Reorder.first->moveBefore(Reorder.second->getIterator());
4287 }
4288
4290
4291 for (Instruction *GuardedI : SPMDCompatibilityTracker) {
4292 BasicBlock *BB = GuardedI->getParent();
4293 auto *CalleeAA = A.lookupAAFor<AAKernelInfo>(
4294 IRPosition::function(*GuardedI->getFunction()), nullptr,
4295 DepClassTy::NONE);
4296 assert(CalleeAA != nullptr && "Expected Callee AAKernelInfo");
4297 auto &CalleeAAFunction = *cast<AAKernelInfoFunction>(CalleeAA);
4298 // Continue if instruction is already guarded.
4299 if (CalleeAAFunction.getGuardedInstructions().contains(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(&I)) {
4307 CalleeAAFunction.getGuardedInstructions().insert(&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 std::make_pair(GuardedRegionStart, 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 KernelInitCB->getNextNode(), "main.thread.user_code");
4348 BasicBlock *ReturnBB =
4349 BasicBlock::Create(Ctx, "exit.threads", Kernel, UserCodeBB);
4350
4351 // Register blocks with attributor:
4352 A.registerManifestAddedBasicBlock(*InitBB);
4353 A.registerManifestAddedBasicBlock(*UserCodeBB);
4354 A.registerManifestAddedBasicBlock(*ReturnBB);
4355
4356 // Debug location:
4357 const DebugLoc &DLoc = KernelInitCB->getDebugLoc();
4358 ReturnInst::Create(Ctx, 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, OMPRTL___kmpc_get_hardware_thread_id_in_block);
4367
4368 // Get thread ID in block.
4369 CallInst *ThreadIdInBlock =
4370 CallInst::Create(ThreadIdInBlockFn, "thread_id.in.block", InitBB);
4371 OMPInfoCache.setCallingConvention(ThreadIdInBlockFn, 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(ICmpInst::ICmp, CmpInst::ICMP_NE, ThreadIdInBlock,
4377 ConstantInt::get(ThreadIdInBlock->getType(), 0),
4378 "thread.is_main", InitBB);
4379 IsMainThread->setDebugLoc(DLoc);
4380 CondBrInst::Create(IsMainThread, ReturnBB, UserCodeBB, 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>(NonCompatibleI))
4393 if (OMPInfoCache.RTLFunctions.contains(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>(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>(NonCompatibleI, "OMP121",
4406 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>(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 =
4428 auto *ExecModeC =
4429 KernelInfo::getExecModeFromKernelEnvironment(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(ExecModeC->getIntegerType(),
4450 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(Kernel);
4457
4458 auto Remark = [&](OptimizationRemark OR) {
4459 return OR << "Transformed generic-mode kernel to SPMD-mode.";
4460 };
4461 A.emitRemark<OptimizationRemark>(KernelInitCB, "OMP120", 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
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({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 =
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 ExistingKernelEnvC);
4491 ConstantInt *ModeC =
4492 KernelInfo::getExecModeFromKernelEnvironment(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() ||
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(UseStateMachineC->getIntegerType(), 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>(KernelInitCB, "OMP130", 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>(KernelInitCB, "OMP131", 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>(KernelInitCB, "OMP132", 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>(UnknownParallelRegionCB,
4550 "OMP133", 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 KernelInitCB->getNextNode(), "thread.user_code.check");
4589 BasicBlock *IsWorkerCheckBB =
4590 BasicBlock::Create(Ctx, "is_worker_check", Kernel, UserCodeEntryBB);
4591 BasicBlock *StateMachineBeginBB = BasicBlock::Create(
4592 Ctx, "worker_state_machine.begin", Kernel, UserCodeEntryBB);
4593 BasicBlock *StateMachineFinishedBB = BasicBlock::Create(
4594 Ctx, "worker_state_machine.finished", Kernel, UserCodeEntryBB);
4595 BasicBlock *StateMachineIsActiveCheckBB = BasicBlock::Create(
4596 Ctx, "worker_state_machine.is_active.check", Kernel, UserCodeEntryBB);
4597 BasicBlock *StateMachineIfCascadeCurrentBB =
4598 BasicBlock::Create(Ctx, "worker_state_machine.parallel_region.check",
4599 Kernel, UserCodeEntryBB);
4600 BasicBlock *StateMachineEndParallelBB =
4601 BasicBlock::Create(Ctx, "worker_state_machine.parallel_region.end",
4602 Kernel, UserCodeEntryBB);
4603 BasicBlock *StateMachineDoneBarrierBB = BasicBlock::Create(
4604 Ctx, "worker_state_machine.done.barrier", Kernel, UserCodeEntryBB);
4605 A.registerManifestAddedBasicBlock(*InitBB);
4606 A.registerManifestAddedBasicBlock(*UserCodeEntryBB);
4607 A.registerManifestAddedBasicBlock(*IsWorkerCheckBB);
4608 A.registerManifestAddedBasicBlock(*StateMachineBeginBB);
4609 A.registerManifestAddedBasicBlock(*StateMachineFinishedBB);
4610 A.registerManifestAddedBasicBlock(*StateMachineIsActiveCheckBB);
4611 A.registerManifestAddedBasicBlock(*StateMachineIfCascadeCurrentBB);
4612 A.registerManifestAddedBasicBlock(*StateMachineEndParallelBB);
4613 A.registerManifestAddedBasicBlock(*StateMachineDoneBarrierBB);
4614
4615 const DebugLoc &DLoc = KernelInitCB->getDebugLoc();
4616 ReturnInst::Create(Ctx, StateMachineFinishedBB)->setDebugLoc(DLoc);
4617 InitBB->getTerminator()->eraseFromParent();
4618
4619 Instruction *IsWorker =
4620 ICmpInst::Create(ICmpInst::ICmp, llvm::CmpInst::ICMP_NE, KernelInitCB,
4621 ConstantInt::getAllOnesValue(KernelInitCB->getType()),
4622 "thread.is_worker", InitBB);
4623 IsWorker->setDebugLoc(DLoc);
4624 CondBrInst::Create(IsWorker, IsWorkerCheckBB, UserCodeEntryBB, 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, OMPRTL___kmpc_get_max_team_threads);
4634 Constant *IsSPMDArg = ConstantInt::get(OMPInfoCache.OMPBuilder.Int32, 0);
4635 CallInst *MaxTeamThreads = CallInst::Create(
4636 MaxTeamThreadsFn, {IsSPMDArg}, "max_team_threads", IsWorkerCheckBB);
4637 OMPInfoCache.setCallingConvention(MaxTeamThreadsFn, MaxTeamThreads);
4638 MaxTeamThreads->setDebugLoc(DLoc);
4639 Instruction *IsMainOrWorker = ICmpInst::Create(
4640 ICmpInst::ICmp, llvm::CmpInst::ICMP_SLT, KernelInitCB, MaxTeamThreads,
4641 "thread.is_main_or_worker", IsWorkerCheckBB);
4642 IsMainOrWorker->setDebugLoc(DLoc);
4643 CondBrInst::Create(IsMainOrWorker, StateMachineBeginBB,
4644 StateMachineFinishedBB, IsWorkerCheckBB);
4645
4646 // Create local storage for the work function pointer.
4647 const DataLayout &DL = M.getDataLayout();
4648 Type *VoidPtrTy = PointerType::getUnqual(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 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, OMPRTL___kmpc_barrier_simple_generic);
4663 CallInst *Barrier =
4664 CallInst::Create(BarrierFn, {Ident, GTid}, "", StateMachineBeginBB);
4665 OMPInfoCache.setCallingConvention(BarrierFn, Barrier);
4666 Barrier->setDebugLoc(DLoc);
4667
4668 if (WorkFnAI->getType()->getPointerAddressSpace() !=
4669 (unsigned int)AddressSpace::Generic) {
4670 WorkFnAI = new AddrSpaceCastInst(
4671 WorkFnAI, PointerType::get(Ctx, (unsigned int)AddressSpace::Generic),
4672 WorkFnAI->getName() + ".generic", StateMachineBeginBB);
4673 WorkFnAI->setDebugLoc(DLoc);
4674 }
4675
4676 FunctionCallee KernelParallelFn =
4677 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4678 M, OMPRTL___kmpc_kernel_parallel);
4679 CallInst *IsActiveWorker = CallInst::Create(
4680 KernelParallelFn, {WorkFnAI}, "worker.is_active", StateMachineBeginBB);
4681 OMPInfoCache.setCallingConvention(KernelParallelFn, 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 Type::getVoidTy(Ctx), {Type::getInt16Ty(Ctx), Type::getInt32Ty(Ctx)},
4689 false);
4690
4691 Instruction *IsDone =
4692 ICmpInst::Create(ICmpInst::ICmp, llvm::CmpInst::ICMP_EQ, WorkFn,
4693 Constant::getNullValue(VoidPtrTy), "worker.is_done",
4694 StateMachineBeginBB);
4695 IsDone->setDebugLoc(DLoc);
4696 CondBrInst::Create(IsDone, StateMachineFinishedBB,
4697 StateMachineIsActiveCheckBB, StateMachineBeginBB)
4698 ->setDebugLoc(DLoc);
4699
4700 CondBrInst::Create(IsActiveWorker, StateMachineIfCascadeCurrentBB,
4701 StateMachineDoneBarrierBB, StateMachineIsActiveCheckBB)
4702 ->setDebugLoc(DLoc);
4703
4704 Value *ZeroArg =
4705 Constant::getNullValue(ParallelRegionFnTy->getParamType(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 CB->getArgOperand(WrapperFunctionArgNo)->stripPointerCasts());
4716 BasicBlock *PRExecuteBB = BasicBlock::Create(
4717 Ctx, "worker_state_machine.parallel_region.execute", Kernel,
4718 StateMachineEndParallelBB);
4719 CallInst::Create(ParallelRegion, {ZeroArg, GTid}, "", PRExecuteBB)
4720 ->setDebugLoc(DLoc);
4721 UncondBrInst::Create(StateMachineEndParallelBB, PRExecuteBB)
4722 ->setDebugLoc(DLoc);
4723
4724 BasicBlock *PRNextBB =
4725 BasicBlock::Create(Ctx, "worker_state_machine.parallel_region.check",
4726 Kernel, StateMachineEndParallelBB);
4727 A.registerManifestAddedBasicBlock(*PRExecuteBB);
4728 A.registerManifestAddedBasicBlock(*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 ICmpInst::ICmp, llvm::CmpInst::ICMP_EQ, WorkFn, ParallelRegion,
4736 "worker.check_parallel_region", StateMachineIfCascadeCurrentBB);
4737 CmpI->setDebugLoc(DLoc);
4738 IsPR = CmpI;
4739 } else {
4740 IsPR = ConstantInt::getTrue(Ctx);
4741 }
4742
4743 CondBrInst::Create(IsPR, PRExecuteBB, PRNextBB,
4744 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(ParallelRegionFnTy, WorkFn, {ZeroArg, GTid}, "",
4756 StateMachineIfCascadeCurrentBB)
4757 ->setDebugLoc(DLoc);
4758 }
4759 UncondBrInst::Create(StateMachineEndParallelBB,
4760 StateMachineIfCascadeCurrentBB)
4761 ->setDebugLoc(DLoc);
4762
4763 FunctionCallee EndParallelFn =
4764 OMPInfoCache.OMPBuilder.getOrCreateRuntimeFunction(
4765 M, OMPRTL___kmpc_kernel_end_parallel);
4766 CallInst *EndParallel =
4767 CallInst::Create(EndParallelFn, {}, "", StateMachineEndParallelBB);
4768 OMPInfoCache.setCallingConvention(EndParallelFn, EndParallel);
4769 EndParallel->setDebugLoc(DLoc);
4770 UncondBrInst::Create(StateMachineDoneBarrierBB, StateMachineEndParallelBB)
4771 ->setDebugLoc(DLoc);
4772
4773 CallInst::Create(BarrierFn, {Ident, GTid}, "", StateMachineDoneBarrierBB)
4774 ->setDebugLoc(DLoc);
4775 UncondBrInst::Create(StateMachineBeginBB, 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 =
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 ExistingKernelEnvC));
4811
4812 if (!AA.SPMDCompatibilityTracker.isValidState())
4813 AA.setExecModeOfKernelEnvironment(
4814 KernelInfo::getExecModeFromKernelEnvironment(ExistingKernelEnvC));
4815
4816 ConstantInt *MayUseNestedParallelismC =
4817 KernelInfo::getMayUseNestedParallelismFromKernelEnvironment(
4818 AA.KernelEnvC);
4819 ConstantInt *NewMayUseNestedParallelismC = ConstantInt::get(
4820 MayUseNestedParallelismC->getIntegerType(), 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>(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>(&I)) {
4835 const auto *UnderlyingObjsAA = A.getAAFor<AAUnderlyingObjects>(
4836 *this, IRPosition::value(*SI->getPointerOperand()),
4837 DepClassTy::OPTIONAL);
4838 auto *HS = A.getAAFor<AAHeapToStack>(
4839 *this, IRPosition::function(*I.getFunction()),
4840 DepClassTy::OPTIONAL);
4841 if (UnderlyingObjsAA &&
4842 UnderlyingObjsAA->forallUnderlyingObjects([&](Value &Obj) {
4843 if (AA::isAssumedThreadLocalObject(A, Obj, *this))
4844 return true;
4845 // Check for AAHeapToStack moved objects which must not be
4846 // guarded.
4847 auto *CB = dyn_cast<CallBase>(&Obj);
4848 return CB && HS && HS->isAssumedHeapToStack(*CB);
4849 }))
4850 return true;
4851 }
4852
4853 // Insert instruction that needs guarding.
4854 SPMDCompatibilityTracker.insert(&I);
4855 return true;
4856 };
4857
4858 bool UsedAssumedInformationInCheckRWInst = false;
4859 if (!SPMDCompatibilityTracker.isAtFixpoint())
4860 if (!A.checkForAllReadWriteInstructions(
4861 CheckRWInst, *this, 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 *this, IRPosition::function(*Kernel), 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>(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 *this, IRPosition::function(*Callback), 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 *this, IRPosition::callsite_function(CB), 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 CheckCallInst, *this, 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 IRPosition::function(*Caller), this, 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(PredCallSite, *this,
4986 true /* RequireAllCallSites */,
4987 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>(IRPosition::function(*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(PredCallSite, *this,
5029 true /* RequireAllCallSites */,
5030 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>(getAssociatedValue());
5047 auto *AssumptionAA = A.getAAFor<AAAssumptionInfo>(
5048 *this, IRPosition::callsite_function(CB), DepClassTy::OPTIONAL);
5049
5050 // Check for SPMD-mode assumptions.
5051 if (AssumptionAA && AssumptionAA->hasAssumption("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>(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(Callee);
5071 if (It == OMPInfoCache.RuntimeFunctionIDMap.end()) {
5072 // Unknown caller or declarations are not analyzable, we give up.
5073 if (!Callee || !A.isFunctionIPOAmendable(*Callee)) {
5074
5075 // Unknown callees might contain parallel regions, except if they have
5076 // an appropriate assumption attached.
5077 if (!AssumptionAA ||
5078 !(AssumptionAA->hasAssumption("omp_no_openmp") ||
5079 AssumptionAA->hasAssumption("omp_no_parallelism")))
5080 ReachedUnknownParallelRegions.insert(&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(&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(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>(CB.getArgOperand(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(&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(&CB);
5190 ReachedUnknownParallelRegions.insert(&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(&CB);
5214 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5215 SPMDCompatibilityTracker.insert(&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(&CB);
5235 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5236 SPMDCompatibilityTracker.insert(&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(&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>(*this, getIRPosition(), 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(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);
5280 auto *FnAA =
5281 A.getAAFor<AAKernelInfo>(*this, FnPos, 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(LLVMContext::MD_callback))
5292 return indicatePessimisticFixpoint();
5293
5294 CallBase &CB = cast<CallBase>(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 *this, IRPosition::function(*CB.getCaller()), DepClassTy::OPTIONAL);
5311 auto *HeapToSharedAA = A.getAAFor<AAHeapToShared>(
5312 *this, IRPosition::function(*CB.getCaller()), 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(&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(&CB);
5330 break;
5331 default:
5332 SPMDCompatibilityTracker.indicatePessimisticFixpoint();
5333 SPMDCompatibilityTracker.insert(&CB);
5334 }
5335 return ChangeStatus::CHANGED;
5336 };
5337
5338 const auto *AACE =
5339 A.getAAFor<AACallEdges>(*this, getIRPosition(), 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 CB.getArgOperand(ParallelRegionOpArgNo)->stripPointerCasts());
5367 if (!ParallelRegion)
5368 return false;
5369
5370 ReachedKnownParallelRegions.insert(&CB);
5371 /// Check nested parallelism
5372 auto *FnAA = A.getAAFor<AAKernelInfo>(
5373 *this, IRPosition::function(*ParallelRegion), 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>(*SimplifiedValue))
5429 return Str + std::to_string(CI->getSExtValue());
5430
5431 return Str + std::string("unknown");
5432 }
5433
5434 void initialize(Attributor &A) override {
5436 indicatePessimisticFixpoint();
5437
5438 Function *Callee = getAssociatedFunction();
5439
5440 auto &OMPInfoCache = static_cast<OMPInformationCache &>(A.getInfoCache());
5441 const auto &It = OMPInfoCache.RuntimeFunctionIDMap.find(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>(getAssociatedValue());
5448 A.registerSimplificationCallback(
5450 [&](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(*this, *AA, 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, "omp_target_thread_limit");
5475 break;
5476 case OMPRTL___kmpc_get_hardware_num_blocks:
5477 Changed = Changed | foldKernelFnAttribute(A, "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(IRPosition::inst(I), **SimplifiedValue);
5492 A.deleteAfterManifest(I);
5493
5494 CallBase *CB = dyn_cast<CallBase>(&I);
5495 auto Remark = [&](OptimizationRemark OR) {
5496 if (auto *C = dyn_cast<ConstantInt>(*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>(CB, "OMP180", 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 *this, IRPosition::function(*getAnchorScope()), 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>(*this, IRPosition::function(*K),
5537 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(Type::getInt8Ty(Ctx), 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(Type::getInt8Ty(Ctx), 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 *this, IRPosition::function(*getAnchorScope()), 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>(*this, IRPosition::function(*K),
5609 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(Type::getInt8Ty(Ctx), 1);
5638 } else {
5639 assert(KnownSPMDCount == 0 && AssumedSPMDCount == 0 &&
5640 "Expected only non-SPMD kernels!");
5641 SimplifiedValue = ConstantInt::get(Type::getInt8Ty(Ctx), 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 *this, IRPosition::function(*getAnchorScope()), 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(Attr, -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(Type::getInt32Ty(Ctx), 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, [&](Use &U, Function &F) {
5693 CallInst *CI = OpenMPOpt::getCallIfRegularCall(U, &RFI);
5694 if (!CI)
5695 return false;
5696 A.getOrCreateAAFor<AAFoldRuntimeCall>(
5697 IRPosition::callsite_returned(*CI), /* QueryingAA */ nullptr,
5698 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 IRPosition::function(Kernel), /* QueryingAA */ nullptr,
5716 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, CreateKernelInfoCB);
5723
5724 registerFoldRuntimeCall(OMPRTL___kmpc_is_spmd_exec_mode);
5725 registerFoldRuntimeCall(OMPRTL___kmpc_parallel_level);
5726 registerFoldRuntimeCall(OMPRTL___kmpc_get_hardware_num_threads_in_block);
5727 registerFoldRuntimeCall(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, &GetterRFI);
5739 if (!CI)
5740 return false;
5741
5742 auto &CB = cast<CallBase>(*CI);
5743
5744 IRPosition CBPos = IRPosition::callsite_function(CB);
5745 A.getOrCreateAAFor<AAICVTracker>(CBPos);
5746 return false;
5747 };
5748
5749 GetterRFI.foreachUse(SCC, 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(F->uses(), [this](const Use &U) {
5767 const auto *CB = dyn_cast<CallBase>(U.getUser());
5768 return CB && CB->isCallee(&U) &&
5769 A.isRunOn(const_cast<Function *>(CB->getCaller()));
5770 }))
5771 continue;
5772 }
5773 registerAAsForFunction(A, *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>(FPos);
5782 if (F.hasFnAttribute(Attribute::Convergent))
5783 A.getOrCreateAAFor<AANonConvergent>(FPos);
5784
5785 bool FunctionUsesSharedAlloc = false;
5787 const OMPInformationCache::RuntimeFunctionInfo::UseVector *SharedAllocUses =
5788 OMPInfoCache.RFIs[OMPRTL___kmpc_alloc_shared].getUseVector(
5789 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>(&I)) {
5797 bool UsedAssumedInformation = false;
5798 A.getAssumedSimplified(IRPosition::value(*LI), /* AA */ nullptr,
5799 UsedAssumedInformation, AA::Interprocedural);
5800 A.getOrCreateAAFor<AAAddressSpace>(
5801 IRPosition::value(*LI->getPointerOperand()));
5802 continue;
5803 }
5804 if (auto *CI = dyn_cast<CallBase>(&I)) {
5805 if (!DisableOpenMPOptDeglobalization && !HasHeapToStackCandidate) {
5806 if (!TLI)
5807 TLI = A.getInfoCache().getTargetLibraryInfoForFunction(F);
5808 HasHeapToStackCandidate =
5809 isRemovableAlloc(CI, TLI) || getFreedOperand(CI, TLI);
5810 }
5811 if (CI->isIndirectCall())
5812 A.getOrCreateAAFor<AAIndirectCallInfo>(
5814 }
5815 if (auto *SI = dyn_cast<StoreInst>(&I)) {
5816 A.getOrCreateAAFor<AAIsDead>(IRPosition::value(*SI));
5817 A.getOrCreateAAFor<AAAddressSpace>(
5818 IRPosition::value(*SI->getPointerOperand()));
5819 continue;
5820 }
5821 if (auto *FI = dyn_cast<FenceInst>(&I)) {
5822 A.getOrCreateAAFor<AAIsDead>(IRPosition::value(*FI));
5823 continue;
5824 }
5825 if (auto *II = dyn_cast<IntrinsicInst>(&I)) {
5826 if (II->getIntrinsicID() == Intrinsic::assume) {
5827 A.getOrCreateAAFor<AAPotentialValues>(
5828 IRPosition::value(*II->getArgOperand(0)));
5829 continue;
5830 }
5831 }
5832 }
5833
5834 if (FunctionUsesSharedAlloc)
5835 A.getOrCreateAAFor<AAHeapToShared>(FPos);
5836 if (HasHeapToStackCandidate)
5837 A.getOrCreateAAFor<AAHeapToStack>(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()) {
5854 llvm_unreachable("ICVTracker can only be created for function position!");
5856 AA = new (A.Allocator) AAICVTrackerFunctionReturned(IRP, A);
5857 break;
5859 AA = new (A.Allocator) AAICVTrackerCallSiteReturned(IRP, A);
5860 break;
5862 AA = new (A.Allocator) AAICVTrackerCallSite(IRP, A);
5863 break;
5865 AA = new (A.Allocator) AAICVTrackerFunction(IRP, A);
5866 break;
5867 }
5868
5869 return *AA;
5870}
5871
5873 Attributor &A) {
5874 AAExecutionDomainFunction *AA = nullptr;
5875 switch (IRP.getPositionKind()) {
5884 "AAExecutionDomain can only be created for function position!");
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()) {
5905 "AAHeapToShared can only be created for function position!");
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()) {
5924 llvm_unreachable("KernelInfo can only be created for function position!");
5926 AA = new (A.Allocator) AAKernelInfoCallSite(IRP, A);
5927 break;
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()) {
5947 llvm_unreachable("KernelInfo can only be created for call site position!");
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.
5967 const AbstractAttribute &,
5968 CallBase &, Function &,
5969 unsigned NumAssumedCallees) {
5970 return NumAssumedCallees <= MaxCalleesForSpecialization;
5971}
5972
5974 if (!containsOpenMP(M))
5975 return PreservedAnalyses::all();
5977 return PreservedAnalyses::all();
5978
5981 KernelSet Kernels = getDeviceKernels(M);
5982
5984 LLVM_DEBUG(dbgs() << TAG << "Module before OpenMPOpt Module Pass:\n" << M);
5985
5986 auto IsCalled = [&](Function &F) {
5987 if (Kernels.contains(&F))
5988 return true;
5989 return !F.use_empty();
5990 };
5991
5992 auto EmitRemark = [&](Function &F) {
5993 auto &ORE = FAM.getResult<OptimizationRemarkEmitterAnalysis>(F);
5994 ORE.emit([&]() {
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(&F) && IsCalled(F) &&
6012 InternalizeFns.insert(&F);
6013 } else if (!F.hasLocalLinkage() && !F.hasFnAttribute(Attribute::Cold)) {
6014 EmitRemark(F);
6015 }
6016 }
6017
6018 Changed |=
6019 Attributor::internalizeFunctions(InternalizeFns, InternalizedMap);
6020 }
6021
6022 // Look at every function in the Module unless it was internalized.
6023 SetVector<Function *> Functions;
6025 for (Function &F : M)
6026 if (!F.isDeclaration() && !InternalizedMap.lookup(&F)) {
6027 SCC.push_back(&F);
6028 Functions.insert(&F);
6029 }
6030
6031 if (SCC.empty())
6033
6034 AnalysisGetter AG(FAM);
6035
6036 auto OREGetter = [&FAM](Function *F) -> OptimizationRemarkEmitter & {
6037 return FAM.getResult<OptimizationRemarkEmitterAnalysis>(*F);
6038 };
6039
6040 BumpPtrAllocator Allocator;
6041 CallGraphUpdater CGUpdater;
6042
6043 bool PostLink = LTOPhase == ThinOrFullLTOPhase::FullLTOPostLink ||
6046 OMPInformationCache InfoCache(M, AG, Allocator, /*CGSCC*/ nullptr, PostLink);
6047
6048 unsigned MaxFixpointIterations =
6050
6051 AttributorConfig AC(CGUpdater);
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;
6060 AC.IPOAmendableCB = [](const Function &F) {
6061 return F.hasFnAttribute("kernel");
6062 };
6063
6064 Attributor A(Functions, InfoCache, AC);
6065
6066 OpenMPOpt OMPOpt(SCC, CGUpdater, OREGetter, InfoCache, A);
6067 Changed |= OMPOpt.run(true);
6068
6069 // Optionally inline device functions for potentially better performance.
6071 for (Function &F : M)
6072 if (!F.isDeclaration() && !Kernels.contains(&F) &&
6073 !F.hasFnAttribute(Attribute::NoInline))
6074 F.addFnAttr(Attribute::AlwaysInline);
6075
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
6087 LazyCallGraph &CG,
6088 CGSCCUpdateResult &UR) {
6089 if (!containsOpenMP(*C.begin()->getFunction().getParent()))
6090 return PreservedAnalyses::all();
6092 return PreservedAnalyses::all();
6093
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(Fn);
6099 }
6100
6101 if (SCC.empty())
6102 return PreservedAnalyses::all();
6103
6104 Module &M = *C.begin()->getFunction().getParent();
6105
6107 LLVM_DEBUG(dbgs() << TAG << "Module before OpenMPOpt CGSCC Pass:\n" << M);
6108
6110 AM.getResult<FunctionAnalysisManagerCGSCCProxy>(C, CG).getManager();
6111
6112 AnalysisGetter AG(FAM);
6113
6114 auto OREGetter = [&FAM](Function *F) -> OptimizationRemarkEmitter & {
6115 return FAM.getResult<OptimizationRemarkEmitterAnalysis>(*F);
6116 };
6117
6118 BumpPtrAllocator Allocator;
6119 CallGraphUpdater CGUpdater;
6120 CGUpdater.initialize(CG, C, AM, UR);
6121
6122 bool PostLink = LTOPhase == ThinOrFullLTOPhase::FullLTOPostLink ||
6126 OMPInformationCache InfoCache(*(Functions.back()->getParent()), AG, Allocator,
6127 /*CGSCC*/ &Functions, PostLink);
6128
6129 unsigned MaxFixpointIterations =
6131
6132 AttributorConfig AC(CGUpdater);
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;
6141
6142 Attributor A(Functions, InfoCache, AC);
6143
6144 OpenMPOpt OMPOpt(SCC, CGUpdater, OREGetter, InfoCache, A);
6145 bool Changed = OMPOpt.run(false);
6146
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
6157 return Fn.hasFnAttribute("kernel");
6158}
6159
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(F)) {
6169 ++NumOpenMPTargetRegionKernels;
6170 Kernels.insert(&F);
6171 } else
6172 ++NumNonOpenMPTargetRegionKernels;
6173 }
6174
6175 return Kernels;
6176}
6177
6179 Metadata *MD = M.getModuleFlag("openmp");
6180 if (!MD)
6181 return false;
6182
6183 return true;
6184}
6185
6187 Metadata *MD = M.getModuleFlag("openmp-device");
6188 if (!MD)
6189 return false;
6190
6191 return true;
6192}
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
amdgpu aa AMDGPU Address space based Alias Analysis Wrapper
unsigned uint64_t
amdgpu next use AMDGPU Next Use Analysis Printer
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
Expand Atomic instructions
static cl::opt< unsigned > SetFixpointIterations("attributor-max-iterations", cl::Hidden, cl::desc("Maximal number of fixpoint iterations."), cl::init(32))
static const Function * getParent(const Value *V)
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static GCRegistry::Add< ErlangGC > A("erlang", "erlang-compatible garbage collector")
static GCRegistry::Add< CoreCLRGC > E("coreclr", "CoreCLR-compatible GC")
This file provides interfaces used to manipulate a call graph, regardless if it is a "old style" Call...
This file provides interfaces used to build and manipulate a call graph, which is a very useful tool ...
This file contains the declarations for the subclasses of Constant, which represent the different fla...
This file defines the DenseSet and SmallDenseSet classes.
This file defines an array type that can be indexed using scoped enum values.
#define DEBUG_TYPE
static void emitRemark(const Function &F, OptimizationRemarkEmitter &ORE, bool Skip)
Loop::LoopBounds::Direction Direction
Definition LoopInfo.cpp:253
#define F(x, y, z)
Definition MD5.cpp:54
#define I(x, y, z)
Definition MD5.cpp:57
Machine Check Debug Module
This file provides utility analysis objects describing memory locations.
#define T
uint64_t IntrinsicInst * II
This file defines constans and helpers used when dealing with OpenMP.
This file defines constans that will be used by both host and device compilation.
static constexpr auto TAG
static cl::opt< bool > HideMemoryTransferLatency("openmp-hide-memory-transfer-latency", cl::desc("[WIP] Tries to hide the latency of host to device memory" " transfers"), cl::Hidden, cl::init(false))
static cl::opt< bool > DisableOpenMPOptStateMachineRewrite("openmp-opt-disable-state-machine-rewrite", cl::desc("Disable OpenMP optimizations that replace the state machine."), cl::Hidden, cl::init(false))
static cl::opt< bool > EnableParallelRegionMerging("openmp-opt-enable-merging", cl::desc("Enable the OpenMP region merging optimization."), cl::Hidden, cl::init(false))
static cl::opt< bool > PrintModuleAfterOptimizations("openmp-opt-print-module-after", cl::desc("Print the current module after OpenMP optimizations."), cl::Hidden, cl::init(false))
#define KERNEL_ENVIRONMENT_CONFIGURATION_GETTER(MEMBER)
#define KERNEL_ENVIRONMENT_CONFIGURATION_IDX(MEMBER, IDX)
#define KERNEL_ENVIRONMENT_CONFIGURATION_SETTER(MEMBER)
static cl::opt< bool > PrintOpenMPKernels("openmp-print-gpu-kernels", cl::init(false), cl::Hidden)
static cl::opt< bool > DisableOpenMPOptFolding("openmp-opt-disable-folding", cl::desc("Disable OpenMP optimizations involving folding."), cl::Hidden, cl::init(false))
static bool shouldSpecializeIndirectCallee(Attributor &, const AbstractAttribute &, CallBase &, Function &, unsigned NumAssumedCallees)
Bound the if-cascade AAIndirectCallInfo builds for an indirect call.
static cl::opt< bool > PrintModuleBeforeOptimizations("openmp-opt-print-module-before", cl::desc("Print the current module before OpenMP optimizations."), cl::Hidden, cl::init(false))
static cl::opt< unsigned > SetFixpointIterations("openmp-opt-max-iterations", cl::Hidden, cl::desc("Maximal number of attributor iterations."), cl::init(256))
static cl::opt< bool > DisableInternalization("openmp-opt-disable-internalization", cl::desc("Disable function internalization."), cl::Hidden, cl::init(false))
static cl::opt< bool > PrintICVValues("openmp-print-icv-values", cl::init(false), cl::Hidden)
static cl::opt< bool > DisableOpenMPOptimizations("openmp-opt-disable", cl::desc("Disable OpenMP specific optimizations."), cl::Hidden, cl::init(false))
static cl::opt< unsigned > SharedMemoryLimit("openmp-opt-shared-limit", cl::Hidden, cl::desc("Maximum amount of shared memory to use."), cl::init(std::numeric_limits< unsigned >::max()))
static cl::opt< bool > EnableVerboseRemarks("openmp-opt-verbose-remarks", cl::desc("Enables more verbose remarks."), cl::Hidden, cl::init(false))
static cl::opt< unsigned > MaxCalleesForSpecialization("openmp-opt-max-callees-for-specialization", cl::Hidden, cl::desc("Number of possible callees above which an indirect call site is " "left alone rather than specialized into an if-cascade."), cl::init(3))
static cl::opt< bool > DisableOpenMPOptDeglobalization("openmp-opt-disable-deglobalization", cl::desc("Disable OpenMP optimizations involving deglobalization."), cl::Hidden, cl::init(false))
static cl::opt< bool > DisableOpenMPOptBarrierElimination("openmp-opt-disable-barrier-elimination", cl::desc("Disable OpenMP optimizations that eliminate barriers."), cl::Hidden, cl::init(false))
#define DEBUG_TYPE
Definition OpenMPOpt.cpp:72
static cl::opt< bool > DeduceICVValues("openmp-deduce-icv-values", cl::init(false), cl::Hidden)
#define KERNEL_ENVIRONMENT_IDX(MEMBER, IDX)
#define KERNEL_ENVIRONMENT_GETTER(MEMBER, RETURNTYPE)
static cl::opt< bool > DisableOpenMPOptSPMDization("openmp-opt-disable-spmdization", cl::desc("Disable OpenMP optimizations involving SPMD-ization."), cl::Hidden, cl::init(false))
static cl::opt< bool > AlwaysInlineDeviceFunctions("openmp-opt-inline-device", cl::desc("Inline all applicable functions on the device."), cl::Hidden, cl::init(false))
#define P(N)
FunctionAnalysisManager FAM
This file builds on the ADT/GraphTraits.h file to build a generic graph post order iterator.
This file contains the declarations for profiling metadata utility functions.
static StringRef getName(Value *V)
R600 Clause Merge
Basic Register Allocator
Remove Loads Into Fake Uses
std::pair< BasicBlock *, BasicBlock * > Edge
static bool contains(SmallPtrSetImpl< ConstantExpr * > &Cache, ConstantExpr *Expr, Constant *C)
Definition Value.cpp:484
This file implements a set that has insertion order iteration characteristics.
This file defines the SmallPtrSet class.
This file defines the SmallVector class.
This file defines the 'Statistic' class, which is designed to be an easy way to expose various metric...
#define STATISTIC(VARNAME, DESC)
Definition Statistic.h:171
This file contains some functions that are useful when dealing with strings.
#define LLVM_DEBUG(...)
Definition Debug.h:119
static void initialize(TargetLibraryInfoImpl &TLI, const Triple &T, const llvm::StringTable &StandardNames, VectorLibrary VecLib)
Initialize the set of available library functions based on the specified target triple.
Value * RHS
static cl::opt< unsigned > MaxThreads("xcore-max-threads", cl::desc("Maximum number of threads (for emulation thread-local storage)"), cl::Hidden, cl::value_desc("number"), cl::init(8))
PassT::Result & getResult(IRUnitT &IR, ExtraArgTs... ExtraArgs)
Get the result of an analysis pass for a given IR unit.
Represent a constant reference to an array (0 or more elements consecutively in memory),...
Definition ArrayRef.h:40
size_t size() const
Get the array size.
Definition ArrayRef.h:141
iterator end()
Definition BasicBlock.h:459
iterator begin()
Instruction iterator methods.
Definition BasicBlock.h:446
LLVM_ABI const_iterator getFirstInsertionPt() const
Returns an iterator to the first instruction in this block that is suitable for inserting a non-PHI i...
LLVM_ABI BasicBlock * splitBasicBlock(iterator I, const Twine &BBName="")
Split the basic block into two basic blocks at the specified instruction.
const Function * getParent() const
Return the enclosing method, or null if none.
Definition BasicBlock.h:213
reverse_iterator rbegin()
Definition BasicBlock.h:462
static BasicBlock * Create(LLVMContext &Context, const Twine &Name="", Function *Parent=nullptr, BasicBlock *InsertBefore=nullptr)
Creates a new BasicBlock.
Definition BasicBlock.h:206
LLVM_ABI const BasicBlock * getUniqueSuccessor() const
Return the successor of this block if it has a unique successor.
InstListType::reverse_iterator reverse_iterator
Definition BasicBlock.h:172
reverse_iterator rend()
Definition BasicBlock.h:464
const Instruction * getTerminator() const LLVM_READONLY
Returns the terminator instruction; assumes that the block is well-formed.
Definition BasicBlock.h:237
Base class for all callable instructions (InvokeInst and CallInst) Holds everything related to callin...
void setCallingConv(CallingConv::ID CC)
bool arg_empty() const
Function * getCalledFunction() const
Returns the function called, or null if this is an indirect function invocation or the function signa...
bool doesNotAccessMemory(unsigned OpNo) const
bool hasFnAttr(Attribute::AttrKind Kind) const
Determine whether this call has the given attribute.
LLVM_ABI bool isIndirectCall() const
Return true if the callsite is an indirect call.
bool isCallee(Value::const_user_iterator UI) const
Determine whether the passed iterator points to the callee operand's Use.
Value * getCalledOperand() const
Value * getArgOperand(unsigned i) const
void setArgOperand(unsigned i, Value *v)
iterator_range< User::op_iterator > args()
Iteration adapter for range-for loops.
unsigned getArgOperandNo(const Use *U) const
Given a use for a arg operand, get the arg operand number that corresponds to it.
unsigned arg_size() const
AttributeList getAttributes() const
Return the attributes for this call.
void addParamAttr(unsigned ArgNo, Attribute::AttrKind Kind)
Adds the attribute to the indicated argument.
bool isArgOperand(const Use *U) const
bool hasOperandBundles() const
Return true if this User has any operand bundles.
LLVM_ABI Function * getCaller()
Helper to get the caller (the parent function).
Wrapper to unify "old style" CallGraph and "new style" LazyCallGraph.
void initialize(LazyCallGraph &LCG, LazyCallGraph::SCC &SCC, CGSCCAnalysisManager &AM, CGSCCUpdateResult &UR)
Initializers for usage outside of a CGSCC pass, inside a CGSCC pass in the old and new pass manager (...
This class represents a function call, abstracting a target machine's calling convention.
static CallInst * Create(FunctionType *Ty, Value *F, const Twine &NameStr="", InsertPosition InsertBefore=nullptr)
@ ICMP_SLT
signed less than
Definition InstrTypes.h:769
@ ICMP_NE
not equal
Definition InstrTypes.h:762
static CondBrInst * Create(Value *Cond, BasicBlock *IfTrue, BasicBlock *IfFalse, InsertPosition InsertBefore=nullptr)
static LLVM_ABI Constant * getPointerCast(Constant *C, Type *Ty)
Create a BitCast, AddrSpaceCast, or a PtrToInt cast constant expression.
static LLVM_ABI Constant * getPointerBitCastOrAddrSpaceCast(Constant *C, Type *Ty)
Create a BitCast or AddrSpaceCast for a pointer type depending on the address space.
This is the shared class of boolean and integer constants.
Definition Constants.h:87
IntegerType * getIntegerType() const
Variant of the getType() method to always return an IntegerType, which reduces the amount of casting ...
Definition Constants.h:198
static LLVM_ABI ConstantInt * getTrue(LLVMContext &Context)
bool isZero() const
This is just a convenience method to make client code smaller for a common code.
Definition Constants.h:219
int64_t getSExtValue() const
Return the constant as a 64-bit integer value after it has been sign extended as appropriate for the ...
Definition Constants.h:174
static LLVM_ABI ConstantPointerNull * get(PointerType *T)
Static factory methods - Return objects of the specified value.
This is an important base class in LLVM.
Definition Constant.h:43
static LLVM_ABI Constant * getNullValue(Type *Ty)
Constructor to create a '0' constant of arbitrary type.
ValueT lookup(const_arg_type_t< KeyT > Val) const
Return the entry for the specified key, or a default constructed value if no such entry exists.
Definition DenseMap.h:794
std::pair< iterator, bool > insert(const std::pair< KeyT, ValueT > &KV)
Definition DenseMap.h:828
LLVM_ABI Instruction * findNearestCommonDominator(Instruction *I1, Instruction *I2) const
Find the nearest instruction I that dominates both I1 and I2, in the sense that a result produced bef...
static ErrorSuccess success()
Create a success value.
Definition Error.h:336
AtomicOrdering getOrdering() const
Returns the ordering constraint of this fence instruction.
A proxy from a FunctionAnalysisManager to an SCC.
const BasicBlock & getEntryBlock() const
Definition Function.h:794
const BasicBlock & front() const
Definition Function.h:845
LLVMContext & getContext() const
getContext - Return a reference to the LLVMContext associated with this function.
Definition Function.cpp:356
void setEntryCount(uint64_t Count, const DenseSet< GlobalValue::GUID > *Imports=nullptr)
Set the entry count for this function.
Argument * getArg(unsigned i) const
Definition Function.h:871
bool hasFnAttribute(Attribute::AttrKind Kind) const
Return true if the function has the attribute.
Definition Function.cpp:734
LLVM_ABI bool isDeclaration() const
Return true if the primary definition of this global value is outside of the current translation unit...
Definition Globals.cpp:408
bool hasLocalLinkage() const
Module * getParent()
Get the module that this global value is contained inside of...
@ PrivateLinkage
Like Internal, but omit from symbol table.
Definition GlobalValue.h:61
@ InternalLinkage
Rename collisions when linking (static functions).
Definition GlobalValue.h:60
const Constant * getInitializer() const
getInitializer - Return the initializer for this global variable.
LLVM_ABI void setInitializer(Constant *InitVal)
setInitializer - Sets the initializer for this global variable, removing any existing initializer if ...
Definition Globals.cpp:613
CondBrInst * CreateCondBr(Value *Cond, BasicBlock *True, BasicBlock *False, MDNode *BranchWeights=nullptr, MDNode *Unpredictable=nullptr)
Create a conditional 'br Cond, TrueDest, FalseDest' instruction.
Definition IRBuilder.h:1221
CallInst * CreateCall(FunctionType *FTy, Value *Callee, ArrayRef< Value * > Args={}, const Twine &Name="", MDNode *FPMathTag=nullptr)
Definition IRBuilder.h:2570
Value * CreateIsNull(Value *Arg, const Twine &Name="")
Return a boolean value testing if Arg == 0.
Definition IRBuilder.h:2767
LLVM_ABI bool isLifetimeStartOrEnd() const LLVM_READONLY
Return true if the instruction is a llvm.lifetime.start or llvm.lifetime.end marker.
LLVM_ABI bool mayWriteToMemory() const LLVM_READONLY
Return true if this instruction may modify memory.
const DebugLoc & getDebugLoc() const
Return the debug location for this node as a DebugLoc.
LLVM_ABI const Module * getModule() const
Return the module owning the function this instruction belongs to or nullptr it the function does not...
LLVM_ABI InstListType::iterator eraseFromParent()
This method unlinks 'this' from the containing basic block and deletes it.
LLVM_ABI const Function * getFunction() const
Return the function this instruction belongs to.
LLVM_ABI bool mayHaveSideEffects() const LLVM_READONLY
Return true if the instruction may have side effects.
LLVM_ABI bool mayReadFromMemory() const LLVM_READONLY
Return true if this instruction may read memory.
iterator_range< user_iterator > users()
void setDebugLoc(DebugLoc Loc)
Set the debug location information for this instruction.
LLVM_ABI void setSuccessor(unsigned Idx, BasicBlock *BB)
Update the specified successor to point at the provided block.
LLVM_ABI const DiagnosticHandler * getDiagHandlerPtr() const
getDiagHandlerPtr - Returns const raw pointer of DiagnosticHandler set by setDiagnosticHandler.
A node in the call graph.
An SCC of the call graph.
A lazily constructed view of the call graph of a module.
const MDOperand & getOperand(unsigned I) const
Definition Metadata.h:1437
static MDTuple * get(LLVMContext &Context, ArrayRef< Metadata * > MDs)
Definition Metadata.h:1579
unsigned getNumOperands() const
Return number of MDNode operands.
Definition Metadata.h:1443
LLVM_ABI void eraseFromParent()
This method unlinks 'this' from the containing function and deletes it.
LLVM_ABI StringRef getName() const
Return the name of the corresponding LLVM basic block, or an empty string.
Root of the metadata hierarchy.
Definition Metadata.h:64
A Module instance is used to store all the information related to an LLVM module.
Definition Module.h:68
const Triple & getTargetTriple() const
Get the target triple which is a string describing the target host.
Definition Module.h:328
LLVM_ABI Constant * getOrCreateIdent(Constant *SrcLocStr, uint32_t SrcLocStrSize, omp::IdentFlag Flags=omp::IdentFlag(0), unsigned Reserve2Flags=0)
Return an ident_t* encoding the source location SrcLocStr and Flags.
LLVM_ABI FunctionCallee getOrCreateRuntimeFunction(Module &M, omp::RuntimeFunction FnID)
Return the function declaration for the runtime function with FnID.
static LLVM_ABI std::pair< int32_t, int32_t > readThreadBoundsForKernel(const Triple &T, Function &Kernel)
}
LLVM_ABI Constant * getOrCreateSrcLocStr(StringRef LocStr, uint32_t &SrcLocStrSize)
Return the (LLVM-IR) string describing the source location LocStr.
IRBuilder<>::InsertPoint InsertPointTy
Type used throughout for insertion points.
IRBuilder Builder
The LLVM-IR Builder used to create IR.
static LLVM_ABI std::pair< int32_t, int32_t > readTeamBoundsForKernel(const Triple &T, Function &Kernel)
Read/write a bounds on teams for Kernel.
bool updateToLocation(const LocationDescription &Loc)
Update the internal location to Loc.
LLVM_ABI PreservedAnalyses run(LazyCallGraph::SCC &C, CGSCCAnalysisManager &AM, LazyCallGraph &CG, CGSCCUpdateResult &UR)
LLVM_ABI PreservedAnalyses run(Module &M, ModuleAnalysisManager &AM)
Diagnostic information for optimization analysis remarks.
The optimization diagnostic interface.
static LLVM_ABI PoisonValue * get(Type *T)
Static factory methods - Return an 'poison' object of the specified type.
A set of analyses that are preserved following a run of a transformation pass.
Definition Analysis.h:112
static PreservedAnalyses none()
Convenience factory function for the empty preserved set.
Definition Analysis.h:115
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
Definition Analysis.h:118
static LLVM_ABI ProfileSummary * getFromMD(Metadata *MD)
Construct profile summary from metdata.
static ReturnInst * Create(LLVMContext &C, Value *retVal=nullptr, InsertPosition InsertBefore=nullptr)
A vector that has set insertion semantics.
Definition SetVector.h:57
size_type size() const
Determine the number of elements in the SetVector.
Definition SetVector.h:103
size_type count(const_arg_type key) const
Count the number of elements of a given key in the SetVector.
Definition SetVector.h:268
bool insert(const value_type &X)
Insert a new element into the SetVector.
Definition SetVector.h:157
size_type size() const
size_type count(ConstPtrType Ptr) const
count - Return 1 if the specified pointer is in the set, 0 otherwise.
std::pair< iterator, bool > insert(PtrType Ptr)
Inserts Ptr if and only if there is no element in the container equal to Ptr.
iterator begin() const
SmallPtrSet - This class implements a set which is optimized for holding SmallSize or less elements.
reference emplace_back(ArgTypes &&... Args)
void append(ItTy in_start, ItTy in_end)
Add the specified range to the end of the SmallVector.
void push_back(const T &Elt)
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
bool starts_with(StringRef Prefix) const
Check if this string starts with the given Prefix.
Definition StringRef.h:258
Triple - Helper class for working with autoconf configuration names.
Definition Triple.h:48
static LLVM_ABI IntegerType * getInt32Ty(LLVMContext &C)
Definition Type.cpp:299
LLVM_ABI unsigned getPointerAddressSpace() const
Get the address space of this pointer or pointer vector type.
static UncondBrInst * Create(BasicBlock *Target, InsertPosition InsertBefore=nullptr)
static LLVM_ABI UndefValue * get(Type *T)
Static factory methods - Return an 'undef' object of the specified type.
A Use represents the edge between a Value definition and its users.
Definition Use.h:35
LLVM_ABI bool replaceUsesOfWith(Value *From, Value *To)
Replace uses of one Value with another.
Definition User.cpp:25
Type * getType() const
All values are typed, get the type of this value.
Definition Value.h:257
LLVM_ABI void setName(const Twine &Name)
Change the name of the value.
Definition Value.cpp:394
bool hasOneUse() const
Return true if there is exactly one use of this value.
Definition Value.h:441
LLVM_ABI void replaceAllUsesWith(Value *V)
Change all uses of this to point to a new Value.
Definition Value.cpp:553
iterator_range< user_iterator > users()
Definition Value.h:428
User * user_back()
Definition Value.h:414
LLVM_ABI const Value * stripPointerCasts() const
Strip off pointer casts, all-zero GEPs and address space casts.
Definition Value.cpp:712
LLVM_ABI StringRef getName() const
Return a constant reference to the value's name.
Definition Value.cpp:319
const ParentTy * getParent() const
Definition ilist_node.h:34
self_iterator getIterator()
Definition ilist_node.h:123
NodeTy * getNextNode()
Get the next node, or nullptr for the list tail.
Definition ilist_node.h:348
Changed
#define UINT64_MAX
Definition DataTypes.h:77
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
GlobalVariable * getKernelEnvironementGVFromKernelInitCB(CallBase *KernelInitCB)
ConstantStruct * getKernelEnvironementFromKernelInitCB(CallBase *KernelInitCB)
Abstract Attribute helper functions.
Definition Attributor.h:165
LLVM_ABI bool isValidAtPosition(const ValueAndContext &VAC, InformationCache &InfoCache)
Return true if the value of VAC is a valid at the position of VAC, that is a constant,...
LLVM_ABI bool isPotentiallyAffectedByBarrier(Attributor &A, const Instruction &I, const AbstractAttribute &QueryingAA)
Return true if I is potentially affected by a barrier.
@ Interprocedural
Definition Attributor.h:188
LLVM_ABI bool isNoSyncInst(Attributor &A, const Instruction &I, const AbstractAttribute &QueryingAA)
Return true if I is a nosync instruction.
constexpr char Args[]
Key for Kernel::Metadata::mArgs.
E & operator^=(E &LHS, E RHS)
@ Entry
Definition COFF.h:862
@ BasicBlock
Various leaf nodes.
Definition ISDOpcodes.h:83
initializer< Ty > init(const Ty &Val)
PointerTypeMap run(const Module &M)
Compute the PointerTypeMap for the module M.
llvm::unique_function< void(llvm::Expected< T >)> Callback
A Callback<T> is a void function that accepts Expected<T>.
Definition Transport.h:132
LLVM_ABI bool isOpenMPDevice(Module &M)
Helper to determine if M is a OpenMP target offloading device module.
LLVM_ABI bool containsOpenMP(Module &M)
Helper to determine if M contains OpenMP.
InternalControlVar
IDs for all Internal Control Variables (ICVs).
RuntimeFunction
IDs for all omp runtime library (RTL) functions.
LLVM_ABI KernelSet getDeviceKernels(Module &M)
Get OpenMP device kernels in M.
@ OMP_TGT_EXEC_MODE_GENERIC_SPMD
SetVector< Kernel > KernelSet
Set of kernels in the module.
Definition OpenMPOpt.h:24
Function * Kernel
Summary of a kernel (=entry point for target offloading).
Definition OpenMPOpt.h:21
LLVM_ABI bool isOpenMPKernel(Function &Fn)
Return true iff Fn is an OpenMP GPU kernel; Fn has the "kernel" attribute.
DiagnosticInfoOptimizationBase::Argument NV
NodeAddr< UseNode * > Use
Definition RDFGraph.h:385
bool empty() const
Definition BasicBlock.h:101
iterator end() const
Definition BasicBlock.h:89
friend class Instruction
Iterator for Instructions in a `BasicBlock.
Definition BasicBlock.h:73
LLVM_ABI iterator begin() const
This is an optimization pass for GlobalISel generic memory operations.
auto drop_begin(T &&RangeOrContainer, size_t N=1)
Return a range covering RangeOrContainer with the first N elements excluded.
Definition STLExtras.h:316
@ Offset
Definition DWP.cpp:577
bool all_of(R &&range, UnaryPredicate P)
Provide wrappers to std::all_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1755
auto size(R &&Range, std::enable_if_t< std::is_base_of< std::random_access_iterator_tag, typename std::iterator_traits< decltype(Range.begin())>::iterator_category >::value, void > *=nullptr)
Get the size of a range.
Definition STLExtras.h:1685
bool succ_empty(const Instruction *I)
Definition CFG.h:141
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
LLVM_ABI bool isRemovableAlloc(const CallBase *V, const TargetLibraryInfo *TLI)
Return true if this is a call to an allocation function that does not have side effects that we are r...
bool operator!=(uint64_t V1, const APInt &V2)
Definition APInt.h:2139
constexpr from_range_t from_range
Value * GetPointerBaseWithConstantOffset(Value *Ptr, int64_t &Offset, const DataLayout &DL, bool AllowNonInbounds=true)
Analyze the specified pointer to see if it can be expressed as a base pointer plus a constant offset.
InnerAnalysisManagerProxy< FunctionAnalysisManager, Module > FunctionAnalysisManagerModuleProxy
Provide the FunctionAnalysisManager to Module proxy.
bool operator==(const AddressRangeValuePair &LHS, const AddressRangeValuePair &RHS)
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Value
Definition InstrProf.h:143
AnalysisManager< LazyCallGraph::SCC, LazyCallGraph & > CGSCCAnalysisManager
The CGSCC analysis manager.
@ ThinLTOPostLink
ThinLTO postlink (backend compile) phase.
Definition Pass.h:83
@ FullLTOPostLink
Full LTO postlink (backend compile) phase.
Definition Pass.h:87
@ ThinLTOPreLink
ThinLTO prelink (summary) phase.
Definition Pass.h:81
auto dyn_cast_or_null(const Y &Val)
Definition Casting.h:753
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
Definition Debug.cpp:209
IRBuilder(LLVMContext &, FolderTy, InserterTy) -> IRBuilder< FolderTy, InserterTy >
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
LLVM_ABI const Value * getUnderlyingObject(const Value *V, unsigned MaxLookup=MaxLookupSearchDepth, bool MustPreserveProvenance=false)
This method strips off any GEP address adjustments, pointer casts or llvm.threadlocal....
bool isa(const From &Val)
isa<X> - Return true if the parameter to the template is an instance of one of the template type argu...
Definition Casting.h:547
MutableArrayRef(T &OneElt) -> MutableArrayRef< T >
void cantFail(Error Err, const char *Msg=nullptr)
Report a fatal error if Err is a failure value.
Definition Error.h:769
bool operator&=(SparseBitVector< ElementSize > *LHS, const SparseBitVector< ElementSize > &RHS)
LLVM_ABI BasicBlock * SplitBlock(BasicBlock *Old, BasicBlock::iterator SplitPt, DominatorTree *DT, LoopInfo *LI=nullptr, MemorySSAUpdater *MSSAU=nullptr, const Twine &BBName="")
Split the specified block at the specified instruction.
auto count(R &&Range, const E &Element)
Wrapper function around std::count to count the number of times an element Element occurs in the give...
Definition STLExtras.h:2028
ArrayRef(const T &OneElt) -> ArrayRef< T >
LLVM_ABI Value * getFreedOperand(const CallBase *CB, const TargetLibraryInfo *TLI)
If this if a call to a free function, return the freed operand.
std::string toString(const APInt &I, unsigned Radix, bool Signed, bool formatAsCLiteral=false, bool UpperCase=true, bool InsertSeparators=false)
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
auto predecessors(const MachineBasicBlock *BB)
ChangeStatus
{
Definition Attributor.h:477
LLVM_ABI Constant * ConstantFoldInsertValueInstruction(Constant *Agg, Constant *Val, ArrayRef< unsigned > Idxs)
Attempt to constant fold an insertvalue instruction with the specified operands and indices.
@ OPTIONAL
The target may be valid if the source is not.
Definition Attributor.h:489
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.
BumpPtrAllocatorImpl<> BumpPtrAllocator
The standard BumpPtrAllocator which just uses the default template parameters.
Definition Allocator.h:391
LLVM_ABI void setFittedBranchWeights(Instruction &I, ArrayRef< uint64_t > Weights, bool IsExpected, bool ElideAllZero=false)
Variant of setBranchWeights where the Weights will be fit first to uint32_t by shifting right.
AnalysisManager< Module > ModuleAnalysisManager
Convenience typedef for the Module analysis manager.
Definition MIRParser.h:39
#define N
static LLVM_ABI AAExecutionDomain & createForPosition(const IRPosition &IRP, Attributor &A)
Create an abstract attribute view for the position IRP.
AAExecutionDomain(const IRPosition &IRP, Attributor &A)
static LLVM_ABI const char ID
Unique ID (due to the unique address)
AccessKind
Simple enum to distinguish read/write/read-write accesses.
StateType::base_t MemoryLocationsKind
static LLVM_ABI bool isAlignedBarrier(const CallBase &CB, bool ExecutedAligned)
Helper function to determine if CB is an aligned (GPU) barrier.
Base struct for all "concrete attribute" deductions.
virtual const char * getIdAddr() const =0
This function should return the address of the ID of the AbstractAttribute.
An interface to query the internal state of an abstract attribute.
Wrapper for FunctionAnalysisManager.
Configuration for the Attributor.
std::function< void(Attributor &A, const Function &F)> InitializationCallback
Callback function to be invoked on internal functions marked live.
std::optional< unsigned > MaxFixpointIterations
Maximum number of iterations to run until fixpoint.
bool RewriteSignatures
Flag to determine if we rewrite function signatures.
const char * PassName
}
OptimizationRemarkGetter OREGetter
IPOAmendableCBTy IPOAmendableCB
bool IsModulePass
Is the user of the Attributor a module pass or not.
std::function< bool(Attributor &A, const AbstractAttribute &AA, CallBase &CB, Function &AssumedCallee, unsigned NumAssumedCallees)> IndirectCalleeSpecializationCallback
Callback function to determine if an indirect call targets should be made direct call targets (with a...
bool DefaultInitializeLiveInternals
Flag to determine if we want to initialize all default AAs for an internal function marked live.
The fixpoint analysis framework that orchestrates the attribute deduction.
static LLVM_ABI bool isInternalizable(Function &F)
Returns true if the function F can be internalized.
std::function< std::optional< Value * >( const IRPosition &, const AbstractAttribute *, bool &)> SimplifictionCallbackTy
Register CB as a simplification callback.
std::function< std::optional< Constant * >( const GlobalVariable &, const AbstractAttribute *, bool &)> GlobalVariableSimplifictionCallbackTy
Register CB as a simplification callback.
std::function< bool(Attributor &, const AbstractAttribute *)> VirtualUseCallbackTy
static LLVM_ABI bool internalizeFunctions(SmallPtrSetImpl< Function * > &FnSet, DenseMap< Function *, Function * > &FnMap)
Make copies of each function in the set FnSet such that the copied version has internal linkage after...
Simple wrapper for a single bit (boolean) state.
Support structure for SCC passes to communicate updates the call graph back to the CGSCC pass manager...
bool isAnyRemarkEnabled(StringRef PassName) const
Return true if any type of remarks are enabled for this pass.
Helper to describe and deal with positions in the LLVM-IR.
Definition Attributor.h:573
static const IRPosition callsite_returned(const CallBase &CB)
Create a position describing the returned value of CB.
Definition Attributor.h:641
static const IRPosition returned(const Function &F, const CallBaseContext *CBContext=nullptr)
Create a position describing the returned value of F.
Definition Attributor.h:623
static const IRPosition value(const Value &V, const CallBaseContext *CBContext=nullptr)
Create a position describing the value of V.
Definition Attributor.h:597
static const IRPosition inst(const Instruction &I, const CallBaseContext *CBContext=nullptr)
Create a position describing the instruction I.
Definition Attributor.h:609
@ IRP_ARGUMENT
An attribute for a function argument.
Definition Attributor.h:587
@ IRP_RETURNED
An attribute for the function return value.
Definition Attributor.h:583
@ IRP_CALL_SITE
An attribute for a call site (function scope).
Definition Attributor.h:586
@ IRP_CALL_SITE_RETURNED
An attribute for a call site return value.
Definition Attributor.h:584
@ IRP_FUNCTION
An attribute for a function (scope).
Definition Attributor.h:585
@ IRP_FLOAT
A position that is not associated with a spot suitable for attributes.
Definition Attributor.h:581
@ IRP_CALL_SITE_ARGUMENT
An attribute for a call site argument.
Definition Attributor.h:588
@ IRP_INVALID
An invalid position.
Definition Attributor.h:580
static const IRPosition function(const Function &F, const CallBaseContext *CBContext=nullptr)
Create a position describing the function scope of F.
Definition Attributor.h:616
Kind getPositionKind() const
Return the associated position kind.
Definition Attributor.h:847
static const IRPosition callsite_function(const CallBase &CB)
Create a position describing the function scope of CB.
Definition Attributor.h:636
Data structure to hold cached (LLVM-IR) information.
Defines various target-specific GPU grid values that must be consistent between host RTL (plugin),...