1 //===- DivergenceAnalysisTest.cpp - DivergenceAnalysis unit tests ---------===//
2 //
3 //                     The LLVM Compiler Infrastructure
4 //
5 // This file is distributed under the University of Illinois Open Source
6 // License. See LICENSE.TXT for details.
7 //
8 //===----------------------------------------------------------------------===//
9 
10 #include "llvm/ADT/SmallVector.h"
11 #include "llvm/Analysis/AssumptionCache.h"
12 #include "llvm/Analysis/DivergenceAnalysis.h"
13 #include "llvm/Analysis/LoopInfo.h"
14 #include "llvm/Analysis/PostDominators.h"
15 #include "llvm/Analysis/SyncDependenceAnalysis.h"
16 #include "llvm/Analysis/TargetLibraryInfo.h"
17 #include "llvm/AsmParser/Parser.h"
18 #include "llvm/IR/Constants.h"
19 #include "llvm/IR/Dominators.h"
20 #include "llvm/IR/GlobalVariable.h"
21 #include "llvm/IR/IRBuilder.h"
22 #include "llvm/IR/InstIterator.h"
23 #include "llvm/IR/LLVMContext.h"
24 #include "llvm/IR/LegacyPassManager.h"
25 #include "llvm/IR/Module.h"
26 #include "llvm/IR/Verifier.h"
27 #include "llvm/Support/SourceMgr.h"
28 #include "gtest/gtest.h"
29 
30 namespace llvm {
31 namespace {
32 
33 BasicBlock *GetBlockByName(StringRef BlockName, Function &F) {
34   for (auto &BB : F) {
35     if (BB.getName() != BlockName)
36       continue;
37     return &BB;
38   }
39   return nullptr;
40 }
41 
42 // We use this fixture to ensure that we clean up DivergenceAnalysis before
43 // deleting the PassManager.
44 class DivergenceAnalysisTest : public testing::Test {
45 protected:
46   LLVMContext Context;
47   Module M;
48   TargetLibraryInfoImpl TLII;
49   TargetLibraryInfo TLI;
50 
51   std::unique_ptr<DominatorTree> DT;
52   std::unique_ptr<PostDominatorTree> PDT;
53   std::unique_ptr<LoopInfo> LI;
54   std::unique_ptr<SyncDependenceAnalysis> SDA;
55 
56   DivergenceAnalysisTest() : M("", Context), TLII(), TLI(TLII) {}
57 
58   DivergenceAnalysis buildDA(Function &F, bool IsLCSSA) {
59     DT.reset(new DominatorTree(F));
60     PDT.reset(new PostDominatorTree(F));
61     LI.reset(new LoopInfo(*DT));
62     SDA.reset(new SyncDependenceAnalysis(*DT, *PDT, *LI));
63     return DivergenceAnalysis(F, nullptr, *DT, *LI, *SDA, IsLCSSA);
64   }
65 
66   void runWithDA(
67       Module &M, StringRef FuncName, bool IsLCSSA,
68       function_ref<void(Function &F, LoopInfo &LI, DivergenceAnalysis &DA)>
69           Test) {
70     auto *F = M.getFunction(FuncName);
71     ASSERT_NE(F, nullptr) << "Could not find " << FuncName;
72     DivergenceAnalysis DA = buildDA(*F, IsLCSSA);
73     Test(*F, *LI, DA);
74   }
75 };
76 
77 // Simple initial state test
78 TEST_F(DivergenceAnalysisTest, DAInitialState) {
79   IntegerType *IntTy = IntegerType::getInt32Ty(Context);
80   FunctionType *FTy =
81       FunctionType::get(Type::getVoidTy(Context), {IntTy}, false);
82   Function *F = cast<Function>(M.getOrInsertFunction("f", FTy));
83   BasicBlock *BB = BasicBlock::Create(Context, "entry", F);
84   ReturnInst::Create(Context, nullptr, BB);
85 
86   DivergenceAnalysis DA = buildDA(*F, false);
87 
88   // Whole function region
89   EXPECT_EQ(DA.getRegionLoop(), nullptr);
90 
91   // No divergence in initial state
92   EXPECT_FALSE(DA.hasDetectedDivergence());
93 
94   // No spurious divergence
95   DA.compute();
96   EXPECT_FALSE(DA.hasDetectedDivergence());
97 
98   // Detected divergence after marking
99   Argument &arg = *F->arg_begin();
100   DA.markDivergent(arg);
101 
102   EXPECT_TRUE(DA.hasDetectedDivergence());
103   EXPECT_TRUE(DA.isDivergent(arg));
104 
105   DA.compute();
106   EXPECT_TRUE(DA.hasDetectedDivergence());
107   EXPECT_TRUE(DA.isDivergent(arg));
108 }
109 
110 TEST_F(DivergenceAnalysisTest, DANoLCSSA) {
111   LLVMContext C;
112   SMDiagnostic Err;
113 
114   std::unique_ptr<Module> M = parseAssemblyString(
115       "target datalayout = \"e-m:e-p:32:32-f64:32:64-f80:32-n8:16:32-S128\" "
116       " "
117       "define i32 @f_1(i8* nocapture %arr, i32 %n, i32* %A, i32* %B) "
118       "    local_unnamed_addr { "
119       "entry: "
120       "  br label %loop.ph "
121       " "
122       "loop.ph: "
123       "  br label %loop "
124       " "
125       "loop: "
126       "  %iv0 = phi i32 [ %iv0.inc, %loop ], [ 0, %loop.ph ] "
127       "  %iv1 = phi i32 [ %iv1.inc, %loop ], [ -2147483648, %loop.ph ] "
128       "  %iv0.inc = add i32 %iv0, 1 "
129       "  %iv1.inc = add i32 %iv1, 3 "
130       "  %cond.cont = icmp slt i32 %iv0, %n "
131       "  br i1 %cond.cont, label %loop, label %for.end.loopexit "
132       " "
133       "for.end.loopexit: "
134       "  ret i32 %iv0 "
135       "} ",
136       Err, C);
137 
138   Function *F = M->getFunction("f_1");
139   DivergenceAnalysis DA = buildDA(*F, false);
140   EXPECT_FALSE(DA.hasDetectedDivergence());
141 
142   auto ItArg = F->arg_begin();
143   ItArg++;
144   auto &NArg = *ItArg;
145 
146   // Seed divergence in argument %n
147   DA.markDivergent(NArg);
148 
149   DA.compute();
150   EXPECT_TRUE(DA.hasDetectedDivergence());
151 
152   // Verify that "ret %iv.0" is divergent
153   auto ItBlock = F->begin();
154   std::advance(ItBlock, 3);
155   auto &ExitBlock = *GetBlockByName("for.end.loopexit", *F);
156   auto &RetInst = *cast<ReturnInst>(ExitBlock.begin());
157   EXPECT_TRUE(DA.isDivergent(RetInst));
158 }
159 
160 TEST_F(DivergenceAnalysisTest, DALCSSA) {
161   LLVMContext C;
162   SMDiagnostic Err;
163 
164   std::unique_ptr<Module> M = parseAssemblyString(
165       "target datalayout = \"e-m:e-p:32:32-f64:32:64-f80:32-n8:16:32-S128\" "
166       " "
167       "define i32 @f_lcssa(i8* nocapture %arr, i32 %n, i32* %A, i32* %B) "
168       "    local_unnamed_addr { "
169       "entry: "
170       "  br label %loop.ph "
171       " "
172       "loop.ph: "
173       "  br label %loop "
174       " "
175       "loop: "
176       "  %iv0 = phi i32 [ %iv0.inc, %loop ], [ 0, %loop.ph ] "
177       "  %iv1 = phi i32 [ %iv1.inc, %loop ], [ -2147483648, %loop.ph ] "
178       "  %iv0.inc = add i32 %iv0, 1 "
179       "  %iv1.inc = add i32 %iv1, 3 "
180       "  %cond.cont = icmp slt i32 %iv0, %n "
181       "  br i1 %cond.cont, label %loop, label %for.end.loopexit "
182       " "
183       "for.end.loopexit: "
184       "  %val.ret = phi i32 [ %iv0, %loop ] "
185       "  br label %detached.return "
186       " "
187       "detached.return: "
188       "  ret i32 %val.ret "
189       "} ",
190       Err, C);
191 
192   Function *F = M->getFunction("f_lcssa");
193   DivergenceAnalysis DA = buildDA(*F, true);
194   EXPECT_FALSE(DA.hasDetectedDivergence());
195 
196   auto ItArg = F->arg_begin();
197   ItArg++;
198   auto &NArg = *ItArg;
199 
200   // Seed divergence in argument %n
201   DA.markDivergent(NArg);
202 
203   DA.compute();
204   EXPECT_TRUE(DA.hasDetectedDivergence());
205 
206   // Verify that "ret %iv.0" is divergent
207   auto ItBlock = F->begin();
208   std::advance(ItBlock, 4);
209   auto &ExitBlock = *GetBlockByName("detached.return", *F);
210   auto &RetInst = *cast<ReturnInst>(ExitBlock.begin());
211   EXPECT_TRUE(DA.isDivergent(RetInst));
212 }
213 
214 TEST_F(DivergenceAnalysisTest, DAJoinDivergence) {
215   LLVMContext C;
216   SMDiagnostic Err;
217 
218   std::unique_ptr<Module> M = parseAssemblyString(
219       "target datalayout = \"e-m:e-p:32:32-f64:32:64-f80:32-n8:16:32-S128\" "
220       " "
221       "define void @f_1(i1 %a, i1 %b, i1 %c) "
222       "    local_unnamed_addr { "
223       "A: "
224       "  br i1 %a, label %B, label %C "
225       " "
226       "B: "
227       "  br i1 %b, label %C, label %D "
228       " "
229       "C: "
230       "  %c.join = phi i32 [ 0, %A ], [ 1, %B ] "
231       "  br i1 %c, label %D, label %E "
232       " "
233       "D: "
234       "  %d.join = phi i32 [ 0, %B ], [ 1, %C ] "
235       "  br label %E "
236       " "
237       "E: "
238       "  %e.join = phi i32 [ 0, %C ], [ 1, %D ] "
239       "  ret void "
240       "} "
241       " "
242       "define void @f_2(i1 %a, i1 %b, i1 %c) "
243       "    local_unnamed_addr { "
244       "A: "
245       "  br i1 %a, label %B, label %E "
246       " "
247       "B: "
248       "  br i1 %b, label %C, label %D "
249       " "
250       "C: "
251       "  br label %D "
252       " "
253       "D: "
254       "  %d.join = phi i32 [ 0, %B ], [ 1, %C ] "
255       "  br label %E "
256       " "
257       "E: "
258       "  %e.join = phi i32 [ 0, %A ], [ 1, %D ] "
259       "  ret void "
260       "} "
261       " "
262       "define void @f_3(i1 %a, i1 %b, i1 %c)"
263       "    local_unnamed_addr { "
264       "A: "
265       "  br i1 %a, label %B, label %C "
266       " "
267       "B: "
268       "  br label %C "
269       " "
270       "C: "
271       "  %c.join = phi i32 [ 0, %A ], [ 1, %B ] "
272       "  br i1 %c, label %D, label %E "
273       " "
274       "D: "
275       "  br label %E "
276       " "
277       "E: "
278       "  %e.join = phi i32 [ 0, %C ], [ 1, %D ] "
279       "  ret void "
280       "} ",
281       Err, C);
282 
283   // Maps divergent conditions to the basic blocks whose Phi nodes become
284   // divergent. Blocks need to be listed in IR order.
285   using SmallBlockVec = SmallVector<const BasicBlock *, 4>;
286   using InducedDivJoinMap = std::map<const Value *, SmallBlockVec>;
287 
288   // Actual function performing the checks.
289   auto CheckDivergenceFunc = [this](Function &F,
290                                     InducedDivJoinMap &ExpectedDivJoins) {
291     for (auto &ItCase : ExpectedDivJoins) {
292       auto *DivVal = ItCase.first;
293       auto DA = buildDA(F, false);
294       DA.markDivergent(*DivVal);
295       DA.compute();
296 
297       // List of basic blocks that shall host divergent Phi nodes.
298       auto ItDivJoins = ItCase.second.begin();
299 
300       for (auto &BB : F) {
301         auto *Phi = dyn_cast<PHINode>(BB.begin());
302         if (!Phi)
303           continue;
304 
305         if (ItDivJoins != ItCase.second.end() && &BB == *ItDivJoins) {
306           EXPECT_TRUE(DA.isDivergent(*Phi));
307           // Advance to next block with expected divergent PHI node.
308           ++ItDivJoins;
309         } else {
310           EXPECT_FALSE(DA.isDivergent(*Phi));
311         }
312       }
313     }
314   };
315 
316   {
317     auto *F = M->getFunction("f_1");
318     auto ItBlocks = F->begin();
319     ItBlocks++; // Skip A
320     ItBlocks++; // Skip B
321     auto *C = &*ItBlocks++;
322     auto *D = &*ItBlocks++;
323     auto *E = &*ItBlocks;
324 
325     auto ItArg = F->arg_begin();
326     auto *AArg = &*ItArg++;
327     auto *BArg = &*ItArg++;
328     auto *CArg = &*ItArg;
329 
330     InducedDivJoinMap DivJoins;
331     DivJoins.emplace(AArg, SmallBlockVec({C, D, E}));
332     DivJoins.emplace(BArg, SmallBlockVec({D, E}));
333     DivJoins.emplace(CArg, SmallBlockVec({E}));
334 
335     CheckDivergenceFunc(*F, DivJoins);
336   }
337 
338   {
339     auto *F = M->getFunction("f_2");
340     auto ItBlocks = F->begin();
341     ItBlocks++; // Skip A
342     ItBlocks++; // Skip B
343     ItBlocks++; // Skip C
344     auto *D = &*ItBlocks++;
345     auto *E = &*ItBlocks;
346 
347     auto ItArg = F->arg_begin();
348     auto *AArg = &*ItArg++;
349     auto *BArg = &*ItArg++;
350     auto *CArg = &*ItArg;
351 
352     InducedDivJoinMap DivJoins;
353     DivJoins.emplace(AArg, SmallBlockVec({E}));
354     DivJoins.emplace(BArg, SmallBlockVec({D}));
355     DivJoins.emplace(CArg, SmallBlockVec({}));
356 
357     CheckDivergenceFunc(*F, DivJoins);
358   }
359 
360   {
361     auto *F = M->getFunction("f_3");
362     auto ItBlocks = F->begin();
363     ItBlocks++; // Skip A
364     ItBlocks++; // Skip B
365     auto *C = &*ItBlocks++;
366     ItBlocks++; // Skip D
367     auto *E = &*ItBlocks;
368 
369     auto ItArg = F->arg_begin();
370     auto *AArg = &*ItArg++;
371     auto *BArg = &*ItArg++;
372     auto *CArg = &*ItArg;
373 
374     InducedDivJoinMap DivJoins;
375     DivJoins.emplace(AArg, SmallBlockVec({C}));
376     DivJoins.emplace(BArg, SmallBlockVec({}));
377     DivJoins.emplace(CArg, SmallBlockVec({E}));
378 
379     CheckDivergenceFunc(*F, DivJoins);
380   }
381 }
382 
383 TEST_F(DivergenceAnalysisTest, DASwitchUnreachableDefault) {
384   LLVMContext C;
385   SMDiagnostic Err;
386 
387   std::unique_ptr<Module> M = parseAssemblyString(
388       "target datalayout = \"e-m:e-p:32:32-f64:32:64-f80:32-n8:16:32-S128\" "
389       " "
390       "define void @switch_unreachable_default(i32 %cond) local_unnamed_addr { "
391       "entry: "
392       "  switch i32 %cond, label %sw.default [ "
393       "    i32 0, label %sw.bb0 "
394       "    i32 1, label %sw.bb1 "
395       "  ] "
396       " "
397       "sw.bb0: "
398       "  br label %sw.epilog "
399       " "
400       "sw.bb1: "
401       "  br label %sw.epilog "
402       " "
403       "sw.default: "
404       "  unreachable "
405       " "
406       "sw.epilog: "
407       "  %div.dbl = phi double [ 0.0, %sw.bb0], [ -1.0, %sw.bb1 ] "
408       "  ret void "
409       "}",
410       Err, C);
411 
412   auto *F = M->getFunction("switch_unreachable_default");
413   auto &CondArg = *F->arg_begin();
414   auto DA = buildDA(*F, false);
415 
416   EXPECT_FALSE(DA.hasDetectedDivergence());
417 
418   DA.markDivergent(CondArg);
419   DA.compute();
420 
421   // Still %CondArg is divergent.
422   EXPECT_TRUE(DA.hasDetectedDivergence());
423 
424   // The join uni.dbl is not divergent (see D52221)
425   auto &ExitBlock = *GetBlockByName("sw.epilog", *F);
426   auto &DivDblPhi = *cast<PHINode>(ExitBlock.begin());
427   EXPECT_TRUE(DA.isDivergent(DivDblPhi));
428 }
429 
430 } // end anonymous namespace
431 } // end namespace llvm
432