1//===- AMDGPULibCalls.cpp -------------------------------------------------===//
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/// \file
10/// This file does AMD library function optimizations.
11//
12//===----------------------------------------------------------------------===//
13
14#include "AMDGPU.h"
15#include "AMDGPULibFunc.h"
16#include "llvm/Analysis/AssumptionCache.h"
17#include "llvm/Analysis/TargetLibraryInfo.h"
18#include "llvm/Analysis/ValueTracking.h"
19#include "llvm/IR/AttributeMask.h"
20#include "llvm/IR/Dominators.h"
21#include "llvm/IR/IRBuilder.h"
22#include "llvm/IR/MDBuilder.h"
23#include "llvm/IR/PatternMatch.h"
24#include <cmath>
25
26#define DEBUG_TYPE "amdgpu-simplifylib"
27
28using namespace llvm;
29using namespace llvm::PatternMatch;
30
31static cl::opt<bool> EnablePreLink("amdgpu-prelink",
32 cl::desc("Enable pre-link mode optimizations"),
33 cl::init(Val: false),
34 cl::Hidden);
35
36static cl::list<std::string> UseNative("amdgpu-use-native",
37 cl::desc("Comma separated list of functions to replace with native, or all"),
38 cl::CommaSeparated, cl::ValueOptional,
39 cl::Hidden);
40
41#define MATH_PI numbers::pi
42#define MATH_E numbers::e
43#define MATH_SQRT2 numbers::sqrt2
44#define MATH_SQRT1_2 numbers::inv_sqrt2
45
46enum class PowKind { Pow, PowR, PowN, RootN };
47
48namespace llvm {
49
50class AMDGPULibCalls {
51private:
52 SimplifyQuery SQ;
53
54 using FuncInfo = llvm::AMDGPULibFunc;
55
56 // -fuse-native.
57 bool AllNative = false;
58
59 bool useNativeFunc(const StringRef F) const;
60
61 // Return a pointer (pointer expr) to the function if function definition with
62 // "FuncName" exists. It may create a new function prototype in pre-link mode.
63 FunctionCallee getFunction(Module *M, const FuncInfo &fInfo);
64
65 /// Wrapper around getFunction which tries to use a faster variant if
66 /// available, and falls back to a less fast option.
67 ///
68 /// Return a replacement function for \p fInfo that has float-typed fast
69 /// variants. \p NewFunc is a base replacement function to use. \p
70 /// NewFuncFastVariant is a faster version to use if the calling context knows
71 /// it's legal. If there is no fast variant to use, \p NewFuncFastVariant
72 /// should be EI_NONE.
73 FunctionCallee getFloatFastVariant(Module *M, const FuncInfo &fInfo,
74 FuncInfo &newInfo,
75 AMDGPULibFunc::EFuncId NewFunc,
76 AMDGPULibFunc::EFuncId NewFuncFastVariant);
77
78 bool parseFunctionName(const StringRef &FMangledName, FuncInfo &FInfo);
79
80 bool TDOFold(CallInst *CI, const FuncInfo &FInfo);
81
82 /* Specialized optimizations */
83
84 // pow/powr/pown
85 bool fold_pow(FPMathOperator *FPOp, IRBuilder<> &B, const FuncInfo &FInfo);
86
87 /// Peform a fast math expansion of pow, powr, pown or rootn.
88 bool expandFastPow(FPMathOperator *FPOp, IRBuilder<> &B, PowKind Kind);
89
90 bool tryOptimizePow(FPMathOperator *FPOp, IRBuilder<> &B,
91 const FuncInfo &FInfo);
92
93 // rootn
94 bool fold_rootn(FPMathOperator *FPOp, IRBuilder<> &B, const FuncInfo &FInfo);
95
96 // -fuse-native for sincos
97 bool sincosUseNative(CallInst *aCI, const FuncInfo &FInfo);
98
99 // evaluate calls if calls' arguments are constants.
100 bool evaluateScalarMathFunc(const FuncInfo &FInfo, APFloat &Res0,
101 APFloat &Res1, Constant *copr0, Constant *copr1);
102 bool evaluateCall(CallInst *aCI, const FuncInfo &FInfo);
103
104 /// Insert a value to sincos function \p Fsincos. Returns (value of sin, value
105 /// of cos, sincos call).
106 std::tuple<Value *, Value *, Value *> insertSinCos(Value *Arg,
107 FastMathFlags FMF,
108 IRBuilder<> &B,
109 FunctionCallee Fsincos);
110
111 // sin/cos
112 bool fold_sincos(FPMathOperator *FPOp, IRBuilder<> &B, const FuncInfo &FInfo);
113
114 // __read_pipe/__write_pipe
115 bool fold_read_write_pipe(CallInst *CI, IRBuilder<> &B,
116 const FuncInfo &FInfo);
117
118 /// Substitute a call to a known libcall with an intrinsic call. If \p
119 /// AllowMinSize is true, allow the replacement in a minsize function.
120 bool shouldReplaceLibcallWithIntrinsic(const CallInst *CI,
121 bool AllowMinSizeF32 = false,
122 bool AllowF64 = false,
123 bool AllowStrictFP = false);
124 void replaceLibCallWithSimpleIntrinsic(IRBuilder<> &B, CallInst *CI,
125 Intrinsic::ID IntrID);
126
127 bool tryReplaceLibcallWithSimpleIntrinsic(IRBuilder<> &B, CallInst *CI,
128 Intrinsic::ID IntrID,
129 bool AllowMinSizeF32 = false,
130 bool AllowF64 = false,
131 bool AllowStrictFP = false);
132
133protected:
134 bool isUnsafeFiniteOnlyMath(const FPMathOperator *FPOp) const;
135
136 bool canIncreasePrecisionOfConstantFold(const FPMathOperator *FPOp) const;
137
138 static void replaceCall(Instruction *I, Value *With) {
139 I->replaceAllUsesWith(V: With);
140 I->eraseFromParent();
141 }
142
143 static void replaceCall(FPMathOperator *I, Value *With) {
144 replaceCall(I: cast<Instruction>(Val: I), With);
145 }
146
147public:
148 AMDGPULibCalls(Function &F, FunctionAnalysisManager &FAM);
149
150 bool fold(CallInst *CI);
151
152 void initNativeFuncs();
153
154 // Replace a normal math function call with that native version
155 bool useNative(CallInst *CI);
156};
157
158} // end namespace llvm
159
160template <typename IRB>
161static CallInst *CreateCallEx(IRB &B, FunctionCallee Callee, Value *Arg,
162 const Twine &Name = "") {
163 CallInst *R = B.CreateCall(Callee, Arg, Name);
164 if (Function *F = dyn_cast<Function>(Val: Callee.getCallee()))
165 R->setCallingConv(F->getCallingConv());
166 return R;
167}
168
169template <typename IRB>
170static CallInst *CreateCallEx2(IRB &B, FunctionCallee Callee, Value *Arg1,
171 Value *Arg2, const Twine &Name = "") {
172 CallInst *R = B.CreateCall(Callee, {Arg1, Arg2}, Name);
173 if (Function *F = dyn_cast<Function>(Val: Callee.getCallee()))
174 R->setCallingConv(F->getCallingConv());
175 return R;
176}
177
178static FunctionType *getPownType(FunctionType *FT) {
179 Type *PowNExpTy = Type::getInt32Ty(C&: FT->getContext());
180 if (VectorType *VecTy = dyn_cast<VectorType>(Val: FT->getReturnType()))
181 PowNExpTy = VectorType::get(ElementType: PowNExpTy, EC: VecTy->getElementCount());
182
183 return FunctionType::get(Result: FT->getReturnType(),
184 Params: {FT->getParamType(i: 0), PowNExpTy}, isVarArg: false);
185}
186
187// Data structures for table-driven optimizations.
188// FuncTbl works for both f32 and f64 functions with 1 input argument
189
190struct TableEntry {
191 double result;
192 double input;
193};
194
195/* a list of {result, input} */
196static const TableEntry tbl_acos[] = {
197 {MATH_PI / 2.0, .input: 0.0},
198 {MATH_PI / 2.0, .input: -0.0},
199 {.result: 0.0, .input: 1.0},
200 {MATH_PI, .input: -1.0}
201};
202static const TableEntry tbl_acosh[] = {
203 {.result: 0.0, .input: 1.0}
204};
205static const TableEntry tbl_acospi[] = {
206 {.result: 0.5, .input: 0.0},
207 {.result: 0.5, .input: -0.0},
208 {.result: 0.0, .input: 1.0},
209 {.result: 1.0, .input: -1.0}
210};
211static const TableEntry tbl_asin[] = {
212 {.result: 0.0, .input: 0.0},
213 {.result: -0.0, .input: -0.0},
214 {MATH_PI / 2.0, .input: 1.0},
215 {.result: -MATH_PI / 2.0, .input: -1.0}
216};
217static const TableEntry tbl_asinh[] = {
218 {.result: 0.0, .input: 0.0},
219 {.result: -0.0, .input: -0.0}
220};
221static const TableEntry tbl_asinpi[] = {
222 {.result: 0.0, .input: 0.0},
223 {.result: -0.0, .input: -0.0},
224 {.result: 0.5, .input: 1.0},
225 {.result: -0.5, .input: -1.0}
226};
227static const TableEntry tbl_atan[] = {
228 {.result: 0.0, .input: 0.0},
229 {.result: -0.0, .input: -0.0},
230 {MATH_PI / 4.0, .input: 1.0},
231 {.result: -MATH_PI / 4.0, .input: -1.0}
232};
233static const TableEntry tbl_atanh[] = {
234 {.result: 0.0, .input: 0.0},
235 {.result: -0.0, .input: -0.0}
236};
237static const TableEntry tbl_atanpi[] = {
238 {.result: 0.0, .input: 0.0},
239 {.result: -0.0, .input: -0.0},
240 {.result: 0.25, .input: 1.0},
241 {.result: -0.25, .input: -1.0}
242};
243static const TableEntry tbl_cbrt[] = {
244 {.result: 0.0, .input: 0.0},
245 {.result: -0.0, .input: -0.0},
246 {.result: 1.0, .input: 1.0},
247 {.result: -1.0, .input: -1.0},
248};
249static const TableEntry tbl_cos[] = {
250 {.result: 1.0, .input: 0.0},
251 {.result: 1.0, .input: -0.0}
252};
253static const TableEntry tbl_cosh[] = {
254 {.result: 1.0, .input: 0.0},
255 {.result: 1.0, .input: -0.0}
256};
257static const TableEntry tbl_cospi[] = {
258 {.result: 1.0, .input: 0.0},
259 {.result: 1.0, .input: -0.0}
260};
261static const TableEntry tbl_erfc[] = {
262 {.result: 1.0, .input: 0.0},
263 {.result: 1.0, .input: -0.0}
264};
265static const TableEntry tbl_erf[] = {
266 {.result: 0.0, .input: 0.0},
267 {.result: -0.0, .input: -0.0}
268};
269static const TableEntry tbl_exp[] = {
270 {.result: 1.0, .input: 0.0},
271 {.result: 1.0, .input: -0.0},
272 {MATH_E, .input: 1.0}
273};
274static const TableEntry tbl_exp2[] = {
275 {.result: 1.0, .input: 0.0},
276 {.result: 1.0, .input: -0.0},
277 {.result: 2.0, .input: 1.0}
278};
279static const TableEntry tbl_exp10[] = {
280 {.result: 1.0, .input: 0.0},
281 {.result: 1.0, .input: -0.0},
282 {.result: 10.0, .input: 1.0}
283};
284static const TableEntry tbl_expm1[] = {
285 {.result: 0.0, .input: 0.0},
286 {.result: -0.0, .input: -0.0}
287};
288static const TableEntry tbl_log[] = {
289 {.result: 0.0, .input: 1.0},
290 {.result: 1.0, MATH_E}
291};
292static const TableEntry tbl_log2[] = {
293 {.result: 0.0, .input: 1.0},
294 {.result: 1.0, .input: 2.0}
295};
296static const TableEntry tbl_log10[] = {
297 {.result: 0.0, .input: 1.0},
298 {.result: 1.0, .input: 10.0}
299};
300static const TableEntry tbl_rsqrt[] = {
301 {.result: 1.0, .input: 1.0},
302 {MATH_SQRT1_2, .input: 2.0}
303};
304static const TableEntry tbl_sin[] = {
305 {.result: 0.0, .input: 0.0},
306 {.result: -0.0, .input: -0.0}
307};
308static const TableEntry tbl_sinh[] = {
309 {.result: 0.0, .input: 0.0},
310 {.result: -0.0, .input: -0.0}
311};
312static const TableEntry tbl_sinpi[] = {
313 {.result: 0.0, .input: 0.0},
314 {.result: -0.0, .input: -0.0}
315};
316static const TableEntry tbl_sqrt[] = {
317 {.result: 0.0, .input: 0.0},
318 {.result: 1.0, .input: 1.0},
319 {MATH_SQRT2, .input: 2.0}
320};
321static const TableEntry tbl_tan[] = {
322 {.result: 0.0, .input: 0.0},
323 {.result: -0.0, .input: -0.0}
324};
325static const TableEntry tbl_tanh[] = {
326 {.result: 0.0, .input: 0.0},
327 {.result: -0.0, .input: -0.0}
328};
329static const TableEntry tbl_tanpi[] = {
330 {.result: 0.0, .input: 0.0},
331 {.result: -0.0, .input: -0.0}
332};
333static const TableEntry tbl_tgamma[] = {
334 {.result: 1.0, .input: 1.0},
335 {.result: 1.0, .input: 2.0},
336 {.result: 2.0, .input: 3.0},
337 {.result: 6.0, .input: 4.0}
338};
339
340static bool HasNative(AMDGPULibFunc::EFuncId id) {
341 switch(id) {
342 case AMDGPULibFunc::EI_DIVIDE:
343 case AMDGPULibFunc::EI_COS:
344 case AMDGPULibFunc::EI_EXP:
345 case AMDGPULibFunc::EI_EXP2:
346 case AMDGPULibFunc::EI_EXP10:
347 case AMDGPULibFunc::EI_LOG:
348 case AMDGPULibFunc::EI_LOG2:
349 case AMDGPULibFunc::EI_LOG10:
350 case AMDGPULibFunc::EI_POWR:
351 case AMDGPULibFunc::EI_RECIP:
352 case AMDGPULibFunc::EI_RSQRT:
353 case AMDGPULibFunc::EI_SIN:
354 case AMDGPULibFunc::EI_SINCOS:
355 case AMDGPULibFunc::EI_SQRT:
356 case AMDGPULibFunc::EI_TAN:
357 return true;
358 default:;
359 }
360 return false;
361}
362
363using TableRef = ArrayRef<TableEntry>;
364
365static TableRef getOptTable(AMDGPULibFunc::EFuncId id) {
366 switch(id) {
367 case AMDGPULibFunc::EI_ACOS: return TableRef(tbl_acos);
368 case AMDGPULibFunc::EI_ACOSH: return TableRef(tbl_acosh);
369 case AMDGPULibFunc::EI_ACOSPI: return TableRef(tbl_acospi);
370 case AMDGPULibFunc::EI_ASIN: return TableRef(tbl_asin);
371 case AMDGPULibFunc::EI_ASINH: return TableRef(tbl_asinh);
372 case AMDGPULibFunc::EI_ASINPI: return TableRef(tbl_asinpi);
373 case AMDGPULibFunc::EI_ATAN: return TableRef(tbl_atan);
374 case AMDGPULibFunc::EI_ATANH: return TableRef(tbl_atanh);
375 case AMDGPULibFunc::EI_ATANPI: return TableRef(tbl_atanpi);
376 case AMDGPULibFunc::EI_CBRT: return TableRef(tbl_cbrt);
377 case AMDGPULibFunc::EI_NCOS:
378 case AMDGPULibFunc::EI_COS: return TableRef(tbl_cos);
379 case AMDGPULibFunc::EI_COSH: return TableRef(tbl_cosh);
380 case AMDGPULibFunc::EI_COSPI: return TableRef(tbl_cospi);
381 case AMDGPULibFunc::EI_ERFC: return TableRef(tbl_erfc);
382 case AMDGPULibFunc::EI_ERF: return TableRef(tbl_erf);
383 case AMDGPULibFunc::EI_EXP: return TableRef(tbl_exp);
384 case AMDGPULibFunc::EI_NEXP2:
385 case AMDGPULibFunc::EI_EXP2: return TableRef(tbl_exp2);
386 case AMDGPULibFunc::EI_EXP10: return TableRef(tbl_exp10);
387 case AMDGPULibFunc::EI_EXPM1: return TableRef(tbl_expm1);
388 case AMDGPULibFunc::EI_LOG: return TableRef(tbl_log);
389 case AMDGPULibFunc::EI_NLOG2:
390 case AMDGPULibFunc::EI_LOG2: return TableRef(tbl_log2);
391 case AMDGPULibFunc::EI_LOG10: return TableRef(tbl_log10);
392 case AMDGPULibFunc::EI_NRSQRT:
393 case AMDGPULibFunc::EI_RSQRT: return TableRef(tbl_rsqrt);
394 case AMDGPULibFunc::EI_NSIN:
395 case AMDGPULibFunc::EI_SIN: return TableRef(tbl_sin);
396 case AMDGPULibFunc::EI_SINH: return TableRef(tbl_sinh);
397 case AMDGPULibFunc::EI_SINPI: return TableRef(tbl_sinpi);
398 case AMDGPULibFunc::EI_NSQRT:
399 case AMDGPULibFunc::EI_SQRT: return TableRef(tbl_sqrt);
400 case AMDGPULibFunc::EI_TAN: return TableRef(tbl_tan);
401 case AMDGPULibFunc::EI_TANH: return TableRef(tbl_tanh);
402 case AMDGPULibFunc::EI_TANPI: return TableRef(tbl_tanpi);
403 case AMDGPULibFunc::EI_TGAMMA: return TableRef(tbl_tgamma);
404 default:;
405 }
406 return TableRef();
407}
408
409static inline int getVecSize(const AMDGPULibFunc& FInfo) {
410 return FInfo.getLeads()[0].VectorSize;
411}
412
413static inline AMDGPULibFunc::EType getArgType(const AMDGPULibFunc& FInfo) {
414 return (AMDGPULibFunc::EType)FInfo.getLeads()[0].ArgType;
415}
416
417FunctionCallee AMDGPULibCalls::getFunction(Module *M, const FuncInfo &fInfo) {
418 // If we are doing PreLinkOpt, the function is external. So it is safe to
419 // use getOrInsertFunction() at this stage.
420
421 return EnablePreLink ? AMDGPULibFunc::getOrInsertFunction(M, fInfo)
422 : AMDGPULibFunc::getFunction(M, fInfo);
423}
424
425FunctionCallee AMDGPULibCalls::getFloatFastVariant(
426 Module *M, const FuncInfo &fInfo, FuncInfo &newInfo,
427 AMDGPULibFunc::EFuncId NewFunc, AMDGPULibFunc::EFuncId FastVariant) {
428 assert(NewFunc != FastVariant);
429
430 if (FastVariant != AMDGPULibFunc::EI_NONE &&
431 getArgType(FInfo: fInfo) == AMDGPULibFunc::F32) {
432 newInfo = AMDGPULibFunc(FastVariant, fInfo);
433 if (FunctionCallee NewCallee = getFunction(M, fInfo: newInfo))
434 return NewCallee;
435 }
436
437 newInfo = AMDGPULibFunc(NewFunc, fInfo);
438 return getFunction(M, fInfo: newInfo);
439}
440
441bool AMDGPULibCalls::parseFunctionName(const StringRef &FMangledName,
442 FuncInfo &FInfo) {
443 return AMDGPULibFunc::parse(MangledName: FMangledName, Ptr&: FInfo);
444}
445
446bool AMDGPULibCalls::isUnsafeFiniteOnlyMath(const FPMathOperator *FPOp) const {
447 return FPOp->hasApproxFunc() && FPOp->hasNoNaNs() && FPOp->hasNoInfs();
448}
449
450bool AMDGPULibCalls::canIncreasePrecisionOfConstantFold(
451 const FPMathOperator *FPOp) const {
452 // TODO: Refine to approxFunc or contract
453 return FPOp->isFast();
454}
455
456AMDGPULibCalls::AMDGPULibCalls(Function &F, FunctionAnalysisManager &FAM)
457 : SQ(F.getDataLayout(), &FAM.getResult<TargetLibraryAnalysis>(IR&: F),
458 FAM.getCachedResult<DominatorTreeAnalysis>(IR&: F),
459 &FAM.getResult<AssumptionAnalysis>(IR&: F)) {}
460
461bool AMDGPULibCalls::useNativeFunc(const StringRef F) const {
462 return AllNative || llvm::is_contained(Range&: UseNative, Element: F);
463}
464
465void AMDGPULibCalls::initNativeFuncs() {
466 AllNative = useNativeFunc(F: "all") ||
467 (UseNative.getNumOccurrences() && UseNative.size() == 1 &&
468 UseNative.begin()->empty());
469}
470
471bool AMDGPULibCalls::sincosUseNative(CallInst *aCI, const FuncInfo &FInfo) {
472 bool native_sin = useNativeFunc(F: "sin");
473 bool native_cos = useNativeFunc(F: "cos");
474
475 if (native_sin && native_cos) {
476 Module *M = aCI->getModule();
477 Value *opr0 = aCI->getArgOperand(i: 0);
478
479 AMDGPULibFunc nf;
480 nf.getLeads()[0].ArgType = FInfo.getLeads()[0].ArgType;
481 nf.getLeads()[0].VectorSize = FInfo.getLeads()[0].VectorSize;
482
483 nf.setPrefix(AMDGPULibFunc::NATIVE);
484 nf.setId(AMDGPULibFunc::EI_SIN);
485 FunctionCallee sinExpr = getFunction(M, fInfo: nf);
486
487 nf.setPrefix(AMDGPULibFunc::NATIVE);
488 nf.setId(AMDGPULibFunc::EI_COS);
489 FunctionCallee cosExpr = getFunction(M, fInfo: nf);
490 if (sinExpr && cosExpr) {
491 Value *sinval =
492 CallInst::Create(Func: sinExpr, Args: opr0, NameStr: "splitsin", InsertBefore: aCI->getIterator());
493 Value *cosval =
494 CallInst::Create(Func: cosExpr, Args: opr0, NameStr: "splitcos", InsertBefore: aCI->getIterator());
495 new StoreInst(cosval, aCI->getArgOperand(i: 1), aCI->getIterator());
496
497 DEBUG_WITH_TYPE("usenative", dbgs() << "<useNative> replace " << *aCI
498 << " with native version of sin/cos");
499
500 replaceCall(I: aCI, With: sinval);
501 return true;
502 }
503 }
504 return false;
505}
506
507bool AMDGPULibCalls::useNative(CallInst *aCI) {
508 Function *Callee = aCI->getCalledFunction();
509 if (!Callee || aCI->isNoBuiltin())
510 return false;
511
512 FuncInfo FInfo;
513 if (!parseFunctionName(FMangledName: Callee->getName(), FInfo) || !FInfo.isMangled() ||
514 FInfo.getPrefix() != AMDGPULibFunc::NOPFX ||
515 getArgType(FInfo) == AMDGPULibFunc::F64 || !HasNative(id: FInfo.getId()) ||
516 !(AllNative || useNativeFunc(F: FInfo.getName()))) {
517 return false;
518 }
519
520 if (FInfo.getId() == AMDGPULibFunc::EI_SINCOS)
521 return sincosUseNative(aCI, FInfo);
522
523 FInfo.setPrefix(AMDGPULibFunc::NATIVE);
524 FunctionCallee F = getFunction(M: aCI->getModule(), fInfo: FInfo);
525 if (!F)
526 return false;
527
528 aCI->setCalledFunction(F);
529 DEBUG_WITH_TYPE("usenative", dbgs() << "<useNative> replace " << *aCI
530 << " with native version");
531 return true;
532}
533
534// Clang emits call of __read_pipe_2 or __read_pipe_4 for OpenCL read_pipe
535// builtin, with appended type size and alignment arguments, where 2 or 4
536// indicates the original number of arguments. The library has optimized version
537// of __read_pipe_2/__read_pipe_4 when the type size and alignment has the same
538// power of 2 value. This function transforms __read_pipe_2 to __read_pipe_2_N
539// for such cases where N is the size in bytes of the type (N = 1, 2, 4, 8, ...,
540// 128). The same for __read_pipe_4, write_pipe_2, and write_pipe_4.
541bool AMDGPULibCalls::fold_read_write_pipe(CallInst *CI, IRBuilder<> &B,
542 const FuncInfo &FInfo) {
543 auto *Callee = CI->getCalledFunction();
544 if (!Callee->isDeclaration())
545 return false;
546
547 assert(Callee->hasName() && "Invalid read_pipe/write_pipe function");
548 auto *M = Callee->getParent();
549 std::string Name = std::string(Callee->getName());
550 auto NumArg = CI->arg_size();
551 if (NumArg != 4 && NumArg != 6)
552 return false;
553 ConstantInt *PacketSize =
554 dyn_cast<ConstantInt>(Val: CI->getArgOperand(i: NumArg - 2));
555 ConstantInt *PacketAlign =
556 dyn_cast<ConstantInt>(Val: CI->getArgOperand(i: NumArg - 1));
557 if (!PacketSize || !PacketAlign)
558 return false;
559
560 unsigned Size = PacketSize->getZExtValue();
561 Align Alignment = PacketAlign->getAlignValue();
562 if (Alignment != Size)
563 return false;
564
565 unsigned PtrArgLoc = CI->arg_size() - 3;
566 Value *PtrArg = CI->getArgOperand(i: PtrArgLoc);
567 Type *PtrTy = PtrArg->getType();
568
569 SmallVector<llvm::Type *, 6> ArgTys;
570 for (unsigned I = 0; I != PtrArgLoc; ++I)
571 ArgTys.push_back(Elt: CI->getArgOperand(i: I)->getType());
572 ArgTys.push_back(Elt: PtrTy);
573
574 Name = Name + "_" + std::to_string(val: Size);
575 auto *FTy = FunctionType::get(Result: Callee->getReturnType(),
576 Params: ArrayRef<Type *>(ArgTys), isVarArg: false);
577 AMDGPULibFunc NewLibFunc(Name, FTy);
578 FunctionCallee F = AMDGPULibFunc::getOrInsertFunction(M, fInfo: NewLibFunc);
579 if (!F)
580 return false;
581
582 SmallVector<Value *, 6> Args;
583 for (unsigned I = 0; I != PtrArgLoc; ++I)
584 Args.push_back(Elt: CI->getArgOperand(i: I));
585 Args.push_back(Elt: PtrArg);
586
587 auto *NCI = B.CreateCall(Callee: F, Args);
588 NCI->setAttributes(CI->getAttributes());
589 CI->replaceAllUsesWith(V: NCI);
590 CI->dropAllReferences();
591 CI->eraseFromParent();
592
593 return true;
594}
595
596// This function returns false if no change; return true otherwise.
597bool AMDGPULibCalls::fold(CallInst *CI) {
598 Function *Callee = CI->getCalledFunction();
599 // Ignore indirect calls.
600 if (!Callee || Callee->isIntrinsic() || CI->isNoBuiltin())
601 return false;
602
603 FuncInfo FInfo;
604 if (!parseFunctionName(FMangledName: Callee->getName(), FInfo))
605 return false;
606
607 // Further check the number of arguments to see if they match.
608 // TODO: Check calling convention matches too
609 if (!FInfo.isCompatibleSignature(M: *Callee->getParent(), FuncTy: CI->getFunctionType()))
610 return false;
611
612 LLVM_DEBUG(dbgs() << "AMDIC: try folding " << *CI << '\n');
613
614 if (TDOFold(CI, FInfo))
615 return true;
616
617 IRBuilder<> B(CI);
618 if (CI->isStrictFP())
619 B.setIsFPConstrained(true);
620
621 if (FPMathOperator *FPOp = dyn_cast<FPMathOperator>(Val: CI)) {
622 // Under unsafe-math, evaluate calls if possible.
623 // According to Brian Sumner, we can do this for all f32 function calls
624 // using host's double function calls.
625 if (canIncreasePrecisionOfConstantFold(FPOp) && evaluateCall(aCI: CI, FInfo))
626 return true;
627
628 // Copy fast flags from the original call.
629 FastMathFlags FMF = FPOp->getFastMathFlags();
630 B.setFastMathFlags(FMF);
631
632 // Specialized optimizations for each function call.
633 //
634 // TODO: Handle native functions
635 switch (FInfo.getId()) {
636 case AMDGPULibFunc::EI_EXP:
637 if (FMF.none())
638 return false;
639 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::exp,
640 AllowMinSizeF32: FMF.approxFunc());
641 case AMDGPULibFunc::EI_EXP2:
642 if (FMF.none())
643 return false;
644 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::exp2,
645 AllowMinSizeF32: FMF.approxFunc());
646 case AMDGPULibFunc::EI_LOG:
647 if (FMF.none())
648 return false;
649 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::log,
650 AllowMinSizeF32: FMF.approxFunc());
651 case AMDGPULibFunc::EI_LOG2:
652 if (FMF.none())
653 return false;
654 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::log2,
655 AllowMinSizeF32: FMF.approxFunc());
656 case AMDGPULibFunc::EI_LOG10:
657 if (FMF.none())
658 return false;
659 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::log10,
660 AllowMinSizeF32: FMF.approxFunc());
661 case AMDGPULibFunc::EI_FMIN:
662 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::minnum,
663 AllowMinSizeF32: true, AllowF64: true);
664 case AMDGPULibFunc::EI_FMAX:
665 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::maxnum,
666 AllowMinSizeF32: true, AllowF64: true);
667 case AMDGPULibFunc::EI_FMA:
668 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::fma, AllowMinSizeF32: true,
669 AllowF64: true);
670 case AMDGPULibFunc::EI_MAD:
671 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::fmuladd,
672 AllowMinSizeF32: true, AllowF64: true);
673 case AMDGPULibFunc::EI_FABS:
674 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::fabs, AllowMinSizeF32: true,
675 AllowF64: true, AllowStrictFP: true);
676 case AMDGPULibFunc::EI_COPYSIGN:
677 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::copysign,
678 AllowMinSizeF32: true, AllowF64: true, AllowStrictFP: true);
679 case AMDGPULibFunc::EI_FLOOR:
680 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::floor, AllowMinSizeF32: true,
681 AllowF64: true);
682 case AMDGPULibFunc::EI_CEIL:
683 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::ceil, AllowMinSizeF32: true,
684 AllowF64: true);
685 case AMDGPULibFunc::EI_TRUNC:
686 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::trunc, AllowMinSizeF32: true,
687 AllowF64: true);
688 case AMDGPULibFunc::EI_RINT:
689 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::rint, AllowMinSizeF32: true,
690 AllowF64: true);
691 case AMDGPULibFunc::EI_ROUND:
692 return tryReplaceLibcallWithSimpleIntrinsic(B, CI, IntrID: Intrinsic::round, AllowMinSizeF32: true,
693 AllowF64: true);
694 case AMDGPULibFunc::EI_LDEXP: {
695 if (!shouldReplaceLibcallWithIntrinsic(CI, AllowMinSizeF32: true, AllowF64: true))
696 return false;
697
698 Value *Arg1 = CI->getArgOperand(i: 1);
699 if (VectorType *VecTy = dyn_cast<VectorType>(Val: CI->getType());
700 VecTy && !isa<VectorType>(Val: Arg1->getType())) {
701 Value *SplatArg1 = B.CreateVectorSplat(EC: VecTy->getElementCount(), V: Arg1);
702 CI->setArgOperand(i: 1, v: SplatArg1);
703 }
704
705 CI->setCalledFunction(Intrinsic::getOrInsertDeclaration(
706 M: CI->getModule(), id: Intrinsic::ldexp,
707 OverloadTys: {CI->getType(), CI->getArgOperand(i: 1)->getType()}));
708 CI->setCallingConv(CallingConv::C);
709 return true;
710 }
711 case AMDGPULibFunc::EI_POW:
712 case AMDGPULibFunc::EI_POW_FAST:
713 return tryOptimizePow(FPOp, B, FInfo);
714 case AMDGPULibFunc::EI_POWR:
715 case AMDGPULibFunc::EI_POWR_FAST: {
716 if (fold_pow(FPOp, B, FInfo))
717 return true;
718 if (!FMF.approxFunc())
719 return false;
720
721 if (FInfo.getId() == AMDGPULibFunc::EI_POWR && FMF.approxFunc() &&
722 getArgType(FInfo) == AMDGPULibFunc::F32) {
723 Module *M = Callee->getParent();
724 AMDGPULibFunc PowrFastInfo(AMDGPULibFunc::EI_POWR_FAST, FInfo);
725 if (FunctionCallee PowrFastFunc = getFunction(M, fInfo: PowrFastInfo)) {
726 CI->setCalledFunction(PowrFastFunc);
727 return true;
728 }
729 }
730
731 if (!shouldReplaceLibcallWithIntrinsic(CI))
732 return false;
733 return expandFastPow(FPOp, B, Kind: PowKind::PowR);
734 }
735 case AMDGPULibFunc::EI_POWN:
736 case AMDGPULibFunc::EI_POWN_FAST: {
737 if (fold_pow(FPOp, B, FInfo))
738 return true;
739 if (!FMF.approxFunc())
740 return false;
741
742 if (FInfo.getId() == AMDGPULibFunc::EI_POWN &&
743 getArgType(FInfo) == AMDGPULibFunc::F32) {
744 Module *M = Callee->getParent();
745 AMDGPULibFunc PownFastInfo(AMDGPULibFunc::EI_POWN_FAST, FInfo);
746 if (FunctionCallee PownFastFunc = getFunction(M, fInfo: PownFastInfo)) {
747 CI->setCalledFunction(PownFastFunc);
748 return true;
749 }
750 }
751
752 if (!shouldReplaceLibcallWithIntrinsic(CI))
753 return false;
754 return expandFastPow(FPOp, B, Kind: PowKind::PowN);
755 }
756 case AMDGPULibFunc::EI_ROOTN:
757 case AMDGPULibFunc::EI_ROOTN_FAST: {
758 if (fold_rootn(FPOp, B, FInfo))
759 return true;
760 if (!FMF.approxFunc())
761 return false;
762
763 if (getArgType(FInfo) == AMDGPULibFunc::F32) {
764 Module *M = Callee->getParent();
765 AMDGPULibFunc RootnFastInfo(AMDGPULibFunc::EI_ROOTN_FAST, FInfo);
766 if (FunctionCallee RootnFastFunc = getFunction(M, fInfo: RootnFastInfo)) {
767 CI->setCalledFunction(RootnFastFunc);
768 return true;
769 }
770 }
771
772 return expandFastPow(FPOp, B, Kind: PowKind::RootN);
773 }
774 case AMDGPULibFunc::EI_SQRT:
775 // TODO: Allow with strictfp + constrained intrinsic
776 return tryReplaceLibcallWithSimpleIntrinsic(
777 B, CI, IntrID: Intrinsic::sqrt, AllowMinSizeF32: true, AllowF64: true, /*AllowStrictFP=*/false);
778 case AMDGPULibFunc::EI_COS:
779 case AMDGPULibFunc::EI_SIN:
780 return fold_sincos(FPOp, B, FInfo);
781 default:
782 break;
783 }
784 } else {
785 // Specialized optimizations for each function call
786 switch (FInfo.getId()) {
787 case AMDGPULibFunc::EI_READ_PIPE_2:
788 case AMDGPULibFunc::EI_READ_PIPE_4:
789 case AMDGPULibFunc::EI_WRITE_PIPE_2:
790 case AMDGPULibFunc::EI_WRITE_PIPE_4:
791 return fold_read_write_pipe(CI, B, FInfo);
792 default:
793 break;
794 }
795 }
796
797 return false;
798}
799
800static Constant *getConstantFloat(const ArrayRef<APFloat> Values,
801 const Type *Ty) {
802
803 assert(Ty->isSingleValueType() &&
804 "Type must either be a scalar or a vector.");
805 assert((!Ty->isVectorTy() || Ty->isScalableTy() ||
806 Values.size() == cast<FixedVectorType>(Ty)->getNumElements()) &&
807 "Unexpected number of constant values.");
808 assert((Ty->isVectorTy() || Values.size() == 1) &&
809 "Expected exactly one constant value");
810
811 Type *ElemTy = Ty->getScalarType();
812 const fltSemantics &FltSem = ElemTy->getFltSemantics();
813
814 SmallVector<Constant *, 4> ConstValues;
815 ConstValues.reserve(N: Values.size());
816 for (APFloat APF : Values) {
817 bool Unused;
818 APF.convert(ToSemantics: FltSem, RM: APFloat::rmNearestTiesToEven, losesInfo: &Unused);
819 ConstValues.push_back(Elt: ConstantFP::get(Ty: ElemTy, V: APF));
820 }
821
822 return Ty->isVectorTy() ? ConstantVector::get(V: ConstValues) : ConstValues[0];
823}
824
825bool AMDGPULibCalls::TDOFold(CallInst *CI, const FuncInfo &FInfo) {
826 // Table-Driven optimization
827 const TableRef tr = getOptTable(id: FInfo.getId());
828 if (tr.empty())
829 return false;
830
831 int const sz = (int)tr.size();
832 Value *opr0 = CI->getArgOperand(i: 0);
833
834 int vecSize = getVecSize(FInfo);
835 if (vecSize > 1) {
836 // Vector version
837 Constant *CV = dyn_cast<Constant>(Val: opr0);
838 if (CV && CV->getType()->isVectorTy()) {
839 SmallVector<APFloat, 4> Values;
840 Values.reserve(N: vecSize);
841 for (int eltNo = 0; eltNo < vecSize; ++eltNo) {
842 // A lane may be undef or poison, in which case there is nothing to
843 // look up in the table.
844 ConstantFP *eltval = dyn_cast_or_null<ConstantFP>(
845 Val: CV->getAggregateElement(Elt: (unsigned)eltNo));
846 if (!eltval)
847 return false;
848 auto MatchingRow = llvm::find_if(Range: tr, P: [eltval](const TableEntry &entry) {
849 return eltval->isExactlyValue(V: entry.input);
850 });
851 if (MatchingRow == tr.end())
852 return false;
853 Values.push_back(Elt: APFloat(MatchingRow->result));
854 }
855 Constant *NewValues = getConstantFloat(Values, Ty: CI->getType());
856 LLVM_DEBUG(errs() << "AMDIC: " << *CI << " ---> " << *NewValues << "\n");
857 replaceCall(I: CI, With: NewValues);
858 return true;
859 }
860 } else {
861 // Scalar version
862 if (ConstantFP *CF = dyn_cast<ConstantFP>(Val: opr0)) {
863 for (int i = 0; i < sz; ++i) {
864 if (CF->isExactlyValue(V: tr[i].input)) {
865 Value *nval = ConstantFP::get(Ty: CF->getType(), V: tr[i].result);
866 LLVM_DEBUG(errs() << "AMDIC: " << *CI << " ---> " << *nval << "\n");
867 replaceCall(I: CI, With: nval);
868 return true;
869 }
870 }
871 }
872 }
873
874 return false;
875}
876
877namespace llvm {
878static double log2(double V) {
879#if _XOPEN_SOURCE >= 600 || defined(_ISOC99_SOURCE) || _POSIX_C_SOURCE >= 200112L
880 return ::log2(x: V);
881#else
882 return log(V) / numbers::ln2;
883#endif
884}
885} // namespace llvm
886
887bool AMDGPULibCalls::fold_pow(FPMathOperator *FPOp, IRBuilder<> &B,
888 const FuncInfo &FInfo) {
889 assert((FInfo.getId() == AMDGPULibFunc::EI_POW ||
890 FInfo.getId() == AMDGPULibFunc::EI_POW_FAST ||
891 FInfo.getId() == AMDGPULibFunc::EI_POWR ||
892 FInfo.getId() == AMDGPULibFunc::EI_POWR_FAST ||
893 FInfo.getId() == AMDGPULibFunc::EI_POWN ||
894 FInfo.getId() == AMDGPULibFunc::EI_POWN_FAST) &&
895 "fold_pow: encounter a wrong function call");
896
897 Module *M = B.getModule();
898 Type *eltType = FPOp->getType()->getScalarType();
899 Value *opr0 = FPOp->getOperand(i: 0);
900 Value *opr1 = FPOp->getOperand(i: 1);
901
902 const APFloat *CF = nullptr;
903 const APInt *CINT = nullptr;
904 if (!match(V: opr1, P: m_APFloatAllowPoison(Res&: CF)))
905 match(V: opr1, P: m_APIntAllowPoison(Res&: CINT));
906
907 // 0x1111111 means that we don't do anything for this call.
908 int ci_opr1 = (CINT ? (int)CINT->getSExtValue() : 0x1111111);
909
910 // OpenCL powr(x<0, y) = NaN, but the folds below would turn it into a
911 // finite number. Skip them unless NaNs are ignored or the base is known
912 // non-negative.
913 bool IsPowr = FInfo.getId() == AMDGPULibFunc::EI_POWR ||
914 FInfo.getId() == AMDGPULibFunc::EI_POWR_FAST;
915 bool SkipConstantFolds =
916 IsPowr && !FPOp->hasNoNaNs() &&
917 !cannotBeOrderedLessThanZero(
918 V: opr0, SQ: SQ.getWithInstruction(I: cast<Instruction>(Val: FPOp)));
919
920 if (CF && (CF->isExactlyValue(V: 0.5) || CF->isExactlyValue(V: -0.5))) {
921 // pow[r](x, [-]0.5) = sqrt(x) / rsqrt(x)
922 //
923 // sqrt/rsqrt and pow disagree on two negative inputs:
924 // pow(-Inf, 0.5) == +Inf but sqrt(-Inf) == NaN (ninf case)
925 // pow(-0.0, 0.5) == +0.0 but sqrt(-0.0) == -0.0 (nsz case)
926 // powr requires x >= 0 by the OpenCL spec, so -Inf is undefined behaviour
927 // and the ninf check can be skipped for powr/powr_fast. -0.0 is a valid
928 // input for powr since -0.0 >= 0 by IEEE comparison, so nsz is still
929 // required for all variants. sqrt/rsqrt already return NaN for a
930 // negative base like powr does, so this fold skips the base-sign check.
931 if (FPOp->hasNoSignedZeros() && (IsPowr || FPOp->hasNoInfs())) {
932 bool issqrt = CF->isExactlyValue(V: 0.5);
933 if (FunctionCallee FPExpr =
934 getFunction(M, fInfo: AMDGPULibFunc(issqrt ? AMDGPULibFunc::EI_SQRT
935 : AMDGPULibFunc::EI_RSQRT,
936 FInfo))) {
937 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> " << FInfo.getName()
938 << '(' << *opr0 << ")\n");
939 Value *nval = CreateCallEx(B, Callee: FPExpr, Arg: opr0,
940 Name: issqrt ? "__pow2sqrt" : "__pow2rsqrt");
941 replaceCall(I: FPOp, With: nval);
942 return true;
943 }
944 }
945 }
946
947 if (!SkipConstantFolds) {
948 if ((CF && CF->isZero()) || (CINT && ci_opr1 == 0)) {
949 // pow/powr/pown(x, 0) == 1
950 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> 1\n");
951 Constant *cnval = ConstantFP::get(Ty: eltType, V: 1.0);
952 if (getVecSize(FInfo) > 1) {
953 cnval = ConstantDataVector::getSplat(NumElts: getVecSize(FInfo), Elt: cnval);
954 }
955 replaceCall(I: FPOp, With: cnval);
956 return true;
957 }
958 if ((CF && CF->isOne()) || (CINT && ci_opr1 == 1)) {
959 // pow/powr/pown(x, 1.0) = x
960 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> " << *opr0 << "\n");
961 replaceCall(I: FPOp, With: opr0);
962 return true;
963 }
964 if ((CF && CF->isExactlyValue(V: 2.0)) || (CINT && ci_opr1 == 2)) {
965 // pow/powr/pown(x, 2.0) = x*x
966 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> " << *opr0 << " * "
967 << *opr0 << "\n");
968 Value *nval = B.CreateFMul(L: opr0, R: opr0, Name: "__pow2");
969 replaceCall(I: FPOp, With: nval);
970 return true;
971 }
972 if ((CF && CF->isMinusOne()) || (CINT && ci_opr1 == -1)) {
973 // pow/powr/pown(x, -1.0) = 1.0/x
974 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> 1 / " << *opr0 << "\n");
975 Constant *cnval = ConstantFP::get(Ty: eltType, V: 1.0);
976 if (getVecSize(FInfo) > 1) {
977 cnval = ConstantDataVector::getSplat(NumElts: getVecSize(FInfo), Elt: cnval);
978 }
979 Value *nval = B.CreateFDiv(L: cnval, R: opr0, Name: "__powrecip");
980 replaceCall(I: FPOp, With: nval);
981 return true;
982 }
983 }
984
985 if (!isUnsafeFiniteOnlyMath(FPOp))
986 return false;
987
988 // Unsafe Math optimization
989
990 // Remember that ci_opr1 is set if opr1 is integral
991 if (CF) {
992 double dval = (getArgType(FInfo) == AMDGPULibFunc::F32)
993 ? (double)CF->convertToFloat()
994 : CF->convertToDouble();
995 int ival = (int)dval;
996 if ((double)ival == dval) {
997 ci_opr1 = ival;
998 } else
999 ci_opr1 = 0x11111111;
1000 }
1001
1002 // pow/powr/pown(x, c) = [1/](x*x*..x); where
1003 // trunc(c) == c && the number of x == c && |c| <= 12
1004 unsigned abs_opr1 = (ci_opr1 < 0) ? -ci_opr1 : ci_opr1;
1005 if (abs_opr1 <= 12) {
1006 Constant *cnval;
1007 Value *nval;
1008 if (abs_opr1 == 0) {
1009 cnval = ConstantFP::get(Ty: eltType, V: 1.0);
1010 if (getVecSize(FInfo) > 1) {
1011 cnval = ConstantDataVector::getSplat(NumElts: getVecSize(FInfo), Elt: cnval);
1012 }
1013 nval = cnval;
1014 } else {
1015 Value *valx2 = nullptr;
1016 nval = nullptr;
1017 while (abs_opr1 > 0) {
1018 valx2 = valx2 ? B.CreateFMul(L: valx2, R: valx2, Name: "__powx2") : opr0;
1019 if (abs_opr1 & 1) {
1020 nval = nval ? B.CreateFMul(L: nval, R: valx2, Name: "__powprod") : valx2;
1021 }
1022 abs_opr1 >>= 1;
1023 }
1024 }
1025
1026 if (ci_opr1 < 0) {
1027 cnval = ConstantFP::get(Ty: eltType, V: 1.0);
1028 if (getVecSize(FInfo) > 1) {
1029 cnval = ConstantDataVector::getSplat(NumElts: getVecSize(FInfo), Elt: cnval);
1030 }
1031 nval = B.CreateFDiv(L: cnval, R: nval, Name: "__1powprod");
1032 }
1033 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> "
1034 << ((ci_opr1 < 0) ? "1/prod(" : "prod(") << *opr0
1035 << ")\n");
1036 replaceCall(I: FPOp, With: nval);
1037 return true;
1038 }
1039
1040 // If we should use the generic intrinsic instead of emitting a libcall
1041 const bool ShouldUseIntrinsic = eltType->isFloatTy() || eltType->isHalfTy();
1042
1043 // powr ---> exp2(y * log2(x))
1044 // pown/pow ---> powr(fabs(x), y) | (x & ((int)y << 31))
1045 FunctionCallee ExpExpr;
1046 if (ShouldUseIntrinsic)
1047 ExpExpr = Intrinsic::getOrInsertDeclaration(M, id: Intrinsic::exp2,
1048 OverloadTys: {FPOp->getType()});
1049 else {
1050 ExpExpr = getFunction(M, fInfo: AMDGPULibFunc(AMDGPULibFunc::EI_EXP2, FInfo));
1051 if (!ExpExpr)
1052 return false;
1053 }
1054
1055 bool needlog = false;
1056 bool needabs = false;
1057 bool needcopysign = false;
1058 Constant *cnval = nullptr;
1059 if (getVecSize(FInfo) == 1) {
1060 CF = nullptr;
1061 match(V: opr0, P: m_APFloatAllowPoison(Res&: CF));
1062
1063 if (CF) {
1064 double V = (getArgType(FInfo) == AMDGPULibFunc::F32)
1065 ? (double)CF->convertToFloat()
1066 : CF->convertToDouble();
1067
1068 V = log2(V: std::abs(x: V));
1069 cnval = ConstantFP::get(Ty: eltType, V);
1070 needcopysign = (FInfo.getId() != AMDGPULibFunc::EI_POWR &&
1071 FInfo.getId() != AMDGPULibFunc::EI_POWR_FAST) &&
1072 CF->isNegative();
1073 } else {
1074 needlog = true;
1075 needcopysign = needabs = FInfo.getId() != AMDGPULibFunc::EI_POWR &&
1076 FInfo.getId() != AMDGPULibFunc::EI_POWR_FAST;
1077 }
1078 } else {
1079 ConstantDataVector *CDV = dyn_cast<ConstantDataVector>(Val: opr0);
1080
1081 if (!CDV) {
1082 needlog = true;
1083 needcopysign = needabs = FInfo.getId() != AMDGPULibFunc::EI_POWR &&
1084 FInfo.getId() != AMDGPULibFunc::EI_POWR_FAST;
1085 } else {
1086 assert ((int)CDV->getNumElements() == getVecSize(FInfo) &&
1087 "Wrong vector size detected");
1088
1089 SmallVector<double, 0> DVal;
1090 for (int i=0; i < getVecSize(FInfo); ++i) {
1091 double V = CDV->getElementAsAPFloat(i).convertToDouble();
1092 if (V < 0.0) needcopysign = true;
1093 V = log2(V: std::abs(x: V));
1094 DVal.push_back(Elt: V);
1095 }
1096 if (getArgType(FInfo) == AMDGPULibFunc::F32) {
1097 SmallVector<float, 0> FVal;
1098 for (double D : DVal)
1099 FVal.push_back(Elt: (float)D);
1100 ArrayRef<float> tmp(FVal);
1101 cnval = ConstantDataVector::get(Context&: M->getContext(), Elts: tmp);
1102 } else {
1103 ArrayRef<double> tmp(DVal);
1104 cnval = ConstantDataVector::get(Context&: M->getContext(), Elts: tmp);
1105 }
1106 }
1107 }
1108
1109 if (needcopysign && (FInfo.getId() == AMDGPULibFunc::EI_POW ||
1110 FInfo.getId() == AMDGPULibFunc::EI_POW_FAST)) {
1111 // We cannot handle corner cases for a general pow() function, give up
1112 // unless y is a constant integral value. Then proceed as if it were pown.
1113 if (!isKnownIntegral(V: opr1, SQ: SQ.getWithInstruction(I: cast<Instruction>(Val: FPOp)),
1114 FMF: FPOp->getFastMathFlags()))
1115 return false;
1116 }
1117
1118 Value *nval;
1119 if (needabs) {
1120 nval = B.CreateFAbs(V: opr0, FMFSource: nullptr, Name: "__fabs");
1121 } else {
1122 nval = cnval ? cnval : opr0;
1123 }
1124 if (needlog) {
1125 FunctionCallee LogExpr;
1126 if (ShouldUseIntrinsic) {
1127 LogExpr = Intrinsic::getOrInsertDeclaration(M, id: Intrinsic::log2,
1128 OverloadTys: {FPOp->getType()});
1129 } else {
1130 LogExpr = getFunction(M, fInfo: AMDGPULibFunc(AMDGPULibFunc::EI_LOG2, FInfo));
1131 if (!LogExpr)
1132 return false;
1133 }
1134
1135 nval = CreateCallEx(B,Callee: LogExpr, Arg: nval, Name: "__log2");
1136 }
1137
1138 if (FInfo.getId() == AMDGPULibFunc::EI_POWN ||
1139 FInfo.getId() == AMDGPULibFunc::EI_POWN_FAST) {
1140 // convert int(32) to fp(f32 or f64)
1141 opr1 = B.CreateSIToFP(V: opr1, DestTy: nval->getType(), Name: "pownI2F");
1142 }
1143 nval = B.CreateFMul(L: opr1, R: nval, Name: "__ylogx");
1144
1145 CallInst *Exp2Call = CreateCallEx(B, Callee: ExpExpr, Arg: nval, Name: "__exp2");
1146
1147 // TODO: Generalized fpclass logic for pow
1148 FPClassTest KnownNot = FPClassTest::fcNegative;
1149 if (FPOp->hasNoNaNs())
1150 KnownNot |= FPClassTest::fcNan;
1151
1152 Exp2Call->addRetAttr(
1153 Attr: Attribute::getWithNoFPClass(Context&: Exp2Call->getContext(), Mask: KnownNot));
1154 nval = Exp2Call;
1155
1156 if (needcopysign) {
1157 Type* nTyS = B.getIntNTy(N: eltType->getPrimitiveSizeInBits());
1158 Type *nTy = FPOp->getType()->getWithNewType(EltTy: nTyS);
1159 Value *opr_n = FPOp->getOperand(i: 1);
1160 if (opr_n->getType()->getScalarType()->isIntegerTy())
1161 opr_n = B.CreateZExtOrTrunc(V: opr_n, DestTy: nTy, Name: "__ytou");
1162 else
1163 opr_n = B.CreateFPToSI(V: opr1, DestTy: nTy, Name: "__ytou");
1164
1165 unsigned size = nTy->getScalarSizeInBits();
1166 Value *sign = B.CreateShl(LHS: opr_n, RHS: size-1, Name: "__yeven");
1167 sign = B.CreateAnd(LHS: B.CreateBitCast(V: opr0, DestTy: nTy), RHS: sign, Name: "__pow_sign");
1168
1169 nval = B.CreateCopySign(LHS: nval, RHS: B.CreateBitCast(V: sign, DestTy: nval->getType()),
1170 FMFSource: nullptr, Name: "__pow_sign");
1171 }
1172
1173 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> "
1174 << "exp2(" << *opr1 << " * log2(" << *opr0 << "))\n");
1175 replaceCall(I: FPOp, With: nval);
1176
1177 return true;
1178}
1179
1180bool AMDGPULibCalls::fold_rootn(FPMathOperator *FPOp, IRBuilder<> &B,
1181 const FuncInfo &FInfo) {
1182 Value *opr0 = FPOp->getOperand(i: 0);
1183 Value *opr1 = FPOp->getOperand(i: 1);
1184
1185 const APInt *CINT = nullptr;
1186 if (!match(V: opr1, P: m_APIntAllowPoison(Res&: CINT)))
1187 return false;
1188
1189 Function *Parent = B.GetInsertBlock()->getParent();
1190
1191 int ci_opr1 = (int)CINT->getSExtValue();
1192 if (ci_opr1 == 1 && !Parent->hasFnAttribute(Kind: Attribute::StrictFP)) {
1193 // rootn(x, 1) = x
1194 //
1195 // TODO: Insert constrained canonicalize for strictfp case.
1196 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> " << *opr0 << '\n');
1197 replaceCall(I: FPOp, With: opr0);
1198 return true;
1199 }
1200
1201 Module *M = B.getModule();
1202
1203 CallInst *CI = cast<CallInst>(Val: FPOp);
1204
1205 // rootn and sqrt disagree on signed-zero / -Inf inputs (e.g. rootn(-0.0, 2)
1206 // is +0.0, sqrt(-0.0) is -0.0), so require nsz/ninf.
1207 bool FMFOkForSqrt = FPOp->hasNoSignedZeros() && FPOp->hasNoInfs();
1208
1209 if (ci_opr1 == 2 && FMFOkForSqrt &&
1210 shouldReplaceLibcallWithIntrinsic(CI,
1211 /*AllowMinSizeF32=*/true,
1212 /*AllowF64=*/true)) {
1213 // rootn(x, 2) = sqrt(x)
1214 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> sqrt(" << *opr0 << ")\n");
1215
1216 Value *NewCall = B.CreateUnaryIntrinsic(ID: Intrinsic::sqrt, Op: opr0, FMFSource: CI);
1217 NewCall->takeName(V: CI);
1218
1219 // OpenCL rootn has a looser ulp of 2 requirement than sqrt, so add some
1220 // metadata.
1221 MDBuilder MDHelper(M->getContext());
1222 MDNode *FPMD = MDHelper.createFPMath(Accuracy: std::max(a: FPOp->getFPAccuracy(), b: 2.0f));
1223 if (auto *NewCallI = dyn_cast<Instruction>(Val: NewCall))
1224 NewCallI->setMetadata(KindID: LLVMContext::MD_fpmath, Node: FPMD);
1225
1226 replaceCall(I: CI, With: NewCall);
1227 return true;
1228 }
1229
1230 if (ci_opr1 == 3) { // rootn(x, 3) = cbrt(x)
1231 if (FunctionCallee FPExpr =
1232 getFunction(M, fInfo: AMDGPULibFunc(AMDGPULibFunc::EI_CBRT, FInfo))) {
1233 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> cbrt(" << *opr0
1234 << ")\n");
1235 Value *nval = CreateCallEx(B,Callee: FPExpr, Arg: opr0, Name: "__rootn2cbrt");
1236 replaceCall(I: FPOp, With: nval);
1237 return true;
1238 }
1239 } else if (ci_opr1 == -1) { // rootn(x, -1) = 1.0/x
1240 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> 1.0 / " << *opr0 << "\n");
1241 Value *nval = B.CreateFDiv(L: ConstantFP::get(Ty: opr0->getType(), V: 1.0),
1242 R: opr0,
1243 Name: "__rootn2div");
1244 replaceCall(I: FPOp, With: nval);
1245 return true;
1246 }
1247
1248 if (ci_opr1 == -2 && FMFOkForSqrt &&
1249 shouldReplaceLibcallWithIntrinsic(CI,
1250 /*AllowMinSizeF32=*/true,
1251 /*AllowF64=*/true)) {
1252 // rootn(x, -2) = rsqrt(x)
1253
1254 // The original rootn had looser ulp requirements than the resultant sqrt
1255 // and fdiv.
1256 MDBuilder MDHelper(M->getContext());
1257 MDNode *FPMD = MDHelper.createFPMath(Accuracy: std::max(a: FPOp->getFPAccuracy(), b: 2.0f));
1258
1259 // TODO: Could handle strictfp but need to fix strict sqrt emission
1260 FastMathFlags FMF = FPOp->getFastMathFlags();
1261 FMF.setAllowContract(true);
1262
1263 Value *Sqrt = B.CreateUnaryIntrinsic(ID: Intrinsic::sqrt, Op: opr0, FMFSource: CI);
1264 Instruction *RSqrt = cast<Instruction>(
1265 Val: B.CreateFDiv(L: ConstantFP::get(Ty: opr0->getType(), V: 1.0), R: Sqrt));
1266 if (auto *SqrtI = dyn_cast<Instruction>(Val: Sqrt))
1267 SqrtI->setFastMathFlags(FMF);
1268 RSqrt->setFastMathFlags(FMF);
1269 RSqrt->setMetadata(KindID: LLVMContext::MD_fpmath, Node: FPMD);
1270
1271 LLVM_DEBUG(errs() << "AMDIC: " << *FPOp << " ---> rsqrt(" << *opr0
1272 << ")\n");
1273 replaceCall(I: CI, With: RSqrt);
1274 return true;
1275 }
1276
1277 return false;
1278}
1279
1280// is_integer(y) => trunc(y) == y
1281static Value *emitIsInteger(IRBuilder<> &B, Value *Y) {
1282 Value *TruncY = B.CreateUnaryIntrinsic(ID: Intrinsic::trunc, Op: Y);
1283 return B.CreateFCmpOEQ(LHS: TruncY, RHS: Y);
1284}
1285
1286static Value *emitIsEvenInteger(IRBuilder<> &B, Value *Y) {
1287 // Even integers are still integers after division by 2.
1288 auto *HalfY = B.CreateFMul(L: Y, R: ConstantFP::get(Ty: Y->getType(), V: 0.5));
1289 return emitIsInteger(B, Y: HalfY);
1290}
1291
1292// is_odd_integer(y) => is_integer(y) && !is_even_integer(y)
1293static Value *emitIsOddInteger(IRBuilder<> &B, Value *Y) {
1294 Value *IsIntY = emitIsInteger(B, Y);
1295 Value *IsEvenY = emitIsEvenInteger(B, Y);
1296 Value *NotEvenY = B.CreateNot(V: IsEvenY);
1297 return B.CreateAnd(LHS: IsIntY, RHS: NotEvenY);
1298}
1299
1300// isinf(val) => fabs(val) == +inf
1301static Value *emitIsInf(IRBuilder<> &B, Value *val) {
1302 auto *fabsVal = B.CreateFAbs(V: val);
1303 return B.CreateFCmpOEQ(LHS: fabsVal, RHS: ConstantFP::getInfinity(Ty: val->getType()));
1304}
1305
1306// y * log2(fabs(x))
1307static Value *emitFastExpYLnx(IRBuilder<> &B, Value *X, Value *Y) {
1308 Value *AbsX = B.CreateFAbs(V: X);
1309 Value *LogAbsX = B.CreateUnaryIntrinsic(ID: Intrinsic::log2, Op: AbsX);
1310 Value *YTimesLogX = B.CreateFMul(L: Y, R: LogAbsX);
1311 return B.CreateUnaryIntrinsic(ID: Intrinsic::exp2, Op: YTimesLogX);
1312}
1313
1314/// Emit special case management epilog code for fast pow, powr, pown, and rootn
1315/// expansions. \p x and \p y should be the arguments to the library call
1316/// (possibly with some values clamped). \p expylnx should be the result to use
1317/// in normal circumstances.
1318static Value *emitPowFixup(IRBuilder<> &B, Value *X, Value *Y, Value *ExpYLnX,
1319 PowKind Kind) {
1320 Constant *Zero = ConstantFP::getZero(Ty: X->getType());
1321 Constant *One = ConstantFP::get(Ty: X->getType(), V: 1.0);
1322 Constant *QNaN = ConstantFP::getQNaN(Ty: X->getType());
1323 Constant *PInf = ConstantFP::getInfinity(Ty: X->getType());
1324
1325 switch (Kind) {
1326 case PowKind::Pow: {
1327 // is_odd_integer(y)
1328 Value *IsOddY = emitIsOddInteger(B, Y);
1329
1330 // ret = copysign(expylnx, is_odd_y ? x : 1.0f)
1331 Value *SelSign = B.CreateSelect(C: IsOddY, True: X, False: One);
1332 Value *Ret = B.CreateCopySign(LHS: ExpYLnX, RHS: SelSign);
1333
1334 // if (x < 0 && !is_integer(y)) ret = QNAN
1335 Value *IsIntY = emitIsInteger(B, Y);
1336 Value *condNegX = B.CreateFCmpOLT(LHS: X, RHS: Zero);
1337 Value *condNotIntY = B.CreateNot(V: IsIntY);
1338 Value *condNaN = B.CreateAnd(LHS: condNegX, RHS: condNotIntY);
1339 Ret = B.CreateSelect(C: condNaN, True: QNaN, False: Ret);
1340
1341 // if (isinf(ay)) { ... }
1342
1343 // FIXME: Missing backend optimization to save on materialization cost of
1344 // mixed sign constant infinities.
1345 Value *YIsInf = emitIsInf(B, val: Y);
1346
1347 Value *AY = B.CreateFAbs(V: Y);
1348 Value *YIsNegInf = B.CreateFCmpUNE(LHS: Y, RHS: AY);
1349
1350 Value *AX = B.CreateFAbs(V: X);
1351 Value *AxEqOne = B.CreateFCmpOEQ(LHS: AX, RHS: One);
1352 Value *AxLtOne = B.CreateFCmpOLT(LHS: AX, RHS: One);
1353 Value *XorCond = B.CreateXor(LHS: AxLtOne, RHS: YIsNegInf);
1354 Value *SelInf =
1355 B.CreateSelect(C: AxEqOne, True: AX, False: B.CreateSelect(C: XorCond, True: Zero, False: AY));
1356 Ret = B.CreateSelect(C: YIsInf, True: SelInf, False: Ret);
1357
1358 // if (isinf(ax) || x == 0.0f) { ... }
1359 Value *XIsInf = emitIsInf(B, val: X);
1360 Value *XEqZero = B.CreateFCmpOEQ(LHS: X, RHS: Zero);
1361 Value *AxInfOrZero = B.CreateOr(LHS: XIsInf, RHS: XEqZero);
1362 Value *YLtZero = B.CreateFCmpOLT(LHS: Y, RHS: Zero);
1363 Value *XorZeroInf = B.CreateXor(LHS: XEqZero, RHS: YLtZero);
1364 Value *SelVal = B.CreateSelect(C: XorZeroInf, True: Zero, False: PInf);
1365 Value *SelSign2 = B.CreateSelect(C: IsOddY, True: X, False: Zero);
1366 Value *Copysign = B.CreateCopySign(LHS: SelVal, RHS: SelSign2);
1367 Ret = B.CreateSelect(C: AxInfOrZero, True: Copysign, False: Ret);
1368
1369 // if (isunordered(x, y)) ret = QNAN
1370 Value *isUnordered = B.CreateFCmpUNO(LHS: X, RHS: Y);
1371 return B.CreateSelect(C: isUnordered, True: QNaN, False: Ret);
1372 }
1373 case PowKind::PowR: {
1374 Value *YIsNeg = B.CreateFCmpOLT(LHS: Y, RHS: Zero);
1375 Value *IZ = B.CreateSelect(C: YIsNeg, True: PInf, False: Zero);
1376 Value *ZI = B.CreateSelect(C: YIsNeg, True: Zero, False: PInf);
1377
1378 Value *YEqZero = B.CreateFCmpOEQ(LHS: Y, RHS: Zero);
1379 Value *SelZeroCase = B.CreateSelect(C: YEqZero, True: QNaN, False: IZ);
1380 Value *XEqZero = B.CreateFCmpOEQ(LHS: X, RHS: Zero);
1381 Value *Ret = B.CreateSelect(C: XEqZero, True: SelZeroCase, False: ExpYLnX);
1382
1383 Value *XEqInf = B.CreateFCmpOEQ(LHS: X, RHS: PInf);
1384 Value *YNeZero = B.CreateFCmpUNE(LHS: Y, RHS: Zero);
1385 Value *CondInfCase = B.CreateAnd(LHS: XEqInf, RHS: YNeZero);
1386 Ret = B.CreateSelect(C: CondInfCase, True: ZI, False: Ret);
1387
1388 Value *IsInfY = emitIsInf(B, val: Y);
1389 Value *XNeOne = B.CreateFCmpUNE(LHS: X, RHS: One);
1390 Value *CondInfY = B.CreateAnd(LHS: IsInfY, RHS: XNeOne);
1391 Value *XLtOne = B.CreateFCmpOLT(LHS: X, RHS: One);
1392 Value *SelInfYCase = B.CreateSelect(C: XLtOne, True: IZ, False: ZI);
1393 Ret = B.CreateSelect(C: CondInfY, True: SelInfYCase, False: Ret);
1394
1395 Value *IsUnordered = B.CreateFCmpUNO(LHS: X, RHS: Y);
1396 return B.CreateSelect(C: IsUnordered, True: QNaN, False: Ret);
1397 }
1398 case PowKind::PowN: {
1399 Constant *ZeroI = ConstantInt::get(Ty: Y->getType(), V: 0);
1400
1401 // is_odd_y = (ny & 1) != 0
1402 Value *OneI = ConstantInt::get(Ty: Y->getType(), V: 1);
1403 Value *YAnd1 = B.CreateAnd(LHS: Y, RHS: OneI);
1404 Value *IsOddY = B.CreateICmpNE(LHS: YAnd1, RHS: ZeroI);
1405
1406 // ret = copysign(expylnx, is_odd_y ? x : 1.0f)
1407 Value *SelSign = B.CreateSelect(C: IsOddY, True: X, False: One);
1408 Value *Ret = B.CreateCopySign(LHS: ExpYLnX, RHS: SelSign);
1409
1410 // if (isinf(x) || x == 0.0f)
1411 Value *FabsX = B.CreateFAbs(V: X);
1412 Value *XIsInf = B.CreateFCmpOEQ(LHS: FabsX, RHS: PInf);
1413 Value *XEqZero = B.CreateFCmpOEQ(LHS: X, RHS: Zero);
1414 Value *InfOrZero = B.CreateOr(LHS: XIsInf, RHS: XEqZero);
1415
1416 // (x == 0.0f) ^ (ny < 0) ? 0.0f : +inf
1417 Value *YLtZero = B.CreateICmpSLT(LHS: Y, RHS: ZeroI);
1418 Value *XorZeroInf = B.CreateXor(LHS: XEqZero, RHS: YLtZero);
1419 Value *SelVal = B.CreateSelect(C: XorZeroInf, True: Zero, False: PInf);
1420
1421 // copysign(selVal, is_odd_y ? x : 0.0f)
1422 Value *SelSign2 = B.CreateSelect(C: IsOddY, True: X, False: Zero);
1423 Value *Copysign = B.CreateCopySign(LHS: SelVal, RHS: SelSign2);
1424
1425 return B.CreateSelect(C: InfOrZero, True: Copysign, False: Ret);
1426 }
1427 case PowKind::RootN: {
1428 Constant *ZeroI = ConstantInt::get(Ty: Y->getType(), V: 0);
1429
1430 // is_odd_y = (ny & 1) != 0
1431 Value *YAnd1 = B.CreateAnd(LHS: Y, RHS: ConstantInt::get(Ty: Y->getType(), V: 1));
1432 Value *IsOddY = B.CreateICmpNE(LHS: YAnd1, RHS: ZeroI);
1433
1434 // ret = copysign(expylnx, is_odd_y ? x : 1.0f)
1435 Value *SelSign = B.CreateSelect(C: IsOddY, True: X, False: One);
1436 Value *Ret = B.CreateCopySign(LHS: ExpYLnX, RHS: SelSign);
1437
1438 // if (isinf(x) || x == 0.0f)
1439 Value *FabsX = B.CreateFAbs(V: X);
1440 Value *IsInfX = B.CreateFCmpOEQ(LHS: FabsX, RHS: PInf);
1441 Value *XEqZero = B.CreateFCmpOEQ(LHS: X, RHS: Zero);
1442 Value *CondInfOrZero = B.CreateOr(LHS: IsInfX, RHS: XEqZero);
1443
1444 // (x == 0.0f) ^ (ny < 0) ? 0.0f : +inf
1445 Value *YLtZero = B.CreateICmpSLT(LHS: Y, RHS: ZeroI);
1446 Value *XorZeroInf = B.CreateXor(LHS: XEqZero, RHS: YLtZero);
1447 Value *SelVal = B.CreateSelect(C: XorZeroInf, True: Zero, False: PInf);
1448
1449 // copysign(selVal, is_odd_y ? x : 0.0f)
1450 Value *SelSign2 = B.CreateSelect(C: IsOddY, True: X, False: Zero);
1451 Value *Copysign = B.CreateCopySign(LHS: SelVal, RHS: SelSign2);
1452
1453 Ret = B.CreateSelect(C: CondInfOrZero, True: Copysign, False: Ret);
1454
1455 // if ((x < 0.0f && !is_odd_y) || ny == 0) ret = QNAN
1456 Value *XIsNeg = B.CreateFCmpOLT(LHS: X, RHS: Zero);
1457 Value *NotOddY = B.CreateNot(V: IsOddY);
1458 Value *CondNegAndNotOdd = B.CreateAnd(LHS: XIsNeg, RHS: NotOddY);
1459 Value *YEqZero = B.CreateICmpEQ(LHS: Y, RHS: ZeroI);
1460 Value *CondBad = B.CreateOr(LHS: CondNegAndNotOdd, RHS: YEqZero);
1461 return B.CreateSelect(C: CondBad, True: QNaN, False: Ret);
1462 }
1463 }
1464
1465 llvm_unreachable("covered switch");
1466}
1467
1468// TODO: Move the fold_pow folding to sqrt/fdiv here
1469bool AMDGPULibCalls::expandFastPow(FPMathOperator *FPOp, IRBuilder<> &B,
1470 PowKind Kind) {
1471 Type *Ty = FPOp->getType();
1472
1473 // There's currently no reason to do this for half. The correct path is
1474 // promote to float and use the fast float expansion.
1475 //
1476 // TODO: We could move this expansion to lowering to get half pow to work.
1477 if (!Ty->getScalarType()->isFloatTy())
1478 return false;
1479
1480 // TODO: Verify optimization for double and bfloat.
1481 Value *X = FPOp->getOperand(i: 0);
1482 Value *Y = FPOp->getOperand(i: 1);
1483
1484 switch (Kind) {
1485 case PowKind::Pow: {
1486 Constant *One = ConstantFP::get(Ty: X->getType(), V: 1.0);
1487
1488 // if (x == 1.0f) y = 1.0f;
1489 Value *XEqOne = B.CreateFCmpOEQ(LHS: X, RHS: One);
1490 Y = B.CreateSelect(C: XEqOne, True: One, False: Y);
1491
1492 // if (y == 0.0f) x = 1.0f;
1493 Value *YEqZero = B.CreateFCmpOEQ(LHS: Y, RHS: ConstantFP::getZero(Ty: X->getType()));
1494 X = B.CreateSelect(C: YEqZero, True: One, False: X);
1495
1496 Value *ExpYLnX = emitFastExpYLnx(B, X, Y);
1497 Value *Fixed = emitPowFixup(B, X, Y, ExpYLnX, Kind);
1498 replaceCall(I: FPOp, With: Fixed);
1499 return true;
1500 }
1501 case PowKind::PowR: {
1502 Value *NegX = B.CreateFCmpOLT(LHS: X, RHS: ConstantFP::getZero(Ty: X->getType()));
1503 X = B.CreateSelect(C: NegX, True: ConstantFP::getQNaN(Ty: X->getType()), False: X);
1504
1505 Value *ExpYLnX = emitFastExpYLnx(B, X, Y);
1506 Value *Fixed = emitPowFixup(B, X, Y, ExpYLnX, Kind);
1507 replaceCall(I: FPOp, With: Fixed);
1508 return true;
1509 }
1510 case PowKind::PowN: {
1511 // ny == 0
1512 Value *YEqZero = B.CreateICmpEQ(LHS: Y, RHS: ConstantInt::get(Ty: Y->getType(), V: 0));
1513
1514 // x = (ny == 0 ? 1.0f : x)
1515 X = B.CreateSelect(C: YEqZero, True: ConstantFP::get(Ty: X->getType(), V: 1.0), False: X);
1516
1517 Value *CastY = B.CreateSIToFP(V: Y, DestTy: X->getType());
1518 Value *ExpYLnX = emitFastExpYLnx(B, X, Y: CastY);
1519 Value *Fixed = emitPowFixup(B, X, Y, ExpYLnX, Kind);
1520 replaceCall(I: FPOp, With: Fixed);
1521 return true;
1522 }
1523 case PowKind::RootN: {
1524 Value *CastY = B.CreateSIToFP(V: Y, DestTy: X->getType());
1525
1526 // This is afn anyway, so we will turn into rcp.
1527 Value *RcpY = B.CreateFDiv(L: ConstantFP::get(Ty: X->getType(), V: 1.0), R: CastY);
1528
1529 Value *ExpYLnX = emitFastExpYLnx(B, X, Y: RcpY);
1530 Value *Fixed = emitPowFixup(B, X, Y, ExpYLnX, Kind);
1531 replaceCall(I: FPOp, With: Fixed);
1532 return true;
1533 }
1534 }
1535 llvm_unreachable("Unhandled PowKind enum");
1536}
1537
1538bool AMDGPULibCalls::tryOptimizePow(FPMathOperator *FPOp, IRBuilder<> &B,
1539 const FuncInfo &FInfo) {
1540 FastMathFlags FMF = FPOp->getFastMathFlags();
1541 CallInst *Call = cast<CallInst>(Val: FPOp);
1542 Module *M = Call->getModule();
1543
1544 FuncInfo PowrInfo;
1545 AMDGPULibFunc::EFuncId FastPowrFuncId =
1546 FMF.approxFunc() || FInfo.getId() == AMDGPULibFunc::EI_POW_FAST
1547 ? AMDGPULibFunc::EI_POWR_FAST
1548 : AMDGPULibFunc::EI_NONE;
1549 FunctionCallee PowrFunc = getFloatFastVariant(
1550 M, fInfo: FInfo, newInfo&: PowrInfo, NewFunc: AMDGPULibFunc::EI_POWR, FastVariant: FastPowrFuncId);
1551
1552 // TODO: Prefer fast pown to fast powr, but slow powr to slow pown.
1553
1554 // pow(x, y) -> powr(x, y) for x >= -0.0
1555 // TODO: Account for flags on current call
1556 if (PowrFunc && cannotBeOrderedLessThanZero(V: FPOp->getOperand(i: 0),
1557 SQ: SQ.getWithInstruction(I: Call))) {
1558 Call->setCalledFunction(PowrFunc);
1559 return fold_pow(FPOp, B, FInfo: PowrInfo) || true;
1560 }
1561
1562 // pow(x, y) -> pown(x, y) for known integral y
1563 if (isKnownIntegral(V: FPOp->getOperand(i: 1), SQ: SQ.getWithInstruction(I: Call),
1564 FMF: FPOp->getFastMathFlags())) {
1565 FunctionType *PownType = getPownType(FT: Call->getFunctionType());
1566
1567 FuncInfo PownInfo;
1568 AMDGPULibFunc::EFuncId FastPownFuncId =
1569 FMF.approxFunc() || FInfo.getId() == AMDGPULibFunc::EI_POW_FAST
1570 ? AMDGPULibFunc::EI_POWN_FAST
1571 : AMDGPULibFunc::EI_NONE;
1572 FunctionCallee PownFunc = getFloatFastVariant(
1573 M, fInfo: FInfo, newInfo&: PownInfo, NewFunc: AMDGPULibFunc::EI_POWN, FastVariant: FastPownFuncId);
1574
1575 if (PownFunc) {
1576 // TODO: If the incoming integral value is an sitofp/uitofp, it won't
1577 // fold out without a known range. We can probably take the source
1578 // value directly.
1579 Value *CastedArg =
1580 B.CreateFPToSI(V: FPOp->getOperand(i: 1), DestTy: PownType->getParamType(i: 1));
1581 // Have to drop any nofpclass attributes on the original call site.
1582 Call->removeParamAttrs(
1583 ArgNo: 1, AttrsToRemove: AttributeFuncs::typeIncompatible(Ty: CastedArg->getType(),
1584 AS: Call->getParamAttributes(ArgNo: 1)));
1585 Call->setCalledFunction(PownFunc);
1586 Call->setArgOperand(i: 1, v: CastedArg);
1587 return fold_pow(FPOp, B, FInfo: PownInfo) || true;
1588 }
1589 }
1590
1591 if (fold_pow(FPOp, B, FInfo))
1592 return true;
1593
1594 if (!FMF.approxFunc())
1595 return false;
1596
1597 if (FInfo.getId() == AMDGPULibFunc::EI_POW && FMF.approxFunc() &&
1598 getArgType(FInfo) == AMDGPULibFunc::F32) {
1599 AMDGPULibFunc PowFastInfo(AMDGPULibFunc::EI_POW_FAST, FInfo);
1600 if (FunctionCallee PowFastFunc = getFunction(M, fInfo: PowFastInfo)) {
1601 Call->setCalledFunction(PowFastFunc);
1602 return fold_pow(FPOp, B, FInfo: PowFastInfo) || true;
1603 }
1604 }
1605
1606 return expandFastPow(FPOp, B, Kind: PowKind::Pow);
1607}
1608
1609// Some library calls are just wrappers around llvm intrinsics, but compiled
1610// conservatively. Preserve the flags from the original call site by
1611// substituting them with direct calls with all the flags.
1612bool AMDGPULibCalls::shouldReplaceLibcallWithIntrinsic(const CallInst *CI,
1613 bool AllowMinSizeF32,
1614 bool AllowF64,
1615 bool AllowStrictFP) {
1616 Type *FltTy = CI->getType()->getScalarType();
1617 const bool IsF32 = FltTy->isFloatTy();
1618
1619 // f64 intrinsics aren't implemented for most operations.
1620 if (!IsF32 && !FltTy->isHalfTy() && (!AllowF64 || !FltTy->isDoubleTy()))
1621 return false;
1622
1623 // We're implicitly inlining by replacing the libcall with the intrinsic, so
1624 // don't do it for noinline call sites.
1625 if (CI->isNoInline())
1626 return false;
1627
1628 const Function *ParentF = CI->getFunction();
1629 // TODO: Handle strictfp
1630 if (!AllowStrictFP && ParentF->hasFnAttribute(Kind: Attribute::StrictFP))
1631 return false;
1632
1633 if (IsF32 && !AllowMinSizeF32 && ParentF->hasMinSize())
1634 return false;
1635 return true;
1636}
1637
1638void AMDGPULibCalls::replaceLibCallWithSimpleIntrinsic(IRBuilder<> &B,
1639 CallInst *CI,
1640 Intrinsic::ID IntrID) {
1641 if (CI->arg_size() == 2) {
1642 Value *Arg0 = CI->getArgOperand(i: 0);
1643 Value *Arg1 = CI->getArgOperand(i: 1);
1644 VectorType *Arg0VecTy = dyn_cast<VectorType>(Val: Arg0->getType());
1645 VectorType *Arg1VecTy = dyn_cast<VectorType>(Val: Arg1->getType());
1646 if (Arg0VecTy && !Arg1VecTy) {
1647 Value *SplatRHS = B.CreateVectorSplat(EC: Arg0VecTy->getElementCount(), V: Arg1);
1648 CI->setArgOperand(i: 1, v: SplatRHS);
1649 } else if (!Arg0VecTy && Arg1VecTy) {
1650 Value *SplatLHS = B.CreateVectorSplat(EC: Arg1VecTy->getElementCount(), V: Arg0);
1651 CI->setArgOperand(i: 0, v: SplatLHS);
1652 }
1653 }
1654
1655 CI->setCalledFunction(Intrinsic::getOrInsertDeclaration(
1656 M: CI->getModule(), id: IntrID, OverloadTys: {CI->getType()}));
1657 CI->setCallingConv(CallingConv::C);
1658}
1659
1660bool AMDGPULibCalls::tryReplaceLibcallWithSimpleIntrinsic(
1661 IRBuilder<> &B, CallInst *CI, Intrinsic::ID IntrID, bool AllowMinSizeF32,
1662 bool AllowF64, bool AllowStrictFP) {
1663 if (!shouldReplaceLibcallWithIntrinsic(CI, AllowMinSizeF32, AllowF64,
1664 AllowStrictFP))
1665 return false;
1666 replaceLibCallWithSimpleIntrinsic(B, CI, IntrID);
1667 return true;
1668}
1669
1670std::tuple<Value *, Value *, Value *>
1671AMDGPULibCalls::insertSinCos(Value *Arg, FastMathFlags FMF, IRBuilder<> &B,
1672 FunctionCallee Fsincos) {
1673 DebugLoc DL = B.getCurrentDebugLocation();
1674 Function *F = B.GetInsertBlock()->getParent();
1675 B.SetInsertPointPastAllocas(F);
1676
1677 AllocaInst *Alloc = B.CreateAlloca(Ty: Arg->getType(), ArraySize: nullptr, Name: "__sincos_");
1678
1679 if (Instruction *ArgInst = dyn_cast<Instruction>(Val: Arg)) {
1680 // If the argument is an instruction, it must dominate all uses so put our
1681 // sincos call there. Otherwise, right after the allocas works well enough
1682 // if it's an argument or constant.
1683
1684 B.SetInsertPoint(*ArgInst->getInsertionPointAfterDef());
1685
1686 // SetInsertPoint unwelcomely always tries to set the debug loc.
1687 B.SetCurrentDebugLocation(DL);
1688 }
1689
1690 Type *CosPtrTy = Fsincos.getFunctionType()->getParamType(i: 1);
1691
1692 // The allocaInst allocates the memory in private address space. This need
1693 // to be addrspacecasted to point to the address space of cos pointer type.
1694 // In OpenCL 2.0 this is generic, while in 1.2 that is private.
1695 Value *CastAlloc = B.CreateAddrSpaceCast(V: Alloc, DestTy: CosPtrTy);
1696
1697 CallInst *SinCos = CreateCallEx2(B, Callee: Fsincos, Arg1: Arg, Arg2: CastAlloc);
1698
1699 // TODO: Is it worth trying to preserve the location for the cos calls for the
1700 // load?
1701
1702 LoadInst *LoadCos = B.CreateLoad(Ty: Arg->getType(), Ptr: Alloc);
1703 return {SinCos, LoadCos, SinCos};
1704}
1705
1706// fold sin, cos -> sincos.
1707bool AMDGPULibCalls::fold_sincos(FPMathOperator *FPOp, IRBuilder<> &B,
1708 const FuncInfo &fInfo) {
1709 assert(fInfo.getId() == AMDGPULibFunc::EI_SIN ||
1710 fInfo.getId() == AMDGPULibFunc::EI_COS);
1711
1712 if ((getArgType(FInfo: fInfo) != AMDGPULibFunc::F32 &&
1713 getArgType(FInfo: fInfo) != AMDGPULibFunc::F64) ||
1714 fInfo.getPrefix() != AMDGPULibFunc::NOPFX)
1715 return false;
1716
1717 bool const isSin = fInfo.getId() == AMDGPULibFunc::EI_SIN;
1718
1719 Value *CArgVal = FPOp->getOperand(i: 0);
1720
1721 // TODO: Constant fold the call
1722 if (isa<ConstantData>(Val: CArgVal))
1723 return false;
1724
1725 CallInst *CI = cast<CallInst>(Val: FPOp);
1726
1727 Function *F = B.GetInsertBlock()->getParent();
1728 Module *M = F->getParent();
1729
1730 // Merge the sin and cos. For OpenCL 2.0, there may only be a generic pointer
1731 // implementation. Prefer the private form if available.
1732 AMDGPULibFunc SinCosLibFuncPrivate(AMDGPULibFunc::EI_SINCOS, fInfo);
1733 SinCosLibFuncPrivate.getLeads()[0].PtrKind =
1734 AMDGPULibFunc::getEPtrKindFromAddrSpace(AS: AMDGPUAS::PRIVATE_ADDRESS);
1735
1736 AMDGPULibFunc SinCosLibFuncGeneric(AMDGPULibFunc::EI_SINCOS, fInfo);
1737 SinCosLibFuncGeneric.getLeads()[0].PtrKind =
1738 AMDGPULibFunc::getEPtrKindFromAddrSpace(AS: AMDGPUAS::FLAT_ADDRESS);
1739
1740 FunctionCallee FSinCosPrivate = getFunction(M, fInfo: SinCosLibFuncPrivate);
1741 FunctionCallee FSinCosGeneric = getFunction(M, fInfo: SinCosLibFuncGeneric);
1742 FunctionCallee FSinCos = FSinCosPrivate ? FSinCosPrivate : FSinCosGeneric;
1743 if (!FSinCos)
1744 return false;
1745
1746 SmallVector<CallInst *> SinCalls;
1747 SmallVector<CallInst *> CosCalls;
1748 SmallVector<CallInst *> SinCosCalls;
1749 FuncInfo PartnerInfo(isSin ? AMDGPULibFunc::EI_COS : AMDGPULibFunc::EI_SIN,
1750 fInfo);
1751 const std::string PairName = PartnerInfo.mangle();
1752
1753 StringRef SinName = isSin ? CI->getCalledFunction()->getName() : PairName;
1754 StringRef CosName = isSin ? PairName : CI->getCalledFunction()->getName();
1755 const std::string SinCosPrivateName = SinCosLibFuncPrivate.mangle();
1756 const std::string SinCosGenericName = SinCosLibFuncGeneric.mangle();
1757
1758 // Intersect the two sets of flags.
1759 FastMathFlags FMF = FPOp->getFastMathFlags();
1760 MDNode *FPMath = CI->getMetadata(KindID: LLVMContext::MD_fpmath);
1761
1762 SmallVector<DILocation *> MergeDbgLocs = {CI->getDebugLoc()};
1763
1764 for (User* U : CArgVal->users()) {
1765 CallInst *XI = dyn_cast<CallInst>(Val: U);
1766 if (!XI || XI->getFunction() != F || XI->isNoBuiltin())
1767 continue;
1768
1769 Function *UCallee = XI->getCalledFunction();
1770 if (!UCallee)
1771 continue;
1772
1773 bool Handled = true;
1774
1775 if (UCallee->getName() == SinName)
1776 SinCalls.push_back(Elt: XI);
1777 else if (UCallee->getName() == CosName)
1778 CosCalls.push_back(Elt: XI);
1779 else if (UCallee->getName() == SinCosPrivateName ||
1780 UCallee->getName() == SinCosGenericName)
1781 SinCosCalls.push_back(Elt: XI);
1782 else
1783 Handled = false;
1784
1785 if (Handled) {
1786 MergeDbgLocs.push_back(Elt: XI->getDebugLoc());
1787 auto *OtherOp = cast<FPMathOperator>(Val: XI);
1788 FMF &= OtherOp->getFastMathFlags();
1789 FPMath = MDNode::getMostGenericFPMath(
1790 A: FPMath, B: XI->getMetadata(KindID: LLVMContext::MD_fpmath));
1791 }
1792 }
1793
1794 if (SinCalls.empty() || CosCalls.empty())
1795 return false;
1796
1797 // insertSinCos needs an insertion point after the argument's def.
1798 if (auto *ArgInst = dyn_cast<Instruction>(Val: CArgVal);
1799 ArgInst && !ArgInst->getInsertionPointAfterDef())
1800 return false;
1801
1802 B.setFastMathFlags(FMF);
1803 B.setDefaultFPMathTag(FPMath);
1804 DILocation *DbgLoc = DILocation::getMergedLocations(Locs: MergeDbgLocs);
1805 B.SetCurrentDebugLocation(DbgLoc);
1806
1807 auto [Sin, Cos, SinCos] = insertSinCos(Arg: CArgVal, FMF, B, Fsincos: FSinCos);
1808
1809 auto replaceTrigInsts = [](ArrayRef<CallInst *> Calls, Value *Res) {
1810 for (CallInst *C : Calls)
1811 C->replaceAllUsesWith(V: Res);
1812
1813 // Leave the other dead instructions to avoid clobbering iterators.
1814 };
1815
1816 replaceTrigInsts(SinCalls, Sin);
1817 replaceTrigInsts(CosCalls, Cos);
1818 replaceTrigInsts(SinCosCalls, SinCos);
1819
1820 // It's safe to delete the original now.
1821 CI->eraseFromParent();
1822 return true;
1823}
1824
1825bool AMDGPULibCalls::evaluateScalarMathFunc(const FuncInfo &FInfo,
1826 APFloat &Res0, APFloat &Res1,
1827 Constant *copr0, Constant *copr1) {
1828 // Every function handled below reads its first operand as a floating-point
1829 // value. Refuse anything else, e.g. a poison vector lane: silently treating
1830 // it as 0.0 misfolds the whole call.
1831 ConstantFP *fpopr0 = dyn_cast_or_null<ConstantFP>(Val: copr0);
1832 if (!fpopr0)
1833 return false;
1834
1835 double opr0 = (getArgType(FInfo) == AMDGPULibFunc::F64)
1836 ? fpopr0->getValueAPF().convertToDouble()
1837 : (double)fpopr0->getValueAPF().convertToFloat();
1838
1839 switch (FInfo.getId()) {
1840 default:
1841 return false;
1842
1843 case AMDGPULibFunc::EI_ACOS:
1844 Res0 = APFloat{acos(x: opr0)};
1845 return true;
1846
1847 case AMDGPULibFunc::EI_ACOSH:
1848 // acosh(x) == log(x + sqrt(x*x - 1))
1849 Res0 = APFloat{log(x: opr0 + sqrt(x: opr0 * opr0 - 1.0))};
1850 return true;
1851
1852 case AMDGPULibFunc::EI_ACOSPI:
1853 Res0 = APFloat{acos(x: opr0) / MATH_PI};
1854 return true;
1855
1856 case AMDGPULibFunc::EI_ASIN:
1857 Res0 = APFloat{asin(x: opr0)};
1858 return true;
1859
1860 case AMDGPULibFunc::EI_ASINH:
1861 // asinh(x) == log(x + sqrt(x*x + 1))
1862 Res0 = APFloat{log(x: opr0 + sqrt(x: opr0 * opr0 + 1.0))};
1863 return true;
1864
1865 case AMDGPULibFunc::EI_ASINPI:
1866 Res0 = APFloat{asin(x: opr0) / MATH_PI};
1867 return true;
1868
1869 case AMDGPULibFunc::EI_ATAN:
1870 Res0 = APFloat{atan(x: opr0)};
1871 return true;
1872
1873 case AMDGPULibFunc::EI_ATANH:
1874 // atanh(x) == (log(x+1) - log(x-1))/2;
1875 Res0 = APFloat{(log(x: opr0 + 1.0) - log(x: opr0 - 1.0)) / 2.0};
1876 return true;
1877
1878 case AMDGPULibFunc::EI_ATANPI:
1879 Res0 = APFloat{atan(x: opr0) / MATH_PI};
1880 return true;
1881
1882 case AMDGPULibFunc::EI_CBRT:
1883 Res0 =
1884 APFloat{(opr0 < 0.0) ? -pow(x: -opr0, y: 1.0 / 3.0) : pow(x: opr0, y: 1.0 / 3.0)};
1885 return true;
1886
1887 case AMDGPULibFunc::EI_COS:
1888 Res0 = APFloat{cos(x: opr0)};
1889 return true;
1890
1891 case AMDGPULibFunc::EI_COSH:
1892 Res0 = APFloat{cosh(x: opr0)};
1893 return true;
1894
1895 case AMDGPULibFunc::EI_COSPI:
1896 Res0 = APFloat{cos(MATH_PI * opr0)};
1897 return true;
1898
1899 case AMDGPULibFunc::EI_EXP:
1900 Res0 = APFloat{std::exp(x: opr0)};
1901 return true;
1902
1903 case AMDGPULibFunc::EI_EXP2:
1904 Res0 = APFloat{pow(x: 2.0, y: opr0)};
1905 return true;
1906
1907 case AMDGPULibFunc::EI_EXP10:
1908 Res0 = APFloat{pow(x: 10.0, y: opr0)};
1909 return true;
1910
1911 case AMDGPULibFunc::EI_LOG:
1912 Res0 = APFloat{log(x: opr0)};
1913 return true;
1914
1915 case AMDGPULibFunc::EI_LOG2:
1916 Res0 = APFloat{log(x: opr0) / log(x: 2.0)};
1917 return true;
1918
1919 case AMDGPULibFunc::EI_LOG10:
1920 Res0 = APFloat{log(x: opr0) / log(x: 10.0)};
1921 return true;
1922
1923 case AMDGPULibFunc::EI_RSQRT:
1924 Res0 = APFloat{1.0 / sqrt(x: opr0)};
1925 return true;
1926
1927 case AMDGPULibFunc::EI_SIN:
1928 Res0 = APFloat{sin(x: opr0)};
1929 return true;
1930
1931 case AMDGPULibFunc::EI_SINH:
1932 Res0 = APFloat{sinh(x: opr0)};
1933 return true;
1934
1935 case AMDGPULibFunc::EI_SINPI:
1936 Res0 = APFloat{sin(MATH_PI * opr0)};
1937 return true;
1938
1939 case AMDGPULibFunc::EI_TAN:
1940 Res0 = APFloat{tan(x: opr0)};
1941 return true;
1942
1943 case AMDGPULibFunc::EI_TANH:
1944 Res0 = APFloat{tanh(x: opr0)};
1945 return true;
1946
1947 case AMDGPULibFunc::EI_TANPI:
1948 Res0 = APFloat{tan(MATH_PI * opr0)};
1949 return true;
1950
1951 // two-arg functions
1952 case AMDGPULibFunc::EI_POW:
1953 case AMDGPULibFunc::EI_POWR: {
1954 ConstantFP *fpopr1 = dyn_cast_or_null<ConstantFP>(Val: copr1);
1955 if (!fpopr1)
1956 return false;
1957 double opr1 = (getArgType(FInfo) == AMDGPULibFunc::F64)
1958 ? fpopr1->getValueAPF().convertToDouble()
1959 : (double)fpopr1->getValueAPF().convertToFloat();
1960 Res0 = APFloat{pow(x: opr0, y: opr1)};
1961 return true;
1962 }
1963
1964 case AMDGPULibFunc::EI_POWN: {
1965 if (ConstantInt *iopr1 = dyn_cast_or_null<ConstantInt>(Val: copr1)) {
1966 double val = (double)iopr1->getSExtValue();
1967 Res0 = APFloat{pow(x: opr0, y: val)};
1968 return true;
1969 }
1970 return false;
1971 }
1972
1973 case AMDGPULibFunc::EI_ROOTN: {
1974 if (ConstantInt *iopr1 = dyn_cast_or_null<ConstantInt>(Val: copr1)) {
1975 double val = (double)iopr1->getSExtValue();
1976 Res0 = APFloat{pow(x: opr0, y: 1.0 / val)};
1977 return true;
1978 }
1979 return false;
1980 }
1981
1982 // with ptr arg
1983 case AMDGPULibFunc::EI_SINCOS:
1984 Res0 = APFloat{sin(x: opr0)};
1985 Res1 = APFloat{cos(x: opr0)};
1986 return true;
1987 }
1988
1989 return false;
1990}
1991
1992bool AMDGPULibCalls::evaluateCall(CallInst *aCI, const FuncInfo &FInfo) {
1993 int numArgs = (int)aCI->arg_size();
1994 if (numArgs > 3)
1995 return false;
1996
1997 Constant *copr0 = nullptr;
1998 Constant *copr1 = nullptr;
1999 if (numArgs > 0) {
2000 if ((copr0 = dyn_cast<Constant>(Val: aCI->getArgOperand(i: 0))) == nullptr)
2001 return false;
2002 }
2003
2004 if (numArgs > 1) {
2005 if ((copr1 = dyn_cast<Constant>(Val: aCI->getArgOperand(i: 1))) == nullptr) {
2006 if (FInfo.getId() != AMDGPULibFunc::EI_SINCOS)
2007 return false;
2008 }
2009 }
2010
2011 // At this point, all arguments to aCI are constants.
2012
2013 // max vector size is 16, and sincos will generate two results.
2014 SmallVector<APFloat, 16> Val0, Val1;
2015 int FuncVecSize = getVecSize(FInfo);
2016 if (FuncVecSize == 1) {
2017 if (!evaluateScalarMathFunc(FInfo, Res0&: Val0.emplace_back(Args: 0.0),
2018 Res1&: Val1.emplace_back(Args: 0.0), copr0, copr1)) {
2019 return false;
2020 }
2021 } else {
2022 // An operand of a vector variant is not necessarily a vector: sincos takes
2023 // a pointer as its second operand, and fmin/fmax/ldexp accept an
2024 // implicitly splatted scalar. Only index into actual vectors.
2025 Constant *CV0 = copr0 && copr0->getType()->isVectorTy() ? copr0 : nullptr;
2026 Constant *CV1 = copr1 && copr1->getType()->isVectorTy() ? copr1 : nullptr;
2027 for (int i = 0; i < FuncVecSize; ++i) {
2028 Constant *celt0 = CV0 ? CV0->getAggregateElement(Elt: (unsigned)i) : nullptr;
2029 Constant *celt1 = CV1 ? CV1->getAggregateElement(Elt: (unsigned)i) : nullptr;
2030 if (!evaluateScalarMathFunc(FInfo, Res0&: Val0.emplace_back(Args: 0.0),
2031 Res1&: Val1.emplace_back(Args: 0.0), copr0: celt0, copr1: celt1)) {
2032 return false;
2033 }
2034 }
2035 }
2036
2037 Constant *nval0 = getConstantFloat(Values: Val0, Ty: aCI->getType());
2038
2039 // sincos
2040 if (FInfo.getId() == AMDGPULibFunc::EI_SINCOS) {
2041 Constant *nval1 = getConstantFloat(Values: Val1, Ty: aCI->getType());
2042 new StoreInst(nval1, aCI->getArgOperand(i: 1), aCI->getIterator());
2043 }
2044
2045 replaceCall(I: aCI, With: nval0);
2046 return true;
2047}
2048
2049PreservedAnalyses AMDGPUSimplifyLibCallsPass::run(Function &F,
2050 FunctionAnalysisManager &AM) {
2051 AMDGPULibCalls Simplifier(F, AM);
2052 Simplifier.initNativeFuncs();
2053
2054 bool Changed = false;
2055
2056 LLVM_DEBUG(dbgs() << "AMDIC: process function ";
2057 F.printAsOperand(dbgs(), false, F.getParent()); dbgs() << '\n';);
2058
2059 for (auto &BB : F) {
2060 for (BasicBlock::iterator I = BB.begin(), E = BB.end(); I != E;) {
2061 // Ignore non-calls.
2062 CallInst *CI = dyn_cast<CallInst>(Val&: I);
2063 ++I;
2064
2065 if (CI) {
2066 if (Simplifier.fold(CI))
2067 Changed = true;
2068 }
2069 }
2070 }
2071 return Changed ? PreservedAnalyses::none() : PreservedAnalyses::all();
2072}
2073
2074PreservedAnalyses AMDGPUUseNativeCallsPass::run(Function &F,
2075 FunctionAnalysisManager &AM) {
2076 if (UseNative.empty())
2077 return PreservedAnalyses::all();
2078
2079 AMDGPULibCalls Simplifier(F, AM);
2080 Simplifier.initNativeFuncs();
2081
2082 bool Changed = false;
2083 for (auto &BB : F) {
2084 for (BasicBlock::iterator I = BB.begin(), E = BB.end(); I != E;) {
2085 // Ignore non-calls.
2086 CallInst *CI = dyn_cast<CallInst>(Val&: I);
2087 ++I;
2088 if (CI && Simplifier.useNative(aCI: CI))
2089 Changed = true;
2090 }
2091 }
2092 return Changed ? PreservedAnalyses::none() : PreservedAnalyses::all();
2093}
2094