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