1// RUN: mlir-opt %s -test-vector-to-vector-conversion="unroll" | FileCheck %s 2 3// CHECK-DAG: #[[MAP1:map[0-9]+]] = affine_map<(d0, d1, d2) -> (d1, d2)> 4 5// CHECK-LABEL: func @add4x2 6// CHECK: %[[ES1:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x2xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>> 7// CHECK-NEXT: %[[ES2:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x2xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>> 8// CHECK-NEXT: %[[TG1:.*]] = vector.tuple_get %[[ES1]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>> 9// CHECK-NEXT: %[[TG2:.*]] = vector.tuple_get %[[ES2]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>> 10// CHECK-NEXT: %[[A1:.*]] = addf %[[TG1]], %[[TG2]] : vector<2x2xf32> 11// CHECK-NEXT: %[[TG3:.*]] = vector.tuple_get %[[ES1]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>> 12// CHECK-NEXT: %[[TG4:.*]] = vector.tuple_get %[[ES2]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>> 13// CHECK-NEXT: %[[A2:.*]] = addf %[[TG3]], %[[TG4]] : vector<2x2xf32> 14// CHECK-NEXT: %[[R1:.*]] = vector.tuple %[[A1]], %[[A2]] : vector<2x2xf32>, vector<2x2xf32> 15// CHECK-NEXT: %[[R2:.*]] = vector.insert_slices %[[R1]], [2, 2], [1, 1] : tuple<vector<2x2xf32>, vector<2x2xf32>> into vector<4x2xf32> 16// CHECK-NEXT: return %[[R2:.*]] : vector<4x2xf32> 17 18func @add4x2(%0: vector<4x2xf32>) -> vector<4x2xf32> { 19 %1 = addf %0, %0: vector<4x2xf32> 20 return %1: vector<4x2xf32> 21} 22 23// CHECK-LABEL: func @add4x4 24// CHECK: %[[ES1:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x4xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 25// CHECK-NEXT: %[[ES2:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x4xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 26 27// CHECK-NEXT: %[[TG1:.*]] = vector.tuple_get %[[ES1]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 28// CHECK-NEXT: %[[TG2:.*]] = vector.tuple_get %[[ES2]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 29// CHECK-NEXT: %[[A1:.*]] = addf %[[TG1]], %[[TG2]] : vector<2x2xf32> 30 31// CHECK-NEXT: %[[TG3:.*]] = vector.tuple_get %[[ES1]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 32// CHECK-NEXT: %[[TG4:.*]] = vector.tuple_get %[[ES2]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 33// CHECK-NEXT: %[[A2:.*]] = addf %[[TG3]], %[[TG4]] : vector<2x2xf32> 34 35// CHECK-NEXT: %[[TG5:.*]] = vector.tuple_get %[[ES1]], 2 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 36// CHECK-NEXT: %[[TG6:.*]] = vector.tuple_get %[[ES2]], 2 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 37// CHECK-NEXT: %[[A3:.*]] = addf %[[TG5]], %[[TG6]] : vector<2x2xf32> 38 39// CHECK-NEXT: %[[TG7:.*]] = vector.tuple_get %[[ES1]], 3 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 40// CHECK-NEXT: %[[TG8:.*]] = vector.tuple_get %[[ES2]], 3 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 41// CHECK-NEXT: %[[A4:.*]] = addf %[[TG7]], %[[TG8]] : vector<2x2xf32> 42 43// CHECK-NEXT: %[[ES3:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x4xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 44 45// CHECK-NEXT: %[[TG9:.*]] = vector.tuple_get %[[ES3]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 46// CHECK-NEXT: %[[A5:.*]] = addf %[[TG9]], %[[A1]] : vector<2x2xf32> 47 48// CHECK-NEXT: %[[TG11:.*]] = vector.tuple_get %[[ES3]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 49// CHECK-NEXT: %[[A6:.*]] = addf %[[TG11]], %[[A2]] : vector<2x2xf32> 50 51// CHECK-NEXT: %[[TG13:.*]] = vector.tuple_get %[[ES3]], 2 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 52// CHECK-NEXT: %[[A7:.*]] = addf %[[TG13]], %[[A3]] : vector<2x2xf32> 53 54// CHECK-NEXT: %[[TG15:.*]] = vector.tuple_get %[[ES3]], 3 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 55// CHECK-NEXT: %[[A8:.*]] = addf %[[TG15]], %[[A4]] : vector<2x2xf32> 56 57// CHECK-NEXT: %[[R3:.*]] = vector.tuple %[[A5]], %[[A6]], %[[A7]], %[[A8]] : vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32> 58// CHECK-NEXT: %[[R4:.*]] = vector.insert_slices %[[R3]], [2, 2], [1, 1] : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> into vector<4x4xf32> 59// CHECK-NEXT: return %[[R4]] : vector<4x4xf32> 60 61func @add4x4(%0: vector<4x4xf32>, %1: vector<4x4xf32>) -> vector<4x4xf32> { 62 %2 = addf %0, %1: vector<4x4xf32> 63 %3 = addf %1, %2: vector<4x4xf32> 64 return %3: vector<4x4xf32> 65} 66 67#contraction_accesses0 = [ 68 affine_map<(i, j, k) -> (i, k)>, 69 affine_map<(i, j, k) -> (k, j)>, 70 affine_map<(i, j, k) -> (i, j)> 71] 72#contraction_trait0 = { 73 indexing_maps = #contraction_accesses0, 74 iterator_types = ["parallel", "parallel", "reduction"] 75} 76 77// CHECK-LABEL: func @contraction4x4_ijk 78 79// CHECK: %[[LMASK:.*]] = vector.constant_mask [4, 6] : vector<4x6xi1> 80// CHECK-NEXT: %[[RMASK:.*]] = vector.constant_mask [6, 4] : vector<6x4xi1> 81 82// Reducing output vector [0, 0] 83 84// CHECK-NEXT: %[[ES1:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x6xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 85// CHECK-NEXT: %[[ES2:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<6x4xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 86// CHECK-NEXT: %[[ES3:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x4xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 87// CHECK-NEXT: %[[ES4:.*]] = vector.extract_slices %[[LMASK]], [2, 2], [1, 1] : vector<4x6xi1> into tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 88// CHECK-NEXT: %[[ES5:.*]] = vector.extract_slices %[[RMASK]], [2, 2], [1, 1] : vector<6x4xi1> into tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 89 90// CHECK-NEXT: %[[TG1:.*]] = vector.tuple_get %[[ES1]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 91// CHECK-NEXT: %[[TG2:.*]] = vector.tuple_get %[[ES2]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 92// CHECK-NEXT: %[[TG3:.*]] = vector.tuple_get %[[ES3]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 93// CHECK-NEXT: %[[TG4:.*]] = vector.tuple_get %[[ES4]], 0 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 94// CHECK-NEXT: %[[TG5:.*]] = vector.tuple_get %[[ES5]], 0 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 95// CHECK-NEXT: %[[R1S00:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG1]], %[[TG2]], %[[TG3]], %[[TG4]], %[[TG5]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 96 97// CHECK-NEXT: %[[TG6:.*]] = vector.tuple_get %[[ES1]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 98// CHECK-NEXT: %[[TG7:.*]] = vector.tuple_get %[[ES2]], 2 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 99// CHECK-NEXT: %[[TG8:.*]] = vector.tuple_get %[[ES4]], 1 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 100// CHECK-NEXT: %[[TG9:.*]] = vector.tuple_get %[[ES5]], 2 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 101// CHECK-NEXT: %[[R2S00:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG6]], %[[TG7]], %[[R1S00]], %[[TG8]], %[[TG9]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 102 103// CHECK-NEXT: %[[TG10:.*]] = vector.tuple_get %[[ES1]], 2 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 104// CHECK-NEXT: %[[TG11:.*]] = vector.tuple_get %[[ES2]], 4 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 105// CHECK-NEXT: %[[TG12:.*]] = vector.tuple_get %[[ES4]], 2 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 106// CHECK-NEXT: %[[TG13:.*]] = vector.tuple_get %[[ES5]], 4 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 107// CHECK-NEXT: %[[R3S00:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG10]], %[[TG11]], %[[R2S00]], %[[TG12]], %[[TG13]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 108 109// Reducing output vector [0, 2] 110 111// CHECK-NEXT: %[[TG14:.*]] = vector.tuple_get %[[ES2]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 112// CHECK-NEXT: %[[TG15:.*]] = vector.tuple_get %[[ES3]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 113// CHECK-NEXT: %[[TG16:.*]] = vector.tuple_get %[[ES5]], 1 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 114// CHECK-NEXT: %[[R1S02:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG1]], %[[TG14]], %[[TG15]], %[[TG4]], %[[TG16]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 115 116// CHECK-NEXT: %[[TG17:.*]] = vector.tuple_get %[[ES2]], 3 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 117// CHECK-NEXT: %[[TG18:.*]] = vector.tuple_get %[[ES5]], 3 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 118// CHECK-NEXT: %[[R2S02:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG6]], %[[TG17]], %[[R1S02]], %[[TG8]], %[[TG18]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 119 120// CHECK-NEXT: %[[TG19:.*]] = vector.tuple_get %[[ES2]], 5 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 121// CHECK-NEXT: %[[TG20:.*]] = vector.tuple_get %[[ES5]], 5 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 122// CHECK-NEXT: %[[R3S02:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG10]], %[[TG19]], %[[R2S02]], %[[TG12]], %[[TG20]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 123 124// Reducing output vector [2, 0] 125 126// CHECK-NEXT: %[[TG21:.*]] = vector.tuple_get %[[ES1]], 3 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 127// CHECK-NEXT: %[[TG22:.*]] = vector.tuple_get %[[ES3]], 2 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 128// CHECK-NEXT: %[[TG23:.*]] = vector.tuple_get %[[ES4]], 3 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 129// CHECK-NEXT: %[[R1S20:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG21]], %[[TG2]], %[[TG22]], %[[TG23]], %[[TG5]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 130 131// CHECK-NEXT: %[[TG24:.*]] = vector.tuple_get %[[ES1]], 4 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 132// CHECK-NEXT: %[[TG25:.*]] = vector.tuple_get %[[ES4]], 4 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 133// CHECK-NEXT: %[[R2S20:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG24]], %[[TG7]], %[[R1S20]], %[[TG25]], %[[TG9]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 134 135// CHECK-NEXT: %[[TG26:.*]] = vector.tuple_get %[[ES1]], 5 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 136// CHECK-NEXT: %[[TG27:.*]] = vector.tuple_get %[[ES4]], 5 : tuple<vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>, vector<2x2xi1>> 137// CHECK-NEXT: %[[R3S20:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG26]], %[[TG11]], %[[R2S20]], %[[TG27]], %[[TG13]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 138 139// Reducing output vector [2, 2] 140 141// CHECK-NEXT: %[[TG28:.*]] = vector.tuple_get %[[ES3]], 3 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 142// CHECK-NEXT: %[[R1S22:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG21]], %[[TG14]], %[[TG28]], %[[TG23]], %[[TG16]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 143// CHECK-NEXT: %[[R2S22:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG24]], %[[TG17]], %[[R1S22]], %[[TG25]], %[[TG18]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 144// CHECK-NEXT: %[[R3S22:.*]] = vector.contract {indexing_maps = [#map0, #map1, #map2], iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[TG26]], %[[TG19]], %[[R2S22]], %[[TG27]], %[[TG20]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 145 146// CHECK-NEXT: %[[RES0:.*]] = vector.tuple %[[R3S00]], %[[R3S02]], %[[R3S20]], %[[R3S22]] : vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32> 147// CHECK-NEXT: %[[RES1:.*]] = vector.insert_slices %[[RES0]], [2, 2], [1, 1] : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> into vector<4x4xf32> 148// CHECK-NEXT: return %[[RES1]] : vector<4x4xf32> 149 150func @contraction4x4_ijk(%arg0 : vector<4x6xf32>, %arg1 : vector<6x4xf32>, 151 %arg2 : vector<4x4xf32>, %arg3 : index) 152 -> (vector<4x4xf32>) { 153 %lhsm = vector.constant_mask [4, 6] : vector<4x6xi1> 154 %rhsm = vector.constant_mask [6, 4] : vector<6x4xi1> 155 %0 = vector.contract #contraction_trait0 %arg0, %arg1, %arg2, %lhsm, %rhsm 156 : vector<4x6xf32>, vector<6x4xf32> into vector<4x4xf32> 157 158 return %0 : vector<4x4xf32> 159} 160 161#contraction_accesses1 = [ 162 affine_map<(i, k, j) -> (i, k)>, 163 affine_map<(i, k, j) -> (k, j)>, 164 affine_map<(i, k, j) -> (i, j)> 165] 166#contraction_trait1 = { 167 indexing_maps = #contraction_accesses1, 168 iterator_types = ["parallel", "reduction", "parallel"] 169} 170 171// CHECK-LABEL: func @contraction4x4_ikj 172 173 174// CHECK: %[[LMASK:.*]] = vector.constant_mask [4, 2] : vector<4x2xi1> 175// CHECK-NEXT: %[[RMASK:.*]] = vector.constant_mask [2, 4] : vector<2x4xi1> 176 177// Reducing output vector [0, 0] 178 179// CHECK-NEXT: %[[ES1:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x2xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>> 180// CHECK-NEXT: %[[ES2:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<2x4xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>> 181// CHECK-NEXT: %[[ES3:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x4xf32> into tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 182// CHECK-NEXT: %[[ES4:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<4x2xi1> into tuple<vector<2x2xi1>, vector<2x2xi1>> 183// CHECK-NEXT: %[[ES5:.*]] = vector.extract_slices %{{.*}}, [2, 2], [1, 1] : vector<2x4xi1> into tuple<vector<2x2xi1>, vector<2x2xi1>> 184 185// CHECK-NEXT: %[[TG1:.*]] = vector.tuple_get %[[ES1]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>> 186// CHECK-NEXT: %[[TG2:.*]] = vector.tuple_get %[[ES2]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>> 187// CHECK-NEXT: %[[TG3:.*]] = vector.tuple_get %[[ES3]], 0 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 188// CHECK-NEXT: %[[TG4:.*]] = vector.tuple_get %[[ES4]], 0 : tuple<vector<2x2xi1>, vector<2x2xi1>> 189// CHECK-NEXT: %[[TG5:.*]] = vector.tuple_get %[[ES5]], 0 : tuple<vector<2x2xi1>, vector<2x2xi1>> 190// CHECK-NEXT: %[[R1S00:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[TG1]], %[[TG2]], %[[TG3]], %[[TG4]], %[[TG5]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 191 192// Reducing output vector [0, 2] 193 194// CHECK-NEXT: %[[TG6:.*]] = vector.tuple_get %[[ES2]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>> 195// CHECK-NEXT: %[[TG7:.*]] = vector.tuple_get %[[ES3]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 196// CHECK-NEXT: %[[TG8:.*]] = vector.tuple_get %[[ES5]], 1 : tuple<vector<2x2xi1>, vector<2x2xi1>> 197// CHECK-NEXT: %[[R1S02:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[TG1]], %[[TG6]], %[[TG7]], %[[TG4]], %[[TG8]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 198 199// Reducing output vector [2, 0] 200 201// CHECK-NEXT: %[[TG9:.*]] = vector.tuple_get %[[ES1]], 1 : tuple<vector<2x2xf32>, vector<2x2xf32>> 202// CHECK-NEXT: %[[TG10:.*]] = vector.tuple_get %[[ES3]], 2 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 203// CHECK-NEXT: %[[TG11:.*]] = vector.tuple_get %[[ES4]], 1 : tuple<vector<2x2xi1>, vector<2x2xi1>> 204// CHECK-NEXT: %[[R1S20:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[TG9]], %[[TG2]], %[[TG10]], %[[TG11]], %[[TG5]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 205 206// Reducing output vector [2, 2] 207 208// CHECK-NEXT: %[[TG12:.*]] = vector.tuple_get %[[ES3]], 3 : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> 209// CHECK-NEXT: %[[R1S22:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[TG9]], %[[TG6]], %[[TG12]], %[[TG11]], %[[TG8]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 210 211// CHECK-NEXT: %[[RES0:.*]] = vector.tuple %[[R1S00]], %[[R1S02]], %[[R1S20]], %[[R1S22]] : vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32> 212// CHECK-NEXT: %[[RES1:.*]] = vector.insert_slices %[[RES0]], [2, 2], [1, 1] : tuple<vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>, vector<2x2xf32>> into vector<4x4xf32> 213// CHECK-NEXT: return %[[RES1]] : vector<4x4xf32> 214 215func @contraction4x4_ikj(%arg0 : vector<4x2xf32>, %arg1 : vector<2x4xf32>, 216 %arg2 : vector<4x4xf32>, %arg3 : index) 217 -> (vector<4x4xf32>) { 218 %lhsm = vector.constant_mask [4, 2] : vector<4x2xi1> 219 %rhsm = vector.constant_mask [2, 4] : vector<2x4xi1> 220 %0 = vector.contract #contraction_trait1 %arg0, %arg1, %arg2, %lhsm, %rhsm 221 : vector<4x2xf32>, vector<2x4xf32> into vector<4x4xf32> 222 223 return %0 : vector<4x4xf32> 224} 225 226// CHECK-LABEL: func @contraction4x4_ikj_xfer_read 227 228// CHECK-DAG: %[[C2:.*]] = constant 2 : index 229// CHECK-DAG: %[[C0:.*]] = constant 0 : index 230 231// Check LHS vector.transfer read is split for each user. 232 233// CHECK: %[[VTR0:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]]], %{{.*}} : memref<4x2xf32>, vector<2x2xf32> 234// CHECK-NEXT: %[[VTR1:.*]] = vector.transfer_read %{{.*}}[%[[C2]], %[[C0]]], %{{.*}} : memref<4x2xf32>, vector<2x2xf32> 235 236// CHECK-NEXT: %[[VTR2:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]]], %{{.*}} : memref<2x4xf32>, vector<2x2xf32> 237// CHECK-NEXT: %[[VTR3:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C2]]], %{{.*}} : memref<2x4xf32>, vector<2x2xf32> 238 239// CHECK-NEXT: %[[VTR4:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]]], %{{.*}} : memref<4x4xf32>, vector<2x2xf32> 240// CHECK-NEXT: %[[VTR5:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C2]]], %{{.*}} : memref<4x4xf32>, vector<2x2xf32> 241// CHECK-NEXT: %[[VTR6:.*]] = vector.transfer_read %{{.*}}[%[[C2]], %[[C0]]], %{{.*}} : memref<4x4xf32>, vector<2x2xf32> 242// CHECK-NEXT: %[[VTR7:.*]] = vector.transfer_read %{{.*}}[%[[C2]], %[[C2]]], %{{.*}} : memref<4x4xf32>, vector<2x2xf32> 243 244// CHECK-NEXT: %[[R0:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR0]], %[[VTR2]], %[[VTR4]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 245// CHECK-NEXT: %[[R1:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR0]], %[[VTR3]], %[[VTR5]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 246// CHECK-NEXT: %[[R2:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR1]], %[[VTR2]], %[[VTR6]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 247// CHECK-NEXT: %[[R3:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR1]], %[[VTR3]], %[[VTR7]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 248 249// CHECK-NEXT: vector.transfer_write %[[R0]], %{{.*}}[%[[C0]], %[[C0]]] {masked = [false, false]} : vector<2x2xf32>, memref<4x4xf32> 250// CHECK-NEXT: vector.transfer_write %[[R1]], %{{.*}}[%[[C0]], %[[C2]]] {masked = [false, false]} : vector<2x2xf32>, memref<4x4xf32> 251// CHECK-NEXT: vector.transfer_write %[[R2]], %{{.*}}[%[[C2]], %[[C0]]] {masked = [false, false]} : vector<2x2xf32>, memref<4x4xf32> 252// CHECK-NEXT: vector.transfer_write %[[R3]], %{{.*}}[%[[C2]], %[[C2]]] {masked = [false, false]} : vector<2x2xf32>, memref<4x4xf32> 253// CHECK-NEXT: return 254 255func @contraction4x4_ikj_xfer_read(%arg0 : memref<4x2xf32>, 256 %arg1 : memref<2x4xf32>, 257 %arg2 : memref<4x4xf32>) { 258 %c0 = constant 0 : index 259 %cf0 = constant 0.0 : f32 260 261 %0 = vector.transfer_read %arg0[%c0, %c0], %cf0 262 { permutation_map = affine_map<(d0, d1) -> (d0, d1)> } 263 : memref<4x2xf32>, vector<4x2xf32> 264 265 %1 = vector.transfer_read %arg1[%c0, %c0], %cf0 266 { permutation_map = affine_map<(d0, d1) -> (d0, d1)> } 267 : memref<2x4xf32>, vector<2x4xf32> 268 269 %2 = vector.transfer_read %arg2[%c0, %c0], %cf0 270 { permutation_map = affine_map<(d0, d1) -> (d0, d1)> } 271 : memref<4x4xf32>, vector<4x4xf32> 272 273 %3 = vector.contract #contraction_trait1 %0, %1, %2 274 : vector<4x2xf32>, vector<2x4xf32> into vector<4x4xf32> 275 276 vector.transfer_write %3, %arg2[%c0, %c0] 277 {permutation_map = affine_map<(d0, d1) -> (d0, d1)>} 278 : vector<4x4xf32>, memref<4x4xf32> 279 return 280} 281 282// TODO: Update test with VTR split transform. 283// CHECK-LABEL: func @vector_transfers 284// CHECK-COUNT-8: vector.transfer_read 285// CHECK-COUNT-4: addf 286// CHECK-COUNT-4: vector.transfer_write 287 288func @vector_transfers(%arg0: index, %arg1: index) { 289 %cst = constant 0.000000e+00 : f32 290 %0 = memref.alloc(%arg0, %arg1) : memref<?x?xf32> 291 %1 = memref.alloc(%arg0, %arg1) : memref<?x?xf32> 292 %2 = memref.alloc(%arg0, %arg1) : memref<?x?xf32> 293 %cst_0 = constant 1.000000e+00 : f32 294 %cst_1 = constant 2.000000e+00 : f32 295 affine.for %arg2 = 0 to %arg0 step 4 { 296 affine.for %arg3 = 0 to %arg1 step 4 { 297 %4 = vector.transfer_read %0[%arg2, %arg3], %cst {permutation_map = affine_map<(d0, d1) -> (d0, d1)>} : memref<?x?xf32>, vector<4x4xf32> 298 %5 = vector.transfer_read %1[%arg2, %arg3], %cst {permutation_map = affine_map<(d0, d1) -> (d0, d1)>} : memref<?x?xf32>, vector<4x4xf32> 299 %6 = addf %4, %5 : vector<4x4xf32> 300 vector.transfer_write %6, %2[%arg2, %arg3] {permutation_map = affine_map<(d0, d1) -> (d0, d1)>} : vector<4x4xf32>, memref<?x?xf32> 301 } 302 } 303 return 304} 305 306// CHECK-LABEL: func @tuple_get(%arg0: vector<4xf32>, %arg1: vector<8xf32>) 307// CHECK: return %arg1 308 309func @tuple_get(%arg0: vector<4xf32>, %arg1: vector<8xf32>) -> vector<8xf32> { 310 %0 = vector.tuple %arg0, %arg1 : vector<4xf32>, vector<8xf32> 311 %1 = vector.tuple_get %0, 1 : tuple<vector<4xf32>, vector<8xf32>> 312 return %1 : vector<8xf32> 313} 314 315// CHECK-LABEL: func @tuple_get_producer_consumer 316// CHECK-SAME: %[[A0:.*0]]: vector<2x4xf32>, 317// CHECK-SAME: %[[A1:.*1]]: vector<2x4xf32>, 318// CHECK-SAME: %[[A2:.*2]]: vector<2x4xf32>, 319// CHECK-SAME: %[[A3:.*3]]: vector<2x4xf32>, 320// CHECK-SAME: %[[A4:.*4]]: vector<2x4xf32>, 321// CHECK-SAME: %[[A5:.*5]]: vector<2x4xf32>, 322// CHECK-SAME: %[[A6:.*6]]: vector<2x4xf32>, 323// CHECK-SAME: %[[A7:.*7]]: vector<2x4xf32> 324// CHECK: return %[[A7]] : vector<2x4xf32> 325 326func @tuple_get_producer_consumer( 327 %arg0 : vector<2x4xf32>, %arg1 : vector<2x4xf32>, 328 %arg2 : vector<2x4xf32>, %arg3 : vector<2x4xf32>, 329 %arg4 : vector<2x4xf32>, %arg5 : vector<2x4xf32>, 330 %arg6 : vector<2x4xf32>, %arg7 : vector<2x4xf32>) -> vector<2x4xf32> { 331 %0 = vector.tuple %arg0, %arg1, %arg2, %arg3, %arg4, %arg5, %arg6, %arg7 332 : vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, 333 vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32> 334 // %arg7 == %0 at tupleIndex = 7, offsets = [0, 0] 335 %1 = vector.insert_slices %0, [2, 4], [1, 1] 336 : tuple<vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, 337 vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>> 338 into vector<4x16xf32> 339 // %arg7 == %1 at tupleIndex = -1, offsets = [2, 12] 340 %2 = vector.extract_slices %1, [4, 8], [1, 1] 341 : vector<4x16xf32> into tuple<vector<4x8xf32>, vector<4x8xf32>> 342 // %arg7 == %2 at tupleIndex = 1, offsets = [2, 4] 343 %3 = vector.shape_cast %2 : tuple<vector<4x8xf32>, vector<4x8xf32>> to 344 tuple<vector<1x1x4x8xf32>, vector<1x1x4x8xf32>> 345 // %arg7 = %3 at tupleIndex = 1, offsets = [0, 0, 2, 4] 346 %4 = vector.tuple_get %3, 1 : tuple<vector<1x1x4x8xf32>, vector<1x1x4x8xf32>> 347 // %arg7 == %4 at tupleIndex = -1, offsets = [0, 0, 2, 4] 348 %5 = vector.shape_cast %4 : vector<1x1x4x8xf32> to vector<4x8xf32> 349 // %arg7 == %5 at tupleIndex = -1, offsets = [2, 4] 350 %6 = vector.extract_slices %5, [2, 4], [1, 1] 351 : vector<4x8xf32> into 352 tuple<vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>> 353 // %arg7 == %6 at tupleIndex = 3, offsets = [0, 0] 354 %7 = vector.tuple_get %6, 3 355 : tuple<vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>> 356 // %arg7 == %7 357 return %7 : vector<2x4xf32> 358} 359 360// CHECK-LABEL: func @tuple_get_producer_consumer_swizzle 361// CHECK-SAME: %[[A0:.*0]]: vector<2x4xf32>, 362// CHECK-SAME: %[[A1:.*1]]: vector<2x4xf32>, 363// CHECK-SAME: %[[A2:.*2]]: vector<2x4xf32>, 364// CHECK-SAME: %[[A3:.*3]]: vector<2x4xf32>, 365// CHECK-SAME: %[[A4:.*4]]: vector<2x4xf32>, 366// CHECK-SAME: %[[A5:.*5]]: vector<2x4xf32>, 367// CHECK-SAME: %[[A6:.*6]]: vector<2x4xf32>, 368// CHECK-SAME: %[[A7:.*7]]: vector<2x4xf32> 369// CHECK: return %[[A7]] : vector<2x4xf32> 370 371func @tuple_get_producer_consumer_swizzle( 372 %arg0 : vector<2x4xf32>, %arg1 : vector<2x4xf32>, 373 %arg2 : vector<2x4xf32>, %arg3 : vector<2x4xf32>, 374 %arg4 : vector<2x4xf32>, %arg5 : vector<2x4xf32>, 375 %arg6 : vector<2x4xf32>, %arg7 : vector<2x4xf32>) -> vector<2x4xf32> { 376 %0 = vector.tuple %arg0, %arg1, %arg2, %arg3, %arg4, %arg5, %arg6, %arg7 377 : vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, 378 vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32> 379 // %arg7 == %0 at tupleIndex = 7, offsets = [0, 0] 380 %1 = vector.insert_slices %0, [2, 4], [1, 1] 381 : tuple<vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, 382 vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>> 383 into vector<4x16xf32> 384 // %arg7 == %1 at tupleIndex = -1, offsets = [2, 12] 385 %2 = vector.extract_slices %1, [4, 8], [1, 1] 386 : vector<4x16xf32> into tuple<vector<4x8xf32>, vector<4x8xf32>> 387 // %arg7 == %2 at tupleIndex = 1, offsets = [2, 4] 388 %3= vector.shape_cast %2 : tuple<vector<4x8xf32>, vector<4x8xf32>> to 389 tuple<vector<1x1x4x8xf32>, vector<1x1x4x8xf32>> 390 // %arg7 = %3 at tupleIndex = 1, offsets = [0, 0, 2, 4] 391 392 // Extract tuple elements. 393 %4 = vector.tuple_get %3, 0 : tuple<vector<1x1x4x8xf32>, vector<1x1x4x8xf32>> 394 %5 = vector.tuple_get %3, 1 : tuple<vector<1x1x4x8xf32>, vector<1x1x4x8xf32>> 395 // %arg7 == %5 at tupleIndex = -1, offsets = [0, 0, 2, 4] 396 397 // Swizzle tuple elements. 398 %6 = vector.tuple %5, %4 : vector<1x1x4x8xf32>, vector<1x1x4x8xf32> 399 // %arg7 == %6 at tupleIndex = 0, offsets = [0, 0, 2, 4] 400 %7 = vector.shape_cast %6 : tuple<vector<1x1x4x8xf32>, vector<1x1x4x8xf32>> to 401 tuple<vector<4x8xf32>, vector<4x8xf32>> 402 // %arg7 = %7 at tupleIndex = 0, offsets = [2, 4] 403 %8 = vector.tuple_get %7, 0 : tuple<vector<4x8xf32>, vector<4x8xf32>> 404 // %arg7 == %8 at tupleIndex = -1, offsets = [2, 4] 405 %9 = vector.extract_slices %8, [2, 4], [1, 1] 406 : vector<4x8xf32> into 407 tuple<vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>> 408 // %arg7 == %9 at tupleIndex = 3, offsets = [0, 0] 409 %10 = vector.tuple_get %9, 3 410 : tuple<vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>, vector<2x4xf32>> 411 // %arg7 == %10 412 return %10 : vector<2x4xf32> 413} 414 415// CHECK-LABEL: func @cancelling_shape_cast_ops 416// CHECK-SAME: %[[A0:.*0]]: vector<2x4xf32> 417// CHECK: return %[[A0]] : vector<2x4xf32> 418func @cancelling_shape_cast_ops(%arg0 : vector<2x4xf32>) -> vector<2x4xf32> { 419 %0 = vector.shape_cast %arg0 : vector<2x4xf32> to vector<8xf32> 420 %1 = vector.shape_cast %0 : vector<8xf32> to vector<2x4xf32> 421 return %1 : vector<2x4xf32> 422} 423 424// CHECK-LABEL: func @vector_transfers_vector_element_type 425// CHECK-DAG: %[[C1:.*]] = constant 1 : index 426// CHECK-DAG: %[[C0:.*]] = constant 0 : index 427// CHECK: %[[VTR0:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]], %[[C0]]], %{{.*}} {masked = [false, false]} : memref<6x2x1xvector<2x4xf32>>, vector<1x1x2x4xf32> 428// CHECK-NEXT: %[[VTR1:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C1]], %[[C0]]], %{{.*}} {masked = [false, false]} : memref<6x2x1xvector<2x4xf32>>, vector<1x1x2x4xf32> 429// CHECK-NEXT: vector.transfer_write %[[VTR0]], %{{.*}}[%[[C0]], %[[C0]], %[[C0]]] {masked = [false, false]} : vector<1x1x2x4xf32>, memref<6x2x1xvector<2x4xf32>> 430// CHECK-NEXT: vector.transfer_write %[[VTR1]], %{{.*}}[%[[C0]], %[[C1]], %[[C0]]] {masked = [false, false]} : vector<1x1x2x4xf32>, memref<6x2x1xvector<2x4xf32>> 431 432func @vector_transfers_vector_element_type() { 433 %c0 = constant 0 : index 434 %cf0 = constant 0.000000e+00 : f32 435 %vf0 = splat %cf0 : vector<2x4xf32> 436 437 %0 = memref.alloc() : memref<6x2x1xvector<2x4xf32>> 438 439 %1 = vector.transfer_read %0[%c0, %c0, %c0], %vf0 440 {permutation_map = affine_map<(d0, d1, d2) -> (d1, d2)>} 441 : memref<6x2x1xvector<2x4xf32>>, vector<2x1x2x4xf32> 442 443 %2 = vector.extract_slices %1, [1, 1, 2, 4], [1, 1, 1, 1] 444 : vector<2x1x2x4xf32> into tuple<vector<1x1x2x4xf32>, vector<1x1x2x4xf32>> 445 %3 = vector.tuple_get %2, 0 : tuple<vector<1x1x2x4xf32>, vector<1x1x2x4xf32>> 446 %4 = vector.tuple_get %2, 1 : tuple<vector<1x1x2x4xf32>, vector<1x1x2x4xf32>> 447 %5 = vector.tuple %3, %4 : vector<1x1x2x4xf32>, vector<1x1x2x4xf32> 448 %6 = vector.insert_slices %5, [1, 1, 2, 4], [1, 1, 1, 1] 449 : tuple<vector<1x1x2x4xf32>, vector<1x1x2x4xf32>> into vector<2x1x2x4xf32> 450 451 vector.transfer_write %6, %0[%c0, %c0, %c0] 452 {permutation_map = affine_map<(d0, d1, d2) -> (d1, d2)>} 453 : vector<2x1x2x4xf32>, memref<6x2x1xvector<2x4xf32>> 454 455 return 456} 457 458// Test that ShapeCastOp on tuple of vectors, decomposes to multiple 459// ShapeCastOps on vectors. 460// CHECK-LABEL: func @shape_cast_decomposition 461// CHECK: %[[V0:.*]] = vector.shape_cast %{{.*}} : vector<5x4x2xf32> to vector<20x2xf32> 462// CHECK-NEXT: %[[V1:.*]] = vector.shape_cast %{{.*}} : vector<3x4x2xf32> to vector<12x2xf32> 463// CHECK-NEXT: return %[[V0]], %[[V1]] : vector<20x2xf32>, vector<12x2xf32> 464 465func @shape_cast_decomposition(%arg0 : vector<5x4x2xf32>, 466 %arg1 : vector<3x4x2xf32>) 467 -> (vector<20x2xf32>, vector<12x2xf32>) { 468 %0 = vector.tuple %arg0, %arg1 : vector<5x4x2xf32>, vector<3x4x2xf32> 469 %1 = vector.shape_cast %0 : tuple<vector<5x4x2xf32>, vector<3x4x2xf32>> to 470 tuple<vector<20x2xf32>, vector<12x2xf32>> 471 %2 = vector.tuple_get %1, 0 : tuple<vector<20x2xf32>, vector<12x2xf32>> 472 %3 = vector.tuple_get %1, 1 : tuple<vector<20x2xf32>, vector<12x2xf32>> 473 return %2, %3 : vector<20x2xf32>, vector<12x2xf32> 474} 475 476// Test that cancelling ShapeCastOps are canonicalized away. 477// EX: 478// 479// The following MLIR with cancelling ShapeCastOps: 480// 481// %0 = source : vector<5x4x2xf32> 482// %1 = shape_cast %0 : vector<5x4x2xf32> to vector<20x2xf32> 483// %2 = shape_cast %1 : vector<20x2xf32> to vector<5x4x2xf32> 484// %3 = user %2 : vector<5x4x2xf32> 485// 486// Should canonicalize to the following: 487// 488// 489// %0 = source : vector<5x4x2xf32> 490// %1 = user %0 : vector<5x4x2xf32> 491// 492 493// ShapeCastOps on vectors. 494// CHECK-LABEL: func @shape_cast_fold 495// CHECK: return %{{.*}}, %{{.*}} : vector<5x4x2xf32>, vector<3x4x2xf32> 496 497func @shape_cast_fold(%arg0 : vector<5x4x2xf32>, %arg1 : vector<3x4x2xf32>) 498 -> (vector<5x4x2xf32>, vector<3x4x2xf32>) { 499 %0 = vector.tuple %arg0, %arg1 : vector<5x4x2xf32>, vector<3x4x2xf32> 500 501 %1 = vector.shape_cast %0 : tuple<vector<5x4x2xf32>, vector<3x4x2xf32>> to 502 tuple<vector<20x2xf32>, vector<12x2xf32>> 503 504 %2 = vector.tuple_get %1, 0 : tuple<vector<20x2xf32>, vector<12x2xf32>> 505 %3 = vector.tuple_get %1, 1 : tuple<vector<20x2xf32>, vector<12x2xf32>> 506 507 %4 = vector.tuple %2, %3 : vector<20x2xf32>, vector<12x2xf32> 508 %5 = vector.shape_cast %4 : tuple<vector<20x2xf32>, vector<12x2xf32>> to 509 tuple<vector<5x4x2xf32>, vector<3x4x2xf32>> 510 511 %6 = vector.tuple_get %5, 0 : tuple<vector<5x4x2xf32>, vector<3x4x2xf32>> 512 %7 = vector.tuple_get %5, 1 : tuple<vector<5x4x2xf32>, vector<3x4x2xf32>> 513 514 return %6, %7 : vector<5x4x2xf32>, vector<3x4x2xf32> 515} 516 517// CHECK-LABEL: func @elementwise_unroll 518// CHECK-SAME: (%[[ARG0:.*]]: memref<4x4xf32>, %[[ARG1:.*]]: memref<4x4xf32>) 519// CHECK-DAG: %[[C2:.*]] = constant 2 : index 520// CHECK-DAG: %[[C0:.*]] = constant 0 : index 521// CHECK: %[[VT0:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 522// CHECK: %[[VT1:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 523// CHECK: %[[VT2:.*]] = vector.transfer_read %[[ARG0]][%[[C2]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 524// CHECK: %[[VT3:.*]] = vector.transfer_read %[[ARG0]][%[[C2]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 525// CHECK: %[[VT4:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 526// CHECK: %[[VT5:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 527// CHECK: %[[VT6:.*]] = vector.transfer_read %[[ARG1]][%[[C2]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 528// CHECK: %[[VT7:.*]] = vector.transfer_read %[[ARG1]][%[[C2]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 529// CHECK: %[[CMP0:.*]] = cmpf ult, %[[VT0]], %[[VT4]] : vector<2x2xf32> 530// CHECK: %[[CMP1:.*]] = cmpf ult, %[[VT1]], %[[VT5]] : vector<2x2xf32> 531// CHECK: %[[CMP2:.*]] = cmpf ult, %[[VT2]], %[[VT6]] : vector<2x2xf32> 532// CHECK: %[[CMP3:.*]] = cmpf ult, %[[VT3]], %[[VT7]] : vector<2x2xf32> 533// CHECK: %[[VT0:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 534// CHECK: %[[VT1:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 535// CHECK: %[[VT2:.*]] = vector.transfer_read %[[ARG0]][%[[C2]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 536// CHECK: %[[VT3:.*]] = vector.transfer_read %[[ARG0]][%[[C2]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 537// CHECK: %[[VT4:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 538// CHECK: %[[VT5:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 539// CHECK: %[[VT6:.*]] = vector.transfer_read %[[ARG1]][%[[C2]], %[[C0]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 540// CHECK: %[[VT7:.*]] = vector.transfer_read %[[ARG1]][%[[C2]], %[[C2]]], {{.*}} : memref<4x4xf32>, vector<2x2xf32> 541// CHECK: %[[SEL0:.*]] = select %[[CMP0]], %[[VT0]], %[[VT4]] : vector<2x2xi1>, vector<2x2xf32> 542// CHECK: %[[SEL1:.*]] = select %[[CMP1]], %[[VT1]], %[[VT5]] : vector<2x2xi1>, vector<2x2xf32> 543// CHECK: %[[SEL2:.*]] = select %[[CMP2]], %[[VT2]], %[[VT6]] : vector<2x2xi1>, vector<2x2xf32> 544// CHECK: %[[SEL3:.*]] = select %[[CMP3]], %[[VT3]], %[[VT7]] : vector<2x2xi1>, vector<2x2xf32> 545// CHECK: vector.transfer_write %[[SEL0]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<2x2xf32>, memref<4x4xf32> 546// CHECK: vector.transfer_write %[[SEL1]], %[[ARG0]][%[[C0]], %[[C2]]] {{.*}} : vector<2x2xf32>, memref<4x4xf32> 547// CHECK: vector.transfer_write %[[SEL2]], %[[ARG0]][%[[C2]], %[[C0]]] {{.*}} : vector<2x2xf32>, memref<4x4xf32> 548// CHECK: vector.transfer_write %[[SEL3]], %[[ARG0]][%[[C2]], %[[C2]]] {{.*}} : vector<2x2xf32>, memref<4x4xf32> 549func @elementwise_unroll(%arg0 : memref<4x4xf32>, %arg1 : memref<4x4xf32>) { 550 %c0 = constant 0 : index 551 %cf0 = constant 0.0 : f32 552 %0 = vector.transfer_read %arg0[%c0, %c0], %cf0 : memref<4x4xf32>, vector<4x4xf32> 553 %1 = vector.transfer_read %arg1[%c0, %c0], %cf0 : memref<4x4xf32>, vector<4x4xf32> 554 %cond = cmpf ult, %0, %1 : vector<4x4xf32> 555 // Vector transfer split pattern only support single user right now. 556 %2 = vector.transfer_read %arg0[%c0, %c0], %cf0 : memref<4x4xf32>, vector<4x4xf32> 557 %3 = vector.transfer_read %arg1[%c0, %c0], %cf0 : memref<4x4xf32>, vector<4x4xf32> 558 %4 = select %cond, %2, %3 : vector<4x4xi1>, vector<4x4xf32> 559 vector.transfer_write %4, %arg0[%c0, %c0] : vector<4x4xf32>, memref<4x4xf32> 560 return 561} 562 563// Check that vector.transfer read/write are split based on contract unrolling. 564// CHECK: %[[VTR0:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]]], %{{.*}} : tensor<4x2xf32>, vector<2x2xf32> 565// CHECK-NEXT: %[[VTR1:.*]] = vector.transfer_read %{{.*}}[%[[C2]], %[[C0]]], %{{.*}} : tensor<4x2xf32>, vector<2x2xf32> 566 567// CHECK-NEXT: %[[VTR2:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]]], %{{.*}} : tensor<2x4xf32>, vector<2x2xf32> 568// CHECK-NEXT: %[[VTR3:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C2]]], %{{.*}} : tensor<2x4xf32>, vector<2x2xf32> 569 570// CHECK-NEXT: %[[VTR4:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]]], %{{.*}} : tensor<4x4xf32>, vector<2x2xf32> 571// CHECK-NEXT: %[[VTR5:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C2]]], %{{.*}} : tensor<4x4xf32>, vector<2x2xf32> 572// CHECK-NEXT: %[[VTR6:.*]] = vector.transfer_read %{{.*}}[%[[C2]], %[[C0]]], %{{.*}} : tensor<4x4xf32>, vector<2x2xf32> 573// CHECK-NEXT: %[[VTR7:.*]] = vector.transfer_read %{{.*}}[%[[C2]], %[[C2]]], %{{.*}} : tensor<4x4xf32>, vector<2x2xf32> 574 575// CHECK-NEXT: %[[R0:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR0]], %[[VTR2]], %[[VTR4]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 576// CHECK-NEXT: %[[R1:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR0]], %[[VTR3]], %[[VTR5]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 577// CHECK-NEXT: %[[R2:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR1]], %[[VTR2]], %[[VTR6]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 578// CHECK-NEXT: %[[R3:.*]] = vector.contract {indexing_maps = [#map2, #map3, #map0], iterator_types = ["parallel", "reduction", "parallel"], kind = #vector.kind<add>} %[[VTR1]], %[[VTR3]], %[[VTR7]] : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 579 580// CHECK-NEXT: %[[VTW0:.*]] = vector.transfer_write %[[R0]], %{{.*}}[%[[C0]], %[[C0]]] {masked = [false, false]} : vector<2x2xf32>, tensor<4x4xf32> 581// CHECK-NEXT: %[[VTW1:.*]] = vector.transfer_write %[[R1]], %[[VTW0]][%[[C0]], %[[C2]]] {masked = [false, false]} : vector<2x2xf32>, tensor<4x4xf32> 582// CHECK-NEXT: %[[VTW2:.*]] = vector.transfer_write %[[R2]], %[[VTW1]][%[[C2]], %[[C0]]] {masked = [false, false]} : vector<2x2xf32>, tensor<4x4xf32> 583// CHECK-NEXT: %[[VTW3:.*]] = vector.transfer_write %[[R3]], %[[VTW2]][%[[C2]], %[[C2]]] {masked = [false, false]} : vector<2x2xf32>, tensor<4x4xf32> 584// CHECK-NEXT: return %[[VTW3]] : tensor<4x4xf32> 585 586func @contraction4x4_ikj_xfer_read_tensor(%arg0 : tensor<4x2xf32>, 587 %arg1 : tensor<2x4xf32>, 588 %arg2 : tensor<4x4xf32>) -> 589 tensor<4x4xf32> { 590 %c0 = constant 0 : index 591 %cf0 = constant 0.0 : f32 592 %0 = vector.transfer_read %arg0[%c0, %c0], %cf0 : 593 tensor<4x2xf32>, vector<4x2xf32> 594 %1 = vector.transfer_read %arg1[%c0, %c0], %cf0 : 595 tensor<2x4xf32>, vector<2x4xf32> 596 %2 = vector.transfer_read %arg2[%c0, %c0], %cf0 : 597 tensor<4x4xf32>, vector<4x4xf32> 598 %3 = vector.contract #contraction_trait1 %0, %1, %2 599 : vector<4x2xf32>, vector<2x4xf32> into vector<4x4xf32> 600 %r = vector.transfer_write %3, %arg2[%c0, %c0] 601 : vector<4x4xf32>, tensor<4x4xf32> 602 return %r : tensor<4x4xf32> 603} 604 605// CHECK-LABEL: func @cast_away_extract_strided_slice_leading_one_dims 606func @cast_away_extract_strided_slice_leading_one_dims(%arg0: vector<1x8x8xf16>) -> vector<1x1x8xf16> { 607 // CHECK: %[[SRC:.+]] = vector.shape_cast %{{.*}} : vector<1x8x8xf16> to vector<8x8xf16> 608 // CHECK: %[[EXTRACT:.+]] = vector.extract_strided_slice %[[SRC]] {offsets = [4], sizes = [1], strides = [1]} : vector<8x8xf16> to vector<1x8xf16> 609 %0 = vector.extract_strided_slice %arg0 {offsets = [0, 4], sizes = [1, 1], strides = [1, 1]} : vector<1x8x8xf16> to vector<1x1x8xf16> 610 // CHECK: %[[RET:.+]] = vector.shape_cast %[[EXTRACT]] : vector<1x8xf16> to vector<1x1x8xf16> 611 // CHECK: return %[[RET]] 612 return %0: vector<1x1x8xf16> 613} 614 615// CHECK-LABEL: func @cast_away_insert_strided_slice_leading_one_dims 616func @cast_away_insert_strided_slice_leading_one_dims(%arg0: vector<1x8xf16>, %arg1: vector<1x8x8xf16>) -> vector<1x8x8xf16> { 617 // CHECK: %[[SRC:.+]] = vector.shape_cast %{{.*}} : vector<1x8xf16> to vector<8xf16> 618 // CHECK: %[[DST:.+]] = vector.shape_cast %{{.*}} : vector<1x8x8xf16> to vector<8x8xf16> 619 // CHECK: %[[INSERT:.+]] = vector.insert_strided_slice %[[SRC]], %[[DST]] {offsets = [0, 0], strides = [1]} : vector<8xf16> into vector<8x8xf16> 620 %0 = vector.insert_strided_slice %arg0, %arg1 {offsets = [0, 0, 0], strides = [1, 1]} : vector<1x8xf16> into vector<1x8x8xf16> 621 // CHECK: %[[RET:.+]] = vector.shape_cast %[[INSERT]] : vector<8x8xf16> to vector<1x8x8xf16> 622 // CHECK: return %[[RET]] 623 return %0: vector<1x8x8xf16> 624} 625 626// CHECK-LABEL: func @cast_away_insert_strided_slice_leading_one_dims_one_element 627func @cast_away_insert_strided_slice_leading_one_dims_one_element(%arg0: vector<1x1xf16>, %arg1: vector<1x1x1xf16>) -> vector<1x1x1xf16> { 628 // CHECK: vector.shape_cast %{{.+}} : vector<1x1xf16> to vector<1xf16> 629 // CHECK: vector.shape_cast %{{.+}} : vector<1x1x1xf16> to vector<1xf16> 630 %0 = vector.insert_strided_slice %arg0, %arg1 {offsets = [0, 0, 0], strides = [1, 1]} : vector<1x1xf16> into vector<1x1x1xf16> 631 return %0: vector<1x1x1xf16> 632} 633 634// CHECK-LABEL: func @cast_away_transfer_read_leading_one_dims 635func @cast_away_transfer_read_leading_one_dims(%arg0: memref<1x4x8x16xf16>) -> vector<1x4xf16> { 636 // CHECK: %[[C0:.+]] = constant 0 : index 637 %c0 = constant 0 : index 638 // CHECK: %[[F0:.+]] = constant 0.000000e+00 : f16 639 %f0 = constant 0. : f16 640 // CHECK: %[[READ:.+]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]], %[[C0]], %[[C0]]], %[[F0]] {masked = [false]} : memref<1x4x8x16xf16>, vector<4xf16> 641 // CHECK: %[[CAST:.+]] = vector.shape_cast %[[READ]] : vector<4xf16> to vector<1x4xf16> 642 %0 = vector.transfer_read %arg0[%c0, %c0, %c0, %c0], %f0 {masked = [false, false]} : memref<1x4x8x16xf16>, vector<1x4xf16> 643 // CHECK: return %[[CAST]] 644 return %0: vector<1x4xf16> 645} 646 647// CHECK-LABEL: func @cast_away_transfer_read_leading_one_dims_one_element 648func @cast_away_transfer_read_leading_one_dims_one_element(%arg0: memref<1x1x1x1xf16>) -> vector<1x1xf16> { 649 %c0 = constant 0 : index 650 %f0 = constant 0. : f16 651 // CHECK: vector.shape_cast %{{.+}} : vector<1xf16> to vector<1x1xf16> 652 %0 = vector.transfer_read %arg0[%c0, %c0, %c0, %c0], %f0 {masked = [false, false]} : memref<1x1x1x1xf16>, vector<1x1xf16> 653 return %0: vector<1x1xf16> 654} 655 656// CHECK-LABEL: func @cast_away_transfer_write_leading_one_dims 657func @cast_away_transfer_write_leading_one_dims(%arg0: memref<1x4x8x16xf16>, %arg1: vector<1x4xf16>) { 658 // CHECK: %[[C0:.+]] = constant 0 : index 659 %c0 = constant 0 : index 660 // CHECK: %[[CAST:.+]] = vector.shape_cast %{{.*}} : vector<1x4xf16> to vector<4xf16> 661 // CHECK: vector.transfer_write %[[CAST]], %{{.*}}[%[[C0]], %[[C0]], %[[C0]], %[[C0]]] {masked = [false]} : vector<4xf16>, memref<1x4x8x16xf16> 662 663 vector.transfer_write %arg1, %arg0[%c0, %c0, %c0, %c0] {masked = [false, false]} : vector<1x4xf16>, memref<1x4x8x16xf16> 664 return 665} 666 667// CHECK-LABEL: func @cast_away_transfer_write_leading_one_dims_one_element 668func @cast_away_transfer_write_leading_one_dims_one_element(%arg0: memref<1x1x1x1xf16>, %arg1: vector<1x1xf16>) { 669 %c0 = constant 0 : index 670 // CHECK: vector.shape_cast %{{.+}} : vector<1x1xf16> to vector<1xf16> 671 vector.transfer_write %arg1, %arg0[%c0, %c0, %c0, %c0] {masked = [false, false]} : vector<1x1xf16>, memref<1x1x1x1xf16> 672 return 673} 674 675// CHECK-LABEL: func @bubble_down_bitcast_in_extract 676// CHECK-SAME: %[[SRC:.+]]: vector<4xf32> 677func @bubble_down_bitcast_in_extract(%src: vector<4xf32>) -> (f16, f16) { 678 %0 = vector.bitcast %src : vector<4xf32> to vector<8xf16> 679 // CHECK: %[[EXTRACT1:.+]] = vector.extract %[[SRC]][1] : vector<4xf32> 680 // CHECK: %[[CAST1:.+]] = vector.bitcast %[[EXTRACT1]] : vector<1xf32> to vector<2xf16> 681 // CHECK: %[[EXTRACT2:.+]] = vector.extract %[[CAST1]][1] : vector<2xf16> 682 %1 = vector.extract %0[3] : vector<8xf16> 683 // CHECK: %[[EXTRACT3:.+]] = vector.extract %[[SRC]][2] : vector<4xf32> 684 // CHECK: %[[CAST2:.+]] = vector.bitcast %[[EXTRACT3]] : vector<1xf32> to vector<2xf16> 685 // CHECK: %[[EXTRACT4:.+]] = vector.extract %[[CAST2]][0] : vector<2xf16> 686 %2 = vector.extract %0[4] : vector<8xf16> 687 // CHECK: return %[[EXTRACT2]], %[[EXTRACT4]] 688 return %1, %2: f16, f16 689} 690 691// CHECK-LABEL: func @bubble_down_bitcast_in_strided_slice_extract 692// CHECK-SAME: %[[SRC:.+]]: vector<4xf32> 693func @bubble_down_bitcast_in_strided_slice_extract(%arg0: vector<4xf32>) -> vector<4xf16> { 694 // CHECK: %[[EXTRACT:.+]] = vector.extract_strided_slice %[[SRC]] {offsets = [2], sizes = [2], strides = [1]} : vector<4xf32> to vector<2xf32> 695 // CHECK: %[[CAST:.+]] = vector.bitcast %[[EXTRACT]] : vector<2xf32> to vector<4xf16> 696 %cast = vector.bitcast %arg0: vector<4xf32> to vector<8xf16> 697 %0 = vector.extract_strided_slice %cast {offsets = [4], sizes = [4], strides = [1]} : vector<8xf16> to vector<4xf16> 698 // CHECK: return %[[CAST]] 699 return %0: vector<4xf16> 700} 701 702// CHECK-LABEL: func @bubble_down_bitcast_in_strided_slice_extract_full_last_dim 703// CHECK-SAME: %[[SRC:.+]]: vector<4x2xf32> 704func @bubble_down_bitcast_in_strided_slice_extract_full_last_dim(%arg0: vector<4x2xf32>) -> vector<2x4xf16> { 705 // CHECK: %[[EXTRACT:.+]] = vector.extract_strided_slice %[[SRC]] {offsets = [1], sizes = [2], strides = [1]} : vector<4x2xf32> to vector<2x2xf32> 706 // CHECK: %[[CAST:.+]] = vector.bitcast %[[EXTRACT]] : vector<2x2xf32> to vector<2x4xf16> 707 %cast = vector.bitcast %arg0: vector<4x2xf32> to vector<4x4xf16> 708 %0 = vector.extract_strided_slice %cast {offsets = [1], sizes = [2], strides = [1]} : vector<4x4xf16> to vector<2x4xf16> 709 // CHECK: return %[[CAST]] 710 return %0: vector<2x4xf16> 711} 712 713// CHECK-LABEL: func @bubble_down_bitcast_in_strided_slice_extract_odd_offset 714func @bubble_down_bitcast_in_strided_slice_extract_odd_offset(%arg0: vector<4xf32>) -> vector<4xf16> { 715 // CHECK: vector.bitcast 716 // CHECK-NEXT: vector.extract_strided_slice 717 %cast = vector.bitcast %arg0: vector<4xf32> to vector<8xf16> 718 %0 = vector.extract_strided_slice %cast {offsets = [3], sizes = [4], strides = [1]} : vector<8xf16> to vector<4xf16> 719 return %0: vector<4xf16> 720} 721 722// CHECK-LABEL: func @bubble_down_bitcast_in_strided_slice_extract_odd_size 723func @bubble_down_bitcast_in_strided_slice_extract_odd_size(%arg0: vector<4xf32>) -> vector<3xf16> { 724 // CHECK: vector.bitcast 725 // CHECK-NEXT: vector.extract_strided_slice 726 %cast = vector.bitcast %arg0: vector<4xf32> to vector<8xf16> 727 %0 = vector.extract_strided_slice %cast {offsets = [0], sizes = [3], strides = [1]} : vector<8xf16> to vector<3xf16> 728 return %0: vector<3xf16> 729} 730 731// CHECK-LABEL: func @bubble_up_bitcast_in_strided_slice_insert 732// CHECK-SAME: (%[[DST:.+]]: vector<8xf16>, %[[SRC1:.+]]: vector<4xf16>, %[[SRC2:.+]]: vector<4xf16>) 733func @bubble_up_bitcast_in_strided_slice_insert(%dst: vector<8xf16>, %src1: vector<4xf16>, %src2: vector<4xf16>) -> vector<4xf32> { 734 // CHECK-DAG: %[[CAST_SRC1:.+]] = vector.bitcast %[[SRC1]] : vector<4xf16> to vector<2xf32> 735 // CHECK-DAG: %[[CAST_SRC2:.+]] = vector.bitcast %[[SRC2]] : vector<4xf16> to vector<2xf32> 736 // CHECK-DAG: %[[CAST_DST:.+]] = vector.bitcast %[[DST]] : vector<8xf16> to vector<4xf32> 737 // CHECK: %[[INSERT1:.+]] = vector.insert_strided_slice %[[CAST_SRC1]], %[[CAST_DST]] {offsets = [0], strides = [1]} : vector<2xf32> into vector<4xf32> 738 // CHECK: %[[INSERT2:.+]] = vector.insert_strided_slice %[[CAST_SRC2]], %[[INSERT1]] {offsets = [2], strides = [1]} : vector<2xf32> into vector<4xf32> 739 %0 = vector.insert_strided_slice %src1, %dst {offsets = [0], strides = [1]} : vector<4xf16> into vector<8xf16> 740 %1 = vector.insert_strided_slice %src2, %0 {offsets = [4], strides = [1]} : vector<4xf16> into vector<8xf16> 741 %cast = vector.bitcast %1: vector<8xf16> to vector<4xf32> 742 // CHECK: return %[[INSERT2]] 743 return %cast: vector<4xf32> 744} 745 746// CHECK-LABEL: func @bubble_up_bitcast_in_strided_slice_insert_odd_offset 747func @bubble_up_bitcast_in_strided_slice_insert_odd_offset(%dst: vector<8xf16>, %src: vector<4xf16>) -> vector<4xf32> { 748 // CHECK: vector.insert_strided_slice 749 // CHECK-NEXT: vector.bitcast 750 %0 = vector.insert_strided_slice %src, %dst {offsets = [3], strides = [1]} : vector<4xf16> into vector<8xf16> 751 %cast = vector.bitcast %0: vector<8xf16> to vector<4xf32> 752 return %cast: vector<4xf32> 753} 754 755// CHECK-LABEL: func @bubble_up_bitcast_in_strided_slice_insert_different_rank 756func @bubble_up_bitcast_in_strided_slice_insert_different_rank(%dst: vector<16x4x8xf16>, %src: vector<2x4xf16>) -> vector<16x4x4xf32> { 757 // CHECK: vector.insert_strided_slice 758 // CHECK-NEXT: vector.bitcast 759 %0 = vector.insert_strided_slice %src, %dst {offsets = [0, 0, 2], strides = [1, 1]} : vector<2x4xf16> into vector<16x4x8xf16> 760 %cast = vector.bitcast %0: vector<16x4x8xf16> to vector<16x4x4xf32> 761 return %cast: vector<16x4x4xf32> 762} 763