1//===----- HipStdPar.cpp - HIP C++ Standard Parallelism Support Passes ----===//
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// This file implements two passes that enable HIP C++ Standard Parallelism
9// Support:
10//
11// 1. AcceleratorCodeSelection (required): Given that only algorithms are
12// accelerated, and that the accelerated implementation exists in the form of
13// a compute kernel, we assume that only the kernel, and all functions
14// reachable from it, constitute code that the user expects the accelerator
15// to execute. Thus, we identify the set of all functions reachable from
16// kernels, and then remove all unreachable ones. This last part is necessary
17// because it is possible for code that the user did not expect to execute on
18// an accelerator to contain constructs that cannot be handled by the target
19// BE, which cannot be provably demonstrated to be dead code in general, and
20// thus can lead to mis-compilation. The degenerate case of this is when a
21// Module contains no kernels (the parent TU had no algorithm invocations fit
22// for acceleration), which we handle by completely emptying said module.
23// **NOTE**: The above does not handle indirectly reachable functions i.e.
24// it is possible to obtain a case where the target of an indirect
25// call is otherwise unreachable and thus is removed; this
26// restriction is aligned with the current `-hipstdpar` limitations
27// and will be relaxed in the future.
28//
29// 2. AllocationInterposition (required only when on-demand paging is
30// unsupported): Some accelerators or operating systems might not support
31// transparent on-demand paging. Thus, they would only be able to access
32// memory that is allocated by an accelerator-aware mechanism. For such cases
33// the user can opt into enabling allocation / deallocation interposition,
34// whereby we replace calls to known allocation / deallocation functions with
35// calls to runtime implemented equivalents that forward the requests to
36// accelerator-aware interfaces. We also support freeing system allocated
37// memory that ends up in one of the runtime equivalents, since this can
38// happen if e.g. a library that was compiled without interposition returns
39// an allocation that can be validly passed to `free`.
40//
41// 3. MathFixup (required): Some accelerators might have an incomplete
42// implementation for the intrinsics used to implement some of the math
43// functions in <cmath> / their corresponding libcall lowerings. Since this
44// can vary quite significantly between accelerators, we replace calls to a
45// set of intrinsics / lib functions known to be problematic with calls to a
46// HIPSTDPAR specific forwarding layer, which gives an uniform interface for
47// accelerators to implement in their own runtime components. This pass
48// should run before AcceleratorCodeSelection so as to prevent the spurious
49// removal of the HIPSTDPAR specific forwarding functions.
50//===----------------------------------------------------------------------===//
51
52#include "llvm/Transforms/HipStdPar/HipStdPar.h"
53
54#include "llvm/ADT/STLExtras.h"
55#include "llvm/ADT/SmallPtrSet.h"
56#include "llvm/ADT/SmallVector.h"
57#include "llvm/Analysis/CallGraph.h"
58#include "llvm/Analysis/OptimizationRemarkEmitter.h"
59#include "llvm/IR/Constants.h"
60#include "llvm/IR/Function.h"
61#include "llvm/IR/IRBuilder.h"
62#include "llvm/IR/Instructions.h"
63#include "llvm/IR/Intrinsics.h"
64#include "llvm/IR/Module.h"
65#include "llvm/Transforms/Utils/ModuleUtils.h"
66
67#include <cassert>
68#include <string>
69#include <utility>
70
71using namespace llvm;
72
73template<typename T>
74static inline void eraseFromModule(T &ToErase) {
75 ToErase.replaceAllUsesWith(PoisonValue::get(T: ToErase.getType()));
76 ToErase.eraseFromParent();
77}
78
79static bool checkIfSupported(GlobalVariable &G) {
80 if (!G.isThreadLocal())
81 return true;
82
83 G.dropDroppableUses();
84
85 if (!G.isConstantUsed())
86 return true;
87
88 std::string W;
89 raw_string_ostream OS(W);
90
91 OS << "Accelerator does not support the thread_local variable "
92 << G.getName();
93
94 Instruction *I = nullptr;
95 SmallVector<User *> Tmp(G.users());
96 SmallPtrSet<User *, 5> Visited;
97 do {
98 auto U = std::move(Tmp.back());
99 Tmp.pop_back();
100
101 if (!Visited.insert(Ptr: U).second)
102 continue;
103
104 if (isa<Instruction>(Val: U))
105 I = cast<Instruction>(Val: U);
106 else
107 Tmp.insert(I: Tmp.end(), From: U->user_begin(), To: U->user_end());
108 } while (!I && !Tmp.empty());
109
110 assert(I && "thread_local global should have at least one non-constant use.");
111
112 G.getContext().diagnose(
113 DI: DiagnosticInfoUnsupported(*I->getParent()->getParent(), W,
114 I->getDebugLoc(), DS_Error));
115
116 return false;
117}
118
119static inline void clearModule(Module &M) { // TODO: simplify.
120 while (!M.functions().empty())
121 eraseFromModule(ToErase&: *M.begin());
122 while (!M.globals().empty())
123 eraseFromModule(ToErase&: *M.globals().begin());
124 while (!M.aliases().empty())
125 eraseFromModule(ToErase&: *M.aliases().begin());
126 while (!M.ifuncs().empty())
127 eraseFromModule(ToErase&: *M.ifuncs().begin());
128}
129
130static SmallVector<std::reference_wrapper<Use>>
131collectIndirectableUses(GlobalVariable *G) {
132 // We are interested only in use chains that end in an Instruction.
133 SmallVector<std::reference_wrapper<Use>> Uses;
134
135 SmallVector<std::reference_wrapper<Use>> Stack(G->use_begin(), G->use_end());
136 while (!Stack.empty()) {
137 Use &U = Stack.pop_back_val();
138 if (isa<Instruction>(Val: U.getUser()))
139 Uses.emplace_back(Args&: U);
140 else
141 transform(Range: U.getUser()->uses(), d_first: std::back_inserter(x&: Stack),
142 F: [](auto &&U) { return std::ref(U); });
143 }
144
145 return Uses;
146}
147
148static inline GlobalVariable *getGlobalForName(GlobalVariable *G) {
149 // Create an anonymous global which stores the variable's name, which will be
150 // used by the HIPSTDPAR runtime to look up the program-wide symbol.
151 LLVMContext &Ctx = G->getContext();
152 auto *CDS = ConstantDataArray::getString(Context&: Ctx, Initializer: G->getName());
153
154 GlobalVariable *N = G->getParent()->getOrInsertGlobal(Name: "", Ty: CDS->getType());
155 N->setInitializer(CDS);
156 N->setLinkage(GlobalValue::LinkageTypes::PrivateLinkage);
157 N->setConstant(true);
158
159 return N;
160}
161
162static inline GlobalVariable *getIndirectionGlobal(Module *M) {
163 // Create an anonymous global which stores a pointer to a pointer, which will
164 // be externally initialised by the HIPSTDPAR runtime with the address of the
165 // program-wide symbol.
166 Type *PtrTy = PointerType::get(
167 C&: M->getContext(), AddressSpace: M->getDataLayout().getDefaultGlobalsAddressSpace());
168 GlobalVariable *NewG = M->getOrInsertGlobal(Name: "", Ty: PtrTy);
169
170 NewG->setInitializer(PoisonValue::get(T: NewG->getValueType()));
171 NewG->setLinkage(GlobalValue::LinkageTypes::PrivateLinkage);
172 NewG->setConstant(true);
173 NewG->setExternallyInitialized(true);
174
175 return NewG;
176}
177
178static Constant *
179appendIndirectedGlobal(const GlobalVariable *IndirectionTable,
180 SmallVector<Constant *> &SymbolIndirections,
181 GlobalVariable *ToIndirect) {
182 Module *M = ToIndirect->getParent();
183
184 auto *InitTy = cast<StructType>(Val: IndirectionTable->getValueType());
185 auto *SymbolListTy = cast<StructType>(Val: InitTy->getStructElementType(N: 2));
186 Type *NameTy = SymbolListTy->getElementType(N: 0);
187 Type *IndirectTy = SymbolListTy->getElementType(N: 1);
188
189 Constant *NameG = getGlobalForName(G: ToIndirect);
190 Constant *IndirectG = getIndirectionGlobal(M);
191 Constant *Entry = ConstantStruct::get(
192 T: SymbolListTy, V: {ConstantExpr::getAddrSpaceCast(C: NameG, Ty: NameTy),
193 ConstantExpr::getAddrSpaceCast(C: IndirectG, Ty: IndirectTy)});
194 SymbolIndirections.push_back(Elt: Entry);
195
196 return IndirectG;
197}
198
199static void fillIndirectionTable(GlobalVariable *IndirectionTable,
200 SmallVector<Constant *> Indirections) {
201 Module *M = IndirectionTable->getParent();
202 size_t SymCnt = Indirections.size();
203
204 auto *InitTy = cast<StructType>(Val: IndirectionTable->getValueType());
205 Type *SymbolListTy = InitTy->getStructElementType(N: 1);
206 auto *SymbolTy = cast<StructType>(Val: InitTy->getStructElementType(N: 2));
207
208 Constant *Count = ConstantInt::get(Ty: InitTy->getStructElementType(N: 0), V: SymCnt);
209 M->removeGlobalVariable(GV: IndirectionTable);
210 GlobalVariable *Symbols =
211 M->getOrInsertGlobal(Name: "", Ty: ArrayType::get(ElementType: SymbolTy, NumElements: SymCnt));
212 Symbols->setLinkage(GlobalValue::LinkageTypes::PrivateLinkage);
213 Symbols->setInitializer(
214 ConstantArray::get(T: ArrayType::get(ElementType: SymbolTy, NumElements: SymCnt), V: {Indirections}));
215 Symbols->setConstant(true);
216
217 Constant *ASCSymbols = ConstantExpr::getAddrSpaceCast(C: Symbols, Ty: SymbolListTy);
218 Constant *Init = ConstantStruct::get(
219 T: InitTy, V: {Count, ASCSymbols, PoisonValue::get(T: SymbolTy)});
220 M->insertGlobalVariable(GV: IndirectionTable);
221 IndirectionTable->setInitializer(Init);
222}
223
224static void replaceWithIndirectUse(const Use &U, const GlobalVariable *G,
225 Constant *IndirectedG) {
226 auto *I = cast<Instruction>(Val: U.getUser());
227
228 IRBuilder<> Builder(I);
229 unsigned OpIdx = U.getOperandNo();
230 Value *Op = I->getOperand(i: OpIdx);
231
232 // We walk back up the use chain, which could be an arbitrarily long sequence
233 // of constexpr AS casts, ptr-to-int and GEP instructions, until we reach the
234 // indirected global.
235 while (auto *CE = dyn_cast<ConstantExpr>(Val: Op)) {
236 assert((CE->getOpcode() == Instruction::GetElementPtr ||
237 CE->getOpcode() == Instruction::AddrSpaceCast ||
238 CE->getOpcode() == Instruction::PtrToInt) &&
239 "Only GEP, ASCAST or PTRTOINT constant uses supported!");
240
241 Instruction *NewI = Builder.Insert(I: CE->getAsInstruction());
242 I->replaceUsesOfWith(From: Op, To: NewI);
243 I = NewI;
244 Op = I->getOperand(i: 0);
245 OpIdx = 0;
246 Builder.SetInsertPoint(I);
247 }
248
249 assert(Op == G && "Must reach indirected global!");
250
251 I->setOperand(i: OpIdx, Val: Builder.CreateLoad(Ty: G->getType(), Ptr: IndirectedG));
252}
253
254static inline bool isValidIndirectionTable(GlobalVariable *IndirectionTable) {
255 std::string W;
256 raw_string_ostream OS(W);
257
258 Type *Ty = IndirectionTable->getValueType();
259 bool Valid = false;
260
261 if (!isa<StructType>(Val: Ty)) {
262 OS << "The Indirection Table must be a struct type; ";
263 Ty->print(O&: OS);
264 OS << " is incorrect.\n";
265 } else if (cast<StructType>(Val: Ty)->getNumElements() != 3u) {
266 OS << "The Indirection Table must have 3 elements; "
267 << cast<StructType>(Val: Ty)->getNumElements() << " is incorrect.\n";
268 } else if (!isa<IntegerType>(Val: cast<StructType>(Val: Ty)->getStructElementType(N: 0))) {
269 OS << "The first element in the Indirection Table must be an integer; ";
270 cast<StructType>(Val: Ty)->getStructElementType(N: 0)->print(O&: OS);
271 OS << " is incorrect.\n";
272 } else if (!isa<PointerType>(Val: cast<StructType>(Val: Ty)->getStructElementType(N: 1))) {
273 OS << "The second element in the Indirection Table must be a pointer; ";
274 cast<StructType>(Val: Ty)->getStructElementType(N: 1)->print(O&: OS);
275 OS << " is incorrect.\n";
276 } else if (!isa<StructType>(Val: cast<StructType>(Val: Ty)->getStructElementType(N: 2))) {
277 OS << "The third element in the Indirection Table must be a struct type; ";
278 cast<StructType>(Val: Ty)->getStructElementType(N: 2)->print(O&: OS);
279 OS << " is incorrect.\n";
280 } else {
281 Valid = true;
282 }
283
284 if (!Valid)
285 IndirectionTable->getContext().diagnose(DI: DiagnosticInfoGeneric(W, DS_Error));
286
287 return Valid;
288}
289
290static void indirectGlobals(GlobalVariable *IndirectionTable,
291 SmallVector<GlobalVariable *> ToIndirect) {
292 // We replace globals with an indirected access via a pointer that will get
293 // set by the HIPSTDPAR runtime, using their accessible, program-wide unique
294 // address as set by the host linker-loader.
295 SmallVector<Constant *> SymbolIndirections;
296 for (auto &&G : ToIndirect) {
297 SmallVector<std::reference_wrapper<Use>> Uses = collectIndirectableUses(G);
298
299 if (Uses.empty())
300 continue;
301
302 Constant *IndirectedGlobal =
303 appendIndirectedGlobal(IndirectionTable, SymbolIndirections, ToIndirect: G);
304
305 for_each(Range&: Uses,
306 F: [=](auto &&U) { replaceWithIndirectUse(U, G, IndirectedGlobal); });
307
308 eraseFromModule(ToErase&: *G);
309 }
310
311 if (SymbolIndirections.empty())
312 return;
313
314 fillIndirectionTable(IndirectionTable, Indirections: std::move(SymbolIndirections));
315}
316
317static inline void maybeHandleGlobals(Module &M) {
318 unsigned GlobAS = M.getDataLayout().getDefaultGlobalsAddressSpace();
319
320 SmallVector<GlobalVariable *> ToIndirect;
321 for (auto &&G : M.globals()) {
322 if (!checkIfSupported(G))
323 return clearModule(M);
324 if (G.getAddressSpace() != GlobAS)
325 continue;
326 if (G.isConstant() && G.hasInitializer() && G.hasAtLeastLocalUnnamedAddr())
327 continue;
328
329 ToIndirect.push_back(Elt: &G);
330 }
331
332 if (ToIndirect.empty())
333 return;
334
335 if (auto *IT = M.getNamedGlobal(Name: "__hipstdpar_symbol_indirection_table")) {
336 if (!isValidIndirectionTable(IndirectionTable: IT))
337 return clearModule(M);
338 return indirectGlobals(IndirectionTable: IT, ToIndirect: std::move(ToIndirect));
339 } else {
340 for (auto &&G : ToIndirect) {
341 // We will internalise these, so we provide a poison initialiser.
342 if (!G->hasInitializer())
343 G->setInitializer(PoisonValue::get(T: G->getValueType()));
344 }
345 }
346}
347
348template<unsigned N>
349static inline void removeUnreachableFunctions(
350 const SmallPtrSet<const Function *, N>& Reachable, Module &M) {
351 removeFromUsedLists(M, [&](Constant *C) {
352 if (auto F = dyn_cast<Function>(Val: C))
353 return !Reachable.contains(F);
354
355 return false;
356 });
357
358 SmallVector<std::reference_wrapper<Function>> ToRemove;
359 copy_if(M, std::back_inserter(x&: ToRemove), [&](auto &&F) {
360 return !F.isIntrinsic() && !Reachable.contains(&F);
361 });
362
363 for_each(Range&: ToRemove, F: eraseFromModule<Function>);
364}
365
366static inline bool isAcceleratorExecutionRoot(const Function *F) {
367 if (!F)
368 return false;
369
370 return F->getCallingConv() == CallingConv::AMDGPU_KERNEL;
371}
372
373static inline bool isCXXExceptionRuntimeFunction(StringRef Name) {
374 return Name == "__cxa_throw" || Name == "__cxa_rethrow" ||
375 Name == "__cxa_bad_cast" || Name == "__cxa_bad_typeid" ||
376 Name == "__cxa_throw_bad_array_new_length" ||
377 Name == "__cxa_rethrow_primary_exception" ||
378 Name == "__cxa_call_unexpected";
379}
380
381static inline bool checkIfExceptionHandlingIsSupported(const Function *F) {
382 for (const BasicBlock &BB : *F) {
383 for (const Instruction &I : BB) {
384 if (!I.isEHPad() &&
385 !isa<InvokeInst, ResumeInst, CatchReturnInst, CleanupReturnInst>(Val: I))
386 continue;
387
388 F->getContext().diagnose(DI: DiagnosticInfoUnsupported(
389 *F, "Accelerator does not support C++ exception handling.",
390 I.getDebugLoc(), DS_Error));
391 return false;
392 }
393 }
394
395 return true;
396}
397
398static inline bool checkIfSupported(const Function *F, const CallBase *CB) {
399 StringRef Name = F->getName();
400 const auto Dx = Name.rfind(Str: "__hipstdpar_unsupported");
401 // HIPStdPar emits unannotated host functions during device compilation and
402 // removes them here when no kernel can reach them. Defer the unsupported
403 // exception diagnostic until this point for the same reason.
404 const bool IsCXXException = isCXXExceptionRuntimeFunction(Name);
405
406 if (Dx == StringRef::npos && !IsCXXException)
407 return true;
408
409 std::string W;
410 raw_string_ostream OS(W);
411
412 if (IsCXXException) {
413 OS << "Accelerator does not support C++ exception handling.";
414 } else {
415 const auto N = Name.substr(Start: 0, N: Dx);
416 if (N == "__CXX_EXCEPTION")
417 OS << "Accelerator does not support C++ exception handling.";
418 else if (N == "__ASM")
419 OS << "Accelerator does not support the ASM block:\n"
420 << cast<ConstantDataArray>(Val: CB->getArgOperand(i: 0))->getAsCString();
421 else
422 OS << "Accelerator does not support the " << N << " function.";
423 }
424
425 auto Caller = CB->getParent()->getParent();
426
427 Caller->getContext().diagnose(
428 DI: DiagnosticInfoUnsupported(*Caller, W, CB->getDebugLoc(), DS_Error));
429
430 return false;
431}
432
433PreservedAnalyses
434 HipStdParAcceleratorCodeSelectionPass::run(Module &M,
435 ModuleAnalysisManager &MAM) {
436 auto &CGA = MAM.getResult<CallGraphAnalysis>(IR&: M);
437
438 SmallPtrSet<const Function *, 32> Reachable;
439 for (auto &&CGN : CGA) {
440 if (!isAcceleratorExecutionRoot(F: CGN.first))
441 continue;
442
443 Reachable.insert(Ptr: CGN.first);
444
445 SmallVector<const Function *> Tmp({CGN.first});
446 do {
447 auto F = std::move(Tmp.back());
448 Tmp.pop_back();
449
450 if (!checkIfExceptionHandlingIsSupported(F))
451 return PreservedAnalyses::none();
452
453 for (auto &&N : *CGA[F]) {
454 if (!N.second)
455 continue;
456 if (!N.second->getFunction())
457 continue;
458 if (Reachable.contains(Ptr: N.second->getFunction()))
459 continue;
460
461 if (!checkIfSupported(F: N.second->getFunction(),
462 CB: dyn_cast<CallBase>(Val&: *N.first)))
463 return PreservedAnalyses::none();
464
465 Reachable.insert(Ptr: N.second->getFunction());
466 Tmp.push_back(Elt: N.second->getFunction());
467 }
468 } while (!std::empty(cont: Tmp));
469 }
470
471 if (std::empty(cont: Reachable))
472 clearModule(M);
473 else
474 removeUnreachableFunctions(Reachable, M);
475
476 maybeHandleGlobals(M);
477
478 return PreservedAnalyses::none();
479}
480
481static constexpr std::pair<StringLiteral, StringLiteral> ReplaceMap[]{
482 {"aligned_alloc", "__hipstdpar_aligned_alloc"},
483 {"calloc", "__hipstdpar_calloc"},
484 {"free", "__hipstdpar_free"},
485 {"malloc", "__hipstdpar_malloc"},
486 {"memalign", "__hipstdpar_aligned_alloc"},
487 {"mmap", "__hipstdpar_mmap"},
488 {"munmap", "__hipstdpar_munmap"},
489 {"posix_memalign", "__hipstdpar_posix_aligned_alloc"},
490 {"realloc", "__hipstdpar_realloc"},
491 {"reallocarray", "__hipstdpar_realloc_array"},
492 {"_ZdaPv", "__hipstdpar_operator_delete"},
493 {"_ZdaPvm", "__hipstdpar_operator_delete_sized"},
494 {"_ZdaPvSt11align_val_t", "__hipstdpar_operator_delete_aligned"},
495 {"_ZdaPvmSt11align_val_t", "__hipstdpar_operator_delete_aligned_sized"},
496 {"_ZdlPv", "__hipstdpar_operator_delete"},
497 {"_ZdlPvm", "__hipstdpar_operator_delete_sized"},
498 {"_ZdlPvSt11align_val_t", "__hipstdpar_operator_delete_aligned"},
499 {"_ZdlPvmSt11align_val_t", "__hipstdpar_operator_delete_aligned_sized"},
500 {"_Znam", "__hipstdpar_operator_new"},
501 {"_ZnamRKSt9nothrow_t", "__hipstdpar_operator_new_nothrow"},
502 {"_ZnamSt11align_val_t", "__hipstdpar_operator_new_aligned"},
503 {"_ZnamSt11align_val_tRKSt9nothrow_t",
504 "__hipstdpar_operator_new_aligned_nothrow"},
505
506 {"_Znwm", "__hipstdpar_operator_new"},
507 {"_ZnwmRKSt9nothrow_t", "__hipstdpar_operator_new_nothrow"},
508 {"_ZnwmSt11align_val_t", "__hipstdpar_operator_new_aligned"},
509 {"_ZnwmSt11align_val_tRKSt9nothrow_t",
510 "__hipstdpar_operator_new_aligned_nothrow"},
511 {"__builtin_calloc", "__hipstdpar_calloc"},
512 {"__builtin_free", "__hipstdpar_free"},
513 {"__builtin_malloc", "__hipstdpar_malloc"},
514 {"__builtin_operator_delete", "__hipstdpar_operator_delete"},
515 {"__builtin_operator_new", "__hipstdpar_operator_new"},
516 {"__builtin_realloc", "__hipstdpar_realloc"},
517 {"__libc_calloc", "__hipstdpar_calloc"},
518 {"__libc_free", "__hipstdpar_free"},
519 {"__libc_malloc", "__hipstdpar_malloc"},
520 {"__libc_memalign", "__hipstdpar_aligned_alloc"},
521 {"__libc_realloc", "__hipstdpar_realloc"}};
522
523static constexpr std::pair<StringLiteral, StringLiteral> HiddenMap[]{
524 // hidden_malloc and hidden_free are only kept for backwards compatibility /
525 // legacy purposes, and we should remove them in the future
526 {"__hipstdpar_hidden_malloc", "__libc_malloc"},
527 {"__hipstdpar_hidden_free", "__libc_free"},
528 {"__hipstdpar_hidden_memalign", "__libc_memalign"},
529 {"__hipstdpar_hidden_mmap", "mmap"},
530 {"__hipstdpar_hidden_munmap", "munmap"}};
531
532PreservedAnalyses
533HipStdParAllocationInterpositionPass::run(Module &M, ModuleAnalysisManager&) {
534 SmallDenseMap<StringRef, StringRef> AllocReplacements(std::cbegin(cont: ReplaceMap),
535 std::cend(cont: ReplaceMap));
536
537 for (auto &&F : M) {
538 if (!F.hasName())
539 continue;
540 auto It = AllocReplacements.find(Val: F.getName());
541 if (It == AllocReplacements.end())
542 continue;
543
544 if (auto R = M.getFunction(Name: It->second)) {
545 F.replaceAllUsesWith(V: R);
546 } else {
547 std::string W;
548 raw_string_ostream OS(W);
549
550 OS << "cannot be interposed, missing: " << AllocReplacements[F.getName()]
551 << ". Tried to run the allocation interposition pass without the "
552 << "replacement functions available.";
553
554 F.getContext().diagnose(DI: DiagnosticInfoUnsupported(F, W,
555 F.getSubprogram(),
556 DS_Warning));
557 }
558 }
559
560 for (auto &&HR : HiddenMap) {
561 if (auto F = M.getFunction(Name: HR.first)) {
562 auto R = M.getOrInsertFunction(Name: HR.second, T: F->getFunctionType(),
563 AttributeList: F->getAttributes());
564 F->replaceAllUsesWith(V: R.getCallee());
565
566 eraseFromModule(ToErase&: *F);
567 }
568 }
569
570 return PreservedAnalyses::none();
571}
572
573static constexpr std::pair<StringLiteral, StringLiteral> MathLibToHipStdPar[]{
574 {"acosh", "__hipstdpar_acosh_f64"},
575 {"acoshf", "__hipstdpar_acosh_f32"},
576 {"asinh", "__hipstdpar_asinh_f64"},
577 {"asinhf", "__hipstdpar_asinh_f32"},
578 {"atanh", "__hipstdpar_atanh_f64"},
579 {"atanhf", "__hipstdpar_atanh_f32"},
580 {"cbrt", "__hipstdpar_cbrt_f64"},
581 {"cbrtf", "__hipstdpar_cbrt_f32"},
582 {"erf", "__hipstdpar_erf_f64"},
583 {"erff", "__hipstdpar_erf_f32"},
584 {"erfc", "__hipstdpar_erfc_f64"},
585 {"erfcf", "__hipstdpar_erfc_f32"},
586 {"fdim", "__hipstdpar_fdim_f64"},
587 {"fdimf", "__hipstdpar_fdim_f32"},
588 {"expm1", "__hipstdpar_expm1_f64"},
589 {"expm1f", "__hipstdpar_expm1_f32"},
590 {"hypot", "__hipstdpar_hypot_f64"},
591 {"hypotf", "__hipstdpar_hypot_f32"},
592 {"ilogb", "__hipstdpar_ilogb_f64"},
593 {"ilogbf", "__hipstdpar_ilogb_f32"},
594 {"lgamma", "__hipstdpar_lgamma_f64"},
595 {"lgammaf", "__hipstdpar_lgamma_f32"},
596 {"log1p", "__hipstdpar_log1p_f64"},
597 {"log1pf", "__hipstdpar_log1p_f32"},
598 {"logb", "__hipstdpar_logb_f64"},
599 {"logbf", "__hipstdpar_logb_f32"},
600 {"nextafter", "__hipstdpar_nextafter_f64"},
601 {"nextafterf", "__hipstdpar_nextafter_f32"},
602 {"nexttoward", "__hipstdpar_nexttoward_f64"},
603 {"nexttowardf", "__hipstdpar_nexttoward_f32"},
604 {"remainder", "__hipstdpar_remainder_f64"},
605 {"remainderf", "__hipstdpar_remainder_f32"},
606 {"remquo", "__hipstdpar_remquo_f64"},
607 {"remquof", "__hipstdpar_remquo_f32"},
608 {"scalbln", "__hipstdpar_scalbln_f64"},
609 {"scalblnf", "__hipstdpar_scalbln_f32"},
610 {"scalbn", "__hipstdpar_scalbn_f64"},
611 {"scalbnf", "__hipstdpar_scalbn_f32"},
612 {"tgamma", "__hipstdpar_tgamma_f64"},
613 {"tgammaf", "__hipstdpar_tgamma_f32"}};
614
615PreservedAnalyses HipStdParMathFixupPass::run(Module &M,
616 ModuleAnalysisManager &) {
617 if (M.empty())
618 return PreservedAnalyses::all();
619
620 SmallVector<std::pair<Function *, std::string>> ToReplace;
621 for (auto &&F : M) {
622 if (!F.hasName())
623 continue;
624
625 StringRef N = F.getName();
626 Intrinsic::ID ID = F.getIntrinsicID();
627
628 switch (ID) {
629 case Intrinsic::not_intrinsic: {
630 auto It =
631 find_if(Range: MathLibToHipStdPar, P: [&](auto &&M) { return M.first == N; });
632 if (It == std::cend(cont: MathLibToHipStdPar))
633 continue;
634 ToReplace.emplace_back(Args: &F, Args: It->second);
635 continue;
636 }
637 case Intrinsic::acos:
638 case Intrinsic::asin:
639 case Intrinsic::atan:
640 case Intrinsic::atan2:
641 case Intrinsic::cosh:
642 case Intrinsic::modf:
643 case Intrinsic::sincos:
644 case Intrinsic::sinh:
645 case Intrinsic::tan:
646 case Intrinsic::tanh:
647 break;
648 default: {
649 if (F.getReturnType()->isDoubleTy()) {
650 switch (ID) {
651 case Intrinsic::cos:
652 case Intrinsic::log:
653 case Intrinsic::log10:
654 case Intrinsic::log2:
655 case Intrinsic::pow:
656 case Intrinsic::sin:
657 break;
658 default:
659 continue;
660 }
661 break;
662 }
663 continue;
664 }
665 }
666
667 ToReplace.emplace_back(Args: &F, Args&: N);
668 llvm::replace(Range&: ToReplace.back().second, OldValue: '.', NewValue: '_');
669 StringRef Prefix = "llvm";
670 ToReplace.back().second.replace(pos: 0, n1: Prefix.size(), s: "__hipstdpar");
671 }
672 for (auto &&[F, NewF] : ToReplace)
673 F->replaceAllUsesWith(
674 V: M.getOrInsertFunction(Name: NewF, T: F->getFunctionType()).getCallee());
675
676 return PreservedAnalyses::none();
677}
678