1 //===- InstCombineNegator.cpp -----------------------------------*- C++ -*-===//
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 file implements sinking of negation into expression trees,
10 // as long as that can be done without increasing instruction count.
11 //
12 //===----------------------------------------------------------------------===//
13 
14 #include "InstCombineInternal.h"
15 #include "llvm/ADT/APInt.h"
16 #include "llvm/ADT/ArrayRef.h"
17 #include "llvm/ADT/None.h"
18 #include "llvm/ADT/Optional.h"
19 #include "llvm/ADT/STLExtras.h"
20 #include "llvm/ADT/SmallVector.h"
21 #include "llvm/ADT/Statistic.h"
22 #include "llvm/ADT/StringRef.h"
23 #include "llvm/ADT/Twine.h"
24 #include "llvm/ADT/iterator_range.h"
25 #include "llvm/Analysis/TargetFolder.h"
26 #include "llvm/Analysis/ValueTracking.h"
27 #include "llvm/IR/Constant.h"
28 #include "llvm/IR/Constants.h"
29 #include "llvm/IR/DebugLoc.h"
30 #include "llvm/IR/DerivedTypes.h"
31 #include "llvm/IR/IRBuilder.h"
32 #include "llvm/IR/Instruction.h"
33 #include "llvm/IR/Instructions.h"
34 #include "llvm/IR/PatternMatch.h"
35 #include "llvm/IR/Type.h"
36 #include "llvm/IR/Use.h"
37 #include "llvm/IR/User.h"
38 #include "llvm/IR/Value.h"
39 #include "llvm/Support/Casting.h"
40 #include "llvm/Support/CommandLine.h"
41 #include "llvm/Support/Compiler.h"
42 #include "llvm/Support/DebugCounter.h"
43 #include "llvm/Support/ErrorHandling.h"
44 #include "llvm/Support/raw_ostream.h"
45 #include <functional>
46 #include <tuple>
47 #include <utility>
48 
49 using namespace llvm;
50 
51 #define DEBUG_TYPE "instcombine"
52 
53 STATISTIC(NegatorTotalNegationsAttempted,
54           "Negator: Number of negations attempted to be sinked");
55 STATISTIC(NegatorNumTreesNegated,
56           "Negator: Number of negations successfully sinked");
57 STATISTIC(NegatorMaxDepthVisited, "Negator: Maximal traversal depth ever "
58                                   "reached while attempting to sink negation");
59 STATISTIC(NegatorTimesDepthLimitReached,
60           "Negator: How many times did the traversal depth limit was reached "
61           "during sinking");
62 STATISTIC(
63     NegatorNumValuesVisited,
64     "Negator: Total number of values visited during attempts to sink negation");
65 STATISTIC(NegatorMaxTotalValuesVisited,
66           "Negator: Maximal number of values ever visited while attempting to "
67           "sink negation");
68 STATISTIC(NegatorNumInstructionsCreatedTotal,
69           "Negator: Number of new negated instructions created, total");
70 STATISTIC(NegatorMaxInstructionsCreated,
71           "Negator: Maximal number of new instructions created during negation "
72           "attempt");
73 STATISTIC(NegatorNumInstructionsNegatedSuccess,
74           "Negator: Number of new negated instructions created in successful "
75           "negation sinking attempts");
76 
77 DEBUG_COUNTER(NegatorCounter, "instcombine-negator",
78               "Controls Negator transformations in InstCombine pass");
79 
80 static cl::opt<bool>
81     NegatorEnabled("instcombine-negator-enabled", cl::init(true),
82                    cl::desc("Should we attempt to sink negations?"));
83 
84 static cl::opt<unsigned>
85     NegatorMaxDepth("instcombine-negator-max-depth",
86                     cl::init(NegatorDefaultMaxDepth),
87                     cl::desc("What is the maximal lookup depth when trying to "
88                              "check for viability of negation sinking."));
89 
90 Negator::Negator(LLVMContext &C, const DataLayout &DL, bool IsTrulyNegation_)
91     : Builder(C, TargetFolder(DL),
92               IRBuilderCallbackInserter([&](Instruction *I) {
93                 ++NegatorNumInstructionsCreatedTotal;
94                 NewInstructions.push_back(I);
95               })),
96       IsTrulyNegation(IsTrulyNegation_) {}
97 
98 #if LLVM_ENABLE_STATS
99 Negator::~Negator() {
100   NegatorMaxTotalValuesVisited.updateMax(NumValuesVisitedInThisNegator);
101 }
102 #endif
103 
104 // FIXME: can this be reworked into a worklist-based algorithm while preserving
105 // the depth-first, early bailout traversal?
106 LLVM_NODISCARD Value *Negator::visit(Value *V, unsigned Depth) {
107   NegatorMaxDepthVisited.updateMax(Depth);
108   ++NegatorNumValuesVisited;
109 
110 #if LLVM_ENABLE_STATS
111   ++NumValuesVisitedInThisNegator;
112 #endif
113 
114   // In i1, negation can simply be ignored.
115   if (V->getType()->isIntOrIntVectorTy(1))
116     return V;
117 
118   Value *X;
119 
120   // -(-(X)) -> X.
121   if (match(V, m_Neg(m_Value(X))))
122     return X;
123 
124   // Integral constants can be freely negated.
125   if (match(V, m_AnyIntegralConstant()))
126     return ConstantExpr::getNeg(cast<Constant>(V), /*HasNUW=*/false,
127                                 /*HasNSW=*/false);
128 
129   // If we have a non-instruction, then give up.
130   if (!isa<Instruction>(V))
131     return nullptr;
132 
133   // If we have started with a true negation (i.e. `sub 0, %y`), then if we've
134   // got instruction that does not require recursive reasoning, we can still
135   // negate it even if it has other uses, without increasing instruction count.
136   if (!V->hasOneUse() && !IsTrulyNegation)
137     return nullptr;
138 
139   auto *I = cast<Instruction>(V);
140   unsigned BitWidth = I->getType()->getScalarSizeInBits();
141 
142   // We must preserve the insertion point and debug info that is set in the
143   // builder at the time this function is called.
144   InstCombiner::BuilderTy::InsertPointGuard Guard(Builder);
145   // And since we are trying to negate instruction I, that tells us about the
146   // insertion point and the debug info that we need to keep.
147   Builder.SetInsertPoint(I);
148 
149   // In some cases we can give the answer without further recursion.
150   switch (I->getOpcode()) {
151   case Instruction::Sub:
152     // `sub` is always negatible.
153     return Builder.CreateSub(I->getOperand(1), I->getOperand(0),
154                              I->getName() + ".neg");
155   case Instruction::Add:
156     // `inc` is always negatible.
157     if (match(I->getOperand(1), m_One()))
158       return Builder.CreateNot(I->getOperand(0), I->getName() + ".neg");
159     break;
160   case Instruction::Xor:
161     // `not` is always negatible.
162     if (match(I, m_Not(m_Value(X))))
163       return Builder.CreateAdd(X, ConstantInt::get(X->getType(), 1),
164                                I->getName() + ".neg");
165     break;
166   case Instruction::AShr:
167   case Instruction::LShr: {
168     // Right-shift sign bit smear is negatible.
169     const APInt *Op1Val;
170     if (match(I->getOperand(1), m_APInt(Op1Val)) && *Op1Val == BitWidth - 1) {
171       Value *BO = I->getOpcode() == Instruction::AShr
172                       ? Builder.CreateLShr(I->getOperand(0), I->getOperand(1))
173                       : Builder.CreateAShr(I->getOperand(0), I->getOperand(1));
174       if (auto *NewInstr = dyn_cast<Instruction>(BO)) {
175         NewInstr->copyIRFlags(I);
176         NewInstr->setName(I->getName() + ".neg");
177       }
178       return BO;
179     }
180     break;
181   }
182   case Instruction::SDiv:
183     // `sdiv` is negatible if divisor is not undef/INT_MIN/1.
184     // While this is normally not behind a use-check,
185     // let's consider division to be special since it's costly.
186     if (!I->hasOneUse())
187       break;
188     if (auto *Op1C = dyn_cast<Constant>(I->getOperand(1))) {
189       if (!Op1C->containsUndefElement() && Op1C->isNotMinSignedValue() &&
190           Op1C->isNotOneValue()) {
191         Value *BO =
192             Builder.CreateSDiv(I->getOperand(0), ConstantExpr::getNeg(Op1C),
193                                I->getName() + ".neg");
194         if (auto *NewInstr = dyn_cast<Instruction>(BO))
195           NewInstr->setIsExact(I->isExact());
196         return BO;
197       }
198     }
199     break;
200   case Instruction::SExt:
201   case Instruction::ZExt:
202     // `*ext` of i1 is always negatible
203     if (I->getOperand(0)->getType()->isIntOrIntVectorTy(1))
204       return I->getOpcode() == Instruction::SExt
205                  ? Builder.CreateZExt(I->getOperand(0), I->getType(),
206                                       I->getName() + ".neg")
207                  : Builder.CreateSExt(I->getOperand(0), I->getType(),
208                                       I->getName() + ".neg");
209     break;
210   default:
211     break; // Other instructions require recursive reasoning.
212   }
213 
214   // Rest of the logic is recursive, and if either the current instruction
215   // has other uses or if it's time to give up then it's time.
216   if (!V->hasOneUse())
217     return nullptr;
218   if (Depth > NegatorMaxDepth) {
219     LLVM_DEBUG(dbgs() << "Negator: reached maximal allowed traversal depth in "
220                       << *V << ". Giving up.\n");
221     ++NegatorTimesDepthLimitReached;
222     return nullptr;
223   }
224 
225   switch (I->getOpcode()) {
226   case Instruction::PHI: {
227     // `phi` is negatible if all the incoming values are negatible.
228     PHINode *PHI = cast<PHINode>(I);
229     SmallVector<Value *, 4> NegatedIncomingValues(PHI->getNumOperands());
230     for (auto I : zip(PHI->incoming_values(), NegatedIncomingValues)) {
231       if (!(std::get<1>(I) = visit(std::get<0>(I), Depth + 1))) // Early return.
232         return nullptr;
233     }
234     // All incoming values are indeed negatible. Create negated PHI node.
235     PHINode *NegatedPHI = Builder.CreatePHI(
236         PHI->getType(), PHI->getNumOperands(), PHI->getName() + ".neg");
237     for (auto I : zip(NegatedIncomingValues, PHI->blocks()))
238       NegatedPHI->addIncoming(std::get<0>(I), std::get<1>(I));
239     return NegatedPHI;
240   }
241   case Instruction::Select: {
242     {
243       // `abs`/`nabs` is always negatible.
244       Value *LHS, *RHS;
245       SelectPatternFlavor SPF =
246           matchSelectPattern(I, LHS, RHS, /*CastOp=*/nullptr, Depth).Flavor;
247       if (SPF == SPF_ABS || SPF == SPF_NABS) {
248         auto *NewSelect = cast<SelectInst>(I->clone());
249         // Just swap the operands of the select.
250         NewSelect->swapValues();
251         // Don't swap prof metadata, we didn't change the branch behavior.
252         NewSelect->setName(I->getName() + ".neg");
253         Builder.Insert(NewSelect);
254         return NewSelect;
255       }
256     }
257     // `select` is negatible if both hands of `select` are negatible.
258     Value *NegOp1 = visit(I->getOperand(1), Depth + 1);
259     if (!NegOp1) // Early return.
260       return nullptr;
261     Value *NegOp2 = visit(I->getOperand(2), Depth + 1);
262     if (!NegOp2)
263       return nullptr;
264     // Do preserve the metadata!
265     return Builder.CreateSelect(I->getOperand(0), NegOp1, NegOp2,
266                                 I->getName() + ".neg", /*MDFrom=*/I);
267   }
268   case Instruction::Trunc: {
269     // `trunc` is negatible if its operand is negatible.
270     Value *NegOp = visit(I->getOperand(0), Depth + 1);
271     if (!NegOp) // Early return.
272       return nullptr;
273     return Builder.CreateTrunc(NegOp, I->getType(), I->getName() + ".neg");
274   }
275   case Instruction::Shl: {
276     // `shl` is negatible if the first operand is negatible.
277     Value *NegOp0 = visit(I->getOperand(0), Depth + 1);
278     if (!NegOp0) // Early return.
279       return nullptr;
280     return Builder.CreateShl(NegOp0, I->getOperand(1), I->getName() + ".neg");
281   }
282   case Instruction::Add: {
283     // `add` is negatible if both of its operands are negatible.
284     Value *NegOp0 = visit(I->getOperand(0), Depth + 1);
285     if (!NegOp0) // Early return.
286       return nullptr;
287     Value *NegOp1 = visit(I->getOperand(1), Depth + 1);
288     if (!NegOp1)
289       return nullptr;
290     return Builder.CreateAdd(NegOp0, NegOp1, I->getName() + ".neg");
291   }
292   case Instruction::Xor:
293     // `xor` is negatible if one of its operands is invertible.
294     // FIXME: InstCombineInverter? But how to connect Inverter and Negator?
295     if (auto *C = dyn_cast<Constant>(I->getOperand(1))) {
296       Value *Xor = Builder.CreateXor(I->getOperand(0), ConstantExpr::getNot(C));
297       return Builder.CreateAdd(Xor, ConstantInt::get(Xor->getType(), 1),
298                                I->getName() + ".neg");
299     }
300     return nullptr;
301   case Instruction::Mul: {
302     // `mul` is negatible if one of its operands is negatible.
303     Value *NegatedOp, *OtherOp;
304     // First try the second operand, in case it's a constant it will be best to
305     // just invert it instead of sinking the `neg` deeper.
306     if (Value *NegOp1 = visit(I->getOperand(1), Depth + 1)) {
307       NegatedOp = NegOp1;
308       OtherOp = I->getOperand(0);
309     } else if (Value *NegOp0 = visit(I->getOperand(0), Depth + 1)) {
310       NegatedOp = NegOp0;
311       OtherOp = I->getOperand(1);
312     } else
313       // Can't negate either of them.
314       return nullptr;
315     return Builder.CreateMul(NegatedOp, OtherOp, I->getName() + ".neg");
316   }
317   default:
318     return nullptr; // Don't know, likely not negatible for free.
319   }
320 
321   llvm_unreachable("Can't get here. We always return from switch.");
322 }
323 
324 LLVM_NODISCARD Optional<Negator::Result> Negator::run(Value *Root) {
325   Value *Negated = visit(Root, /*Depth=*/0);
326   if (!Negated) {
327     // We must cleanup newly-inserted instructions, to avoid any potential
328     // endless combine looping.
329     llvm::for_each(llvm::reverse(NewInstructions),
330                    [&](Instruction *I) { I->eraseFromParent(); });
331     return llvm::None;
332   }
333   return std::make_pair(ArrayRef<Instruction *>(NewInstructions), Negated);
334 }
335 
336 LLVM_NODISCARD Value *Negator::Negate(bool LHSIsZero, Value *Root,
337                                       InstCombiner &IC) {
338   ++NegatorTotalNegationsAttempted;
339   LLVM_DEBUG(dbgs() << "Negator: attempting to sink negation into " << *Root
340                     << "\n");
341 
342   if (!NegatorEnabled || !DebugCounter::shouldExecute(NegatorCounter))
343     return nullptr;
344 
345   Negator N(Root->getContext(), IC.getDataLayout(), LHSIsZero);
346   Optional<Result> Res = N.run(Root);
347   if (!Res) { // Negation failed.
348     LLVM_DEBUG(dbgs() << "Negator: failed to sink negation into " << *Root
349                       << "\n");
350     return nullptr;
351   }
352 
353   LLVM_DEBUG(dbgs() << "Negator: successfully sunk negation into " << *Root
354                     << "\n         NEW: " << *Res->second << "\n");
355   ++NegatorNumTreesNegated;
356 
357   // We must temporarily unset the 'current' insertion point and DebugLoc of the
358   // InstCombine's IRBuilder so that it won't interfere with the ones we have
359   // already specified when producing negated instructions.
360   InstCombiner::BuilderTy::InsertPointGuard Guard(IC.Builder);
361   IC.Builder.ClearInsertionPoint();
362   IC.Builder.SetCurrentDebugLocation(DebugLoc());
363 
364   // And finally, we must add newly-created instructions into the InstCombine's
365   // worklist (in a proper order!) so it can attempt to combine them.
366   LLVM_DEBUG(dbgs() << "Negator: Propagating " << Res->first.size()
367                     << " instrs to InstCombine\n");
368   NegatorMaxInstructionsCreated.updateMax(Res->first.size());
369   NegatorNumInstructionsNegatedSuccess += Res->first.size();
370 
371   // They are in def-use order, so nothing fancy, just insert them in order.
372   llvm::for_each(Res->first,
373                  [&](Instruction *I) { IC.Builder.Insert(I, I->getName()); });
374 
375   // And return the new root.
376   return Res->second;
377 }
378