1 //===- ShapedTypeTest.cpp - ShapedType 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-exception
6 //
7 //===----------------------------------------------------------------------===//
8 
9 #include "mlir/IR/AffineMap.h"
10 #include "mlir/IR/BuiltinTypes.h"
11 #include "mlir/IR/Dialect.h"
12 #include "mlir/IR/DialectInterface.h"
13 #include "llvm/ADT/SmallVector.h"
14 #include "gtest/gtest.h"
15 #include <cstdint>
16 
17 using namespace mlir;
18 using namespace mlir::detail;
19 
20 namespace {
21 TEST(ShapedTypeTest, CloneMemref) {
22   MLIRContext context;
23 
24   Type i32 = IntegerType::get(&context, 32);
25   Type f32 = FloatType::getF32(&context);
26   int memSpace = 7;
27   Type memrefOriginalType = i32;
28   llvm::SmallVector<int64_t> memrefOriginalShape({10, 20});
29   AffineMap map = makeStridedLinearLayoutMap({2, 3}, 5, &context);
30 
31   ShapedType memrefType =
32       MemRefType::Builder(memrefOriginalShape, memrefOriginalType)
33           .setMemorySpace(memSpace)
34           .setAffineMaps(map);
35   // Update shape.
36   llvm::SmallVector<int64_t> memrefNewShape({30, 40});
37   ASSERT_NE(memrefOriginalShape, memrefNewShape);
38   ASSERT_EQ(memrefType.clone(memrefNewShape),
39             (MemRefType)MemRefType::Builder(memrefNewShape, memrefOriginalType)
40                 .setMemorySpace(memSpace)
41                 .setAffineMaps(map));
42   // Update type.
43   Type memrefNewType = f32;
44   ASSERT_NE(memrefOriginalType, memrefNewType);
45   ASSERT_EQ(memrefType.clone(memrefNewType),
46             (MemRefType)MemRefType::Builder(memrefOriginalShape, memrefNewType)
47                 .setMemorySpace(memSpace)
48                 .setAffineMaps(map));
49   // Update both.
50   ASSERT_EQ(memrefType.clone(memrefNewShape, memrefNewType),
51             (MemRefType)MemRefType::Builder(memrefNewShape, memrefNewType)
52                 .setMemorySpace(memSpace)
53                 .setAffineMaps(map));
54 
55   // Test unranked memref cloning.
56   ShapedType unrankedTensorType =
57       UnrankedMemRefType::get(memrefOriginalType, memSpace);
58   ASSERT_EQ(unrankedTensorType.clone(memrefNewShape),
59             (MemRefType)MemRefType::Builder(memrefNewShape, memrefOriginalType)
60                 .setMemorySpace(memSpace));
61   ASSERT_EQ(unrankedTensorType.clone(memrefNewType),
62             UnrankedMemRefType::get(memrefNewType, memSpace));
63   ASSERT_EQ(unrankedTensorType.clone(memrefNewShape, memrefNewType),
64             (MemRefType)MemRefType::Builder(memrefNewShape, memrefNewType)
65                 .setMemorySpace(memSpace));
66 }
67 
68 TEST(ShapedTypeTest, CloneTensor) {
69   MLIRContext context;
70 
71   Type i32 = IntegerType::get(&context, 32);
72   Type f32 = FloatType::getF32(&context);
73 
74   Type tensorOriginalType = i32;
75   llvm::SmallVector<int64_t> tensorOriginalShape({10, 20});
76 
77   // Test ranked tensor cloning.
78   ShapedType tensorType =
79       RankedTensorType::get(tensorOriginalShape, tensorOriginalType);
80   // Update shape.
81   llvm::SmallVector<int64_t> tensorNewShape({30, 40});
82   ASSERT_NE(tensorOriginalShape, tensorNewShape);
83   ASSERT_EQ(tensorType.clone(tensorNewShape),
84             RankedTensorType::get(tensorNewShape, tensorOriginalType));
85   // Update type.
86   Type tensorNewType = f32;
87   ASSERT_NE(tensorOriginalType, tensorNewType);
88   ASSERT_EQ(tensorType.clone(tensorNewType),
89             RankedTensorType::get(tensorOriginalShape, tensorNewType));
90   // Update both.
91   ASSERT_EQ(tensorType.clone(tensorNewShape, tensorNewType),
92             RankedTensorType::get(tensorNewShape, tensorNewType));
93 
94   // Test unranked tensor cloning.
95   ShapedType unrankedTensorType = UnrankedTensorType::get(tensorOriginalType);
96   ASSERT_EQ(unrankedTensorType.clone(tensorNewShape),
97             RankedTensorType::get(tensorNewShape, tensorOriginalType));
98   ASSERT_EQ(unrankedTensorType.clone(tensorNewType),
99             UnrankedTensorType::get(tensorNewType));
100   ASSERT_EQ(unrankedTensorType.clone(tensorNewShape),
101             RankedTensorType::get(tensorNewShape, tensorOriginalType));
102 }
103 
104 TEST(ShapedTypeTest, CloneVector) {
105   MLIRContext context;
106 
107   Type i32 = IntegerType::get(&context, 32);
108   Type f32 = FloatType::getF32(&context);
109 
110   Type vectorOriginalType = i32;
111   llvm::SmallVector<int64_t> vectorOriginalShape({10, 20});
112   ShapedType vectorType =
113       VectorType::get(vectorOriginalShape, vectorOriginalType);
114   // Update shape.
115   llvm::SmallVector<int64_t> vectorNewShape({30, 40});
116   ASSERT_NE(vectorOriginalShape, vectorNewShape);
117   ASSERT_EQ(vectorType.clone(vectorNewShape),
118             VectorType::get(vectorNewShape, vectorOriginalType));
119   // Update type.
120   Type vectorNewType = f32;
121   ASSERT_NE(vectorOriginalType, vectorNewType);
122   ASSERT_EQ(vectorType.clone(vectorNewType),
123             VectorType::get(vectorOriginalShape, vectorNewType));
124   // Update both.
125   ASSERT_EQ(vectorType.clone(vectorNewShape, vectorNewType),
126             VectorType::get(vectorNewShape, vectorNewType));
127 }
128 
129 } // end namespace
130