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