brintos

brintos / llvm-project-archived public Read only

0
0
Text · 15.8 KiB · 49373d3 Raw
506 lines · cpp
1//===- FunctionSpecializationTest.cpp - Cost model unit tests -------------===//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-exception6//7//===----------------------------------------------------------------------===//8 9#include "llvm/Analysis/AssumptionCache.h"10#include "llvm/Analysis/BlockFrequencyInfo.h"11#include "llvm/Analysis/BranchProbabilityInfo.h"12#include "llvm/Analysis/LoopInfo.h"13#include "llvm/Analysis/PostDominators.h"14#include "llvm/Analysis/TargetLibraryInfo.h"15#include "llvm/Analysis/TargetTransformInfo.h"16#include "llvm/AsmParser/Parser.h"17#include "llvm/IR/Constants.h"18#include "llvm/IR/PassInstrumentation.h"19#include "llvm/Support/SourceMgr.h"20#include "llvm/Transforms/IPO/FunctionSpecialization.h"21#include "llvm/Transforms/Utils/SCCPSolver.h"22#include "gtest/gtest.h"23#include <memory>24 25namespace llvm {26 27static void removeSSACopy(Function &F) {28  for (BasicBlock &BB : F) {29    for (Instruction &Inst : llvm::make_early_inc_range(BB)) {30      auto *BC = dyn_cast<BitCastInst>(&Inst);31      if (!BC || BC->getType() != BC->getOperand(0)->getType())32        continue;33      Inst.replaceAllUsesWith(BC->getOperand(0));34      Inst.eraseFromParent();35    }36  }37}38 39class FunctionSpecializationTest : public testing::Test {40protected:41  LLVMContext Ctx;42  FunctionAnalysisManager FAM;43  std::unique_ptr<Module> M;44  std::unique_ptr<SCCPSolver> Solver;45  SmallVector<Instruction *, 8> KnownConstants;46 47  FunctionSpecializationTest() {48    FAM.registerPass([&] { return TargetLibraryAnalysis(); });49    FAM.registerPass([&] { return TargetIRAnalysis(); });50    FAM.registerPass([&] { return BlockFrequencyAnalysis(); });51    FAM.registerPass([&] { return BranchProbabilityAnalysis(); });52    FAM.registerPass([&] { return LoopAnalysis(); });53    FAM.registerPass([&] { return AssumptionAnalysis(); });54    FAM.registerPass([&] { return DominatorTreeAnalysis(); });55    FAM.registerPass([&] { return PostDominatorTreeAnalysis(); });56    FAM.registerPass([&] { return PassInstrumentationAnalysis(); });57  }58 59  Module &parseModule(const char *ModuleString) {60    SMDiagnostic Err;61    M = parseAssemblyString(ModuleString, Err, Ctx);62    EXPECT_TRUE(M);63    return *M;64  }65 66  FunctionSpecializer getSpecializerFor(Function *F) {67    auto GetTLI = [this](Function &F) -> const TargetLibraryInfo & {68      return FAM.getResult<TargetLibraryAnalysis>(F);69    };70    auto GetTTI = [this](Function &F) -> TargetTransformInfo & {71      return FAM.getResult<TargetIRAnalysis>(F);72    };73    auto GetAC = [this](Function &F) -> AssumptionCache & {74      return FAM.getResult<AssumptionAnalysis>(F);75    };76    auto GetDT = [this](Function &F) -> DominatorTree & {77      return FAM.getResult<DominatorTreeAnalysis>(F);78    };79    auto GetBFI = [this](Function &F) -> BlockFrequencyInfo & {80      return FAM.getResult<BlockFrequencyAnalysis>(F);81    };82 83    Solver = std::make_unique<SCCPSolver>(M->getDataLayout(), GetTLI, Ctx);84 85    DominatorTree &DT = GetDT(*F);86    AssumptionCache &AC = GetAC(*F);87    Solver->addPredicateInfo(*F, DT, AC);88 89    Solver->markBlockExecutable(&F->front());90    for (Argument &Arg : F->args())91      Solver->markOverdefined(&Arg);92    Solver->solveWhileResolvedUndefsIn(*M);93 94    removeSSACopy(*F);95 96    return FunctionSpecializer(*Solver, *M, &FAM, GetBFI, GetTLI, GetTTI,97                               GetAC);98  }99 100  Cost getCodeSizeSavings(Instruction &I, bool HasLatencySavings = true) {101    auto &TTI = FAM.getResult<TargetIRAnalysis>(*I.getFunction());102 103    Cost CodeSize =104        TTI.getInstructionCost(&I, TargetTransformInfo::TCK_CodeSize);105 106    if (HasLatencySavings)107      KnownConstants.push_back(&I);108 109    return CodeSize;110  }111 112  Cost getLatencySavings(Function *F) {113    auto &TTI = FAM.getResult<TargetIRAnalysis>(*F);114    auto &BFI = FAM.getResult<BlockFrequencyAnalysis>(*F);115 116    Cost Latency = 0;117    for (const Instruction *I : KnownConstants)118      Latency += BFI.getBlockFreq(I->getParent()).getFrequency() /119                 BFI.getEntryFreq().getFrequency() *120                 TTI.getInstructionCost(I, TargetTransformInfo::TCK_Latency);121 122    return Latency;123  }124};125 126} // namespace llvm127 128using namespace llvm;129 130TEST_F(FunctionSpecializationTest, SwitchInst) {131  const char *ModuleString = R"(132    define void @foo(i32 %a, i32 %b, i32 %i) {133    entry:134      br label %loop135    loop:136      switch i32 %i, label %default137      [ i32 1, label %case1138        i32 2, label %case2 ]139    case1:140      %0 = mul i32 %a, 2141      %1 = sub i32 6, 5142      br label %bb1143    case2:144      %2 = and i32 %b, 3145      %3 = sdiv i32 8, 2146      br label %bb2147    bb1:148      %4 = add i32 %0, %b149      br label %loop150    bb2:151      %5 = or i32 %2, %a152      br label %loop153    default:154      ret void155    }156  )";157 158  Module &M = parseModule(ModuleString);159  Function *F = M.getFunction("foo");160  FunctionSpecializer Specializer = getSpecializerFor(F);161  InstCostVisitor Visitor = Specializer.getInstCostVisitorFor(F);162 163  Constant *One = ConstantInt::get(IntegerType::getInt32Ty(M.getContext()), 1);164 165  auto FuncIter = F->begin();166  BasicBlock &Loop = *++FuncIter;167  BasicBlock &Case1 = *++FuncIter;168  BasicBlock &Case2 = *++FuncIter;169  BasicBlock &BB1 = *++FuncIter;170  BasicBlock &BB2 = *++FuncIter;171 172  Instruction &Switch = Loop.front();173  Instruction &Mul = Case1.front();174  Instruction &And = Case2.front();175  Instruction &Sdiv = *++Case2.begin();176  Instruction &BrBB2 = Case2.back();177  Instruction &Add = BB1.front();178  Instruction &Or = BB2.front();179  Instruction &BrLoop = BB2.back();180 181  // mul182  Cost Ref = getCodeSizeSavings(Mul);183  Cost Test = Visitor.getCodeSizeSavingsForArg(F->getArg(0), One);184  EXPECT_EQ(Test, Ref);185  EXPECT_TRUE(Test > 0);186 187  // and + or + add188  Ref = getCodeSizeSavings(And) + getCodeSizeSavings(Or) +189        getCodeSizeSavings(Add);190  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(1), One);191  EXPECT_EQ(Test, Ref);192  EXPECT_TRUE(Test > 0);193 194  // switch + sdiv + br + br195  Ref = getCodeSizeSavings(Switch) +196        getCodeSizeSavings(Sdiv, /*HasLatencySavings=*/false) +197        getCodeSizeSavings(BrBB2, /*HasLatencySavings=*/false) +198        getCodeSizeSavings(BrLoop, /*HasLatencySavings=*/false);199  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(2), One);200  EXPECT_EQ(Test, Ref);201  EXPECT_TRUE(Test > 0);202 203  // Latency.204  Ref = getLatencySavings(F);205  Test = Visitor.getLatencySavingsForKnownConstants();206  EXPECT_EQ(Test, Ref);207  EXPECT_TRUE(Test > 0);208}209 210TEST_F(FunctionSpecializationTest, BranchInst) {211  const char *ModuleString = R"(212    define void @foo(i32 %a, i32 %b, i1 %cond) {213    entry:214      br label %loop215    loop:216      br i1 %cond, label %bb0, label %bb3217    bb0:218      %0 = mul i32 %a, 2219      %1 = sub i32 6, 5220      br i1 %cond, label %bb1, label %bb2221    bb1:222      %2 = add i32 %0, %b223      %3 = sdiv i32 8, 2224      br label %bb2225    bb2:226      br label %loop227    bb3:228      ret void229    }230  )";231 232  Module &M = parseModule(ModuleString);233  Function *F = M.getFunction("foo");234  FunctionSpecializer Specializer = getSpecializerFor(F);235  InstCostVisitor Visitor = Specializer.getInstCostVisitorFor(F);236 237  Constant *One = ConstantInt::get(IntegerType::getInt32Ty(M.getContext()), 1);238  Constant *False = ConstantInt::getFalse(M.getContext());239 240  auto FuncIter = F->begin();241  BasicBlock &Loop = *++FuncIter;242  BasicBlock &BB0 = *++FuncIter;243  BasicBlock &BB1 = *++FuncIter;244  BasicBlock &BB2 = *++FuncIter;245 246  Instruction &Branch = Loop.front();247  Instruction &Mul = BB0.front();248  Instruction &Sub = *++BB0.begin();249  Instruction &BrBB1BB2 = BB0.back();250  Instruction &Add = BB1.front();251  Instruction &Sdiv = *++BB1.begin();252  Instruction &BrBB2 = BB1.back();253  Instruction &BrLoop = BB2.front();254 255  // mul256  Cost Ref = getCodeSizeSavings(Mul);257  Cost Test = Visitor.getCodeSizeSavingsForArg(F->getArg(0), One);258  EXPECT_EQ(Test, Ref);259  EXPECT_TRUE(Test > 0);260 261  // add262  Ref = getCodeSizeSavings(Add);263  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(1), One);264  EXPECT_EQ(Test, Ref);265  EXPECT_TRUE(Test > 0);266 267  // branch + sub + br + sdiv + br268  Ref = getCodeSizeSavings(Branch) +269        getCodeSizeSavings(Sub, /*HasLatencySavings=*/false) +270        getCodeSizeSavings(BrBB1BB2) +271        getCodeSizeSavings(Sdiv, /*HasLatencySavings=*/false) +272        getCodeSizeSavings(BrBB2, /*HasLatencySavings=*/false) +273        getCodeSizeSavings(BrLoop, /*HasLatencySavings=*/false);274  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(2), False);275  EXPECT_EQ(Test, Ref);276  EXPECT_TRUE(Test > 0);277 278  // Latency.279  Ref = getLatencySavings(F);280  Test = Visitor.getLatencySavingsForKnownConstants();281  EXPECT_EQ(Test, Ref);282  EXPECT_TRUE(Test > 0);283}284 285TEST_F(FunctionSpecializationTest, SelectInst) {286  const char *ModuleString = R"(287    define i32 @foo(i1 %cond, i32 %a, i32 %b) {288      %sel = select i1 %cond, i32 %a, i32 %b289      ret i32 %sel290    }291  )";292 293  Module &M = parseModule(ModuleString);294  Function *F = M.getFunction("foo");295  FunctionSpecializer Specializer = getSpecializerFor(F);296  InstCostVisitor Visitor = Specializer.getInstCostVisitorFor(F);297 298  Constant *One = ConstantInt::get(IntegerType::getInt32Ty(M.getContext()), 1);299  Constant *Zero = ConstantInt::get(IntegerType::getInt32Ty(M.getContext()), 0);300  Constant *False = ConstantInt::getFalse(M.getContext());301  Instruction &Select = *F->front().begin();302 303  Cost RefCodeSize = getCodeSizeSavings(Select);304  Cost RefLatency = getLatencySavings(F);305 306  Cost TestCodeSize = Visitor.getCodeSizeSavingsForArg(F->getArg(0), False);307  EXPECT_TRUE(TestCodeSize == 0);308  TestCodeSize = Visitor.getCodeSizeSavingsForArg(F->getArg(1), One);309  EXPECT_TRUE(TestCodeSize == 0);310  Cost TestLatency = Visitor.getLatencySavingsForKnownConstants();311  EXPECT_TRUE(TestLatency == 0);312 313  TestCodeSize = Visitor.getCodeSizeSavingsForArg(F->getArg(2), Zero);314  EXPECT_EQ(TestCodeSize, RefCodeSize);315  EXPECT_TRUE(TestCodeSize > 0);316  TestLatency = Visitor.getLatencySavingsForKnownConstants();317  EXPECT_EQ(TestLatency, RefLatency);318  EXPECT_TRUE(TestLatency > 0);319}320 321TEST_F(FunctionSpecializationTest, Misc) {322  const char *ModuleString = R"(323    %struct_t = type { [8 x i16], [8 x i16], i32, i32, i32, ptr, [8 x i8] }324    @g = constant %struct_t zeroinitializer, align 16325 326    declare i32 @llvm.smax.i32(i32, i32)327    declare i32 @bar(i32)328 329    define i32 @foo(i8 %a, i1 %cond, ptr %b, i32 %c) {330      %cmp = icmp eq i8 %a, 10331      %ext = zext i1 %cmp to i64332      %sel = select i1 %cond, i64 %ext, i64 1333      %gep = getelementptr inbounds %struct_t, ptr %b, i64 %sel, i32 4334      %ld = load i32, ptr %gep335      %fr = freeze i32 %ld336      %smax = call i32 @llvm.smax.i32(i32 %fr, i32 1)337      %call = call i32 @bar(i32 %smax)338      %fr2 = freeze i32 %c339      %add = add i32 %call, %fr2340      ret i32 %add341    }342  )";343 344  Module &M = parseModule(ModuleString);345  Function *F = M.getFunction("foo");346  FunctionSpecializer Specializer = getSpecializerFor(F);347  InstCostVisitor Visitor = Specializer.getInstCostVisitorFor(F);348 349  GlobalVariable *GV = M.getGlobalVariable("g");350  Constant *One = ConstantInt::get(IntegerType::getInt8Ty(M.getContext()), 1);351  Constant *True = ConstantInt::getTrue(M.getContext());352  Constant *Undef = UndefValue::get(IntegerType::getInt32Ty(M.getContext()));353 354  auto BlockIter = F->front().begin();355  Instruction &Icmp = *BlockIter++;356  Instruction &Zext = *BlockIter++;357  Instruction &Select = *BlockIter++;358  Instruction &Gep = *BlockIter++;359  Instruction &Load = *BlockIter++;360  Instruction &Freeze = *BlockIter++;361  Instruction &Smax = *BlockIter++;362 363  // icmp + zext364  Cost Ref = getCodeSizeSavings(Icmp) + getCodeSizeSavings(Zext);365  Cost Test = Visitor.getCodeSizeSavingsForArg(F->getArg(0), One);366  EXPECT_EQ(Test, Ref);367  EXPECT_TRUE(Test > 0);368 369  // select370  Ref = getCodeSizeSavings(Select);371  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(1), True);372  EXPECT_EQ(Test, Ref);373  EXPECT_TRUE(Test > 0);374 375  // gep + load + freeze + smax376  Ref = getCodeSizeSavings(Gep) + getCodeSizeSavings(Load) +377        getCodeSizeSavings(Freeze) + getCodeSizeSavings(Smax);378  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(2), GV);379  EXPECT_EQ(Test, Ref);380  EXPECT_TRUE(Test > 0);381 382  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(3), Undef);383  EXPECT_TRUE(Test == 0);384 385  // Latency.386  Ref = getLatencySavings(F);387  Test = Visitor.getLatencySavingsForKnownConstants();388  EXPECT_EQ(Test, Ref);389  EXPECT_TRUE(Test > 0);390}391 392TEST_F(FunctionSpecializationTest, PhiNode) {393  const char *ModuleString = R"(394    define void @foo(i32 %a, i32 %b, i32 %i) {395    entry:396      br label %loop397    loop:398      %0 = phi i32 [ %a, %entry ], [ %3, %bb ]399      switch i32 %i, label %default400      [ i32 1, label %case1401        i32 2, label %case2 ]402    case1:403      %1 = add i32 %0, 1404      br label %bb405    case2:406      %2 = phi i32 [ %a, %entry ], [ %0, %loop ]407      br label %bb408    bb:409      %3 = phi i32 [ %b, %case1 ], [ %2, %case2 ], [ %3, %bb ]410      %4 = icmp eq i32 %3, 1411      br i1 %4, label %bb, label %loop412    default:413      ret void414    }415  )";416 417  Module &M = parseModule(ModuleString);418  Function *F = M.getFunction("foo");419  FunctionSpecializer Specializer = getSpecializerFor(F);420  InstCostVisitor Visitor = Specializer.getInstCostVisitorFor(F);421 422  Constant *One = ConstantInt::get(IntegerType::getInt32Ty(M.getContext()), 1);423 424  auto FuncIter = F->begin();425  BasicBlock &Loop = *++FuncIter;426  BasicBlock &Case1 = *++FuncIter;427  BasicBlock &Case2 = *++FuncIter;428  BasicBlock &BB = *++FuncIter;429 430  Instruction &PhiLoop = Loop.front();431  Instruction &Switch = Loop.back();432  Instruction &Add = Case1.front();433  Instruction &PhiCase2 = Case2.front();434  Instruction &BrBB = Case2.back();435  Instruction &PhiBB = BB.front();436  Instruction &Icmp = *++BB.begin();437  Instruction &Branch = BB.back();438 439  Cost Test = Visitor.getCodeSizeSavingsForArg(F->getArg(0), One);440  EXPECT_TRUE(Test == 0);441 442  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(1), One);443  EXPECT_TRUE(Test == 0);444 445  Test = Visitor.getLatencySavingsForKnownConstants();446  EXPECT_TRUE(Test == 0);447 448  // switch + phi + br449  Cost Ref = getCodeSizeSavings(Switch) +450             getCodeSizeSavings(PhiCase2, /*HasLatencySavings=*/false) +451             getCodeSizeSavings(BrBB, /*HasLatencySavings=*/false);452  Test = Visitor.getCodeSizeSavingsForArg(F->getArg(2), One);453  EXPECT_EQ(Test, Ref);454  EXPECT_TRUE(Test > 0 && Test > 0);455 456  // phi + phi + add + icmp + branch457  Ref = getCodeSizeSavings(PhiBB) + getCodeSizeSavings(PhiLoop) +458        getCodeSizeSavings(Add) + getCodeSizeSavings(Icmp) +459        getCodeSizeSavings(Branch);460  Test = Visitor.getCodeSizeSavingsFromPendingPHIs();461  EXPECT_EQ(Test, Ref);462  EXPECT_TRUE(Test > 0);463 464  // Latency.465  Ref = getLatencySavings(F);466  Test = Visitor.getLatencySavingsForKnownConstants();467  EXPECT_EQ(Test, Ref);468  EXPECT_TRUE(Test > 0);469}470 471TEST_F(FunctionSpecializationTest, BinOp) {472  // Verify that we can handle binary operators even when only one operand is473  // constant.474  const char *ModuleString = R"(475    define i32 @foo(i1 %a, i1 %b) {476      %and1 = and i1 %a, %b477      %and2 = and i1 %b, %and1478      %sel = select i1 %and2, i32 1, i32 0479      ret i32 %sel480    }481  )";482 483  Module &M = parseModule(ModuleString);484  Function *F = M.getFunction("foo");485  FunctionSpecializer Specializer = getSpecializerFor(F);486  InstCostVisitor Visitor = Specializer.getInstCostVisitorFor(F);487 488  Constant *False = ConstantInt::getFalse(M.getContext());489  BasicBlock &BB = F->front();490  Instruction &And1 = BB.front();491  Instruction &And2 = *++BB.begin();492  Instruction &Select = *++BB.begin();493 494  Cost RefCodeSize = getCodeSizeSavings(And1) + getCodeSizeSavings(And2) +495                     getCodeSizeSavings(Select);496  Cost RefLatency = getLatencySavings(F);497 498  Cost TestCodeSize = Visitor.getCodeSizeSavingsForArg(F->getArg(0), False);499  Cost TestLatency = Visitor.getLatencySavingsForKnownConstants();500 501  EXPECT_EQ(TestCodeSize, RefCodeSize);502  EXPECT_TRUE(TestCodeSize > 0);503  EXPECT_EQ(TestLatency, RefLatency);504  EXPECT_TRUE(TestLatency > 0);505}506