1// RUN: mlir-opt %s -test-linalg-transform-patterns=test-linalg-to-vector-patterns -split-input-file | FileCheck %s 2 3// ----- 4 5// CHECK-LABEL: contraction_dot 6func @contraction_dot(%A: memref<1584xf32>, %B: memref<1584xf32>, %C: memref<f32>) { 7 // CHECK: vector.contract 8 // CHECK-SAME: vector<1584xf32>, vector<1584xf32> into f32 9 linalg.dot ins(%A, %B: memref<1584xf32>, memref<1584xf32>) 10 outs(%C: memref<f32>) 11 return 12} 13 14// ----- 15 16// CHECK-LABEL: contraction_matvec 17func @contraction_matvec(%A: memref<1584x1584xf32>, %B: memref<1584xf32>, %C: memref<1584xf32>) { 18 // CHECK: vector.contract 19 // CHECK-SAME: vector<1584x1584xf32>, vector<1584xf32> into vector<1584xf32> 20 linalg.matvec ins(%A, %B: memref<1584x1584xf32>, memref<1584xf32>) 21 outs(%C: memref<1584xf32>) 22 return 23} 24 25// ----- 26 27// CHECK-LABEL: contraction_matmul 28func @contraction_matmul(%A: memref<1584x1584xf32>, %B: memref<1584x1584xf32>, %C: memref<1584x1584xf32>) { 29 // CHECK: vector.contract 30 // CHECK-SAME: vector<1584x1584xf32>, vector<1584x1584xf32> into vector<1584x1584xf32> 31 linalg.matmul ins(%A, %B: memref<1584x1584xf32>, memref<1584x1584xf32>) 32 outs(%C: memref<1584x1584xf32>) 33 return 34} 35 36// ----- 37 38// CHECK-LABEL: contraction_batch_matmul 39func @contraction_batch_matmul(%A: memref<1584x1584x1584xf32>, %B: memref<1584x1584x1584xf32>, %C: memref<1584x1584x1584xf32>) { 40 // CHECK: vector.contract 41 // CHECK-SAME: vector<1584x1584x1584xf32>, vector<1584x1584x1584xf32> into vector<1584x1584x1584xf32> 42 linalg.batch_matmul 43 ins(%A, %B: memref<1584x1584x1584xf32>, memref<1584x1584x1584xf32>) 44 outs(%C: memref<1584x1584x1584xf32>) 45 return 46} 47 48// ----- 49 50#matmul_trait = { 51 args_in = 2, 52 args_out = 1, 53 indexing_maps = [ 54 affine_map<(m, n, k) -> (m, k)>, 55 affine_map<(m, n, k) -> (k, n)>, 56 affine_map<(m, n, k) -> (m, n)> 57 ], 58 iterator_types = ["parallel", "parallel", "reduction"] 59} 60 61// CHECK-DAG: #[[$trans_2d:.*]] = affine_map<(d0, d1) -> (d1, d0)> 62// CHECK-DAG: #[[$mk:.*]] = affine_map<(d0, d1, d2) -> (d0, d2)> 63// CHECK-DAG: #[[$nk:.*]] = affine_map<(d0, d1, d2) -> (d1, d2)> 64// CHECK-DAG: #[[$mn:.*]] = affine_map<(d0, d1, d2) -> (d0, d1)> 65 66// CHECK-LABEL: func @vectorization_test 67func @vectorization_test(%A: memref<8x16xf32>, %B: memref<16x32xf32>, 68 %C: memref<8x32xf32>) { 69 // CHECK: vector.transfer_read %{{.*}} : memref<8x16xf32>, vector<8x16xf32> 70 // CHECK: vector.transfer_read %{{.*}} : memref<16x32xf32>, vector<32x16xf32> 71 // CHECK: vector.transfer_read %{{.*}} : memref<8x32xf32>, vector<8x32xf32> 72 // CHECK: vector.contract {indexing_maps = [#[[$mk]], #[[$nk]], #[[$mn]]] 73 // CHECK-SAME: vector<8x16xf32>, vector<32x16xf32> into vector<8x32xf32> 74 // CHECK: vector.transfer_write %{{.*}}, %{{.*}} : vector<8x32xf32>, memref<8x32xf32> 75 linalg.generic #matmul_trait 76 ins(%A, %B : memref<8x16xf32>, memref<16x32xf32>) 77 outs(%C : memref<8x32xf32>) { 78 ^bb(%a: f32, %b: f32, %c: f32) : 79 %d = mulf %a, %b: f32 80 %e = addf %c, %d: f32 81 linalg.yield %e : f32 82 } 83 return 84} 85 86// ----- 87 88#matmul_trait = { 89 args_in = 2, 90 args_out = 1, 91 indexing_maps = [ 92 affine_map<(m, n, k) -> (m, k)>, 93 affine_map<(m, n, k) -> (k, n)>, 94 affine_map<(m, n, k) -> (m, n)> 95 ], 96 iterator_types = ["parallel", "parallel", "reduction"] 97} 98 99// CHECK-DAG: #[[$trans_2d:.*]] = affine_map<(d0, d1) -> (d1, d0)> 100// CHECK-DAG: #[[$mk:.*]] = affine_map<(d0, d1, d2) -> (d0, d2)> 101// CHECK-DAG: #[[$nk:.*]] = affine_map<(d0, d1, d2) -> (d1, d2)> 102// CHECK-DAG: #[[$mn:.*]] = affine_map<(d0, d1, d2) -> (d0, d1)> 103 104// CHECK-LABEL: func @vectorization_test_integer 105func @vectorization_test_integer(%A: memref<8x16xi32>, %B: memref<16x32xi32>, 106 %C: memref<8x32xi32>) { 107 // CHECK: vector.transfer_read %{{.*}} : memref<8x16xi32>, vector<8x16xi32> 108 // CHECK: vector.transfer_read %{{.*}} : memref<16x32xi32>, vector<32x16xi32> 109 // CHECK: vector.transfer_read %{{.*}} : memref<8x32xi32>, vector<8x32xi32> 110 // CHECK: vector.contract {indexing_maps = [#[[$mk]], #[[$nk]], #[[$mn]]], 111 // CHECK-SAME: vector<8x16xi32>, vector<32x16xi32> into vector<8x32xi32> 112 // CHECK: vector.transfer_write %{{.*}}, %{{.*}} : vector<8x32xi32>, memref<8x32xi32> 113 linalg.generic #matmul_trait 114 ins(%A, %B : memref<8x16xi32>, memref<16x32xi32>) 115 outs(%C : memref<8x32xi32>) { 116 ^bb(%a: i32, %b: i32, %c: i32) : 117 %d = muli %a, %b: i32 118 %e = addi %c, %d: i32 119 linalg.yield %e : i32 120 } 121 return 122} 123 124// ----- 125 126// CHECK-LABEL: func @vectorization_test_2 127func @vectorization_test_2(%A: memref<8x16xf32>, %B: memref<16x32xf32>, 128 %C: memref<8x32xf32>) { 129 // CHECK: vector.contract {{.*}} : 130 // vector<8x16xf32>, vector<16x32xf32> into vector<8x32xf32> 131 linalg.matmul 132 ins(%A, %B: memref<8x16xf32>, memref<16x32xf32>) 133 outs(%C: memref<8x32xf32>) 134 return 135} 136 137// ----- 138 139// CHECK-LABEL: func @test_vectorize_fill 140func @test_vectorize_fill(%A : memref<8x16xf32>, %arg0 : f32) { 141 // CHECK: %[[V:.*]] = vector.broadcast {{.*}} : f32 to vector<8x16xf32> 142 // CHECK: vector.transfer_write %[[V]], {{.*}} : vector<8x16xf32>, memref<8x16xf32> 143 linalg.fill(%A, %arg0) : memref<8x16xf32>, f32 144 return 145} 146 147// ----- 148 149// CHECK-LABEL: func @test_vectorize_fill 150func @test_vectorize_fill_scalar(%A : memref<f32>, %arg0 : f32) { 151 // CHECK-SAME: (%[[M:.*]]: memref<f32>, %[[V:.*]]: f32) 152 // CHECK: store %[[V]], %[[M]][] : memref<f32> 153 linalg.fill(%A, %arg0) : memref<f32>, f32 154 return 155} 156 157// ----- 158 159// CHECK-LABEL: func @test_vectorize_copy 160func @test_vectorize_copy(%A : memref<8x16xf32>, %B : memref<8x16xf32>) { 161 // CHECK: %[[V:.*]] = vector.transfer_read {{.*}} : memref<8x16xf32>, vector<8x16xf32> 162 // CHECK: vector.transfer_write %[[V]], {{.*}} : vector<8x16xf32>, memref<8x16xf32> 163 linalg.copy(%A, %B) : memref<8x16xf32>, memref<8x16xf32> 164 return 165} 166 167// ----- 168 169// CHECK-LABEL: func @test_vectorize_copy_scalar 170func @test_vectorize_copy_scalar(%A : memref<f32>, %B : memref<f32>) { 171 // CHECK: %[[V:.*]] = memref.load {{.*}} : memref<f32> 172 // CHECK: store %[[V]], {{.*}} : memref<f32> 173 linalg.copy(%A, %B) : memref<f32>, memref<f32> 174 return 175} 176 177// ----- 178 179// CHECK-LABEL: func @test_vectorize_trailing_index 180 // CHECK-SAME: (%[[ARG0:.*]]: memref<1x2x4x8xindex>) 181func @test_vectorize_trailing_index(%arg0: memref<1x2x4x8xindex>) { 182 // CHECK-DAG: %[[CST0:.*]] = constant dense<[0, 1, 2, 3, 4, 5, 6, 7]> : vector<8xindex> 183 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 184 linalg.generic { 185 indexing_maps = [ 186 affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>], 187 iterator_types = ["parallel", "parallel", "parallel", "parallel"]} 188 outs(%arg0: memref<1x2x4x8xindex>) { 189 ^bb0(%arg1: index): 190 // CHECK: %[[BCST:.*]] = vector.broadcast %[[CST0]] : vector<8xindex> to vector<1x2x4x8xindex> 191 // CHECK: vector.transfer_write %[[BCST]], %[[ARG0]][%[[C0]], %[[C0]], %[[C0]], %[[C0]]] {{.*}} : vector<1x2x4x8xindex>, memref<1x2x4x8xindex> 192 %0 = linalg.index 3 : index 193 linalg.yield %0 : index 194 } 195 return 196} 197 198// ----- 199 200// CHECK-LABEL: func @test_vectorize_inner_index 201 // CHECK-SAME: (%[[ARG0:.*]]: memref<1x2x4x8xindex>) 202func @test_vectorize_inner_index(%arg0: memref<1x2x4x8xindex>) { 203 // CHECK-DAG: %[[CST0:.*]] = constant dense<[0, 1]> : vector<2xindex> 204 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 205 linalg.generic { 206 indexing_maps = [ 207 affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>], 208 iterator_types = ["parallel", "parallel", "parallel", "parallel"]} 209 outs(%arg0: memref<1x2x4x8xindex>) { 210 ^bb0(%arg1: index): 211 // CHECK: %[[BCST:.*]] = vector.broadcast %[[CST0]] : vector<2xindex> to vector<1x8x4x2xindex> 212 // CHECK: %[[TRAN:.*]] = vector.transpose %[[BCST]], [0, 3, 2, 1] : vector<1x8x4x2xindex> to vector<1x2x4x8xindex> 213 // CHECK: vector.transfer_write %[[TRAN]], %[[ARG0]][%[[C0]], %[[C0]], %[[C0]], %[[C0]]] {{.*}} : vector<1x2x4x8xindex>, memref<1x2x4x8xindex> 214 %0 = linalg.index 1 : index 215 linalg.yield %0 : index 216 } 217 return 218} 219 220// ----- 221 222// CHECK-LABEL: func @generic_vectorize 223 // CHECK-SAME: (%[[ARG0:.*]]: memref<4x256xf32>, %[[ARG1:.*]]: memref<4x256xf32>, 224 // CHECK-SAME: %[[ARG2:.*]]: memref<256xf32>, %[[ARG3:.*]]: f32) 225func @generic_vectorize(%arg0: memref<4x256xf32>, 226 %arg1: memref<4x256xf32>, 227 %arg2: memref<256xf32>, %i: f32) { 228 // CHECK-DAG: %[[CST0:.*]] = constant dense<2.000000e+00> : vector<4x256xf32> 229 // CHECK-DAG: %[[CST1:.*]] = constant dense<1.000000e+00> : vector<4x256xf32> 230 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 231 %c1_f32 = constant 1.0 : f32 232 linalg.generic { 233 args_in = 0 : i64, 234 args_out = 10 : i64, 235 indexing_maps = [ 236 affine_map<(d0, d1) -> (d0, d1)>, 237 affine_map<(d0, d1) -> (d1)>, 238 affine_map<(d0, d1) -> (d0, d1)>, 239 affine_map<(d0, d1) -> (d0, d1)>, 240 affine_map<(d0, d1) -> (d0, d1)>, 241 affine_map<(d0, d1) -> (d0, d1)>, 242 affine_map<(d0, d1) -> (d0, d1)>, 243 affine_map<(d0, d1) -> (d0, d1)>, 244 affine_map<(d0, d1) -> (d0, d1)>, 245 affine_map<(d0, d1) -> (d0, d1)>, 246 affine_map<(d0, d1) -> (d0, d1)>, 247 affine_map<(d0, d1) -> (d0, d1)>], 248 iterator_types = ["parallel", "parallel"]} 249 ins(%arg1, %arg2: memref<4x256xf32>, memref<256xf32>) 250 outs( 251 %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0 : 252 memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, 253 memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, 254 memref<4x256xf32>, memref<4x256xf32>) { 255 ^bb0(%arg3 : f32, %arg4 : f32, %arg5: f32, %arg6: f32, %arg7: f32, %arg8: f32, 256 // CHECK: %[[V2:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : memref<4x256xf32>, vector<4x256xf32> 257 // CHECK: %[[V0:.*]] = vector.transfer_read %[[ARG2]][%[[C0]]], {{.*}} : memref<256xf32>, vector<4x256xf32> 258 // CHECK: %[[V3:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x256xf32>, vector<4x256xf32> 259 // CHECK: %[[V1:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x256xf32>, vector<4x256xf32> 260 %arg9 : f32, %arg10 : f32, %arg11 : f32, %arg12 : f32, %arg13 : f32, 261 %arg14 : f32): 262 // CHECK: %[[ADD:.*]] = addf %[[V0]], %[[V1]] : vector<4x256xf32> 263 %6 = addf %arg4, %arg6 : f32 264 // CHECK: %[[CMP:.*]] = cmpf ogt, %[[V2]], %[[V1]] : vector<4x256xf32> 265 %7 = cmpf ogt, %arg3, %arg6 : f32 266 // CHECK: %[[ARG3B:.*]] = vector.broadcast %[[ARG3]] : f32 to vector<4x256xf32> 267 %8 = constant 2.0 : f32 268 // CHECK: %[[DIV:.*]] = divf %[[V3]], %[[ARG3B]] : vector<4x256xf32> 269 %9 = divf %arg5, %i : f32 270 // CHECK: %[[EXP:.*]] = math.exp2 %[[V3]] : vector<4x256xf32> 271 %10 = math.exp2 %arg5 : f32 272 // CHECK: %[[MUL:.*]] = mulf %[[V3]], %[[CST0]] : vector<4x256xf32> 273 %11 = mulf %arg5, %8 : f32 274 // CHECK: %[[RSQRT:.*]] = math.rsqrt %[[V3]] : vector<4x256xf32> 275 %12 = math.rsqrt %arg5 : f32 276 // CHECK: %[[SEL:.*]] = select %[[CMP]], %[[V3]], %[[V1]] : vector<4x256xi1>, vector<4x256xf32> 277 %13 = select %7, %arg5, %arg6 : f32 278 // CHECK: %[[SUB:.*]] = subf %[[V3]], %[[V0]] : vector<4x256xf32> 279 %14 = subf %arg5, %arg4 : f32 280 // CHECK: %[[TAN:.*]] = math.tanh %[[V3]] : vector<4x256xf32> 281 %15 = math.tanh %arg5 : f32 282 // CHECK: vector.transfer_write %[[ADD]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 283 // CHECK: vector.transfer_write %[[CST0]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 284 // CHECK: vector.transfer_write %[[CST1]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 285 // CHECK: vector.transfer_write %[[DIV]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 286 // CHECK: vector.transfer_write %[[EXP]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 287 // CHECK: vector.transfer_write %[[MUL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 288 // CHECK: vector.transfer_write %[[RSQRT]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 289 // CHECK: vector.transfer_write %[[SEL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 290 // CHECK: vector.transfer_write %[[SUB]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 291 // CHECK: vector.transfer_write %[[TAN]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 292 linalg.yield %6, %8, %c1_f32, %9, %10, %11, %12, %13, %14, %15 : f32, f32, 293 f32, f32, f32, f32, f32, f32, f32, f32 294 } 295 return 296} 297 298// ----- 299 300// CHECK-LABEL: func @generic_vectorize_tensor 301// CHECK-SAME: (%[[ARG0:.*]]: tensor<4x256xf32>, %[[ARG1:.*]]: tensor<4x256xf32>, 302// CHECK-SAME: %[[ARG2:.*]]: tensor<256xf32>, %[[ARG3:.*]]: f32) 303func @generic_vectorize_tensor(%arg0: tensor<4x256xf32>, 304 %arg1: tensor<4x256xf32>, %arg2: tensor<256xf32>, 305 %i: f32) -> (tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 306 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 307 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>) { 308 %c1_f32 = constant 1.0 : f32 309 %r:10 = linalg.generic { 310 indexing_maps = [ 311 affine_map<(d0, d1) -> (d0, d1)>, 312 affine_map<(d0, d1) -> (d1)>, 313 affine_map<(d0, d1) -> (d0, d1)>, 314 affine_map<(d0, d1) -> (d0, d1)>, 315 affine_map<(d0, d1) -> (d0, d1)>, 316 affine_map<(d0, d1) -> (d0, d1)>, 317 affine_map<(d0, d1) -> (d0, d1)>, 318 affine_map<(d0, d1) -> (d0, d1)>, 319 affine_map<(d0, d1) -> (d0, d1)>, 320 affine_map<(d0, d1) -> (d0, d1)>, 321 affine_map<(d0, d1) -> (d0, d1)>, 322 affine_map<(d0, d1) -> (d0, d1)>], 323 iterator_types = ["parallel", "parallel"]} 324 ins(%arg1, %arg2: tensor<4x256xf32>, tensor<256xf32>) 325 outs( 326 %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0 : 327 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 328 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 329 tensor<4x256xf32>, tensor<4x256xf32>) { 330 ^bb0(%arg3 : f32, %arg4 : f32, %arg5: f32, %arg6: f32, %arg7: f32, %arg8: f32, 331 %arg9 : f32, %arg10 : f32, %arg11 : f32, %arg12 : f32, %arg13 : f32, 332 %arg14 : f32): 333 // CHECK-DAG: %[[CST0:.*]] = constant dense<2.000000e+00> : vector<4x256xf32> 334 // CHECK-DAG: %[[CST1:.*]] = constant dense<1.000000e+00> : vector<4x256xf32> 335 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 336 // CHECK: %[[V2:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x256xf32>, vector<4x256xf32> 337 // CHECK: %[[V0:.*]] = vector.transfer_read %[[ARG2]][%[[C0]]], {{.*}} : tensor<256xf32>, vector<4x256xf32> 338 // CHECK: %[[V3:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x256xf32>, vector<4x256xf32> 339 // CHECK: %[[V1:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x256xf32>, vector<4x256xf32> 340 // CHECK: %[[ADD:.*]] = addf %[[V0]], %[[V1]] : vector<4x256xf32> 341 %6 = addf %arg4, %arg6 : f32 342 // CHECK: %[[CMP:.*]] = cmpf ogt, %[[V2]], %[[V1]] : vector<4x256xf32> 343 %7 = cmpf ogt, %arg3, %arg6 : f32 344 // CHECK: %[[ARG3B:.*]] = vector.broadcast %[[ARG3]] : f32 to vector<4x256xf32> 345 %8 = constant 2.0 : f32 346 // CHECK: %[[DIV:.*]] = divf %[[V3]], %[[ARG3B]] : vector<4x256xf32> 347 %9 = divf %arg5, %i : f32 348 // CHECK: %[[EXP:.*]] = math.exp2 %[[V3]] : vector<4x256xf32> 349 %10 = math.exp2 %arg5 : f32 350 // CHECK: %[[MUL:.*]] = mulf %[[V3]], %[[CST0]] : vector<4x256xf32> 351 %11 = mulf %arg5, %8 : f32 352 // CHECK: %[[RSQRT:.*]] = math.rsqrt %[[V3]] : vector<4x256xf32> 353 %12 = math.rsqrt %arg5 : f32 354 // CHECK: %[[SEL:.*]] = select %[[CMP]], %[[V3]], %[[V1]] : vector<4x256xi1>, vector<4x256xf32> 355 %13 = select %7, %arg5, %arg6 : f32 356 // CHECK: %[[SUB:.*]] = subf %[[V3]], %[[V0]] : vector<4x256xf32> 357 %14 = subf %arg5, %arg4 : f32 358 // CHECK: %[[TAN:.*]] = math.tanh %[[V3]] : vector<4x256xf32> 359 %15 = math.tanh %arg5 : f32 360 // CHECK: %[[R0:.*]] = vector.transfer_write %[[ADD]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 361 // CHECK: %[[R1:.*]] = vector.transfer_write %[[CST0]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 362 // CHECK: %[[R2:.*]] = vector.transfer_write %[[CST1]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 363 // CHECK: %[[R3:.*]] = vector.transfer_write %[[DIV]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 364 // CHECK: %[[R4:.*]] = vector.transfer_write %[[EXP]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 365 // CHECK: %[[R5:.*]] = vector.transfer_write %[[MUL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 366 // CHECK: %[[R6:.*]] = vector.transfer_write %[[RSQRT]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 367 // CHECK: %[[R7:.*]] = vector.transfer_write %[[SEL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 368 // CHECK: %[[R8:.*]] = vector.transfer_write %[[SUB]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 369 // CHECK: %[[R9:.*]] = vector.transfer_write %[[TAN]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 370 linalg.yield %6, %8, %c1_f32, %9, %10, %11, %12, %13, %14, %15 : f32, f32, 371 f32, f32, f32, f32, f32, f32, f32, f32 372 } -> tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 373 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 374 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32> 375 // CHECK: return %[[R0]], %[[R1]], %[[R2]], %[[R3]], %[[R4]], %[[R5]], %[[R6]], %[[R7]], %[[R8]], %[[R9]] : tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32> 376 return %r#0, %r#1, %r#2, %r#3, %r#4, %r#5, %r#6, %r#7, %r#8, %r#9: 377 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 378 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 379 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32> 380} 381 382// ----- 383 384// Test different input maps. 385#matmul_trait = { 386 indexing_maps = [ 387 affine_map<(d0, d1, d2, d3) -> (d1, d0)>, 388 affine_map<(d0, d1, d2, d3) -> (d3, d1)>, 389 affine_map<(d0, d1, d2, d3) -> (d3, d1, d0, d2)>, 390 affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)> 391 ], 392 iterator_types = ["parallel", "parallel", "parallel", "parallel"] 393} 394 395// CHECK-DAG: #[[MAP0:.*]] = affine_map<(d0, d1) -> (d1, d0, 0, 0)> 396// CHECK-DAG: #[[MAP1:.*]] = affine_map<(d0, d1) -> (0, d1, 0, d0)> 397// CHECK-DAG: #[[MAP2:.*]] = affine_map<(d0, d1, d2, d3) -> (d2, d1, d3, d0)> 398// CHECK: func @vectorization_transpose 399// CHECK: vector.transfer_read {{.*}}{permutation_map = #[[MAP0]]} : memref<14x7xf32>, vector<7x14x8x16xf32> 400// CHECK: vector.transfer_read {{.*}}{permutation_map = #[[MAP1]]} : memref<16x14xf32>, vector<7x14x8x16xf32> 401// CHECK: vector.transfer_read {{.*}}{permutation_map = #[[MAP2]]} : memref<16x14x7x8xf32>, vector<7x14x8x16xf32> 402// CHECK: addf {{.*}} : vector<7x14x8x16xf32> 403// CHECK: addf {{.*}} : vector<7x14x8x16xf32> 404// CHECK: vector.transfer_write {{.*}} : vector<7x14x8x16xf32>, memref<7x14x8x16xf32> 405func @vectorization_transpose(%A: memref<14x7xf32>, %B: memref<16x14xf32>, 406 %C: memref<16x14x7x8xf32>, %D: memref<7x14x8x16xf32>) { 407 linalg.generic #matmul_trait 408 ins(%A, %B, %C : memref<14x7xf32>, memref<16x14xf32>, memref<16x14x7x8xf32>) 409 outs(%D : memref<7x14x8x16xf32>) { 410 ^bb(%a: f32, %b: f32, %c: f32, %d: f32) : 411 %e = addf %a, %b: f32 412 %f = addf %e, %c: f32 413 linalg.yield %f : f32 414 } 415 return 416} 417 418// ----- 419 420// CHECK-LABEL: func @matmul_tensors 421// CHECK-SAME: (%[[ARG0:.*]]: tensor<8x4xf32>, %[[ARG1:.*]]: tensor<4x12xf32>, 422// CHECK-SAME: %[[ARG2:.*]]: tensor<8x12xf32>) -> tensor<8x12xf32> 423func @matmul_tensors( 424 %arg0: tensor<8x4xf32>, %arg1: tensor<4x12xf32>, %arg2: tensor<8x12xf32>) 425 -> tensor<8x12xf32> { 426 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 427 // CHECK-DAG: %[[VEC_C0:.*]] = constant dense<0.000000e+00> : vector<8x12xf32> 428 // CHECK-DAG: %[[V0:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : tensor<8x4xf32>, vector<8x4xf32> 429 // CHECK-DAG: %[[V1:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x12xf32>, vector<12x4xf32> 430 // CHECK-DAG: %[[V2:.*]] = vector.transfer_read %[[ARG2]][%[[C0]], %[[C0]]], {{.*}} : tensor<8x12xf32>, vector<8x12xf32> 431 // 432 // linalg contraction lowers to %tmp = vector.contract %a, %b, %c0 followed by addf %c, %tmp. 433 // a later canonicalization fuses the add into vector.contract. 434 // CHECK: %[[C:.*]] = vector.contract 435 // CHECK-SAME: iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} 436 // CHECK-SAME: %[[V0]], %[[V1]], %[[VEC_C0]] : 437 // CHECK-SAME: vector<8x4xf32>, vector<12x4xf32> into vector<8x12xf32> 438 // CHECK: %[[C2:.*]] = addf %[[V2]], %[[C]] : vector<8x12xf32> 439 // CHECK: %[[W:.*]] = vector.transfer_write %[[C2]], %[[ARG2]][%[[C0]], %[[C0]]] {in_bounds = [true, true]} : vector<8x12xf32>, tensor<8x12xf32> 440 %0 = linalg.matmul ins(%arg0, %arg1: tensor<8x4xf32>, tensor<4x12xf32>) 441 outs(%arg2: tensor<8x12xf32>) 442 -> tensor<8x12xf32> 443 // CHECK: return %[[W]] : tensor<8x12xf32> 444 return %0 : tensor<8x12xf32> 445} 446 447// ----- 448 449// CHECK-LABEL: func @matmul_i8_i8_i32 450// CHECK-SAME: %[[ARG0:[a-z0-9]+]]: memref<4x6xi8> 451// CHECK-SAME: %[[ARG1:[a-z0-9]+]]: memref<6x12xi8> 452// CHECK-SAME: %[[ARG2:[a-z0-9]+]]: memref<4x12xi32> 453func @matmul_i8_i8_i32(%a: memref<4x6xi8>, %b: memref<6x12xi8>, %c: memref<4x12xi32>) { 454 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 455 // CHECK-DAG: %[[VEC_C0:.*]] = constant dense<0> : vector<4x12xi32> 456 // CHECK-DAG: %[[V0:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x6xi8>, vector<4x6xi8> 457 // CHECK-DAG: %[[V1:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : memref<6x12xi8>, vector<12x6xi8> 458 // CHECK-DAG: %[[V2:.*]] = vector.transfer_read %[[ARG2]][%[[C0]], %[[C0]]], {{.*}} : memref<4x12xi32>, vector<4x12xi32> 459 // CHECK-DAG: %[[V0_32:.*]] = sexti %[[V0]] : vector<4x6xi8> to vector<4x6xi32> 460 // CHECK-DAG: %[[V1_32:.*]] = sexti %[[V1]] : vector<12x6xi8> to vector<12x6xi32> 461 // 462 // linalg contraction lowers to %tmp = vector.contract %a, %b, %c0 followed by addf %c, %tmp. 463 // a later canonicalization fuses the add into vector.contract. 464 // CHECK: %[[C:.*]] = vector.contract 465 // CHECK-SAME: iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} 466 // CHECK-SAME: %[[V0_32]], %[[V1_32]], %[[VEC_C0]] 467 // CHECK-SAME: vector<4x6xi32>, vector<12x6xi32> into vector<4x12xi32> 468 // CHECK: %[[RES:.*]] = addi %[[V2]], %[[C]] : vector<4x12xi32> 469 // CHECK: vector.transfer_write %[[RES]], %[[ARG2]][%[[C0]], %[[C0]]] {in_bounds = [true, true]} 470 // CHECK-SAME: vector<4x12xi32>, memref<4x12xi32> 471 linalg.matmul_i8_i8_i32 ins(%a, %b : memref<4x6xi8>, memref<6x12xi8>) 472 outs(%c: memref<4x12xi32>) 473 return 474} 475 476// ----- 477 478// CHECK-LABEL: func @pad_static 479// CHECK-NOT: linalg.pad_tensor 480func @pad_static(%arg0: tensor<?x?x?xf32>, %pad_value: f32) -> tensor<2x3x4xf32> { 481 // CHECK: %[[C0:.*]] = constant 0 : index 482 // CHECK: %[[READ:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]], %[[C0]]] 483 // CHECK-SAME: : tensor<?x?x?xf32>, vector<2x3x4xf32> 484 // CHECK: %[[INIT:.*]] = linalg.init_tensor [2, 3, 4] : tensor<2x3x4xf32> 485 // CHECK: %[[WRITTEN:.*]] = vector.transfer_write %[[READ]], %[[INIT]][%[[C0]], %[[C0]], %[[C0]]] 486 // CHECK-SAME: {in_bounds = [true, true, true]} : vector<2x3x4xf32>, tensor<2x3x4xf32> 487 %c0 = constant 0 : index 488 %0 = linalg.pad_tensor %arg0 low[0, %c0, 0] high[0, 0, %c0] { 489 ^bb0(%arg1: index, %arg2: index, %arg3: index): 490 linalg.yield %pad_value : f32 491 } : tensor<?x?x?xf32> to tensor<2x3x4xf32> 492 493 // CHECK: return %[[WRITTEN]] : tensor<2x3x4xf32> 494 return %0 : tensor<2x3x4xf32> 495} 496 497// ----- 498 499// CHECK-LABEL: func @pad_static_high_padding 500// CHECK: linalg.pad_tensor 501func @pad_static_high_padding(%arg0: tensor<?x?x?xf32>, %pad_value: f32) -> tensor<2x3x4xf32> { 502 %0 = linalg.pad_tensor %arg0 low[0, 0, 0] high[0, 1, 0] { 503 ^bb0(%arg1: index, %arg2: index, %arg3: index): 504 linalg.yield %pad_value : f32 505 } : tensor<?x?x?xf32> to tensor<2x3x4xf32> 506 return %0 : tensor<2x3x4xf32> 507} 508 509// ----- 510 511// CHECK-LABEL: func @pad_dynamic 512// CHECK: linalg.pad_tensor 513func @pad_dynamic(%arg0: tensor<1x2x2x?xf32>, %low: index, %high: index, 514 %pad_value: f32) -> tensor<6x?x?x?xf32> { 515 %0 = linalg.pad_tensor %arg0 low[2, %low, 3, 3] high[3, 3, %high, 2] { 516 ^bb0(%arg1: index, %arg2: index, %arg3: index, %arg4: index): 517 linalg.yield %pad_value : f32 518 } : tensor<1x2x2x?xf32> to tensor<6x?x?x?xf32> 519 return %0 : tensor<6x?x?x?xf32> 520} 521 522// ----- 523 524// CHECK-DAG: #[[$M0:.*]] = affine_map<(d0, d1) -> (d0, d1, 0)> 525 526// CHECK-LABEL: func @sum_exp 527func @sum_exp(%input: tensor<4x16x8xf32>, %output: tensor<4x16xf32>) 528 -> tensor<4x16xf32> 529{ 530 // CHECK: vector.transfer_read {{.*}} : tensor<4x16x8xf32>, vector<4x16x8xf32> 531 // CHECK: vector.transfer_read {{.*}} {permutation_map = #[[$M0]]} : tensor<4x16xf32>, vector<4x16x8xf32> 532 // CHECK: math.exp {{.*}} : vector<4x16x8xf32> 533 // CHECK: addf {{.*}} : vector<4x16x8xf32> 534 // CHECK: vector.multi_reduction #vector.kind<add>, %{{.*}} [2] : vector<4x16x8xf32> to vector<4x16xf32> 535 // CHECK: vector.transfer_write {{.*}} : vector<4x16xf32>, tensor<4x16xf32> 536 // CHECK: return {{.*}} : tensor<4x16xf32> 537 %0 = linalg.generic { 538 indexing_maps = [ 539 affine_map<(d0, d1, d2) -> (d0, d1, d2)>, 540 affine_map<(d0, d1, d2) -> (d0, d1)> 541 ], 542 iterator_types = ["parallel", "parallel", "reduction"] 543 } ins(%input : tensor<4x16x8xf32>) outs(%output : tensor<4x16xf32>) { 544 ^bb0(%arg0: f32, %arg1: f32): // no predecessors 545 %1 = math.exp %arg0 : f32 546 %2 = addf %1, %arg1 : f32 547 linalg.yield %2 : f32 548 } -> tensor<4x16xf32> 549 return %0 : tensor<4x16xf32> 550} 551 552// ----- 553 554// CHECK-DAG: #[[$M1:.*]] = affine_map<(d0, d1) -> (d1, d0, 0, 0)> 555// CHECK-DAG: #[[$M2:.*]] = affine_map<(d0, d1) -> (0, 0, d1, d0)> 556// CHECK-DAG: #[[$M3:.*]] = affine_map<(d0, d1) -> (d1, 0, 0, d0)> 557// CHECK-DAG: #[[$M4:.*]] = affine_map<(d0, d1) -> (d1, d0)> 558 559// CHECK-LABEL: func @sum_exp_2 560func @sum_exp_2(%input: tensor<3x2xf32>, %input_2: tensor<5x4xf32>, %output: tensor<5x2xf32>) 561 -> tensor<5x2xf32> 562{ 563 // CHECK: vector.transfer_read {{.*}} {permutation_map = #[[$M1]]} : tensor<3x2xf32>, vector<2x3x4x5xf32> 564 // CHECK: vector.transfer_read {{.*}} {permutation_map = #[[$M2]]} : tensor<5x4xf32>, vector<2x3x4x5xf32> 565 // CHECK: vector.transfer_read {{.*}} {permutation_map = #[[$M3]]} : tensor<5x2xf32>, vector<2x3x4x5xf32> 566 // CHECK: math.exp {{.*}} : vector<2x3x4x5xf32> 567 // CHECK: math.exp {{.*}} : vector<2x3x4x5xf32> 568 // CHECK: addf {{.*}} : vector<2x3x4x5xf32> 569 // CHECK: addf {{.*}} : vector<2x3x4x5xf32> 570 // CHECK: vector.multi_reduction #vector.kind<add>, {{.*}} [1, 2] : vector<2x3x4x5xf32> to vector<2x5xf32> 571 // CHECK: vector.transfer_write {{.*}} {permutation_map = #[[$M4]]} : vector<2x5xf32>, tensor<5x2xf32> 572 // CHECK: return {{.*}} : tensor<5x2xf32> 573 %0 = linalg.generic { 574 indexing_maps = [ 575 affine_map<(d0, d1, d2, d3) -> (d1, d0)>, 576 affine_map<(d0, d1, d2, d3) -> (d3, d2)>, 577 affine_map<(d0, d1, d2, d3) -> (d3, d0)> 578 ], 579 iterator_types = ["parallel", "reduction", "reduction", "parallel"] 580 } ins(%input, %input_2 : tensor<3x2xf32>, tensor<5x4xf32>) outs(%output : tensor<5x2xf32>) { 581 ^bb0(%arg0: f32, %arg1: f32, %arg2: f32): // no predecessors 582 %1 = math.exp %arg0 : f32 583 %2 = math.exp %arg1 : f32 584 %3 = addf %1, %2 : f32 585 %4 = addf %3, %arg2 : f32 586 linalg.yield %4 : f32 587 } -> tensor<5x2xf32> 588 return %0 : tensor<5x2xf32> 589} 590