1 //===- llvm/unittest/CodeGen/AArch64SelectionDAGTest.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 "llvm/CodeGen/SelectionDAG.h"
10 #include "llvm/Analysis/OptimizationRemarkEmitter.h"
11 #include "llvm/AsmParser/Parser.h"
12 #include "llvm/CodeGen/MachineModuleInfo.h"
13 #include "llvm/CodeGen/TargetLowering.h"
14 #include "llvm/Support/SourceMgr.h"
15 #include "llvm/Support/TargetRegistry.h"
16 #include "llvm/Support/TargetSelect.h"
17 #include "llvm/Target/TargetMachine.h"
18 #include "gtest/gtest.h"
19 
20 namespace llvm {
21 
22 class AArch64SelectionDAGTest : public testing::Test {
23 protected:
24   static void SetUpTestCase() {
25     InitializeAllTargets();
26     InitializeAllTargetMCs();
27   }
28 
29   void SetUp() override {
30     StringRef Assembly = "define void @f() { ret void }";
31 
32     Triple TargetTriple("aarch64--");
33     std::string Error;
34     const Target *T = TargetRegistry::lookupTarget("", TargetTriple, Error);
35     // FIXME: These tests do not depend on AArch64 specifically, but we have to
36     // initialize a target. A skeleton Target for unittests would allow us to
37     // always run these tests.
38     if (!T)
39       return;
40 
41     TargetOptions Options;
42     TM = std::unique_ptr<LLVMTargetMachine>(static_cast<LLVMTargetMachine *>(
43         T->createTargetMachine("AArch64", "", "+sve", Options, None, None,
44                                CodeGenOpt::Aggressive)));
45     if (!TM)
46       return;
47 
48     SMDiagnostic SMError;
49     M = parseAssemblyString(Assembly, SMError, Context);
50     if (!M)
51       report_fatal_error(SMError.getMessage());
52     M->setDataLayout(TM->createDataLayout());
53 
54     F = M->getFunction("f");
55     if (!F)
56       report_fatal_error("F?");
57 
58     MachineModuleInfo MMI(TM.get());
59 
60     MF = std::make_unique<MachineFunction>(*F, *TM, *TM->getSubtargetImpl(*F), 0,
61                                       MMI);
62 
63     DAG = std::make_unique<SelectionDAG>(*TM, CodeGenOpt::None);
64     if (!DAG)
65       report_fatal_error("DAG?");
66     OptimizationRemarkEmitter ORE(F);
67     DAG->init(*MF, ORE, nullptr, nullptr, nullptr, nullptr, nullptr);
68   }
69 
70   TargetLoweringBase::LegalizeTypeAction getTypeAction(EVT VT) {
71     return DAG->getTargetLoweringInfo().getTypeAction(Context, VT);
72   }
73 
74   EVT getTypeToTransformTo(EVT VT) {
75     return DAG->getTargetLoweringInfo().getTypeToTransformTo(Context, VT);
76   }
77 
78   LLVMContext Context;
79   std::unique_ptr<LLVMTargetMachine> TM;
80   std::unique_ptr<Module> M;
81   Function *F;
82   std::unique_ptr<MachineFunction> MF;
83   std::unique_ptr<SelectionDAG> DAG;
84 };
85 
86 TEST_F(AArch64SelectionDAGTest, computeKnownBits_ZERO_EXTEND_VECTOR_INREG) {
87   if (!TM)
88     return;
89   SDLoc Loc;
90   auto Int8VT = EVT::getIntegerVT(Context, 8);
91   auto Int16VT = EVT::getIntegerVT(Context, 16);
92   auto InVecVT = EVT::getVectorVT(Context, Int8VT, 4);
93   auto OutVecVT = EVT::getVectorVT(Context, Int16VT, 2);
94   auto InVec = DAG->getConstant(0, Loc, InVecVT);
95   auto Op = DAG->getNode(ISD::ZERO_EXTEND_VECTOR_INREG, Loc, OutVecVT, InVec);
96   auto DemandedElts = APInt(2, 3);
97   KnownBits Known = DAG->computeKnownBits(Op, DemandedElts);
98   EXPECT_TRUE(Known.isZero());
99 }
100 
101 TEST_F(AArch64SelectionDAGTest, computeKnownBitsSVE_ZERO_EXTEND_VECTOR_INREG) {
102   if (!TM)
103     return;
104   SDLoc Loc;
105   auto Int8VT = EVT::getIntegerVT(Context, 8);
106   auto Int16VT = EVT::getIntegerVT(Context, 16);
107   auto InVecVT = EVT::getVectorVT(Context, Int8VT, 4, true);
108   auto OutVecVT = EVT::getVectorVT(Context, Int16VT, 2, true);
109   auto InVec = DAG->getConstant(0, Loc, InVecVT);
110   auto Op = DAG->getNode(ISD::ZERO_EXTEND_VECTOR_INREG, Loc, OutVecVT, InVec);
111   auto DemandedElts = APInt(2, 3);
112   KnownBits Known = DAG->computeKnownBits(Op, DemandedElts);
113 
114   // We don't know anything for SVE at the moment.
115   EXPECT_EQ(Known.Zero, APInt(16, 0u));
116   EXPECT_EQ(Known.One, APInt(16, 0u));
117   EXPECT_FALSE(Known.isZero());
118 }
119 
120 TEST_F(AArch64SelectionDAGTest, computeKnownBits_EXTRACT_SUBVECTOR) {
121   if (!TM)
122     return;
123   SDLoc Loc;
124   auto IntVT = EVT::getIntegerVT(Context, 8);
125   auto VecVT = EVT::getVectorVT(Context, IntVT, 3);
126   auto IdxVT = EVT::getIntegerVT(Context, 64);
127   auto Vec = DAG->getConstant(0, Loc, VecVT);
128   auto ZeroIdx = DAG->getConstant(0, Loc, IdxVT);
129   auto Op = DAG->getNode(ISD::EXTRACT_SUBVECTOR, Loc, VecVT, Vec, ZeroIdx);
130   auto DemandedElts = APInt(3, 7);
131   KnownBits Known = DAG->computeKnownBits(Op, DemandedElts);
132   EXPECT_TRUE(Known.isZero());
133 }
134 
135 TEST_F(AArch64SelectionDAGTest, ComputeNumSignBits_SIGN_EXTEND_VECTOR_INREG) {
136   if (!TM)
137     return;
138   SDLoc Loc;
139   auto Int8VT = EVT::getIntegerVT(Context, 8);
140   auto Int16VT = EVT::getIntegerVT(Context, 16);
141   auto InVecVT = EVT::getVectorVT(Context, Int8VT, 4);
142   auto OutVecVT = EVT::getVectorVT(Context, Int16VT, 2);
143   auto InVec = DAG->getConstant(1, Loc, InVecVT);
144   auto Op = DAG->getNode(ISD::SIGN_EXTEND_VECTOR_INREG, Loc, OutVecVT, InVec);
145   auto DemandedElts = APInt(2, 3);
146   EXPECT_EQ(DAG->ComputeNumSignBits(Op, DemandedElts), 15u);
147 }
148 
149 TEST_F(AArch64SelectionDAGTest, ComputeNumSignBits_EXTRACT_SUBVECTOR) {
150   if (!TM)
151     return;
152   SDLoc Loc;
153   auto IntVT = EVT::getIntegerVT(Context, 8);
154   auto VecVT = EVT::getVectorVT(Context, IntVT, 3);
155   auto IdxVT = EVT::getIntegerVT(Context, 64);
156   auto Vec = DAG->getConstant(1, Loc, VecVT);
157   auto ZeroIdx = DAG->getConstant(0, Loc, IdxVT);
158   auto Op = DAG->getNode(ISD::EXTRACT_SUBVECTOR, Loc, VecVT, Vec, ZeroIdx);
159   auto DemandedElts = APInt(3, 7);
160   EXPECT_EQ(DAG->ComputeNumSignBits(Op, DemandedElts), 7u);
161 }
162 
163 TEST_F(AArch64SelectionDAGTest, SimplifyDemandedVectorElts_EXTRACT_SUBVECTOR) {
164   if (!TM)
165     return;
166 
167   TargetLowering TL(*TM);
168 
169   SDLoc Loc;
170   auto IntVT = EVT::getIntegerVT(Context, 8);
171   auto VecVT = EVT::getVectorVT(Context, IntVT, 3);
172   auto IdxVT = EVT::getIntegerVT(Context, 64);
173   auto Vec = DAG->getConstant(1, Loc, VecVT);
174   auto ZeroIdx = DAG->getConstant(0, Loc, IdxVT);
175   auto Op = DAG->getNode(ISD::EXTRACT_SUBVECTOR, Loc, VecVT, Vec, ZeroIdx);
176   auto DemandedElts = APInt(3, 7);
177   auto KnownUndef = APInt(3, 0);
178   auto KnownZero = APInt(3, 0);
179   TargetLowering::TargetLoweringOpt TLO(*DAG, false, false);
180   EXPECT_EQ(TL.SimplifyDemandedVectorElts(Op, DemandedElts, KnownUndef,
181                                           KnownZero, TLO),
182             false);
183 }
184 
185 // Piggy-backing on the AArch64 tests to verify SelectionDAG::computeKnownBits.
186 TEST_F(AArch64SelectionDAGTest, ComputeKnownBits_ADD) {
187   if (!TM)
188     return;
189   SDLoc Loc;
190   auto IntVT = EVT::getIntegerVT(Context, 8);
191   auto UnknownOp = DAG->getRegister(0, IntVT);
192   auto Mask = DAG->getConstant(0x8A, Loc, IntVT);
193   auto N0 = DAG->getNode(ISD::AND, Loc, IntVT, Mask, UnknownOp);
194   auto N1 = DAG->getConstant(0x55, Loc, IntVT);
195   auto Op = DAG->getNode(ISD::ADD, Loc, IntVT, N0, N1);
196   // N0 = ?000?0?0
197   // N1 = 01010101
198   //  =>
199   // Known.One  = 01010101 (0x55)
200   // Known.Zero = 00100000 (0x20)
201   KnownBits Known = DAG->computeKnownBits(Op);
202   EXPECT_EQ(Known.Zero, APInt(8, 0x20));
203   EXPECT_EQ(Known.One, APInt(8, 0x55));
204 }
205 
206 // Piggy-backing on the AArch64 tests to verify SelectionDAG::computeKnownBits.
207 TEST_F(AArch64SelectionDAGTest, ComputeKnownBits_SUB) {
208   if (!TM)
209     return;
210   SDLoc Loc;
211   auto IntVT = EVT::getIntegerVT(Context, 8);
212   auto N0 = DAG->getConstant(0x55, Loc, IntVT);
213   auto UnknownOp = DAG->getRegister(0, IntVT);
214   auto Mask = DAG->getConstant(0x2e, Loc, IntVT);
215   auto N1 = DAG->getNode(ISD::AND, Loc, IntVT, Mask, UnknownOp);
216   auto Op = DAG->getNode(ISD::SUB, Loc, IntVT, N0, N1);
217   // N0 = 01010101
218   // N1 = 00?0???0
219   //  =>
220   // Known.One  = 00000001 (0x1)
221   // Known.Zero = 10000000 (0x80)
222   KnownBits Known = DAG->computeKnownBits(Op);
223   EXPECT_EQ(Known.Zero, APInt(8, 0x80));
224   EXPECT_EQ(Known.One, APInt(8, 0x1));
225 }
226 
227 TEST_F(AArch64SelectionDAGTest, isSplatValue_Fixed_BUILD_VECTOR) {
228   if (!TM)
229     return;
230 
231   TargetLowering TL(*TM);
232 
233   SDLoc Loc;
234   auto IntVT = EVT::getIntegerVT(Context, 8);
235   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, false);
236   // Create a BUILD_VECTOR
237   SDValue Op = DAG->getConstant(1, Loc, VecVT);
238   EXPECT_EQ(Op->getOpcode(), ISD::BUILD_VECTOR);
239   EXPECT_TRUE(DAG->isSplatValue(Op, /*AllowUndefs=*/false));
240 
241   APInt UndefElts;
242   APInt DemandedElts;
243   EXPECT_FALSE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
244 
245   // Width=16, Mask=3
246   DemandedElts = APInt(16, 3);
247   EXPECT_TRUE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
248 }
249 
250 TEST_F(AArch64SelectionDAGTest, isSplatValue_Fixed_ADD_of_BUILD_VECTOR) {
251   if (!TM)
252     return;
253 
254   TargetLowering TL(*TM);
255 
256   SDLoc Loc;
257   auto IntVT = EVT::getIntegerVT(Context, 8);
258   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, false);
259 
260   // Should create BUILD_VECTORs
261   SDValue Val1 = DAG->getConstant(1, Loc, VecVT);
262   SDValue Val2 = DAG->getConstant(3, Loc, VecVT);
263   EXPECT_EQ(Val1->getOpcode(), ISD::BUILD_VECTOR);
264   SDValue Op = DAG->getNode(ISD::ADD, Loc, VecVT, Val1, Val2);
265 
266   EXPECT_TRUE(DAG->isSplatValue(Op, /*AllowUndefs=*/false));
267 
268   APInt UndefElts;
269   APInt DemandedElts;
270   EXPECT_FALSE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
271 
272   // Width=16, Mask=3
273   DemandedElts = APInt(16, 3);
274   EXPECT_TRUE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
275 }
276 
277 TEST_F(AArch64SelectionDAGTest, isSplatValue_Scalable_SPLAT_VECTOR) {
278   if (!TM)
279     return;
280 
281   TargetLowering TL(*TM);
282 
283   SDLoc Loc;
284   auto IntVT = EVT::getIntegerVT(Context, 8);
285   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, true);
286   // Create a SPLAT_VECTOR
287   SDValue Op = DAG->getConstant(1, Loc, VecVT);
288   EXPECT_EQ(Op->getOpcode(), ISD::SPLAT_VECTOR);
289   EXPECT_TRUE(DAG->isSplatValue(Op, /*AllowUndefs=*/false));
290 
291   APInt UndefElts;
292   APInt DemandedElts;
293   EXPECT_TRUE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
294 
295   // Width=16, Mask=3. These bits should be ignored.
296   DemandedElts = APInt(16, 3);
297   EXPECT_TRUE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
298 }
299 
300 TEST_F(AArch64SelectionDAGTest, isSplatValue_Scalable_ADD_of_SPLAT_VECTOR) {
301   if (!TM)
302     return;
303 
304   TargetLowering TL(*TM);
305 
306   SDLoc Loc;
307   auto IntVT = EVT::getIntegerVT(Context, 8);
308   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, true);
309 
310   // Should create SPLAT_VECTORS
311   SDValue Val1 = DAG->getConstant(1, Loc, VecVT);
312   SDValue Val2 = DAG->getConstant(3, Loc, VecVT);
313   EXPECT_EQ(Val1->getOpcode(), ISD::SPLAT_VECTOR);
314   SDValue Op = DAG->getNode(ISD::ADD, Loc, VecVT, Val1, Val2);
315 
316   EXPECT_TRUE(DAG->isSplatValue(Op, /*AllowUndefs=*/false));
317 
318   APInt UndefElts;
319   APInt DemandedElts;
320   EXPECT_TRUE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
321 
322   // Width=16, Mask=3. These bits should be ignored.
323   DemandedElts = APInt(16, 3);
324   EXPECT_TRUE(DAG->isSplatValue(Op, DemandedElts, UndefElts));
325 }
326 
327 TEST_F(AArch64SelectionDAGTest, getSplatSourceVector_Fixed_BUILD_VECTOR) {
328   if (!TM)
329     return;
330 
331   TargetLowering TL(*TM);
332 
333   SDLoc Loc;
334   auto IntVT = EVT::getIntegerVT(Context, 8);
335   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, false);
336   // Create a BUILD_VECTOR
337   SDValue Op = DAG->getConstant(1, Loc, VecVT);
338   EXPECT_EQ(Op->getOpcode(), ISD::BUILD_VECTOR);
339 
340   int SplatIdx = -1;
341   EXPECT_EQ(DAG->getSplatSourceVector(Op, SplatIdx), Op);
342   EXPECT_EQ(SplatIdx, 0);
343 }
344 
345 TEST_F(AArch64SelectionDAGTest, getSplatSourceVector_Fixed_ADD_of_BUILD_VECTOR) {
346   if (!TM)
347     return;
348 
349   TargetLowering TL(*TM);
350 
351   SDLoc Loc;
352   auto IntVT = EVT::getIntegerVT(Context, 8);
353   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, false);
354 
355   // Should create BUILD_VECTORs
356   SDValue Val1 = DAG->getConstant(1, Loc, VecVT);
357   SDValue Val2 = DAG->getConstant(3, Loc, VecVT);
358   EXPECT_EQ(Val1->getOpcode(), ISD::BUILD_VECTOR);
359   SDValue Op = DAG->getNode(ISD::ADD, Loc, VecVT, Val1, Val2);
360 
361   int SplatIdx = -1;
362   EXPECT_EQ(DAG->getSplatSourceVector(Op, SplatIdx), Op);
363   EXPECT_EQ(SplatIdx, 0);
364 }
365 
366 TEST_F(AArch64SelectionDAGTest, getSplatSourceVector_Scalable_SPLAT_VECTOR) {
367   if (!TM)
368     return;
369 
370   TargetLowering TL(*TM);
371 
372   SDLoc Loc;
373   auto IntVT = EVT::getIntegerVT(Context, 8);
374   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, true);
375   // Create a SPLAT_VECTOR
376   SDValue Op = DAG->getConstant(1, Loc, VecVT);
377   EXPECT_EQ(Op->getOpcode(), ISD::SPLAT_VECTOR);
378 
379   int SplatIdx = -1;
380   EXPECT_EQ(DAG->getSplatSourceVector(Op, SplatIdx), Op);
381   EXPECT_EQ(SplatIdx, 0);
382 }
383 
384 TEST_F(AArch64SelectionDAGTest, getSplatSourceVector_Scalable_ADD_of_SPLAT_VECTOR) {
385   if (!TM)
386     return;
387 
388   TargetLowering TL(*TM);
389 
390   SDLoc Loc;
391   auto IntVT = EVT::getIntegerVT(Context, 8);
392   auto VecVT = EVT::getVectorVT(Context, IntVT, 16, true);
393 
394   // Should create SPLAT_VECTORS
395   SDValue Val1 = DAG->getConstant(1, Loc, VecVT);
396   SDValue Val2 = DAG->getConstant(3, Loc, VecVT);
397   EXPECT_EQ(Val1->getOpcode(), ISD::SPLAT_VECTOR);
398   SDValue Op = DAG->getNode(ISD::ADD, Loc, VecVT, Val1, Val2);
399 
400   int SplatIdx = -1;
401   EXPECT_EQ(DAG->getSplatSourceVector(Op, SplatIdx), Op);
402   EXPECT_EQ(SplatIdx, 0);
403 }
404 
405 TEST_F(AArch64SelectionDAGTest, getTypeConversion_SplitScalableMVT) {
406   if (!TM)
407     return;
408 
409   MVT VT = MVT::nxv4i64;
410   EXPECT_EQ(getTypeAction(VT), TargetLoweringBase::TypeSplitVector);
411   ASSERT_TRUE(getTypeToTransformTo(VT).isScalableVector());
412 }
413 
414 TEST_F(AArch64SelectionDAGTest, getTypeConversion_PromoteScalableMVT) {
415   if (!TM)
416     return;
417 
418   MVT VT = MVT::nxv2i32;
419   EXPECT_EQ(getTypeAction(VT), TargetLoweringBase::TypePromoteInteger);
420   ASSERT_TRUE(getTypeToTransformTo(VT).isScalableVector());
421 }
422 
423 TEST_F(AArch64SelectionDAGTest, getTypeConversion_NoScalarizeMVT_nxv1f32) {
424   if (!TM)
425     return;
426 
427   MVT VT = MVT::nxv1f32;
428   EXPECT_NE(getTypeAction(VT), TargetLoweringBase::TypeScalarizeVector);
429   ASSERT_TRUE(getTypeToTransformTo(VT).isScalableVector());
430 }
431 
432 TEST_F(AArch64SelectionDAGTest, getTypeConversion_SplitScalableEVT) {
433   if (!TM)
434     return;
435 
436   EVT VT = EVT::getVectorVT(Context, MVT::i64, 256, true);
437   EXPECT_EQ(getTypeAction(VT), TargetLoweringBase::TypeSplitVector);
438   EXPECT_EQ(getTypeToTransformTo(VT), VT.getHalfNumVectorElementsVT(Context));
439 }
440 
441 TEST_F(AArch64SelectionDAGTest, getTypeConversion_WidenScalableEVT) {
442   if (!TM)
443     return;
444 
445   EVT FromVT = EVT::getVectorVT(Context, MVT::i64, 6, true);
446   EVT ToVT = EVT::getVectorVT(Context, MVT::i64, 8, true);
447 
448   EXPECT_EQ(getTypeAction(FromVT), TargetLoweringBase::TypeWidenVector);
449   EXPECT_EQ(getTypeToTransformTo(FromVT), ToVT);
450 }
451 
452 TEST_F(AArch64SelectionDAGTest, getTypeConversion_NoScalarizeEVT_nxv1f128) {
453   if (!TM)
454     return;
455 
456   EVT FromVT = EVT::getVectorVT(Context, MVT::f128, 1, true);
457   EXPECT_DEATH(getTypeAction(FromVT), "Cannot legalize this vector");
458 }
459 
460 } // end namespace llvm
461