1//===- LowerExpectIntrinsic.cpp - Lower expect intrinsic ------------------===//
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// This pass lowers the 'expect' intrinsic to LLVM metadata.
10//
11//===----------------------------------------------------------------------===//
12
13#include "llvm/Transforms/Scalar/LowerExpectIntrinsic.h"
14#include "ScalarOptions.h"
15#include "llvm/ADT/SmallVector.h"
16#include "llvm/ADT/Statistic.h"
17#include "llvm/IR/BasicBlock.h"
18#include "llvm/IR/Constants.h"
19#include "llvm/IR/Function.h"
20#include "llvm/IR/Instructions.h"
21#include "llvm/IR/Intrinsics.h"
22#include "llvm/IR/LLVMContext.h"
23#include "llvm/IR/MDBuilder.h"
24#include "llvm/IR/ProfDataUtils.h"
25#include "llvm/Transforms/Utils/MisExpect.h"
26
27#include <cmath>
28
29using namespace llvm;
30
31#define DEBUG_TYPE "lower-expect-intrinsic"
32
33STATISTIC(ExpectIntrinsicsHandled,
34 "Number of 'expect' intrinsic instructions handled");
35
36// These default values are chosen to represent an extremely skewed outcome for
37// a condition, but they leave some room for interpretation by later passes.
38//
39// If the documentation for __builtin_expect() was made explicit that it should
40// only be used in extreme cases, we could make this ratio higher. As it stands,
41// programmers may be using __builtin_expect() / llvm.expect to annotate that a
42// branch is likely or unlikely to be taken.
43
44static std::tuple<uint32_t, uint32_t>
45getBranchWeight(Intrinsic::ID IntrinsicID, CallInst *CI, int BranchCount) {
46 const ScalarOptions &Opts = ScalarOptions::Global;
47 if (IntrinsicID == Intrinsic::expect) {
48 // __builtin_expect
49 return std::make_tuple(args: Opts.likely_branch_weight,
50 args: Opts.unlikely_branch_weight);
51 } else {
52 // __builtin_expect_with_probability
53 assert(CI->getNumOperands() >= 3 &&
54 "expect with probability must have 3 arguments");
55 auto *Confidence = cast<ConstantFP>(Val: CI->getArgOperand(i: 2));
56 double TrueProb = Confidence->getValueAPF().convertToDouble();
57 assert((TrueProb >= 0.0 && TrueProb <= 1.0) &&
58 "probability value must be in the range [0.0, 1.0]");
59 double FalseProb = (1.0 - TrueProb) / (BranchCount - 1);
60 uint32_t LikelyBW = ceil(x: (TrueProb * (double)(INT32_MAX - 1)) + 1.0);
61 uint32_t UnlikelyBW = ceil(x: (FalseProb * (double)(INT32_MAX - 1)) + 1.0);
62 return std::make_tuple(args&: LikelyBW, args&: UnlikelyBW);
63 }
64}
65
66static bool handleSwitchExpect(SwitchInst &SI) {
67 CallInst *CI = dyn_cast<CallInst>(Val: SI.getCondition());
68 if (!CI)
69 return false;
70
71 Function *Fn = CI->getCalledFunction();
72 if (!Fn || (Fn->getIntrinsicID() != Intrinsic::expect &&
73 Fn->getIntrinsicID() != Intrinsic::expect_with_probability))
74 return false;
75
76 Value *ArgValue = CI->getArgOperand(i: 0);
77 ConstantInt *ExpectedValue = dyn_cast<ConstantInt>(Val: CI->getArgOperand(i: 1));
78 if (!ExpectedValue)
79 return false;
80
81 SwitchInst::CaseHandle Case = *SI.findCaseValue(C: ExpectedValue);
82 unsigned n = SI.getNumCases(); // +1 for default case.
83 uint32_t LikelyBranchWeightVal, UnlikelyBranchWeightVal;
84 std::tie(args&: LikelyBranchWeightVal, args&: UnlikelyBranchWeightVal) =
85 getBranchWeight(IntrinsicID: Fn->getIntrinsicID(), CI, BranchCount: n + 1);
86
87 SmallVector<uint32_t, 16> Weights(n + 1, UnlikelyBranchWeightVal);
88
89 uint64_t Index = (Case == *SI.case_default()) ? 0 : Case.getCaseIndex() + 1;
90 Weights[Index] = LikelyBranchWeightVal;
91
92 misexpect::checkExpectAnnotations(I: SI, ExistingWeights: Weights, /*IsFrontend=*/true);
93
94 SI.setCondition(ArgValue);
95 setBranchWeights(I&: SI, Weights, /*IsExpected=*/true);
96 return true;
97}
98
99/// Handler for PHINodes that define the value argument to an
100/// @llvm.expect call.
101///
102/// If the operand of the phi has a constant value and it 'contradicts'
103/// with the expected value of phi def, then the corresponding incoming
104/// edge of the phi is unlikely to be taken. Using that information,
105/// the branch probability info for the originating branch can be inferred.
106static void handlePhiDef(CallInst *Expect) {
107 Value &Arg = *Expect->getArgOperand(i: 0);
108 ConstantInt *ExpectedValue = dyn_cast<ConstantInt>(Val: Expect->getArgOperand(i: 1));
109 if (!ExpectedValue)
110 return;
111 const APInt &ExpectedPhiValue = ExpectedValue->getValue();
112 bool ExpectedValueIsLikely = true;
113 Function *Fn = Expect->getCalledFunction();
114 // If the function is expect_with_probability, then we need to take the
115 // probability into consideration. For example, in
116 // expect.with.probability.i64(i64 %a, i64 1, double 0.0), the
117 // "ExpectedValue" 1 is unlikely. This affects probability propagation later.
118 if (Fn->getIntrinsicID() == Intrinsic::expect_with_probability) {
119 auto *Confidence = cast<ConstantFP>(Val: Expect->getArgOperand(i: 2));
120 double TrueProb = Confidence->getValueAPF().convertToDouble();
121 ExpectedValueIsLikely = (TrueProb > 0.5);
122 }
123
124 // Walk up in backward a list of instructions that
125 // have 'copy' semantics by 'stripping' the copies
126 // until a PHI node or an instruction of unknown kind
127 // is reached. Negation via xor is also handled.
128 //
129 // C = PHI(...);
130 // B = C;
131 // A = B;
132 // D = __builtin_expect(A, 0);
133 //
134 Value *V = &Arg;
135 SmallVector<Instruction *, 4> Operations;
136 while (!isa<PHINode>(Val: V)) {
137 if (ZExtInst *ZExt = dyn_cast<ZExtInst>(Val: V)) {
138 V = ZExt->getOperand(i_nocapture: 0);
139 Operations.push_back(Elt: ZExt);
140 continue;
141 }
142
143 if (SExtInst *SExt = dyn_cast<SExtInst>(Val: V)) {
144 V = SExt->getOperand(i_nocapture: 0);
145 Operations.push_back(Elt: SExt);
146 continue;
147 }
148
149 BinaryOperator *BinOp = dyn_cast<BinaryOperator>(Val: V);
150 if (!BinOp || BinOp->getOpcode() != Instruction::Xor)
151 return;
152
153 ConstantInt *CInt = dyn_cast<ConstantInt>(Val: BinOp->getOperand(i_nocapture: 1));
154 if (!CInt)
155 return;
156
157 V = BinOp->getOperand(i_nocapture: 0);
158 Operations.push_back(Elt: BinOp);
159 }
160
161 // Executes the recorded operations on input 'Value'.
162 auto ApplyOperations = [&](const APInt &Value) {
163 APInt Result = Value;
164 for (auto *Op : llvm::reverse(C&: Operations)) {
165 switch (Op->getOpcode()) {
166 case Instruction::Xor:
167 Result ^= cast<ConstantInt>(Val: Op->getOperand(i: 1))->getValue();
168 break;
169 case Instruction::ZExt:
170 Result = Result.zext(width: Op->getType()->getIntegerBitWidth());
171 break;
172 case Instruction::SExt:
173 Result = Result.sext(width: Op->getType()->getIntegerBitWidth());
174 break;
175 default:
176 llvm_unreachable("Unexpected operation");
177 }
178 }
179 return Result;
180 };
181
182 auto *PhiDef = cast<PHINode>(Val: V);
183
184 // Get the first dominating conditional branch of the operand
185 // i's incoming block.
186 auto GetDomConditional = [&](unsigned i) -> CondBrInst * {
187 BasicBlock *BB = PhiDef->getIncomingBlock(i);
188 if (CondBrInst *BI = dyn_cast<CondBrInst>(Val: BB->getTerminator()))
189 return BI;
190 BB = BB->getSinglePredecessor();
191 if (!BB)
192 return nullptr;
193 return dyn_cast<CondBrInst>(Val: BB->getTerminator());
194 };
195
196 // Now walk through all Phi operands to find phi oprerands with values
197 // conflicting with the expected phi output value. Any such operand
198 // indicates the incoming edge to that operand is unlikely.
199 for (unsigned i = 0, e = PhiDef->getNumIncomingValues(); i != e; ++i) {
200
201 Value *PhiOpnd = PhiDef->getIncomingValue(i);
202 ConstantInt *CI = dyn_cast<ConstantInt>(Val: PhiOpnd);
203 if (!CI)
204 continue;
205
206 // Not an interesting case when IsUnlikely is false -- we can not infer
207 // anything useful when:
208 // (1) We expect some phi output and the operand value matches it, or
209 // (2) We don't expect some phi output (i.e. the "ExpectedValue" has low
210 // probability) and the operand value doesn't match that.
211 const APInt &CurrentPhiValue = ApplyOperations(CI->getValue());
212 if (ExpectedValueIsLikely == (ExpectedPhiValue == CurrentPhiValue))
213 continue;
214
215 CondBrInst *BI = GetDomConditional(i);
216 if (!BI)
217 continue;
218
219 MDBuilder MDB(PhiDef->getContext());
220
221 // There are two situations in which an operand of the PhiDef comes
222 // from a given successor of a branch instruction BI.
223 // 1) When the incoming block of the operand is the successor block;
224 // 2) When the incoming block is BI's enclosing block and the
225 // successor is the PhiDef's enclosing block.
226 //
227 // Returns true if the operand which comes from OpndIncomingBB
228 // comes from outgoing edge of BI that leads to Succ block.
229 auto *OpndIncomingBB = PhiDef->getIncomingBlock(i);
230 auto IsOpndComingFromSuccessor = [&](BasicBlock *Succ) {
231 if (OpndIncomingBB == Succ)
232 // If this successor is the incoming block for this
233 // Phi operand, then this successor does lead to the Phi.
234 return true;
235 if (OpndIncomingBB == BI->getParent() && Succ == PhiDef->getParent())
236 // Otherwise, if the edge is directly from the branch
237 // to the Phi, this successor is the one feeding this
238 // Phi operand.
239 return true;
240 return false;
241 };
242 uint32_t LikelyBranchWeightVal, UnlikelyBranchWeightVal;
243 std::tie(args&: LikelyBranchWeightVal, args&: UnlikelyBranchWeightVal) = getBranchWeight(
244 IntrinsicID: Expect->getCalledFunction()->getIntrinsicID(), CI: Expect, BranchCount: 2);
245 if (!ExpectedValueIsLikely)
246 std::swap(a&: LikelyBranchWeightVal, b&: UnlikelyBranchWeightVal);
247
248 if (IsOpndComingFromSuccessor(BI->getSuccessor(i: 1)))
249 BI->setMetadata(KindID: LLVMContext::MD_prof,
250 Node: MDB.createBranchWeights(TrueWeight: LikelyBranchWeightVal,
251 FalseWeight: UnlikelyBranchWeightVal,
252 /*IsExpected=*/true));
253 else if (IsOpndComingFromSuccessor(BI->getSuccessor(i: 0)))
254 BI->setMetadata(KindID: LLVMContext::MD_prof,
255 Node: MDB.createBranchWeights(TrueWeight: UnlikelyBranchWeightVal,
256 FalseWeight: LikelyBranchWeightVal,
257 /*IsExpected=*/true));
258 }
259}
260
261// Handle both CondBrInst and SelectInst.
262template <class BrSelInst> static bool handleBrSelExpect(BrSelInst &BSI) {
263
264 // Handle non-optimized IR code like:
265 // %expval = call i64 @llvm.expect.i64(i64 %conv1, i64 1)
266 // %tobool = icmp ne i64 %expval, 0
267 // br i1 %tobool, label %if.then, label %if.end
268 //
269 // Or the following simpler case:
270 // %expval = call i1 @llvm.expect.i1(i1 %cmp, i1 1)
271 // br i1 %expval, label %if.then, label %if.end
272
273 CallInst *CI;
274
275 ICmpInst *CmpI = dyn_cast<ICmpInst>(BSI.getCondition());
276 CmpInst::Predicate Predicate;
277 ConstantInt *CmpConstOperand = nullptr;
278 if (!CmpI) {
279 CI = dyn_cast<CallInst>(BSI.getCondition());
280 Predicate = CmpInst::ICMP_NE;
281 } else {
282 Predicate = CmpI->getPredicate();
283 if (Predicate != CmpInst::ICMP_NE && Predicate != CmpInst::ICMP_EQ)
284 return false;
285
286 CmpConstOperand = dyn_cast<ConstantInt>(Val: CmpI->getOperand(i_nocapture: 1));
287 if (!CmpConstOperand)
288 return false;
289 CI = dyn_cast<CallInst>(Val: CmpI->getOperand(i_nocapture: 0));
290 }
291
292 if (!CI)
293 return false;
294
295 uint64_t ValueComparedTo = 0;
296 if (CmpConstOperand) {
297 if (CmpConstOperand->getBitWidth() > 64)
298 return false;
299 ValueComparedTo = CmpConstOperand->getZExtValue();
300 }
301
302 Function *Fn = CI->getCalledFunction();
303 if (!Fn || (Fn->getIntrinsicID() != Intrinsic::expect &&
304 Fn->getIntrinsicID() != Intrinsic::expect_with_probability))
305 return false;
306
307 Value *ArgValue = CI->getArgOperand(i: 0);
308 ConstantInt *ExpectedValue = dyn_cast<ConstantInt>(Val: CI->getArgOperand(i: 1));
309 if (!ExpectedValue)
310 return false;
311
312 MDBuilder MDB(CI->getContext());
313 MDNode *Node;
314
315 uint32_t LikelyBranchWeightVal, UnlikelyBranchWeightVal;
316 std::tie(args&: LikelyBranchWeightVal, args&: UnlikelyBranchWeightVal) =
317 getBranchWeight(IntrinsicID: Fn->getIntrinsicID(), CI, BranchCount: 2);
318
319 SmallVector<uint32_t, 4> ExpectedWeights;
320 if ((ExpectedValue->getZExtValue() == ValueComparedTo) ==
321 (Predicate == CmpInst::ICMP_EQ)) {
322 Node = MDB.createBranchWeights(
323 TrueWeight: LikelyBranchWeightVal, FalseWeight: UnlikelyBranchWeightVal, /*IsExpected=*/true);
324 ExpectedWeights = {LikelyBranchWeightVal, UnlikelyBranchWeightVal};
325 } else {
326 Node = MDB.createBranchWeights(TrueWeight: UnlikelyBranchWeightVal,
327 FalseWeight: LikelyBranchWeightVal, /*IsExpected=*/true);
328 ExpectedWeights = {UnlikelyBranchWeightVal, LikelyBranchWeightVal};
329 }
330
331 if (CmpI)
332 CmpI->setOperand(i_nocapture: 0, Val_nocapture: ArgValue);
333 else
334 BSI.setCondition(ArgValue);
335
336 misexpect::checkFrontendInstrumentation(I: BSI, ExpectedWeights);
337
338 BSI.setMetadata(LLVMContext::MD_prof, Node);
339
340 return true;
341}
342
343static bool lowerExpectIntrinsic(Function &F) {
344 bool Changed = false;
345
346 for (BasicBlock &BB : F) {
347 // Create "block_weights" metadata.
348 if (CondBrInst *BI = dyn_cast<CondBrInst>(Val: BB.getTerminator())) {
349 if (handleBrSelExpect<CondBrInst>(BSI&: *BI))
350 ExpectIntrinsicsHandled++;
351 } else if (SwitchInst *SI = dyn_cast<SwitchInst>(Val: BB.getTerminator())) {
352 if (handleSwitchExpect(SI&: *SI))
353 ExpectIntrinsicsHandled++;
354 }
355
356 // Remove llvm.expect intrinsics. Iterate backwards in order
357 // to process select instructions before the intrinsic gets
358 // removed.
359 for (Instruction &Inst : llvm::make_early_inc_range(Range: llvm::reverse(C&: BB))) {
360 CallInst *CI = dyn_cast<CallInst>(Val: &Inst);
361 if (!CI) {
362 if (SelectInst *SI = dyn_cast<SelectInst>(Val: &Inst)) {
363 if (handleBrSelExpect(BSI&: *SI))
364 ExpectIntrinsicsHandled++;
365 }
366 continue;
367 }
368
369 Function *Fn = CI->getCalledFunction();
370 if (Fn && (Fn->getIntrinsicID() == Intrinsic::expect ||
371 Fn->getIntrinsicID() == Intrinsic::expect_with_probability)) {
372 // Before erasing the llvm.expect, walk backward to find
373 // phi that define llvm.expect's first arg, and
374 // infer branch probability:
375 handlePhiDef(Expect: CI);
376 Value *Exp = CI->getArgOperand(i: 0);
377 CI->replaceAllUsesWith(V: Exp);
378 CI->eraseFromParent();
379 Changed = true;
380 }
381 }
382 }
383
384 return Changed;
385}
386
387PreservedAnalyses LowerExpectIntrinsicPass::run(Function &F,
388 FunctionAnalysisManager &) {
389 if (lowerExpectIntrinsic(F))
390 return PreservedAnalyses::none();
391
392 return PreservedAnalyses::all();
393}
394