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