1 //===- HexagonLoopIdiomRecognition.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 #include "HexagonLoopIdiomRecognition.h" 10 #include "llvm/ADT/APInt.h" 11 #include "llvm/ADT/DenseMap.h" 12 #include "llvm/ADT/SetVector.h" 13 #include "llvm/ADT/SmallPtrSet.h" 14 #include "llvm/ADT/SmallSet.h" 15 #include "llvm/ADT/SmallVector.h" 16 #include "llvm/ADT/StringRef.h" 17 #include "llvm/ADT/Triple.h" 18 #include "llvm/Analysis/AliasAnalysis.h" 19 #include "llvm/Analysis/InstructionSimplify.h" 20 #include "llvm/Analysis/LoopAnalysisManager.h" 21 #include "llvm/Analysis/LoopInfo.h" 22 #include "llvm/Analysis/LoopPass.h" 23 #include "llvm/Analysis/MemoryLocation.h" 24 #include "llvm/Analysis/ScalarEvolution.h" 25 #include "llvm/Analysis/ScalarEvolutionExpressions.h" 26 #include "llvm/Analysis/TargetLibraryInfo.h" 27 #include "llvm/Analysis/ValueTracking.h" 28 #include "llvm/IR/Attributes.h" 29 #include "llvm/IR/BasicBlock.h" 30 #include "llvm/IR/Constant.h" 31 #include "llvm/IR/Constants.h" 32 #include "llvm/IR/DataLayout.h" 33 #include "llvm/IR/DebugLoc.h" 34 #include "llvm/IR/DerivedTypes.h" 35 #include "llvm/IR/Dominators.h" 36 #include "llvm/IR/Function.h" 37 #include "llvm/IR/IRBuilder.h" 38 #include "llvm/IR/InstrTypes.h" 39 #include "llvm/IR/Instruction.h" 40 #include "llvm/IR/Instructions.h" 41 #include "llvm/IR/IntrinsicInst.h" 42 #include "llvm/IR/Intrinsics.h" 43 #include "llvm/IR/IntrinsicsHexagon.h" 44 #include "llvm/IR/Module.h" 45 #include "llvm/IR/PassManager.h" 46 #include "llvm/IR/PatternMatch.h" 47 #include "llvm/IR/Type.h" 48 #include "llvm/IR/User.h" 49 #include "llvm/IR/Value.h" 50 #include "llvm/InitializePasses.h" 51 #include "llvm/Pass.h" 52 #include "llvm/Support/Casting.h" 53 #include "llvm/Support/CommandLine.h" 54 #include "llvm/Support/Compiler.h" 55 #include "llvm/Support/Debug.h" 56 #include "llvm/Support/ErrorHandling.h" 57 #include "llvm/Support/KnownBits.h" 58 #include "llvm/Support/raw_ostream.h" 59 #include "llvm/Transforms/Scalar.h" 60 #include "llvm/Transforms/Utils.h" 61 #include "llvm/Transforms/Utils/Local.h" 62 #include "llvm/Transforms/Utils/ScalarEvolutionExpander.h" 63 #include <algorithm> 64 #include <array> 65 #include <cassert> 66 #include <cstdint> 67 #include <cstdlib> 68 #include <deque> 69 #include <functional> 70 #include <iterator> 71 #include <map> 72 #include <set> 73 #include <utility> 74 #include <vector> 75 76 #define DEBUG_TYPE "hexagon-lir" 77 78 using namespace llvm; 79 80 static cl::opt<bool> DisableMemcpyIdiom("disable-memcpy-idiom", 81 cl::Hidden, cl::init(false), 82 cl::desc("Disable generation of memcpy in loop idiom recognition")); 83 84 static cl::opt<bool> DisableMemmoveIdiom("disable-memmove-idiom", 85 cl::Hidden, cl::init(false), 86 cl::desc("Disable generation of memmove in loop idiom recognition")); 87 88 static cl::opt<unsigned> RuntimeMemSizeThreshold("runtime-mem-idiom-threshold", 89 cl::Hidden, cl::init(0), cl::desc("Threshold (in bytes) for the runtime " 90 "check guarding the memmove.")); 91 92 static cl::opt<unsigned> CompileTimeMemSizeThreshold( 93 "compile-time-mem-idiom-threshold", cl::Hidden, cl::init(64), 94 cl::desc("Threshold (in bytes) to perform the transformation, if the " 95 "runtime loop count (mem transfer size) is known at compile-time.")); 96 97 static cl::opt<bool> OnlyNonNestedMemmove("only-nonnested-memmove-idiom", 98 cl::Hidden, cl::init(true), 99 cl::desc("Only enable generating memmove in non-nested loops")); 100 101 static cl::opt<bool> HexagonVolatileMemcpy( 102 "disable-hexagon-volatile-memcpy", cl::Hidden, cl::init(false), 103 cl::desc("Enable Hexagon-specific memcpy for volatile destination.")); 104 105 static cl::opt<unsigned> SimplifyLimit("hlir-simplify-limit", cl::init(10000), 106 cl::Hidden, cl::desc("Maximum number of simplification steps in HLIR")); 107 108 static const char *HexagonVolatileMemcpyName 109 = "hexagon_memcpy_forward_vp4cp4n2"; 110 111 112 namespace llvm { 113 114 void initializeHexagonLoopIdiomRecognizeLegacyPassPass(PassRegistry &); 115 Pass *createHexagonLoopIdiomPass(); 116 117 } // end namespace llvm 118 119 namespace { 120 121 class HexagonLoopIdiomRecognize { 122 public: 123 explicit HexagonLoopIdiomRecognize(AliasAnalysis *AA, DominatorTree *DT, 124 LoopInfo *LF, const TargetLibraryInfo *TLI, 125 ScalarEvolution *SE) 126 : AA(AA), DT(DT), LF(LF), TLI(TLI), SE(SE) {} 127 128 bool run(Loop *L); 129 130 private: 131 int getSCEVStride(const SCEVAddRecExpr *StoreEv); 132 bool isLegalStore(Loop *CurLoop, StoreInst *SI); 133 void collectStores(Loop *CurLoop, BasicBlock *BB, 134 SmallVectorImpl<StoreInst *> &Stores); 135 bool processCopyingStore(Loop *CurLoop, StoreInst *SI, const SCEV *BECount); 136 bool coverLoop(Loop *L, SmallVectorImpl<Instruction *> &Insts) const; 137 bool runOnLoopBlock(Loop *CurLoop, BasicBlock *BB, const SCEV *BECount, 138 SmallVectorImpl<BasicBlock *> &ExitBlocks); 139 bool runOnCountableLoop(Loop *L); 140 141 AliasAnalysis *AA; 142 const DataLayout *DL; 143 DominatorTree *DT; 144 LoopInfo *LF; 145 const TargetLibraryInfo *TLI; 146 ScalarEvolution *SE; 147 bool HasMemcpy, HasMemmove; 148 }; 149 150 class HexagonLoopIdiomRecognizeLegacyPass : public LoopPass { 151 public: 152 static char ID; 153 154 explicit HexagonLoopIdiomRecognizeLegacyPass() : LoopPass(ID) { 155 initializeHexagonLoopIdiomRecognizeLegacyPassPass( 156 *PassRegistry::getPassRegistry()); 157 } 158 159 StringRef getPassName() const override { 160 return "Recognize Hexagon-specific loop idioms"; 161 } 162 163 void getAnalysisUsage(AnalysisUsage &AU) const override { 164 AU.addRequired<LoopInfoWrapperPass>(); 165 AU.addRequiredID(LoopSimplifyID); 166 AU.addRequiredID(LCSSAID); 167 AU.addRequired<AAResultsWrapperPass>(); 168 AU.addPreserved<AAResultsWrapperPass>(); 169 AU.addRequired<ScalarEvolutionWrapperPass>(); 170 AU.addRequired<DominatorTreeWrapperPass>(); 171 AU.addRequired<TargetLibraryInfoWrapperPass>(); 172 AU.addPreserved<TargetLibraryInfoWrapperPass>(); 173 } 174 175 bool runOnLoop(Loop *L, LPPassManager &LPM) override; 176 }; 177 178 struct Simplifier { 179 struct Rule { 180 using FuncType = std::function<Value *(Instruction *, LLVMContext &)>; 181 Rule(StringRef N, FuncType F) : Name(N), Fn(F) {} 182 StringRef Name; // For debugging. 183 FuncType Fn; 184 }; 185 186 void addRule(StringRef N, const Rule::FuncType &F) { 187 Rules.push_back(Rule(N, F)); 188 } 189 190 private: 191 struct WorkListType { 192 WorkListType() = default; 193 194 void push_back(Value *V) { 195 // Do not push back duplicates. 196 if (!S.count(V)) { 197 Q.push_back(V); 198 S.insert(V); 199 } 200 } 201 202 Value *pop_front_val() { 203 Value *V = Q.front(); 204 Q.pop_front(); 205 S.erase(V); 206 return V; 207 } 208 209 bool empty() const { return Q.empty(); } 210 211 private: 212 std::deque<Value *> Q; 213 std::set<Value *> S; 214 }; 215 216 using ValueSetType = std::set<Value *>; 217 218 std::vector<Rule> Rules; 219 220 public: 221 struct Context { 222 using ValueMapType = DenseMap<Value *, Value *>; 223 224 Value *Root; 225 ValueSetType Used; // The set of all cloned values used by Root. 226 ValueSetType Clones; // The set of all cloned values. 227 LLVMContext &Ctx; 228 229 Context(Instruction *Exp) 230 : Ctx(Exp->getParent()->getParent()->getContext()) { 231 initialize(Exp); 232 } 233 234 ~Context() { cleanup(); } 235 236 void print(raw_ostream &OS, const Value *V) const; 237 Value *materialize(BasicBlock *B, BasicBlock::iterator At); 238 239 private: 240 friend struct Simplifier; 241 242 void initialize(Instruction *Exp); 243 void cleanup(); 244 245 template <typename FuncT> void traverse(Value *V, FuncT F); 246 void record(Value *V); 247 void use(Value *V); 248 void unuse(Value *V); 249 250 bool equal(const Instruction *I, const Instruction *J) const; 251 Value *find(Value *Tree, Value *Sub) const; 252 Value *subst(Value *Tree, Value *OldV, Value *NewV); 253 void replace(Value *OldV, Value *NewV); 254 void link(Instruction *I, BasicBlock *B, BasicBlock::iterator At); 255 }; 256 257 Value *simplify(Context &C); 258 }; 259 260 struct PE { 261 PE(const Simplifier::Context &c, Value *v = nullptr) : C(c), V(v) {} 262 263 const Simplifier::Context &C; 264 const Value *V; 265 }; 266 267 LLVM_ATTRIBUTE_USED 268 raw_ostream &operator<<(raw_ostream &OS, const PE &P) { 269 P.C.print(OS, P.V ? P.V : P.C.Root); 270 return OS; 271 } 272 273 } // end anonymous namespace 274 275 char HexagonLoopIdiomRecognizeLegacyPass::ID = 0; 276 277 INITIALIZE_PASS_BEGIN(HexagonLoopIdiomRecognizeLegacyPass, "hexagon-loop-idiom", 278 "Recognize Hexagon-specific loop idioms", false, false) 279 INITIALIZE_PASS_DEPENDENCY(LoopInfoWrapperPass) 280 INITIALIZE_PASS_DEPENDENCY(LoopSimplify) 281 INITIALIZE_PASS_DEPENDENCY(LCSSAWrapperPass) 282 INITIALIZE_PASS_DEPENDENCY(ScalarEvolutionWrapperPass) 283 INITIALIZE_PASS_DEPENDENCY(DominatorTreeWrapperPass) 284 INITIALIZE_PASS_DEPENDENCY(TargetLibraryInfoWrapperPass) 285 INITIALIZE_PASS_DEPENDENCY(AAResultsWrapperPass) 286 INITIALIZE_PASS_END(HexagonLoopIdiomRecognizeLegacyPass, "hexagon-loop-idiom", 287 "Recognize Hexagon-specific loop idioms", false, false) 288 289 template <typename FuncT> 290 void Simplifier::Context::traverse(Value *V, FuncT F) { 291 WorkListType Q; 292 Q.push_back(V); 293 294 while (!Q.empty()) { 295 Instruction *U = dyn_cast<Instruction>(Q.pop_front_val()); 296 if (!U || U->getParent()) 297 continue; 298 if (!F(U)) 299 continue; 300 for (Value *Op : U->operands()) 301 Q.push_back(Op); 302 } 303 } 304 305 void Simplifier::Context::print(raw_ostream &OS, const Value *V) const { 306 const auto *U = dyn_cast<const Instruction>(V); 307 if (!U) { 308 OS << V << '(' << *V << ')'; 309 return; 310 } 311 312 if (U->getParent()) { 313 OS << U << '('; 314 U->printAsOperand(OS, true); 315 OS << ')'; 316 return; 317 } 318 319 unsigned N = U->getNumOperands(); 320 if (N != 0) 321 OS << U << '('; 322 OS << U->getOpcodeName(); 323 for (const Value *Op : U->operands()) { 324 OS << ' '; 325 print(OS, Op); 326 } 327 if (N != 0) 328 OS << ')'; 329 } 330 331 void Simplifier::Context::initialize(Instruction *Exp) { 332 // Perform a deep clone of the expression, set Root to the root 333 // of the clone, and build a map from the cloned values to the 334 // original ones. 335 ValueMapType M; 336 BasicBlock *Block = Exp->getParent(); 337 WorkListType Q; 338 Q.push_back(Exp); 339 340 while (!Q.empty()) { 341 Value *V = Q.pop_front_val(); 342 if (M.find(V) != M.end()) 343 continue; 344 if (Instruction *U = dyn_cast<Instruction>(V)) { 345 if (isa<PHINode>(U) || U->getParent() != Block) 346 continue; 347 for (Value *Op : U->operands()) 348 Q.push_back(Op); 349 M.insert({U, U->clone()}); 350 } 351 } 352 353 for (std::pair<Value*,Value*> P : M) { 354 Instruction *U = cast<Instruction>(P.second); 355 for (unsigned i = 0, n = U->getNumOperands(); i != n; ++i) { 356 auto F = M.find(U->getOperand(i)); 357 if (F != M.end()) 358 U->setOperand(i, F->second); 359 } 360 } 361 362 auto R = M.find(Exp); 363 assert(R != M.end()); 364 Root = R->second; 365 366 record(Root); 367 use(Root); 368 } 369 370 void Simplifier::Context::record(Value *V) { 371 auto Record = [this](Instruction *U) -> bool { 372 Clones.insert(U); 373 return true; 374 }; 375 traverse(V, Record); 376 } 377 378 void Simplifier::Context::use(Value *V) { 379 auto Use = [this](Instruction *U) -> bool { 380 Used.insert(U); 381 return true; 382 }; 383 traverse(V, Use); 384 } 385 386 void Simplifier::Context::unuse(Value *V) { 387 if (!isa<Instruction>(V) || cast<Instruction>(V)->getParent() != nullptr) 388 return; 389 390 auto Unuse = [this](Instruction *U) -> bool { 391 if (!U->use_empty()) 392 return false; 393 Used.erase(U); 394 return true; 395 }; 396 traverse(V, Unuse); 397 } 398 399 Value *Simplifier::Context::subst(Value *Tree, Value *OldV, Value *NewV) { 400 if (Tree == OldV) 401 return NewV; 402 if (OldV == NewV) 403 return Tree; 404 405 WorkListType Q; 406 Q.push_back(Tree); 407 while (!Q.empty()) { 408 Instruction *U = dyn_cast<Instruction>(Q.pop_front_val()); 409 // If U is not an instruction, or it's not a clone, skip it. 410 if (!U || U->getParent()) 411 continue; 412 for (unsigned i = 0, n = U->getNumOperands(); i != n; ++i) { 413 Value *Op = U->getOperand(i); 414 if (Op == OldV) { 415 U->setOperand(i, NewV); 416 unuse(OldV); 417 } else { 418 Q.push_back(Op); 419 } 420 } 421 } 422 return Tree; 423 } 424 425 void Simplifier::Context::replace(Value *OldV, Value *NewV) { 426 if (Root == OldV) { 427 Root = NewV; 428 use(Root); 429 return; 430 } 431 432 // NewV may be a complex tree that has just been created by one of the 433 // transformation rules. We need to make sure that it is commoned with 434 // the existing Root to the maximum extent possible. 435 // Identify all subtrees of NewV (including NewV itself) that have 436 // equivalent counterparts in Root, and replace those subtrees with 437 // these counterparts. 438 WorkListType Q; 439 Q.push_back(NewV); 440 while (!Q.empty()) { 441 Value *V = Q.pop_front_val(); 442 Instruction *U = dyn_cast<Instruction>(V); 443 if (!U || U->getParent()) 444 continue; 445 if (Value *DupV = find(Root, V)) { 446 if (DupV != V) 447 NewV = subst(NewV, V, DupV); 448 } else { 449 for (Value *Op : U->operands()) 450 Q.push_back(Op); 451 } 452 } 453 454 // Now, simply replace OldV with NewV in Root. 455 Root = subst(Root, OldV, NewV); 456 use(Root); 457 } 458 459 void Simplifier::Context::cleanup() { 460 for (Value *V : Clones) { 461 Instruction *U = cast<Instruction>(V); 462 if (!U->getParent()) 463 U->dropAllReferences(); 464 } 465 466 for (Value *V : Clones) { 467 Instruction *U = cast<Instruction>(V); 468 if (!U->getParent()) 469 U->deleteValue(); 470 } 471 } 472 473 bool Simplifier::Context::equal(const Instruction *I, 474 const Instruction *J) const { 475 if (I == J) 476 return true; 477 if (!I->isSameOperationAs(J)) 478 return false; 479 if (isa<PHINode>(I)) 480 return I->isIdenticalTo(J); 481 482 for (unsigned i = 0, n = I->getNumOperands(); i != n; ++i) { 483 Value *OpI = I->getOperand(i), *OpJ = J->getOperand(i); 484 if (OpI == OpJ) 485 continue; 486 auto *InI = dyn_cast<const Instruction>(OpI); 487 auto *InJ = dyn_cast<const Instruction>(OpJ); 488 if (InI && InJ) { 489 if (!equal(InI, InJ)) 490 return false; 491 } else if (InI != InJ || !InI) 492 return false; 493 } 494 return true; 495 } 496 497 Value *Simplifier::Context::find(Value *Tree, Value *Sub) const { 498 Instruction *SubI = dyn_cast<Instruction>(Sub); 499 WorkListType Q; 500 Q.push_back(Tree); 501 502 while (!Q.empty()) { 503 Value *V = Q.pop_front_val(); 504 if (V == Sub) 505 return V; 506 Instruction *U = dyn_cast<Instruction>(V); 507 if (!U || U->getParent()) 508 continue; 509 if (SubI && equal(SubI, U)) 510 return U; 511 assert(!isa<PHINode>(U)); 512 for (Value *Op : U->operands()) 513 Q.push_back(Op); 514 } 515 return nullptr; 516 } 517 518 void Simplifier::Context::link(Instruction *I, BasicBlock *B, 519 BasicBlock::iterator At) { 520 if (I->getParent()) 521 return; 522 523 for (Value *Op : I->operands()) { 524 if (Instruction *OpI = dyn_cast<Instruction>(Op)) 525 link(OpI, B, At); 526 } 527 528 B->getInstList().insert(At, I); 529 } 530 531 Value *Simplifier::Context::materialize(BasicBlock *B, 532 BasicBlock::iterator At) { 533 if (Instruction *RootI = dyn_cast<Instruction>(Root)) 534 link(RootI, B, At); 535 return Root; 536 } 537 538 Value *Simplifier::simplify(Context &C) { 539 WorkListType Q; 540 Q.push_back(C.Root); 541 unsigned Count = 0; 542 const unsigned Limit = SimplifyLimit; 543 544 while (!Q.empty()) { 545 if (Count++ >= Limit) 546 break; 547 Instruction *U = dyn_cast<Instruction>(Q.pop_front_val()); 548 if (!U || U->getParent() || !C.Used.count(U)) 549 continue; 550 bool Changed = false; 551 for (Rule &R : Rules) { 552 Value *W = R.Fn(U, C.Ctx); 553 if (!W) 554 continue; 555 Changed = true; 556 C.record(W); 557 C.replace(U, W); 558 Q.push_back(C.Root); 559 break; 560 } 561 if (!Changed) { 562 for (Value *Op : U->operands()) 563 Q.push_back(Op); 564 } 565 } 566 return Count < Limit ? C.Root : nullptr; 567 } 568 569 //===----------------------------------------------------------------------===// 570 // 571 // Implementation of PolynomialMultiplyRecognize 572 // 573 //===----------------------------------------------------------------------===// 574 575 namespace { 576 577 class PolynomialMultiplyRecognize { 578 public: 579 explicit PolynomialMultiplyRecognize(Loop *loop, const DataLayout &dl, 580 const DominatorTree &dt, const TargetLibraryInfo &tli, 581 ScalarEvolution &se) 582 : CurLoop(loop), DL(dl), DT(dt), TLI(tli), SE(se) {} 583 584 bool recognize(); 585 586 private: 587 using ValueSeq = SetVector<Value *>; 588 589 IntegerType *getPmpyType() const { 590 LLVMContext &Ctx = CurLoop->getHeader()->getParent()->getContext(); 591 return IntegerType::get(Ctx, 32); 592 } 593 594 bool isPromotableTo(Value *V, IntegerType *Ty); 595 void promoteTo(Instruction *In, IntegerType *DestTy, BasicBlock *LoopB); 596 bool promoteTypes(BasicBlock *LoopB, BasicBlock *ExitB); 597 598 Value *getCountIV(BasicBlock *BB); 599 bool findCycle(Value *Out, Value *In, ValueSeq &Cycle); 600 void classifyCycle(Instruction *DivI, ValueSeq &Cycle, ValueSeq &Early, 601 ValueSeq &Late); 602 bool classifyInst(Instruction *UseI, ValueSeq &Early, ValueSeq &Late); 603 bool commutesWithShift(Instruction *I); 604 bool highBitsAreZero(Value *V, unsigned IterCount); 605 bool keepsHighBitsZero(Value *V, unsigned IterCount); 606 bool isOperandShifted(Instruction *I, Value *Op); 607 bool convertShiftsToLeft(BasicBlock *LoopB, BasicBlock *ExitB, 608 unsigned IterCount); 609 void cleanupLoopBody(BasicBlock *LoopB); 610 611 struct ParsedValues { 612 ParsedValues() = default; 613 614 Value *M = nullptr; 615 Value *P = nullptr; 616 Value *Q = nullptr; 617 Value *R = nullptr; 618 Value *X = nullptr; 619 Instruction *Res = nullptr; 620 unsigned IterCount = 0; 621 bool Left = false; 622 bool Inv = false; 623 }; 624 625 bool matchLeftShift(SelectInst *SelI, Value *CIV, ParsedValues &PV); 626 bool matchRightShift(SelectInst *SelI, ParsedValues &PV); 627 bool scanSelect(SelectInst *SI, BasicBlock *LoopB, BasicBlock *PrehB, 628 Value *CIV, ParsedValues &PV, bool PreScan); 629 unsigned getInverseMxN(unsigned QP); 630 Value *generate(BasicBlock::iterator At, ParsedValues &PV); 631 632 void setupPreSimplifier(Simplifier &S); 633 void setupPostSimplifier(Simplifier &S); 634 635 Loop *CurLoop; 636 const DataLayout &DL; 637 const DominatorTree &DT; 638 const TargetLibraryInfo &TLI; 639 ScalarEvolution &SE; 640 }; 641 642 } // end anonymous namespace 643 644 Value *PolynomialMultiplyRecognize::getCountIV(BasicBlock *BB) { 645 pred_iterator PI = pred_begin(BB), PE = pred_end(BB); 646 if (std::distance(PI, PE) != 2) 647 return nullptr; 648 BasicBlock *PB = (*PI == BB) ? *std::next(PI) : *PI; 649 650 for (auto I = BB->begin(), E = BB->end(); I != E && isa<PHINode>(I); ++I) { 651 auto *PN = cast<PHINode>(I); 652 Value *InitV = PN->getIncomingValueForBlock(PB); 653 if (!isa<ConstantInt>(InitV) || !cast<ConstantInt>(InitV)->isZero()) 654 continue; 655 Value *IterV = PN->getIncomingValueForBlock(BB); 656 auto *BO = dyn_cast<BinaryOperator>(IterV); 657 if (!BO) 658 continue; 659 if (BO->getOpcode() != Instruction::Add) 660 continue; 661 Value *IncV = nullptr; 662 if (BO->getOperand(0) == PN) 663 IncV = BO->getOperand(1); 664 else if (BO->getOperand(1) == PN) 665 IncV = BO->getOperand(0); 666 if (IncV == nullptr) 667 continue; 668 669 if (auto *T = dyn_cast<ConstantInt>(IncV)) 670 if (T->getZExtValue() == 1) 671 return PN; 672 } 673 return nullptr; 674 } 675 676 static void replaceAllUsesOfWithIn(Value *I, Value *J, BasicBlock *BB) { 677 for (auto UI = I->user_begin(), UE = I->user_end(); UI != UE;) { 678 Use &TheUse = UI.getUse(); 679 ++UI; 680 if (auto *II = dyn_cast<Instruction>(TheUse.getUser())) 681 if (BB == II->getParent()) 682 II->replaceUsesOfWith(I, J); 683 } 684 } 685 686 bool PolynomialMultiplyRecognize::matchLeftShift(SelectInst *SelI, 687 Value *CIV, ParsedValues &PV) { 688 // Match the following: 689 // select (X & (1 << i)) != 0 ? R ^ (Q << i) : R 690 // select (X & (1 << i)) == 0 ? R : R ^ (Q << i) 691 // The condition may also check for equality with the masked value, i.e 692 // select (X & (1 << i)) == (1 << i) ? R ^ (Q << i) : R 693 // select (X & (1 << i)) != (1 << i) ? R : R ^ (Q << i); 694 695 Value *CondV = SelI->getCondition(); 696 Value *TrueV = SelI->getTrueValue(); 697 Value *FalseV = SelI->getFalseValue(); 698 699 using namespace PatternMatch; 700 701 CmpInst::Predicate P; 702 Value *A = nullptr, *B = nullptr, *C = nullptr; 703 704 if (!match(CondV, m_ICmp(P, m_And(m_Value(A), m_Value(B)), m_Value(C))) && 705 !match(CondV, m_ICmp(P, m_Value(C), m_And(m_Value(A), m_Value(B))))) 706 return false; 707 if (P != CmpInst::ICMP_EQ && P != CmpInst::ICMP_NE) 708 return false; 709 // Matched: select (A & B) == C ? ... : ... 710 // select (A & B) != C ? ... : ... 711 712 Value *X = nullptr, *Sh1 = nullptr; 713 // Check (A & B) for (X & (1 << i)): 714 if (match(A, m_Shl(m_One(), m_Specific(CIV)))) { 715 Sh1 = A; 716 X = B; 717 } else if (match(B, m_Shl(m_One(), m_Specific(CIV)))) { 718 Sh1 = B; 719 X = A; 720 } else { 721 // TODO: Could also check for an induction variable containing single 722 // bit shifted left by 1 in each iteration. 723 return false; 724 } 725 726 bool TrueIfZero; 727 728 // Check C against the possible values for comparison: 0 and (1 << i): 729 if (match(C, m_Zero())) 730 TrueIfZero = (P == CmpInst::ICMP_EQ); 731 else if (C == Sh1) 732 TrueIfZero = (P == CmpInst::ICMP_NE); 733 else 734 return false; 735 736 // So far, matched: 737 // select (X & (1 << i)) ? ... : ... 738 // including variations of the check against zero/non-zero value. 739 740 Value *ShouldSameV = nullptr, *ShouldXoredV = nullptr; 741 if (TrueIfZero) { 742 ShouldSameV = TrueV; 743 ShouldXoredV = FalseV; 744 } else { 745 ShouldSameV = FalseV; 746 ShouldXoredV = TrueV; 747 } 748 749 Value *Q = nullptr, *R = nullptr, *Y = nullptr, *Z = nullptr; 750 Value *T = nullptr; 751 if (match(ShouldXoredV, m_Xor(m_Value(Y), m_Value(Z)))) { 752 // Matched: select +++ ? ... : Y ^ Z 753 // select +++ ? Y ^ Z : ... 754 // where +++ denotes previously checked matches. 755 if (ShouldSameV == Y) 756 T = Z; 757 else if (ShouldSameV == Z) 758 T = Y; 759 else 760 return false; 761 R = ShouldSameV; 762 // Matched: select +++ ? R : R ^ T 763 // select +++ ? R ^ T : R 764 // depending on TrueIfZero. 765 766 } else if (match(ShouldSameV, m_Zero())) { 767 // Matched: select +++ ? 0 : ... 768 // select +++ ? ... : 0 769 if (!SelI->hasOneUse()) 770 return false; 771 T = ShouldXoredV; 772 // Matched: select +++ ? 0 : T 773 // select +++ ? T : 0 774 775 Value *U = *SelI->user_begin(); 776 if (!match(U, m_Xor(m_Specific(SelI), m_Value(R))) && 777 !match(U, m_Xor(m_Value(R), m_Specific(SelI)))) 778 return false; 779 // Matched: xor (select +++ ? 0 : T), R 780 // xor (select +++ ? T : 0), R 781 } else 782 return false; 783 784 // The xor input value T is isolated into its own match so that it could 785 // be checked against an induction variable containing a shifted bit 786 // (todo). 787 // For now, check against (Q << i). 788 if (!match(T, m_Shl(m_Value(Q), m_Specific(CIV))) && 789 !match(T, m_Shl(m_ZExt(m_Value(Q)), m_ZExt(m_Specific(CIV))))) 790 return false; 791 // Matched: select +++ ? R : R ^ (Q << i) 792 // select +++ ? R ^ (Q << i) : R 793 794 PV.X = X; 795 PV.Q = Q; 796 PV.R = R; 797 PV.Left = true; 798 return true; 799 } 800 801 bool PolynomialMultiplyRecognize::matchRightShift(SelectInst *SelI, 802 ParsedValues &PV) { 803 // Match the following: 804 // select (X & 1) != 0 ? (R >> 1) ^ Q : (R >> 1) 805 // select (X & 1) == 0 ? (R >> 1) : (R >> 1) ^ Q 806 // The condition may also check for equality with the masked value, i.e 807 // select (X & 1) == 1 ? (R >> 1) ^ Q : (R >> 1) 808 // select (X & 1) != 1 ? (R >> 1) : (R >> 1) ^ Q 809 810 Value *CondV = SelI->getCondition(); 811 Value *TrueV = SelI->getTrueValue(); 812 Value *FalseV = SelI->getFalseValue(); 813 814 using namespace PatternMatch; 815 816 Value *C = nullptr; 817 CmpInst::Predicate P; 818 bool TrueIfZero; 819 820 if (match(CondV, m_ICmp(P, m_Value(C), m_Zero())) || 821 match(CondV, m_ICmp(P, m_Zero(), m_Value(C)))) { 822 if (P != CmpInst::ICMP_EQ && P != CmpInst::ICMP_NE) 823 return false; 824 // Matched: select C == 0 ? ... : ... 825 // select C != 0 ? ... : ... 826 TrueIfZero = (P == CmpInst::ICMP_EQ); 827 } else if (match(CondV, m_ICmp(P, m_Value(C), m_One())) || 828 match(CondV, m_ICmp(P, m_One(), m_Value(C)))) { 829 if (P != CmpInst::ICMP_EQ && P != CmpInst::ICMP_NE) 830 return false; 831 // Matched: select C == 1 ? ... : ... 832 // select C != 1 ? ... : ... 833 TrueIfZero = (P == CmpInst::ICMP_NE); 834 } else 835 return false; 836 837 Value *X = nullptr; 838 if (!match(C, m_And(m_Value(X), m_One())) && 839 !match(C, m_And(m_One(), m_Value(X)))) 840 return false; 841 // Matched: select (X & 1) == +++ ? ... : ... 842 // select (X & 1) != +++ ? ... : ... 843 844 Value *R = nullptr, *Q = nullptr; 845 if (TrueIfZero) { 846 // The select's condition is true if the tested bit is 0. 847 // TrueV must be the shift, FalseV must be the xor. 848 if (!match(TrueV, m_LShr(m_Value(R), m_One()))) 849 return false; 850 // Matched: select +++ ? (R >> 1) : ... 851 if (!match(FalseV, m_Xor(m_Specific(TrueV), m_Value(Q))) && 852 !match(FalseV, m_Xor(m_Value(Q), m_Specific(TrueV)))) 853 return false; 854 // Matched: select +++ ? (R >> 1) : (R >> 1) ^ Q 855 // with commuting ^. 856 } else { 857 // The select's condition is true if the tested bit is 1. 858 // TrueV must be the xor, FalseV must be the shift. 859 if (!match(FalseV, m_LShr(m_Value(R), m_One()))) 860 return false; 861 // Matched: select +++ ? ... : (R >> 1) 862 if (!match(TrueV, m_Xor(m_Specific(FalseV), m_Value(Q))) && 863 !match(TrueV, m_Xor(m_Value(Q), m_Specific(FalseV)))) 864 return false; 865 // Matched: select +++ ? (R >> 1) ^ Q : (R >> 1) 866 // with commuting ^. 867 } 868 869 PV.X = X; 870 PV.Q = Q; 871 PV.R = R; 872 PV.Left = false; 873 return true; 874 } 875 876 bool PolynomialMultiplyRecognize::scanSelect(SelectInst *SelI, 877 BasicBlock *LoopB, BasicBlock *PrehB, Value *CIV, ParsedValues &PV, 878 bool PreScan) { 879 using namespace PatternMatch; 880 881 // The basic pattern for R = P.Q is: 882 // for i = 0..31 883 // R = phi (0, R') 884 // if (P & (1 << i)) ; test-bit(P, i) 885 // R' = R ^ (Q << i) 886 // 887 // Similarly, the basic pattern for R = (P/Q).Q - P 888 // for i = 0..31 889 // R = phi(P, R') 890 // if (R & (1 << i)) 891 // R' = R ^ (Q << i) 892 893 // There exist idioms, where instead of Q being shifted left, P is shifted 894 // right. This produces a result that is shifted right by 32 bits (the 895 // non-shifted result is 64-bit). 896 // 897 // For R = P.Q, this would be: 898 // for i = 0..31 899 // R = phi (0, R') 900 // if ((P >> i) & 1) 901 // R' = (R >> 1) ^ Q ; R is cycled through the loop, so it must 902 // else ; be shifted by 1, not i. 903 // R' = R >> 1 904 // 905 // And for the inverse: 906 // for i = 0..31 907 // R = phi (P, R') 908 // if (R & 1) 909 // R' = (R >> 1) ^ Q 910 // else 911 // R' = R >> 1 912 913 // The left-shifting idioms share the same pattern: 914 // select (X & (1 << i)) ? R ^ (Q << i) : R 915 // Similarly for right-shifting idioms: 916 // select (X & 1) ? (R >> 1) ^ Q 917 918 if (matchLeftShift(SelI, CIV, PV)) { 919 // If this is a pre-scan, getting this far is sufficient. 920 if (PreScan) 921 return true; 922 923 // Need to make sure that the SelI goes back into R. 924 auto *RPhi = dyn_cast<PHINode>(PV.R); 925 if (!RPhi) 926 return false; 927 if (SelI != RPhi->getIncomingValueForBlock(LoopB)) 928 return false; 929 PV.Res = SelI; 930 931 // If X is loop invariant, it must be the input polynomial, and the 932 // idiom is the basic polynomial multiply. 933 if (CurLoop->isLoopInvariant(PV.X)) { 934 PV.P = PV.X; 935 PV.Inv = false; 936 } else { 937 // X is not loop invariant. If X == R, this is the inverse pmpy. 938 // Otherwise, check for an xor with an invariant value. If the 939 // variable argument to the xor is R, then this is still a valid 940 // inverse pmpy. 941 PV.Inv = true; 942 if (PV.X != PV.R) { 943 Value *Var = nullptr, *Inv = nullptr, *X1 = nullptr, *X2 = nullptr; 944 if (!match(PV.X, m_Xor(m_Value(X1), m_Value(X2)))) 945 return false; 946 auto *I1 = dyn_cast<Instruction>(X1); 947 auto *I2 = dyn_cast<Instruction>(X2); 948 if (!I1 || I1->getParent() != LoopB) { 949 Var = X2; 950 Inv = X1; 951 } else if (!I2 || I2->getParent() != LoopB) { 952 Var = X1; 953 Inv = X2; 954 } else 955 return false; 956 if (Var != PV.R) 957 return false; 958 PV.M = Inv; 959 } 960 // The input polynomial P still needs to be determined. It will be 961 // the entry value of R. 962 Value *EntryP = RPhi->getIncomingValueForBlock(PrehB); 963 PV.P = EntryP; 964 } 965 966 return true; 967 } 968 969 if (matchRightShift(SelI, PV)) { 970 // If this is an inverse pattern, the Q polynomial must be known at 971 // compile time. 972 if (PV.Inv && !isa<ConstantInt>(PV.Q)) 973 return false; 974 if (PreScan) 975 return true; 976 // There is no exact matching of right-shift pmpy. 977 return false; 978 } 979 980 return false; 981 } 982 983 bool PolynomialMultiplyRecognize::isPromotableTo(Value *Val, 984 IntegerType *DestTy) { 985 IntegerType *T = dyn_cast<IntegerType>(Val->getType()); 986 if (!T || T->getBitWidth() > DestTy->getBitWidth()) 987 return false; 988 if (T->getBitWidth() == DestTy->getBitWidth()) 989 return true; 990 // Non-instructions are promotable. The reason why an instruction may not 991 // be promotable is that it may produce a different result if its operands 992 // and the result are promoted, for example, it may produce more non-zero 993 // bits. While it would still be possible to represent the proper result 994 // in a wider type, it may require adding additional instructions (which 995 // we don't want to do). 996 Instruction *In = dyn_cast<Instruction>(Val); 997 if (!In) 998 return true; 999 // The bitwidth of the source type is smaller than the destination. 1000 // Check if the individual operation can be promoted. 1001 switch (In->getOpcode()) { 1002 case Instruction::PHI: 1003 case Instruction::ZExt: 1004 case Instruction::And: 1005 case Instruction::Or: 1006 case Instruction::Xor: 1007 case Instruction::LShr: // Shift right is ok. 1008 case Instruction::Select: 1009 case Instruction::Trunc: 1010 return true; 1011 case Instruction::ICmp: 1012 if (CmpInst *CI = cast<CmpInst>(In)) 1013 return CI->isEquality() || CI->isUnsigned(); 1014 llvm_unreachable("Cast failed unexpectedly"); 1015 case Instruction::Add: 1016 return In->hasNoSignedWrap() && In->hasNoUnsignedWrap(); 1017 } 1018 return false; 1019 } 1020 1021 void PolynomialMultiplyRecognize::promoteTo(Instruction *In, 1022 IntegerType *DestTy, BasicBlock *LoopB) { 1023 Type *OrigTy = In->getType(); 1024 assert(!OrigTy->isVoidTy() && "Invalid instruction to promote"); 1025 1026 // Leave boolean values alone. 1027 if (!In->getType()->isIntegerTy(1)) 1028 In->mutateType(DestTy); 1029 unsigned DestBW = DestTy->getBitWidth(); 1030 1031 // Handle PHIs. 1032 if (PHINode *P = dyn_cast<PHINode>(In)) { 1033 unsigned N = P->getNumIncomingValues(); 1034 for (unsigned i = 0; i != N; ++i) { 1035 BasicBlock *InB = P->getIncomingBlock(i); 1036 if (InB == LoopB) 1037 continue; 1038 Value *InV = P->getIncomingValue(i); 1039 IntegerType *Ty = cast<IntegerType>(InV->getType()); 1040 // Do not promote values in PHI nodes of type i1. 1041 if (Ty != P->getType()) { 1042 // If the value type does not match the PHI type, the PHI type 1043 // must have been promoted. 1044 assert(Ty->getBitWidth() < DestBW); 1045 InV = IRBuilder<>(InB->getTerminator()).CreateZExt(InV, DestTy); 1046 P->setIncomingValue(i, InV); 1047 } 1048 } 1049 } else if (ZExtInst *Z = dyn_cast<ZExtInst>(In)) { 1050 Value *Op = Z->getOperand(0); 1051 if (Op->getType() == Z->getType()) 1052 Z->replaceAllUsesWith(Op); 1053 Z->eraseFromParent(); 1054 return; 1055 } 1056 if (TruncInst *T = dyn_cast<TruncInst>(In)) { 1057 IntegerType *TruncTy = cast<IntegerType>(OrigTy); 1058 Value *Mask = ConstantInt::get(DestTy, (1u << TruncTy->getBitWidth()) - 1); 1059 Value *And = IRBuilder<>(In).CreateAnd(T->getOperand(0), Mask); 1060 T->replaceAllUsesWith(And); 1061 T->eraseFromParent(); 1062 return; 1063 } 1064 1065 // Promote immediates. 1066 for (unsigned i = 0, n = In->getNumOperands(); i != n; ++i) { 1067 if (ConstantInt *CI = dyn_cast<ConstantInt>(In->getOperand(i))) 1068 if (CI->getType()->getBitWidth() < DestBW) 1069 In->setOperand(i, ConstantInt::get(DestTy, CI->getZExtValue())); 1070 } 1071 } 1072 1073 bool PolynomialMultiplyRecognize::promoteTypes(BasicBlock *LoopB, 1074 BasicBlock *ExitB) { 1075 assert(LoopB); 1076 // Skip loops where the exit block has more than one predecessor. The values 1077 // coming from the loop block will be promoted to another type, and so the 1078 // values coming into the exit block from other predecessors would also have 1079 // to be promoted. 1080 if (!ExitB || (ExitB->getSinglePredecessor() != LoopB)) 1081 return false; 1082 IntegerType *DestTy = getPmpyType(); 1083 // Check if the exit values have types that are no wider than the type 1084 // that we want to promote to. 1085 unsigned DestBW = DestTy->getBitWidth(); 1086 for (PHINode &P : ExitB->phis()) { 1087 if (P.getNumIncomingValues() != 1) 1088 return false; 1089 assert(P.getIncomingBlock(0) == LoopB); 1090 IntegerType *T = dyn_cast<IntegerType>(P.getType()); 1091 if (!T || T->getBitWidth() > DestBW) 1092 return false; 1093 } 1094 1095 // Check all instructions in the loop. 1096 for (Instruction &In : *LoopB) 1097 if (!In.isTerminator() && !isPromotableTo(&In, DestTy)) 1098 return false; 1099 1100 // Perform the promotion. 1101 std::vector<Instruction*> LoopIns; 1102 std::transform(LoopB->begin(), LoopB->end(), std::back_inserter(LoopIns), 1103 [](Instruction &In) { return &In; }); 1104 for (Instruction *In : LoopIns) 1105 if (!In->isTerminator()) 1106 promoteTo(In, DestTy, LoopB); 1107 1108 // Fix up the PHI nodes in the exit block. 1109 Instruction *EndI = ExitB->getFirstNonPHI(); 1110 BasicBlock::iterator End = EndI ? EndI->getIterator() : ExitB->end(); 1111 for (auto I = ExitB->begin(); I != End; ++I) { 1112 PHINode *P = dyn_cast<PHINode>(I); 1113 if (!P) 1114 break; 1115 Type *Ty0 = P->getIncomingValue(0)->getType(); 1116 Type *PTy = P->getType(); 1117 if (PTy != Ty0) { 1118 assert(Ty0 == DestTy); 1119 // In order to create the trunc, P must have the promoted type. 1120 P->mutateType(Ty0); 1121 Value *T = IRBuilder<>(ExitB, End).CreateTrunc(P, PTy); 1122 // In order for the RAUW to work, the types of P and T must match. 1123 P->mutateType(PTy); 1124 P->replaceAllUsesWith(T); 1125 // Final update of the P's type. 1126 P->mutateType(Ty0); 1127 cast<Instruction>(T)->setOperand(0, P); 1128 } 1129 } 1130 1131 return true; 1132 } 1133 1134 bool PolynomialMultiplyRecognize::findCycle(Value *Out, Value *In, 1135 ValueSeq &Cycle) { 1136 // Out = ..., In, ... 1137 if (Out == In) 1138 return true; 1139 1140 auto *BB = cast<Instruction>(Out)->getParent(); 1141 bool HadPhi = false; 1142 1143 for (auto U : Out->users()) { 1144 auto *I = dyn_cast<Instruction>(&*U); 1145 if (I == nullptr || I->getParent() != BB) 1146 continue; 1147 // Make sure that there are no multi-iteration cycles, e.g. 1148 // p1 = phi(p2) 1149 // p2 = phi(p1) 1150 // The cycle p1->p2->p1 would span two loop iterations. 1151 // Check that there is only one phi in the cycle. 1152 bool IsPhi = isa<PHINode>(I); 1153 if (IsPhi && HadPhi) 1154 return false; 1155 HadPhi |= IsPhi; 1156 if (Cycle.count(I)) 1157 return false; 1158 Cycle.insert(I); 1159 if (findCycle(I, In, Cycle)) 1160 break; 1161 Cycle.remove(I); 1162 } 1163 return !Cycle.empty(); 1164 } 1165 1166 void PolynomialMultiplyRecognize::classifyCycle(Instruction *DivI, 1167 ValueSeq &Cycle, ValueSeq &Early, ValueSeq &Late) { 1168 // All the values in the cycle that are between the phi node and the 1169 // divider instruction will be classified as "early", all other values 1170 // will be "late". 1171 1172 bool IsE = true; 1173 unsigned I, N = Cycle.size(); 1174 for (I = 0; I < N; ++I) { 1175 Value *V = Cycle[I]; 1176 if (DivI == V) 1177 IsE = false; 1178 else if (!isa<PHINode>(V)) 1179 continue; 1180 // Stop if found either. 1181 break; 1182 } 1183 // "I" is the index of either DivI or the phi node, whichever was first. 1184 // "E" is "false" or "true" respectively. 1185 ValueSeq &First = !IsE ? Early : Late; 1186 for (unsigned J = 0; J < I; ++J) 1187 First.insert(Cycle[J]); 1188 1189 ValueSeq &Second = IsE ? Early : Late; 1190 Second.insert(Cycle[I]); 1191 for (++I; I < N; ++I) { 1192 Value *V = Cycle[I]; 1193 if (DivI == V || isa<PHINode>(V)) 1194 break; 1195 Second.insert(V); 1196 } 1197 1198 for (; I < N; ++I) 1199 First.insert(Cycle[I]); 1200 } 1201 1202 bool PolynomialMultiplyRecognize::classifyInst(Instruction *UseI, 1203 ValueSeq &Early, ValueSeq &Late) { 1204 // Select is an exception, since the condition value does not have to be 1205 // classified in the same way as the true/false values. The true/false 1206 // values do have to be both early or both late. 1207 if (UseI->getOpcode() == Instruction::Select) { 1208 Value *TV = UseI->getOperand(1), *FV = UseI->getOperand(2); 1209 if (Early.count(TV) || Early.count(FV)) { 1210 if (Late.count(TV) || Late.count(FV)) 1211 return false; 1212 Early.insert(UseI); 1213 } else if (Late.count(TV) || Late.count(FV)) { 1214 if (Early.count(TV) || Early.count(FV)) 1215 return false; 1216 Late.insert(UseI); 1217 } 1218 return true; 1219 } 1220 1221 // Not sure what would be the example of this, but the code below relies 1222 // on having at least one operand. 1223 if (UseI->getNumOperands() == 0) 1224 return true; 1225 1226 bool AE = true, AL = true; 1227 for (auto &I : UseI->operands()) { 1228 if (Early.count(&*I)) 1229 AL = false; 1230 else if (Late.count(&*I)) 1231 AE = false; 1232 } 1233 // If the operands appear "all early" and "all late" at the same time, 1234 // then it means that none of them are actually classified as either. 1235 // This is harmless. 1236 if (AE && AL) 1237 return true; 1238 // Conversely, if they are neither "all early" nor "all late", then 1239 // we have a mixture of early and late operands that is not a known 1240 // exception. 1241 if (!AE && !AL) 1242 return false; 1243 1244 // Check that we have covered the two special cases. 1245 assert(AE != AL); 1246 1247 if (AE) 1248 Early.insert(UseI); 1249 else 1250 Late.insert(UseI); 1251 return true; 1252 } 1253 1254 bool PolynomialMultiplyRecognize::commutesWithShift(Instruction *I) { 1255 switch (I->getOpcode()) { 1256 case Instruction::And: 1257 case Instruction::Or: 1258 case Instruction::Xor: 1259 case Instruction::LShr: 1260 case Instruction::Shl: 1261 case Instruction::Select: 1262 case Instruction::ICmp: 1263 case Instruction::PHI: 1264 break; 1265 default: 1266 return false; 1267 } 1268 return true; 1269 } 1270 1271 bool PolynomialMultiplyRecognize::highBitsAreZero(Value *V, 1272 unsigned IterCount) { 1273 auto *T = dyn_cast<IntegerType>(V->getType()); 1274 if (!T) 1275 return false; 1276 1277 KnownBits Known(T->getBitWidth()); 1278 computeKnownBits(V, Known, DL); 1279 return Known.countMinLeadingZeros() >= IterCount; 1280 } 1281 1282 bool PolynomialMultiplyRecognize::keepsHighBitsZero(Value *V, 1283 unsigned IterCount) { 1284 // Assume that all inputs to the value have the high bits zero. 1285 // Check if the value itself preserves the zeros in the high bits. 1286 if (auto *C = dyn_cast<ConstantInt>(V)) 1287 return C->getValue().countLeadingZeros() >= IterCount; 1288 1289 if (auto *I = dyn_cast<Instruction>(V)) { 1290 switch (I->getOpcode()) { 1291 case Instruction::And: 1292 case Instruction::Or: 1293 case Instruction::Xor: 1294 case Instruction::LShr: 1295 case Instruction::Select: 1296 case Instruction::ICmp: 1297 case Instruction::PHI: 1298 case Instruction::ZExt: 1299 return true; 1300 } 1301 } 1302 1303 return false; 1304 } 1305 1306 bool PolynomialMultiplyRecognize::isOperandShifted(Instruction *I, Value *Op) { 1307 unsigned Opc = I->getOpcode(); 1308 if (Opc == Instruction::Shl || Opc == Instruction::LShr) 1309 return Op != I->getOperand(1); 1310 return true; 1311 } 1312 1313 bool PolynomialMultiplyRecognize::convertShiftsToLeft(BasicBlock *LoopB, 1314 BasicBlock *ExitB, unsigned IterCount) { 1315 Value *CIV = getCountIV(LoopB); 1316 if (CIV == nullptr) 1317 return false; 1318 auto *CIVTy = dyn_cast<IntegerType>(CIV->getType()); 1319 if (CIVTy == nullptr) 1320 return false; 1321 1322 ValueSeq RShifts; 1323 ValueSeq Early, Late, Cycled; 1324 1325 // Find all value cycles that contain logical right shifts by 1. 1326 for (Instruction &I : *LoopB) { 1327 using namespace PatternMatch; 1328 1329 Value *V = nullptr; 1330 if (!match(&I, m_LShr(m_Value(V), m_One()))) 1331 continue; 1332 ValueSeq C; 1333 if (!findCycle(&I, V, C)) 1334 continue; 1335 1336 // Found a cycle. 1337 C.insert(&I); 1338 classifyCycle(&I, C, Early, Late); 1339 Cycled.insert(C.begin(), C.end()); 1340 RShifts.insert(&I); 1341 } 1342 1343 // Find the set of all values affected by the shift cycles, i.e. all 1344 // cycled values, and (recursively) all their users. 1345 ValueSeq Users(Cycled.begin(), Cycled.end()); 1346 for (unsigned i = 0; i < Users.size(); ++i) { 1347 Value *V = Users[i]; 1348 if (!isa<IntegerType>(V->getType())) 1349 return false; 1350 auto *R = cast<Instruction>(V); 1351 // If the instruction does not commute with shifts, the loop cannot 1352 // be unshifted. 1353 if (!commutesWithShift(R)) 1354 return false; 1355 for (auto I = R->user_begin(), E = R->user_end(); I != E; ++I) { 1356 auto *T = cast<Instruction>(*I); 1357 // Skip users from outside of the loop. They will be handled later. 1358 // Also, skip the right-shifts and phi nodes, since they mix early 1359 // and late values. 1360 if (T->getParent() != LoopB || RShifts.count(T) || isa<PHINode>(T)) 1361 continue; 1362 1363 Users.insert(T); 1364 if (!classifyInst(T, Early, Late)) 1365 return false; 1366 } 1367 } 1368 1369 if (Users.empty()) 1370 return false; 1371 1372 // Verify that high bits remain zero. 1373 ValueSeq Internal(Users.begin(), Users.end()); 1374 ValueSeq Inputs; 1375 for (unsigned i = 0; i < Internal.size(); ++i) { 1376 auto *R = dyn_cast<Instruction>(Internal[i]); 1377 if (!R) 1378 continue; 1379 for (Value *Op : R->operands()) { 1380 auto *T = dyn_cast<Instruction>(Op); 1381 if (T && T->getParent() != LoopB) 1382 Inputs.insert(Op); 1383 else 1384 Internal.insert(Op); 1385 } 1386 } 1387 for (Value *V : Inputs) 1388 if (!highBitsAreZero(V, IterCount)) 1389 return false; 1390 for (Value *V : Internal) 1391 if (!keepsHighBitsZero(V, IterCount)) 1392 return false; 1393 1394 // Finally, the work can be done. Unshift each user. 1395 IRBuilder<> IRB(LoopB); 1396 std::map<Value*,Value*> ShiftMap; 1397 1398 using CastMapType = std::map<std::pair<Value *, Type *>, Value *>; 1399 1400 CastMapType CastMap; 1401 1402 auto upcast = [] (CastMapType &CM, IRBuilder<> &IRB, Value *V, 1403 IntegerType *Ty) -> Value* { 1404 auto H = CM.find(std::make_pair(V, Ty)); 1405 if (H != CM.end()) 1406 return H->second; 1407 Value *CV = IRB.CreateIntCast(V, Ty, false); 1408 CM.insert(std::make_pair(std::make_pair(V, Ty), CV)); 1409 return CV; 1410 }; 1411 1412 for (auto I = LoopB->begin(), E = LoopB->end(); I != E; ++I) { 1413 using namespace PatternMatch; 1414 1415 if (isa<PHINode>(I) || !Users.count(&*I)) 1416 continue; 1417 1418 // Match lshr x, 1. 1419 Value *V = nullptr; 1420 if (match(&*I, m_LShr(m_Value(V), m_One()))) { 1421 replaceAllUsesOfWithIn(&*I, V, LoopB); 1422 continue; 1423 } 1424 // For each non-cycled operand, replace it with the corresponding 1425 // value shifted left. 1426 for (auto &J : I->operands()) { 1427 Value *Op = J.get(); 1428 if (!isOperandShifted(&*I, Op)) 1429 continue; 1430 if (Users.count(Op)) 1431 continue; 1432 // Skip shifting zeros. 1433 if (isa<ConstantInt>(Op) && cast<ConstantInt>(Op)->isZero()) 1434 continue; 1435 // Check if we have already generated a shift for this value. 1436 auto F = ShiftMap.find(Op); 1437 Value *W = (F != ShiftMap.end()) ? F->second : nullptr; 1438 if (W == nullptr) { 1439 IRB.SetInsertPoint(&*I); 1440 // First, the shift amount will be CIV or CIV+1, depending on 1441 // whether the value is early or late. Instead of creating CIV+1, 1442 // do a single shift of the value. 1443 Value *ShAmt = CIV, *ShVal = Op; 1444 auto *VTy = cast<IntegerType>(ShVal->getType()); 1445 auto *ATy = cast<IntegerType>(ShAmt->getType()); 1446 if (Late.count(&*I)) 1447 ShVal = IRB.CreateShl(Op, ConstantInt::get(VTy, 1)); 1448 // Second, the types of the shifted value and the shift amount 1449 // must match. 1450 if (VTy != ATy) { 1451 if (VTy->getBitWidth() < ATy->getBitWidth()) 1452 ShVal = upcast(CastMap, IRB, ShVal, ATy); 1453 else 1454 ShAmt = upcast(CastMap, IRB, ShAmt, VTy); 1455 } 1456 // Ready to generate the shift and memoize it. 1457 W = IRB.CreateShl(ShVal, ShAmt); 1458 ShiftMap.insert(std::make_pair(Op, W)); 1459 } 1460 I->replaceUsesOfWith(Op, W); 1461 } 1462 } 1463 1464 // Update the users outside of the loop to account for having left 1465 // shifts. They would normally be shifted right in the loop, so shift 1466 // them right after the loop exit. 1467 // Take advantage of the loop-closed SSA form, which has all the post- 1468 // loop values in phi nodes. 1469 IRB.SetInsertPoint(ExitB, ExitB->getFirstInsertionPt()); 1470 for (auto P = ExitB->begin(), Q = ExitB->end(); P != Q; ++P) { 1471 if (!isa<PHINode>(P)) 1472 break; 1473 auto *PN = cast<PHINode>(P); 1474 Value *U = PN->getIncomingValueForBlock(LoopB); 1475 if (!Users.count(U)) 1476 continue; 1477 Value *S = IRB.CreateLShr(PN, ConstantInt::get(PN->getType(), IterCount)); 1478 PN->replaceAllUsesWith(S); 1479 // The above RAUW will create 1480 // S = lshr S, IterCount 1481 // so we need to fix it back into 1482 // S = lshr PN, IterCount 1483 cast<User>(S)->replaceUsesOfWith(S, PN); 1484 } 1485 1486 return true; 1487 } 1488 1489 void PolynomialMultiplyRecognize::cleanupLoopBody(BasicBlock *LoopB) { 1490 for (auto &I : *LoopB) 1491 if (Value *SV = SimplifyInstruction(&I, {DL, &TLI, &DT})) 1492 I.replaceAllUsesWith(SV); 1493 1494 for (auto I = LoopB->begin(), N = I; I != LoopB->end(); I = N) { 1495 N = std::next(I); 1496 RecursivelyDeleteTriviallyDeadInstructions(&*I, &TLI); 1497 } 1498 } 1499 1500 unsigned PolynomialMultiplyRecognize::getInverseMxN(unsigned QP) { 1501 // Arrays of coefficients of Q and the inverse, C. 1502 // Q[i] = coefficient at x^i. 1503 std::array<char,32> Q, C; 1504 1505 for (unsigned i = 0; i < 32; ++i) { 1506 Q[i] = QP & 1; 1507 QP >>= 1; 1508 } 1509 assert(Q[0] == 1); 1510 1511 // Find C, such that 1512 // (Q[n]*x^n + ... + Q[1]*x + Q[0]) * (C[n]*x^n + ... + C[1]*x + C[0]) = 1 1513 // 1514 // For it to have a solution, Q[0] must be 1. Since this is Z2[x], the 1515 // operations * and + are & and ^ respectively. 1516 // 1517 // Find C[i] recursively, by comparing i-th coefficient in the product 1518 // with 0 (or 1 for i=0). 1519 // 1520 // C[0] = 1, since C[0] = Q[0], and Q[0] = 1. 1521 C[0] = 1; 1522 for (unsigned i = 1; i < 32; ++i) { 1523 // Solve for C[i] in: 1524 // C[0]Q[i] ^ C[1]Q[i-1] ^ ... ^ C[i-1]Q[1] ^ C[i]Q[0] = 0 1525 // This is equivalent to 1526 // C[0]Q[i] ^ C[1]Q[i-1] ^ ... ^ C[i-1]Q[1] ^ C[i] = 0 1527 // which is 1528 // C[0]Q[i] ^ C[1]Q[i-1] ^ ... ^ C[i-1]Q[1] = C[i] 1529 unsigned T = 0; 1530 for (unsigned j = 0; j < i; ++j) 1531 T = T ^ (C[j] & Q[i-j]); 1532 C[i] = T; 1533 } 1534 1535 unsigned QV = 0; 1536 for (unsigned i = 0; i < 32; ++i) 1537 if (C[i]) 1538 QV |= (1 << i); 1539 1540 return QV; 1541 } 1542 1543 Value *PolynomialMultiplyRecognize::generate(BasicBlock::iterator At, 1544 ParsedValues &PV) { 1545 IRBuilder<> B(&*At); 1546 Module *M = At->getParent()->getParent()->getParent(); 1547 Function *PMF = Intrinsic::getDeclaration(M, Intrinsic::hexagon_M4_pmpyw); 1548 1549 Value *P = PV.P, *Q = PV.Q, *P0 = P; 1550 unsigned IC = PV.IterCount; 1551 1552 if (PV.M != nullptr) 1553 P0 = P = B.CreateXor(P, PV.M); 1554 1555 // Create a bit mask to clear the high bits beyond IterCount. 1556 auto *BMI = ConstantInt::get(P->getType(), APInt::getLowBitsSet(32, IC)); 1557 1558 if (PV.IterCount != 32) 1559 P = B.CreateAnd(P, BMI); 1560 1561 if (PV.Inv) { 1562 auto *QI = dyn_cast<ConstantInt>(PV.Q); 1563 assert(QI && QI->getBitWidth() <= 32); 1564 1565 // Again, clearing bits beyond IterCount. 1566 unsigned M = (1 << PV.IterCount) - 1; 1567 unsigned Tmp = (QI->getZExtValue() | 1) & M; 1568 unsigned QV = getInverseMxN(Tmp) & M; 1569 auto *QVI = ConstantInt::get(QI->getType(), QV); 1570 P = B.CreateCall(PMF, {P, QVI}); 1571 P = B.CreateTrunc(P, QI->getType()); 1572 if (IC != 32) 1573 P = B.CreateAnd(P, BMI); 1574 } 1575 1576 Value *R = B.CreateCall(PMF, {P, Q}); 1577 1578 if (PV.M != nullptr) 1579 R = B.CreateXor(R, B.CreateIntCast(P0, R->getType(), false)); 1580 1581 return R; 1582 } 1583 1584 static bool hasZeroSignBit(const Value *V) { 1585 if (const auto *CI = dyn_cast<const ConstantInt>(V)) 1586 return (CI->getType()->getSignBit() & CI->getSExtValue()) == 0; 1587 const Instruction *I = dyn_cast<const Instruction>(V); 1588 if (!I) 1589 return false; 1590 switch (I->getOpcode()) { 1591 case Instruction::LShr: 1592 if (const auto SI = dyn_cast<const ConstantInt>(I->getOperand(1))) 1593 return SI->getZExtValue() > 0; 1594 return false; 1595 case Instruction::Or: 1596 case Instruction::Xor: 1597 return hasZeroSignBit(I->getOperand(0)) && 1598 hasZeroSignBit(I->getOperand(1)); 1599 case Instruction::And: 1600 return hasZeroSignBit(I->getOperand(0)) || 1601 hasZeroSignBit(I->getOperand(1)); 1602 } 1603 return false; 1604 } 1605 1606 void PolynomialMultiplyRecognize::setupPreSimplifier(Simplifier &S) { 1607 S.addRule("sink-zext", 1608 // Sink zext past bitwise operations. 1609 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1610 if (I->getOpcode() != Instruction::ZExt) 1611 return nullptr; 1612 Instruction *T = dyn_cast<Instruction>(I->getOperand(0)); 1613 if (!T) 1614 return nullptr; 1615 switch (T->getOpcode()) { 1616 case Instruction::And: 1617 case Instruction::Or: 1618 case Instruction::Xor: 1619 break; 1620 default: 1621 return nullptr; 1622 } 1623 IRBuilder<> B(Ctx); 1624 return B.CreateBinOp(cast<BinaryOperator>(T)->getOpcode(), 1625 B.CreateZExt(T->getOperand(0), I->getType()), 1626 B.CreateZExt(T->getOperand(1), I->getType())); 1627 }); 1628 S.addRule("xor/and -> and/xor", 1629 // (xor (and x a) (and y a)) -> (and (xor x y) a) 1630 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1631 if (I->getOpcode() != Instruction::Xor) 1632 return nullptr; 1633 Instruction *And0 = dyn_cast<Instruction>(I->getOperand(0)); 1634 Instruction *And1 = dyn_cast<Instruction>(I->getOperand(1)); 1635 if (!And0 || !And1) 1636 return nullptr; 1637 if (And0->getOpcode() != Instruction::And || 1638 And1->getOpcode() != Instruction::And) 1639 return nullptr; 1640 if (And0->getOperand(1) != And1->getOperand(1)) 1641 return nullptr; 1642 IRBuilder<> B(Ctx); 1643 return B.CreateAnd(B.CreateXor(And0->getOperand(0), And1->getOperand(0)), 1644 And0->getOperand(1)); 1645 }); 1646 S.addRule("sink binop into select", 1647 // (Op (select c x y) z) -> (select c (Op x z) (Op y z)) 1648 // (Op x (select c y z)) -> (select c (Op x y) (Op x z)) 1649 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1650 BinaryOperator *BO = dyn_cast<BinaryOperator>(I); 1651 if (!BO) 1652 return nullptr; 1653 Instruction::BinaryOps Op = BO->getOpcode(); 1654 if (SelectInst *Sel = dyn_cast<SelectInst>(BO->getOperand(0))) { 1655 IRBuilder<> B(Ctx); 1656 Value *X = Sel->getTrueValue(), *Y = Sel->getFalseValue(); 1657 Value *Z = BO->getOperand(1); 1658 return B.CreateSelect(Sel->getCondition(), 1659 B.CreateBinOp(Op, X, Z), 1660 B.CreateBinOp(Op, Y, Z)); 1661 } 1662 if (SelectInst *Sel = dyn_cast<SelectInst>(BO->getOperand(1))) { 1663 IRBuilder<> B(Ctx); 1664 Value *X = BO->getOperand(0); 1665 Value *Y = Sel->getTrueValue(), *Z = Sel->getFalseValue(); 1666 return B.CreateSelect(Sel->getCondition(), 1667 B.CreateBinOp(Op, X, Y), 1668 B.CreateBinOp(Op, X, Z)); 1669 } 1670 return nullptr; 1671 }); 1672 S.addRule("fold select-select", 1673 // (select c (select c x y) z) -> (select c x z) 1674 // (select c x (select c y z)) -> (select c x z) 1675 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1676 SelectInst *Sel = dyn_cast<SelectInst>(I); 1677 if (!Sel) 1678 return nullptr; 1679 IRBuilder<> B(Ctx); 1680 Value *C = Sel->getCondition(); 1681 if (SelectInst *Sel0 = dyn_cast<SelectInst>(Sel->getTrueValue())) { 1682 if (Sel0->getCondition() == C) 1683 return B.CreateSelect(C, Sel0->getTrueValue(), Sel->getFalseValue()); 1684 } 1685 if (SelectInst *Sel1 = dyn_cast<SelectInst>(Sel->getFalseValue())) { 1686 if (Sel1->getCondition() == C) 1687 return B.CreateSelect(C, Sel->getTrueValue(), Sel1->getFalseValue()); 1688 } 1689 return nullptr; 1690 }); 1691 S.addRule("or-signbit -> xor-signbit", 1692 // (or (lshr x 1) 0x800.0) -> (xor (lshr x 1) 0x800.0) 1693 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1694 if (I->getOpcode() != Instruction::Or) 1695 return nullptr; 1696 ConstantInt *Msb = dyn_cast<ConstantInt>(I->getOperand(1)); 1697 if (!Msb || Msb->getZExtValue() != Msb->getType()->getSignBit()) 1698 return nullptr; 1699 if (!hasZeroSignBit(I->getOperand(0))) 1700 return nullptr; 1701 return IRBuilder<>(Ctx).CreateXor(I->getOperand(0), Msb); 1702 }); 1703 S.addRule("sink lshr into binop", 1704 // (lshr (BitOp x y) c) -> (BitOp (lshr x c) (lshr y c)) 1705 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1706 if (I->getOpcode() != Instruction::LShr) 1707 return nullptr; 1708 BinaryOperator *BitOp = dyn_cast<BinaryOperator>(I->getOperand(0)); 1709 if (!BitOp) 1710 return nullptr; 1711 switch (BitOp->getOpcode()) { 1712 case Instruction::And: 1713 case Instruction::Or: 1714 case Instruction::Xor: 1715 break; 1716 default: 1717 return nullptr; 1718 } 1719 IRBuilder<> B(Ctx); 1720 Value *S = I->getOperand(1); 1721 return B.CreateBinOp(BitOp->getOpcode(), 1722 B.CreateLShr(BitOp->getOperand(0), S), 1723 B.CreateLShr(BitOp->getOperand(1), S)); 1724 }); 1725 S.addRule("expose bitop-const", 1726 // (BitOp1 (BitOp2 x a) b) -> (BitOp2 x (BitOp1 a b)) 1727 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1728 auto IsBitOp = [](unsigned Op) -> bool { 1729 switch (Op) { 1730 case Instruction::And: 1731 case Instruction::Or: 1732 case Instruction::Xor: 1733 return true; 1734 } 1735 return false; 1736 }; 1737 BinaryOperator *BitOp1 = dyn_cast<BinaryOperator>(I); 1738 if (!BitOp1 || !IsBitOp(BitOp1->getOpcode())) 1739 return nullptr; 1740 BinaryOperator *BitOp2 = dyn_cast<BinaryOperator>(BitOp1->getOperand(0)); 1741 if (!BitOp2 || !IsBitOp(BitOp2->getOpcode())) 1742 return nullptr; 1743 ConstantInt *CA = dyn_cast<ConstantInt>(BitOp2->getOperand(1)); 1744 ConstantInt *CB = dyn_cast<ConstantInt>(BitOp1->getOperand(1)); 1745 if (!CA || !CB) 1746 return nullptr; 1747 IRBuilder<> B(Ctx); 1748 Value *X = BitOp2->getOperand(0); 1749 return B.CreateBinOp(BitOp2->getOpcode(), X, 1750 B.CreateBinOp(BitOp1->getOpcode(), CA, CB)); 1751 }); 1752 } 1753 1754 void PolynomialMultiplyRecognize::setupPostSimplifier(Simplifier &S) { 1755 S.addRule("(and (xor (and x a) y) b) -> (and (xor x y) b), if b == b&a", 1756 [](Instruction *I, LLVMContext &Ctx) -> Value* { 1757 if (I->getOpcode() != Instruction::And) 1758 return nullptr; 1759 Instruction *Xor = dyn_cast<Instruction>(I->getOperand(0)); 1760 ConstantInt *C0 = dyn_cast<ConstantInt>(I->getOperand(1)); 1761 if (!Xor || !C0) 1762 return nullptr; 1763 if (Xor->getOpcode() != Instruction::Xor) 1764 return nullptr; 1765 Instruction *And0 = dyn_cast<Instruction>(Xor->getOperand(0)); 1766 Instruction *And1 = dyn_cast<Instruction>(Xor->getOperand(1)); 1767 // Pick the first non-null and. 1768 if (!And0 || And0->getOpcode() != Instruction::And) 1769 std::swap(And0, And1); 1770 ConstantInt *C1 = dyn_cast<ConstantInt>(And0->getOperand(1)); 1771 if (!C1) 1772 return nullptr; 1773 uint32_t V0 = C0->getZExtValue(); 1774 uint32_t V1 = C1->getZExtValue(); 1775 if (V0 != (V0 & V1)) 1776 return nullptr; 1777 IRBuilder<> B(Ctx); 1778 return B.CreateAnd(B.CreateXor(And0->getOperand(0), And1), C0); 1779 }); 1780 } 1781 1782 bool PolynomialMultiplyRecognize::recognize() { 1783 LLVM_DEBUG(dbgs() << "Starting PolynomialMultiplyRecognize on loop\n" 1784 << *CurLoop << '\n'); 1785 // Restrictions: 1786 // - The loop must consist of a single block. 1787 // - The iteration count must be known at compile-time. 1788 // - The loop must have an induction variable starting from 0, and 1789 // incremented in each iteration of the loop. 1790 BasicBlock *LoopB = CurLoop->getHeader(); 1791 LLVM_DEBUG(dbgs() << "Loop header:\n" << *LoopB); 1792 1793 if (LoopB != CurLoop->getLoopLatch()) 1794 return false; 1795 BasicBlock *ExitB = CurLoop->getExitBlock(); 1796 if (ExitB == nullptr) 1797 return false; 1798 BasicBlock *EntryB = CurLoop->getLoopPreheader(); 1799 if (EntryB == nullptr) 1800 return false; 1801 1802 unsigned IterCount = 0; 1803 const SCEV *CT = SE.getBackedgeTakenCount(CurLoop); 1804 if (isa<SCEVCouldNotCompute>(CT)) 1805 return false; 1806 if (auto *CV = dyn_cast<SCEVConstant>(CT)) 1807 IterCount = CV->getValue()->getZExtValue() + 1; 1808 1809 Value *CIV = getCountIV(LoopB); 1810 ParsedValues PV; 1811 Simplifier PreSimp; 1812 PV.IterCount = IterCount; 1813 LLVM_DEBUG(dbgs() << "Loop IV: " << *CIV << "\nIterCount: " << IterCount 1814 << '\n'); 1815 1816 setupPreSimplifier(PreSimp); 1817 1818 // Perform a preliminary scan of select instructions to see if any of them 1819 // looks like a generator of the polynomial multiply steps. Assume that a 1820 // loop can only contain a single transformable operation, so stop the 1821 // traversal after the first reasonable candidate was found. 1822 // XXX: Currently this approach can modify the loop before being 100% sure 1823 // that the transformation can be carried out. 1824 bool FoundPreScan = false; 1825 auto FeedsPHI = [LoopB](const Value *V) -> bool { 1826 for (const Value *U : V->users()) { 1827 if (const auto *P = dyn_cast<const PHINode>(U)) 1828 if (P->getParent() == LoopB) 1829 return true; 1830 } 1831 return false; 1832 }; 1833 for (Instruction &In : *LoopB) { 1834 SelectInst *SI = dyn_cast<SelectInst>(&In); 1835 if (!SI || !FeedsPHI(SI)) 1836 continue; 1837 1838 Simplifier::Context C(SI); 1839 Value *T = PreSimp.simplify(C); 1840 SelectInst *SelI = (T && isa<SelectInst>(T)) ? cast<SelectInst>(T) : SI; 1841 LLVM_DEBUG(dbgs() << "scanSelect(pre-scan): " << PE(C, SelI) << '\n'); 1842 if (scanSelect(SelI, LoopB, EntryB, CIV, PV, true)) { 1843 FoundPreScan = true; 1844 if (SelI != SI) { 1845 Value *NewSel = C.materialize(LoopB, SI->getIterator()); 1846 SI->replaceAllUsesWith(NewSel); 1847 RecursivelyDeleteTriviallyDeadInstructions(SI, &TLI); 1848 } 1849 break; 1850 } 1851 } 1852 1853 if (!FoundPreScan) { 1854 LLVM_DEBUG(dbgs() << "Have not found candidates for pmpy\n"); 1855 return false; 1856 } 1857 1858 if (!PV.Left) { 1859 // The right shift version actually only returns the higher bits of 1860 // the result (each iteration discards the LSB). If we want to convert it 1861 // to a left-shifting loop, the working data type must be at least as 1862 // wide as the target's pmpy instruction. 1863 if (!promoteTypes(LoopB, ExitB)) 1864 return false; 1865 // Run post-promotion simplifications. 1866 Simplifier PostSimp; 1867 setupPostSimplifier(PostSimp); 1868 for (Instruction &In : *LoopB) { 1869 SelectInst *SI = dyn_cast<SelectInst>(&In); 1870 if (!SI || !FeedsPHI(SI)) 1871 continue; 1872 Simplifier::Context C(SI); 1873 Value *T = PostSimp.simplify(C); 1874 SelectInst *SelI = dyn_cast_or_null<SelectInst>(T); 1875 if (SelI != SI) { 1876 Value *NewSel = C.materialize(LoopB, SI->getIterator()); 1877 SI->replaceAllUsesWith(NewSel); 1878 RecursivelyDeleteTriviallyDeadInstructions(SI, &TLI); 1879 } 1880 break; 1881 } 1882 1883 if (!convertShiftsToLeft(LoopB, ExitB, IterCount)) 1884 return false; 1885 cleanupLoopBody(LoopB); 1886 } 1887 1888 // Scan the loop again, find the generating select instruction. 1889 bool FoundScan = false; 1890 for (Instruction &In : *LoopB) { 1891 SelectInst *SelI = dyn_cast<SelectInst>(&In); 1892 if (!SelI) 1893 continue; 1894 LLVM_DEBUG(dbgs() << "scanSelect: " << *SelI << '\n'); 1895 FoundScan = scanSelect(SelI, LoopB, EntryB, CIV, PV, false); 1896 if (FoundScan) 1897 break; 1898 } 1899 assert(FoundScan); 1900 1901 LLVM_DEBUG({ 1902 StringRef PP = (PV.M ? "(P+M)" : "P"); 1903 if (!PV.Inv) 1904 dbgs() << "Found pmpy idiom: R = " << PP << ".Q\n"; 1905 else 1906 dbgs() << "Found inverse pmpy idiom: R = (" << PP << "/Q).Q) + " 1907 << PP << "\n"; 1908 dbgs() << " Res:" << *PV.Res << "\n P:" << *PV.P << "\n"; 1909 if (PV.M) 1910 dbgs() << " M:" << *PV.M << "\n"; 1911 dbgs() << " Q:" << *PV.Q << "\n"; 1912 dbgs() << " Iteration count:" << PV.IterCount << "\n"; 1913 }); 1914 1915 BasicBlock::iterator At(EntryB->getTerminator()); 1916 Value *PM = generate(At, PV); 1917 if (PM == nullptr) 1918 return false; 1919 1920 if (PM->getType() != PV.Res->getType()) 1921 PM = IRBuilder<>(&*At).CreateIntCast(PM, PV.Res->getType(), false); 1922 1923 PV.Res->replaceAllUsesWith(PM); 1924 PV.Res->eraseFromParent(); 1925 return true; 1926 } 1927 1928 int HexagonLoopIdiomRecognize::getSCEVStride(const SCEVAddRecExpr *S) { 1929 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(S->getOperand(1))) 1930 return SC->getAPInt().getSExtValue(); 1931 return 0; 1932 } 1933 1934 bool HexagonLoopIdiomRecognize::isLegalStore(Loop *CurLoop, StoreInst *SI) { 1935 // Allow volatile stores if HexagonVolatileMemcpy is enabled. 1936 if (!(SI->isVolatile() && HexagonVolatileMemcpy) && !SI->isSimple()) 1937 return false; 1938 1939 Value *StoredVal = SI->getValueOperand(); 1940 Value *StorePtr = SI->getPointerOperand(); 1941 1942 // Reject stores that are so large that they overflow an unsigned. 1943 uint64_t SizeInBits = DL->getTypeSizeInBits(StoredVal->getType()); 1944 if ((SizeInBits & 7) || (SizeInBits >> 32) != 0) 1945 return false; 1946 1947 // See if the pointer expression is an AddRec like {base,+,1} on the current 1948 // loop, which indicates a strided store. If we have something else, it's a 1949 // random store we can't handle. 1950 auto *StoreEv = dyn_cast<SCEVAddRecExpr>(SE->getSCEV(StorePtr)); 1951 if (!StoreEv || StoreEv->getLoop() != CurLoop || !StoreEv->isAffine()) 1952 return false; 1953 1954 // Check to see if the stride matches the size of the store. If so, then we 1955 // know that every byte is touched in the loop. 1956 int Stride = getSCEVStride(StoreEv); 1957 if (Stride == 0) 1958 return false; 1959 unsigned StoreSize = DL->getTypeStoreSize(SI->getValueOperand()->getType()); 1960 if (StoreSize != unsigned(std::abs(Stride))) 1961 return false; 1962 1963 // The store must be feeding a non-volatile load. 1964 LoadInst *LI = dyn_cast<LoadInst>(SI->getValueOperand()); 1965 if (!LI || !LI->isSimple()) 1966 return false; 1967 1968 // See if the pointer expression is an AddRec like {base,+,1} on the current 1969 // loop, which indicates a strided load. If we have something else, it's a 1970 // random load we can't handle. 1971 Value *LoadPtr = LI->getPointerOperand(); 1972 auto *LoadEv = dyn_cast<SCEVAddRecExpr>(SE->getSCEV(LoadPtr)); 1973 if (!LoadEv || LoadEv->getLoop() != CurLoop || !LoadEv->isAffine()) 1974 return false; 1975 1976 // The store and load must share the same stride. 1977 if (StoreEv->getOperand(1) != LoadEv->getOperand(1)) 1978 return false; 1979 1980 // Success. This store can be converted into a memcpy. 1981 return true; 1982 } 1983 1984 /// mayLoopAccessLocation - Return true if the specified loop might access the 1985 /// specified pointer location, which is a loop-strided access. The 'Access' 1986 /// argument specifies what the verboten forms of access are (read or write). 1987 static bool 1988 mayLoopAccessLocation(Value *Ptr, ModRefInfo Access, Loop *L, 1989 const SCEV *BECount, unsigned StoreSize, 1990 AliasAnalysis &AA, 1991 SmallPtrSetImpl<Instruction *> &Ignored) { 1992 // Get the location that may be stored across the loop. Since the access 1993 // is strided positively through memory, we say that the modified location 1994 // starts at the pointer and has infinite size. 1995 LocationSize AccessSize = LocationSize::unknown(); 1996 1997 // If the loop iterates a fixed number of times, we can refine the access 1998 // size to be exactly the size of the memset, which is (BECount+1)*StoreSize 1999 if (const SCEVConstant *BECst = dyn_cast<SCEVConstant>(BECount)) 2000 AccessSize = LocationSize::precise((BECst->getValue()->getZExtValue() + 1) * 2001 StoreSize); 2002 2003 // TODO: For this to be really effective, we have to dive into the pointer 2004 // operand in the store. Store to &A[i] of 100 will always return may alias 2005 // with store of &A[100], we need to StoreLoc to be "A" with size of 100, 2006 // which will then no-alias a store to &A[100]. 2007 MemoryLocation StoreLoc(Ptr, AccessSize); 2008 2009 for (auto *B : L->blocks()) 2010 for (auto &I : *B) 2011 if (Ignored.count(&I) == 0 && 2012 isModOrRefSet( 2013 intersectModRef(AA.getModRefInfo(&I, StoreLoc), Access))) 2014 return true; 2015 2016 return false; 2017 } 2018 2019 void HexagonLoopIdiomRecognize::collectStores(Loop *CurLoop, BasicBlock *BB, 2020 SmallVectorImpl<StoreInst*> &Stores) { 2021 Stores.clear(); 2022 for (Instruction &I : *BB) 2023 if (StoreInst *SI = dyn_cast<StoreInst>(&I)) 2024 if (isLegalStore(CurLoop, SI)) 2025 Stores.push_back(SI); 2026 } 2027 2028 bool HexagonLoopIdiomRecognize::processCopyingStore(Loop *CurLoop, 2029 StoreInst *SI, const SCEV *BECount) { 2030 assert((SI->isSimple() || (SI->isVolatile() && HexagonVolatileMemcpy)) && 2031 "Expected only non-volatile stores, or Hexagon-specific memcpy" 2032 "to volatile destination."); 2033 2034 Value *StorePtr = SI->getPointerOperand(); 2035 auto *StoreEv = cast<SCEVAddRecExpr>(SE->getSCEV(StorePtr)); 2036 unsigned Stride = getSCEVStride(StoreEv); 2037 unsigned StoreSize = DL->getTypeStoreSize(SI->getValueOperand()->getType()); 2038 if (Stride != StoreSize) 2039 return false; 2040 2041 // See if the pointer expression is an AddRec like {base,+,1} on the current 2042 // loop, which indicates a strided load. If we have something else, it's a 2043 // random load we can't handle. 2044 auto *LI = cast<LoadInst>(SI->getValueOperand()); 2045 auto *LoadEv = cast<SCEVAddRecExpr>(SE->getSCEV(LI->getPointerOperand())); 2046 2047 // The trip count of the loop and the base pointer of the addrec SCEV is 2048 // guaranteed to be loop invariant, which means that it should dominate the 2049 // header. This allows us to insert code for it in the preheader. 2050 BasicBlock *Preheader = CurLoop->getLoopPreheader(); 2051 Instruction *ExpPt = Preheader->getTerminator(); 2052 IRBuilder<> Builder(ExpPt); 2053 SCEVExpander Expander(*SE, *DL, "hexagon-loop-idiom"); 2054 2055 Type *IntPtrTy = Builder.getIntPtrTy(*DL, SI->getPointerAddressSpace()); 2056 2057 // Okay, we have a strided store "p[i]" of a loaded value. We can turn 2058 // this into a memcpy/memmove in the loop preheader now if we want. However, 2059 // this would be unsafe to do if there is anything else in the loop that may 2060 // read or write the memory region we're storing to. For memcpy, this 2061 // includes the load that feeds the stores. Check for an alias by generating 2062 // the base address and checking everything. 2063 Value *StoreBasePtr = Expander.expandCodeFor(StoreEv->getStart(), 2064 Builder.getInt8PtrTy(SI->getPointerAddressSpace()), ExpPt); 2065 Value *LoadBasePtr = nullptr; 2066 2067 bool Overlap = false; 2068 bool DestVolatile = SI->isVolatile(); 2069 Type *BECountTy = BECount->getType(); 2070 2071 if (DestVolatile) { 2072 // The trip count must fit in i32, since it is the type of the "num_words" 2073 // argument to hexagon_memcpy_forward_vp4cp4n2. 2074 if (StoreSize != 4 || DL->getTypeSizeInBits(BECountTy) > 32) { 2075 CleanupAndExit: 2076 // If we generated new code for the base pointer, clean up. 2077 Expander.clear(); 2078 if (StoreBasePtr && (LoadBasePtr != StoreBasePtr)) { 2079 RecursivelyDeleteTriviallyDeadInstructions(StoreBasePtr, TLI); 2080 StoreBasePtr = nullptr; 2081 } 2082 if (LoadBasePtr) { 2083 RecursivelyDeleteTriviallyDeadInstructions(LoadBasePtr, TLI); 2084 LoadBasePtr = nullptr; 2085 } 2086 return false; 2087 } 2088 } 2089 2090 SmallPtrSet<Instruction*, 2> Ignore1; 2091 Ignore1.insert(SI); 2092 if (mayLoopAccessLocation(StoreBasePtr, ModRefInfo::ModRef, CurLoop, BECount, 2093 StoreSize, *AA, Ignore1)) { 2094 // Check if the load is the offending instruction. 2095 Ignore1.insert(LI); 2096 if (mayLoopAccessLocation(StoreBasePtr, ModRefInfo::ModRef, CurLoop, 2097 BECount, StoreSize, *AA, Ignore1)) { 2098 // Still bad. Nothing we can do. 2099 goto CleanupAndExit; 2100 } 2101 // It worked with the load ignored. 2102 Overlap = true; 2103 } 2104 2105 if (!Overlap) { 2106 if (DisableMemcpyIdiom || !HasMemcpy) 2107 goto CleanupAndExit; 2108 } else { 2109 // Don't generate memmove if this function will be inlined. This is 2110 // because the caller will undergo this transformation after inlining. 2111 Function *Func = CurLoop->getHeader()->getParent(); 2112 if (Func->hasFnAttribute(Attribute::AlwaysInline)) 2113 goto CleanupAndExit; 2114 2115 // In case of a memmove, the call to memmove will be executed instead 2116 // of the loop, so we need to make sure that there is nothing else in 2117 // the loop than the load, store and instructions that these two depend 2118 // on. 2119 SmallVector<Instruction*,2> Insts; 2120 Insts.push_back(SI); 2121 Insts.push_back(LI); 2122 if (!coverLoop(CurLoop, Insts)) 2123 goto CleanupAndExit; 2124 2125 if (DisableMemmoveIdiom || !HasMemmove) 2126 goto CleanupAndExit; 2127 bool IsNested = CurLoop->getParentLoop() != nullptr; 2128 if (IsNested && OnlyNonNestedMemmove) 2129 goto CleanupAndExit; 2130 } 2131 2132 // For a memcpy, we have to make sure that the input array is not being 2133 // mutated by the loop. 2134 LoadBasePtr = Expander.expandCodeFor(LoadEv->getStart(), 2135 Builder.getInt8PtrTy(LI->getPointerAddressSpace()), ExpPt); 2136 2137 SmallPtrSet<Instruction*, 2> Ignore2; 2138 Ignore2.insert(SI); 2139 if (mayLoopAccessLocation(LoadBasePtr, ModRefInfo::Mod, CurLoop, BECount, 2140 StoreSize, *AA, Ignore2)) 2141 goto CleanupAndExit; 2142 2143 // Check the stride. 2144 bool StridePos = getSCEVStride(LoadEv) >= 0; 2145 2146 // Currently, the volatile memcpy only emulates traversing memory forward. 2147 if (!StridePos && DestVolatile) 2148 goto CleanupAndExit; 2149 2150 bool RuntimeCheck = (Overlap || DestVolatile); 2151 2152 BasicBlock *ExitB; 2153 if (RuntimeCheck) { 2154 // The runtime check needs a single exit block. 2155 SmallVector<BasicBlock*, 8> ExitBlocks; 2156 CurLoop->getUniqueExitBlocks(ExitBlocks); 2157 if (ExitBlocks.size() != 1) 2158 goto CleanupAndExit; 2159 ExitB = ExitBlocks[0]; 2160 } 2161 2162 // The # stored bytes is (BECount+1)*Size. Expand the trip count out to 2163 // pointer size if it isn't already. 2164 LLVMContext &Ctx = SI->getContext(); 2165 BECount = SE->getTruncateOrZeroExtend(BECount, IntPtrTy); 2166 DebugLoc DLoc = SI->getDebugLoc(); 2167 2168 const SCEV *NumBytesS = 2169 SE->getAddExpr(BECount, SE->getOne(IntPtrTy), SCEV::FlagNUW); 2170 if (StoreSize != 1) 2171 NumBytesS = SE->getMulExpr(NumBytesS, SE->getConstant(IntPtrTy, StoreSize), 2172 SCEV::FlagNUW); 2173 Value *NumBytes = Expander.expandCodeFor(NumBytesS, IntPtrTy, ExpPt); 2174 if (Instruction *In = dyn_cast<Instruction>(NumBytes)) 2175 if (Value *Simp = SimplifyInstruction(In, {*DL, TLI, DT})) 2176 NumBytes = Simp; 2177 2178 CallInst *NewCall; 2179 2180 if (RuntimeCheck) { 2181 unsigned Threshold = RuntimeMemSizeThreshold; 2182 if (ConstantInt *CI = dyn_cast<ConstantInt>(NumBytes)) { 2183 uint64_t C = CI->getZExtValue(); 2184 if (Threshold != 0 && C < Threshold) 2185 goto CleanupAndExit; 2186 if (C < CompileTimeMemSizeThreshold) 2187 goto CleanupAndExit; 2188 } 2189 2190 BasicBlock *Header = CurLoop->getHeader(); 2191 Function *Func = Header->getParent(); 2192 Loop *ParentL = LF->getLoopFor(Preheader); 2193 StringRef HeaderName = Header->getName(); 2194 2195 // Create a new (empty) preheader, and update the PHI nodes in the 2196 // header to use the new preheader. 2197 BasicBlock *NewPreheader = BasicBlock::Create(Ctx, HeaderName+".rtli.ph", 2198 Func, Header); 2199 if (ParentL) 2200 ParentL->addBasicBlockToLoop(NewPreheader, *LF); 2201 IRBuilder<>(NewPreheader).CreateBr(Header); 2202 for (auto &In : *Header) { 2203 PHINode *PN = dyn_cast<PHINode>(&In); 2204 if (!PN) 2205 break; 2206 int bx = PN->getBasicBlockIndex(Preheader); 2207 if (bx >= 0) 2208 PN->setIncomingBlock(bx, NewPreheader); 2209 } 2210 DT->addNewBlock(NewPreheader, Preheader); 2211 DT->changeImmediateDominator(Header, NewPreheader); 2212 2213 // Check for safe conditions to execute memmove. 2214 // If stride is positive, copying things from higher to lower addresses 2215 // is equivalent to memmove. For negative stride, it's the other way 2216 // around. Copying forward in memory with positive stride may not be 2217 // same as memmove since we may be copying values that we just stored 2218 // in some previous iteration. 2219 Value *LA = Builder.CreatePtrToInt(LoadBasePtr, IntPtrTy); 2220 Value *SA = Builder.CreatePtrToInt(StoreBasePtr, IntPtrTy); 2221 Value *LowA = StridePos ? SA : LA; 2222 Value *HighA = StridePos ? LA : SA; 2223 Value *CmpA = Builder.CreateICmpULT(LowA, HighA); 2224 Value *Cond = CmpA; 2225 2226 // Check for distance between pointers. Since the case LowA < HighA 2227 // is checked for above, assume LowA >= HighA. 2228 Value *Dist = Builder.CreateSub(LowA, HighA); 2229 Value *CmpD = Builder.CreateICmpSLE(NumBytes, Dist); 2230 Value *CmpEither = Builder.CreateOr(Cond, CmpD); 2231 Cond = CmpEither; 2232 2233 if (Threshold != 0) { 2234 Type *Ty = NumBytes->getType(); 2235 Value *Thr = ConstantInt::get(Ty, Threshold); 2236 Value *CmpB = Builder.CreateICmpULT(Thr, NumBytes); 2237 Value *CmpBoth = Builder.CreateAnd(Cond, CmpB); 2238 Cond = CmpBoth; 2239 } 2240 BasicBlock *MemmoveB = BasicBlock::Create(Ctx, Header->getName()+".rtli", 2241 Func, NewPreheader); 2242 if (ParentL) 2243 ParentL->addBasicBlockToLoop(MemmoveB, *LF); 2244 Instruction *OldT = Preheader->getTerminator(); 2245 Builder.CreateCondBr(Cond, MemmoveB, NewPreheader); 2246 OldT->eraseFromParent(); 2247 Preheader->setName(Preheader->getName()+".old"); 2248 DT->addNewBlock(MemmoveB, Preheader); 2249 // Find the new immediate dominator of the exit block. 2250 BasicBlock *ExitD = Preheader; 2251 for (auto PI = pred_begin(ExitB), PE = pred_end(ExitB); PI != PE; ++PI) { 2252 BasicBlock *PB = *PI; 2253 ExitD = DT->findNearestCommonDominator(ExitD, PB); 2254 if (!ExitD) 2255 break; 2256 } 2257 // If the prior immediate dominator of ExitB was dominated by the 2258 // old preheader, then the old preheader becomes the new immediate 2259 // dominator. Otherwise don't change anything (because the newly 2260 // added blocks are dominated by the old preheader). 2261 if (ExitD && DT->dominates(Preheader, ExitD)) { 2262 DomTreeNode *BN = DT->getNode(ExitB); 2263 DomTreeNode *DN = DT->getNode(ExitD); 2264 BN->setIDom(DN); 2265 } 2266 2267 // Add a call to memmove to the conditional block. 2268 IRBuilder<> CondBuilder(MemmoveB); 2269 CondBuilder.CreateBr(ExitB); 2270 CondBuilder.SetInsertPoint(MemmoveB->getTerminator()); 2271 2272 if (DestVolatile) { 2273 Type *Int32Ty = Type::getInt32Ty(Ctx); 2274 Type *Int32PtrTy = Type::getInt32PtrTy(Ctx); 2275 Type *VoidTy = Type::getVoidTy(Ctx); 2276 Module *M = Func->getParent(); 2277 FunctionCallee Fn = M->getOrInsertFunction( 2278 HexagonVolatileMemcpyName, VoidTy, Int32PtrTy, Int32PtrTy, Int32Ty); 2279 2280 const SCEV *OneS = SE->getConstant(Int32Ty, 1); 2281 const SCEV *BECount32 = SE->getTruncateOrZeroExtend(BECount, Int32Ty); 2282 const SCEV *NumWordsS = SE->getAddExpr(BECount32, OneS, SCEV::FlagNUW); 2283 Value *NumWords = Expander.expandCodeFor(NumWordsS, Int32Ty, 2284 MemmoveB->getTerminator()); 2285 if (Instruction *In = dyn_cast<Instruction>(NumWords)) 2286 if (Value *Simp = SimplifyInstruction(In, {*DL, TLI, DT})) 2287 NumWords = Simp; 2288 2289 Value *Op0 = (StoreBasePtr->getType() == Int32PtrTy) 2290 ? StoreBasePtr 2291 : CondBuilder.CreateBitCast(StoreBasePtr, Int32PtrTy); 2292 Value *Op1 = (LoadBasePtr->getType() == Int32PtrTy) 2293 ? LoadBasePtr 2294 : CondBuilder.CreateBitCast(LoadBasePtr, Int32PtrTy); 2295 NewCall = CondBuilder.CreateCall(Fn, {Op0, Op1, NumWords}); 2296 } else { 2297 NewCall = CondBuilder.CreateMemMove( 2298 StoreBasePtr, SI->getAlign(), LoadBasePtr, LI->getAlign(), NumBytes); 2299 } 2300 } else { 2301 NewCall = Builder.CreateMemCpy(StoreBasePtr, SI->getAlign(), LoadBasePtr, 2302 LI->getAlign(), NumBytes); 2303 // Okay, the memcpy has been formed. Zap the original store and 2304 // anything that feeds into it. 2305 RecursivelyDeleteTriviallyDeadInstructions(SI, TLI); 2306 } 2307 2308 NewCall->setDebugLoc(DLoc); 2309 2310 LLVM_DEBUG(dbgs() << " Formed " << (Overlap ? "memmove: " : "memcpy: ") 2311 << *NewCall << "\n" 2312 << " from load ptr=" << *LoadEv << " at: " << *LI << "\n" 2313 << " from store ptr=" << *StoreEv << " at: " << *SI 2314 << "\n"); 2315 2316 return true; 2317 } 2318 2319 // Check if the instructions in Insts, together with their dependencies 2320 // cover the loop in the sense that the loop could be safely eliminated once 2321 // the instructions in Insts are removed. 2322 bool HexagonLoopIdiomRecognize::coverLoop(Loop *L, 2323 SmallVectorImpl<Instruction*> &Insts) const { 2324 SmallSet<BasicBlock*,8> LoopBlocks; 2325 for (auto *B : L->blocks()) 2326 LoopBlocks.insert(B); 2327 2328 SetVector<Instruction*> Worklist(Insts.begin(), Insts.end()); 2329 2330 // Collect all instructions from the loop that the instructions in Insts 2331 // depend on (plus their dependencies, etc.). These instructions will 2332 // constitute the expression trees that feed those in Insts, but the trees 2333 // will be limited only to instructions contained in the loop. 2334 for (unsigned i = 0; i < Worklist.size(); ++i) { 2335 Instruction *In = Worklist[i]; 2336 for (auto I = In->op_begin(), E = In->op_end(); I != E; ++I) { 2337 Instruction *OpI = dyn_cast<Instruction>(I); 2338 if (!OpI) 2339 continue; 2340 BasicBlock *PB = OpI->getParent(); 2341 if (!LoopBlocks.count(PB)) 2342 continue; 2343 Worklist.insert(OpI); 2344 } 2345 } 2346 2347 // Scan all instructions in the loop, if any of them have a user outside 2348 // of the loop, or outside of the expressions collected above, then either 2349 // the loop has a side-effect visible outside of it, or there are 2350 // instructions in it that are not involved in the original set Insts. 2351 for (auto *B : L->blocks()) { 2352 for (auto &In : *B) { 2353 if (isa<BranchInst>(In) || isa<DbgInfoIntrinsic>(In)) 2354 continue; 2355 if (!Worklist.count(&In) && In.mayHaveSideEffects()) 2356 return false; 2357 for (auto K : In.users()) { 2358 Instruction *UseI = dyn_cast<Instruction>(K); 2359 if (!UseI) 2360 continue; 2361 BasicBlock *UseB = UseI->getParent(); 2362 if (LF->getLoopFor(UseB) != L) 2363 return false; 2364 } 2365 } 2366 } 2367 2368 return true; 2369 } 2370 2371 /// runOnLoopBlock - Process the specified block, which lives in a counted loop 2372 /// with the specified backedge count. This block is known to be in the current 2373 /// loop and not in any subloops. 2374 bool HexagonLoopIdiomRecognize::runOnLoopBlock(Loop *CurLoop, BasicBlock *BB, 2375 const SCEV *BECount, SmallVectorImpl<BasicBlock*> &ExitBlocks) { 2376 // We can only promote stores in this block if they are unconditionally 2377 // executed in the loop. For a block to be unconditionally executed, it has 2378 // to dominate all the exit blocks of the loop. Verify this now. 2379 auto DominatedByBB = [this,BB] (BasicBlock *EB) -> bool { 2380 return DT->dominates(BB, EB); 2381 }; 2382 if (!all_of(ExitBlocks, DominatedByBB)) 2383 return false; 2384 2385 bool MadeChange = false; 2386 // Look for store instructions, which may be optimized to memset/memcpy. 2387 SmallVector<StoreInst*,8> Stores; 2388 collectStores(CurLoop, BB, Stores); 2389 2390 // Optimize the store into a memcpy, if it feeds an similarly strided load. 2391 for (auto &SI : Stores) 2392 MadeChange |= processCopyingStore(CurLoop, SI, BECount); 2393 2394 return MadeChange; 2395 } 2396 2397 bool HexagonLoopIdiomRecognize::runOnCountableLoop(Loop *L) { 2398 PolynomialMultiplyRecognize PMR(L, *DL, *DT, *TLI, *SE); 2399 if (PMR.recognize()) 2400 return true; 2401 2402 if (!HasMemcpy && !HasMemmove) 2403 return false; 2404 2405 const SCEV *BECount = SE->getBackedgeTakenCount(L); 2406 assert(!isa<SCEVCouldNotCompute>(BECount) && 2407 "runOnCountableLoop() called on a loop without a predictable" 2408 "backedge-taken count"); 2409 2410 SmallVector<BasicBlock *, 8> ExitBlocks; 2411 L->getUniqueExitBlocks(ExitBlocks); 2412 2413 bool Changed = false; 2414 2415 // Scan all the blocks in the loop that are not in subloops. 2416 for (auto *BB : L->getBlocks()) { 2417 // Ignore blocks in subloops. 2418 if (LF->getLoopFor(BB) != L) 2419 continue; 2420 Changed |= runOnLoopBlock(L, BB, BECount, ExitBlocks); 2421 } 2422 2423 return Changed; 2424 } 2425 2426 bool HexagonLoopIdiomRecognize::run(Loop *L) { 2427 const Module &M = *L->getHeader()->getParent()->getParent(); 2428 if (Triple(M.getTargetTriple()).getArch() != Triple::hexagon) 2429 return false; 2430 2431 // If the loop could not be converted to canonical form, it must have an 2432 // indirectbr in it, just give up. 2433 if (!L->getLoopPreheader()) 2434 return false; 2435 2436 // Disable loop idiom recognition if the function's name is a common idiom. 2437 StringRef Name = L->getHeader()->getParent()->getName(); 2438 if (Name == "memset" || Name == "memcpy" || Name == "memmove") 2439 return false; 2440 2441 DL = &L->getHeader()->getModule()->getDataLayout(); 2442 2443 HasMemcpy = TLI->has(LibFunc_memcpy); 2444 HasMemmove = TLI->has(LibFunc_memmove); 2445 2446 if (SE->hasLoopInvariantBackedgeTakenCount(L)) 2447 return runOnCountableLoop(L); 2448 return false; 2449 } 2450 2451 bool HexagonLoopIdiomRecognizeLegacyPass::runOnLoop(Loop *L, 2452 LPPassManager &LPM) { 2453 if (skipLoop(L)) 2454 return false; 2455 2456 auto *AA = &getAnalysis<AAResultsWrapperPass>().getAAResults(); 2457 auto *DT = &getAnalysis<DominatorTreeWrapperPass>().getDomTree(); 2458 auto *LF = &getAnalysis<LoopInfoWrapperPass>().getLoopInfo(); 2459 auto *TLI = &getAnalysis<TargetLibraryInfoWrapperPass>().getTLI( 2460 *L->getHeader()->getParent()); 2461 auto *SE = &getAnalysis<ScalarEvolutionWrapperPass>().getSE(); 2462 return HexagonLoopIdiomRecognize(AA, DT, LF, TLI, SE).run(L); 2463 } 2464 2465 Pass *llvm::createHexagonLoopIdiomPass() { 2466 return new HexagonLoopIdiomRecognizeLegacyPass(); 2467 } 2468 2469 PreservedAnalyses 2470 HexagonLoopIdiomRecognitionPass::run(Loop &L, LoopAnalysisManager &AM, 2471 LoopStandardAnalysisResults &AR, 2472 LPMUpdater &U) { 2473 return HexagonLoopIdiomRecognize(&AR.AA, &AR.DT, &AR.LI, &AR.TLI, &AR.SE) 2474 .run(&L) 2475 ? getLoopPassPreservedAnalyses() 2476 : PreservedAnalyses::all(); 2477 } 2478