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: #[[$mk:.*]] = affine_map<(d0, d1, d2) -> (d0, d2)> 62// CHECK-DAG: #[[$kn:.*]] = affine_map<(d0, d1, d2) -> (d2, d1)> 63// CHECK-DAG: #[[$mn:.*]] = affine_map<(d0, d1, d2) -> (d0, d1)> 64 65// CHECK-LABEL: func @vectorization_test 66func @vectorization_test(%A: memref<8x16xf32>, %B: memref<16x32xf32>, 67 %C: memref<8x32xf32>) { 68 // CHECK: vector.transfer_read %{{.*}} : memref<8x16xf32>, vector<8x16xf32> 69 // CHECK: vector.transfer_read %{{.*}} : memref<16x32xf32>, vector<16x32xf32> 70 // CHECK: vector.transfer_read %{{.*}} : memref<8x32xf32>, vector<8x32xf32> 71 // CHECK: vector.contract {indexing_maps = [#[[$mk]], #[[$kn]], #[[$mn]]] 72 // CHECK-SAME: vector<8x16xf32>, vector<16x32xf32> into vector<8x32xf32> 73 // CHECK: vector.transfer_write %{{.*}}, %{{.*}} : vector<8x32xf32>, memref<8x32xf32> 74 linalg.generic #matmul_trait 75 ins(%A, %B : memref<8x16xf32>, memref<16x32xf32>) 76 outs(%C : memref<8x32xf32>) { 77 ^bb(%a: f32, %b: f32, %c: f32) : 78 %d = mulf %a, %b: f32 79 %e = addf %c, %d: f32 80 linalg.yield %e : f32 81 } 82 return 83} 84 85// ----- 86 87#matmul_trait = { 88 args_in = 2, 89 args_out = 1, 90 indexing_maps = [ 91 affine_map<(m, n, k) -> (m, k)>, 92 affine_map<(m, n, k) -> (k, n)>, 93 affine_map<(m, n, k) -> (m, n)> 94 ], 95 iterator_types = ["parallel", "parallel", "reduction"] 96} 97 98// CHECK-DAG: #[[$mk:.*]] = affine_map<(d0, d1, d2) -> (d0, d2)> 99// CHECK-DAG: #[[$kn:.*]] = affine_map<(d0, d1, d2) -> (d2, d1)> 100// CHECK-DAG: #[[$mn:.*]] = affine_map<(d0, d1, d2) -> (d0, d1)> 101 102// CHECK-LABEL: func @vectorization_test_integer 103func @vectorization_test_integer(%A: memref<8x16xi32>, %B: memref<16x32xi32>, 104 %C: memref<8x32xi32>) { 105 // CHECK: vector.transfer_read %{{.*}} : memref<8x16xi32>, vector<8x16xi32> 106 // CHECK: vector.transfer_read %{{.*}} : memref<16x32xi32>, vector<16x32xi32> 107 // CHECK: vector.transfer_read %{{.*}} : memref<8x32xi32>, vector<8x32xi32> 108 // CHECK: vector.contract {indexing_maps = [#[[$mk]], #[[$kn]], #[[$mn]]], 109 // CHECK-SAME: vector<8x16xi32>, vector<16x32xi32> into vector<8x32xi32> 110 // CHECK: vector.transfer_write %{{.*}}, %{{.*}} : vector<8x32xi32>, memref<8x32xi32> 111 linalg.generic #matmul_trait 112 ins(%A, %B : memref<8x16xi32>, memref<16x32xi32>) 113 outs(%C : memref<8x32xi32>) { 114 ^bb(%a: i32, %b: i32, %c: i32) : 115 %d = muli %a, %b: i32 116 %e = addi %c, %d: i32 117 linalg.yield %e : i32 118 } 119 return 120} 121 122// ----- 123 124// CHECK-LABEL: func @vectorization_test_2 125func @vectorization_test_2(%A: memref<8x16xf32>, %B: memref<16x32xf32>, 126 %C: memref<8x32xf32>) { 127 // CHECK: vector.contract {{.*}} : 128 // vector<8x16xf32>, vector<16x32xf32> into vector<8x32xf32> 129 linalg.matmul 130 ins(%A, %B: memref<8x16xf32>, memref<16x32xf32>) 131 outs(%C: memref<8x32xf32>) 132 return 133} 134 135// ----- 136 137// CHECK-LABEL: func @test_vectorize_fill 138func @test_vectorize_fill(%A : memref<8x16xf32>, %arg0 : f32) { 139 // CHECK: %[[V:.*]] = vector.broadcast {{.*}} : f32 to vector<8x16xf32> 140 // CHECK: vector.transfer_write %[[V]], {{.*}} : vector<8x16xf32>, memref<8x16xf32> 141 linalg.fill(%A, %arg0) : memref<8x16xf32>, f32 142 return 143} 144 145// ----- 146 147// CHECK-LABEL: func @test_vectorize_fill 148func @test_vectorize_fill_scalar(%A : memref<f32>, %arg0 : f32) { 149 // CHECK-SAME: (%[[M:.*]]: memref<f32>, %[[V:.*]]: f32) 150 // CHECK: store %[[V]], %[[M]][] : memref<f32> 151 linalg.fill(%A, %arg0) : memref<f32>, f32 152 return 153} 154 155// ----- 156 157// CHECK-LABEL: func @test_vectorize_copy 158func @test_vectorize_copy(%A : memref<8x16xf32>, %B : memref<8x16xf32>) { 159 // CHECK: %[[V:.*]] = vector.transfer_read {{.*}} : memref<8x16xf32>, vector<8x16xf32> 160 // CHECK: vector.transfer_write %[[V]], {{.*}} : vector<8x16xf32>, memref<8x16xf32> 161 linalg.copy(%A, %B) : memref<8x16xf32>, memref<8x16xf32> 162 return 163} 164 165// ----- 166 167// CHECK-LABEL: func @test_vectorize_copy_scalar 168func @test_vectorize_copy_scalar(%A : memref<f32>, %B : memref<f32>) { 169 // CHECK: %[[V:.*]] = load {{.*}} : memref<f32> 170 // CHECK: store %[[V]], {{.*}} : memref<f32> 171 linalg.copy(%A, %B) : memref<f32>, memref<f32> 172 return 173} 174 175// ----- 176 177// CHECK-LABEL: func @generic_vectorize 178 // CHECK-SAME: (%[[ARG0:.*]]: memref<4x256xf32>, %[[ARG1:.*]]: memref<4x256xf32>, 179 // CHECK-SAME: %[[ARG2:.*]]: memref<256xf32>, %[[ARG3:.*]]: f32) 180func @generic_vectorize(%arg0: memref<4x256xf32>, 181 %arg1: memref<4x256xf32>, 182 %arg2: memref<256xf32>, %i: f32) { 183 // CHECK-DAG: %[[CST0:.*]] = constant dense<2.000000e+00> : vector<4x256xf32> 184 // CHECK-DAG: %[[CST1:.*]] = constant dense<1.000000e+00> : vector<4x256xf32> 185 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 186 %c1_f32 = constant 1.0 : f32 187 linalg.generic { 188 args_in = 0 : i64, 189 args_out = 10 : i64, 190 indexing_maps = [ 191 affine_map<(d0, d1) -> (d0, d1)>, 192 affine_map<(d0, d1) -> (d1)>, 193 affine_map<(d0, d1) -> (d0, d1)>, 194 affine_map<(d0, d1) -> (d0, d1)>, 195 affine_map<(d0, d1) -> (d0, d1)>, 196 affine_map<(d0, d1) -> (d0, d1)>, 197 affine_map<(d0, d1) -> (d0, d1)>, 198 affine_map<(d0, d1) -> (d0, d1)>, 199 affine_map<(d0, d1) -> (d0, d1)>, 200 affine_map<(d0, d1) -> (d0, d1)>, 201 affine_map<(d0, d1) -> (d0, d1)>, 202 affine_map<(d0, d1) -> (d0, d1)>], 203 iterator_types = ["parallel", "parallel"]} 204 ins(%arg1, %arg2: memref<4x256xf32>, memref<256xf32>) 205 outs( 206 %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0 : 207 memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, 208 memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, memref<4x256xf32>, 209 memref<4x256xf32>, memref<4x256xf32>) { 210 ^bb0(%arg3 : f32, %arg4 : f32, %arg5: f32, %arg6: f32, %arg7: f32, %arg8: f32, 211 // CHECK: %[[V2:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : memref<4x256xf32>, vector<4x256xf32> 212 // CHECK: %[[V0:.*]] = vector.transfer_read %[[ARG2]][%[[C0]]], {{.*}} : memref<256xf32>, vector<256xf32> 213 // CHECK: %[[V3:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x256xf32>, vector<4x256xf32> 214 // CHECK: %[[V1:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x256xf32>, vector<4x256xf32> 215 %arg9 : f32, %arg10 : f32, %arg11 : f32, %arg12 : f32, %arg13 : f32, 216 %arg14 : f32): 217 // CHECK: %[[V0B:.*]] = vector.broadcast %[[V0]] : vector<256xf32> to vector<4x256xf32> 218 // CHECK: %[[ADD:.*]] = addf %[[V0B]], %[[V1]] : vector<4x256xf32> 219 %6 = addf %arg4, %arg6 : f32 220 // CHECK: %[[CMP:.*]] = cmpf ogt, %[[V2]], %[[V1]] : vector<4x256xf32> 221 %7 = cmpf ogt, %arg3, %arg6 : f32 222 // CHECK: %[[ARG3B:.*]] = vector.broadcast %[[ARG3]] : f32 to vector<4x256xf32> 223 %8 = constant 2.0 : f32 224 // CHECK: %[[DIV:.*]] = divf %[[V3]], %[[ARG3B]] : vector<4x256xf32> 225 %9 = divf %arg5, %i : f32 226 // CHECK: %[[EXP:.*]] = math.exp2 %[[V3]] : vector<4x256xf32> 227 %10 = math.exp2 %arg5 : f32 228 // CHECK: %[[MUL:.*]] = mulf %[[V3]], %[[CST0]] : vector<4x256xf32> 229 %11 = mulf %arg5, %8 : f32 230 // CHECK: %[[RSQRT:.*]] = math.rsqrt %[[V3]] : vector<4x256xf32> 231 %12 = math.rsqrt %arg5 : f32 232 // CHECK: %[[SEL:.*]] = select %[[CMP]], %[[V3]], %[[V1]] : vector<4x256xi1>, vector<4x256xf32> 233 %13 = select %7, %arg5, %arg6 : f32 234 // CHECK: %[[V0B:.*]] = vector.broadcast %[[V0]] : vector<256xf32> to vector<4x256xf32> 235 // CHECK: %[[SUB:.*]] = subf %[[V3]], %[[V0B]] : vector<4x256xf32> 236 %14 = subf %arg5, %arg4 : f32 237 // CHECK: %[[TAN:.*]] = math.tanh %[[V3]] : vector<4x256xf32> 238 %15 = math.tanh %arg5 : f32 239 // CHECK: vector.transfer_write %[[ADD]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 240 // CHECK: vector.transfer_write %[[CST0]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 241 // CHECK: vector.transfer_write %[[CST1]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 242 // CHECK: vector.transfer_write %[[DIV]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 243 // CHECK: vector.transfer_write %[[EXP]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 244 // CHECK: vector.transfer_write %[[MUL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 245 // CHECK: vector.transfer_write %[[RSQRT]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 246 // CHECK: vector.transfer_write %[[SEL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 247 // CHECK: vector.transfer_write %[[SUB]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 248 // CHECK: vector.transfer_write %[[TAN]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, memref<4x256xf32> 249 linalg.yield %6, %8, %c1_f32, %9, %10, %11, %12, %13, %14, %15 : f32, f32, 250 f32, f32, f32, f32, f32, f32, f32, f32 251 } 252 return 253} 254 255 256// ----- 257 258// CHECK-LABEL: func @generic_vectorize_tensor 259// CHECK-SAME: (%[[ARG0:.*]]: tensor<4x256xf32>, %[[ARG1:.*]]: tensor<4x256xf32>, 260// CHECK-SAME: %[[ARG2:.*]]: tensor<256xf32>, %[[ARG3:.*]]: f32) 261func @generic_vectorize_tensor(%arg0: tensor<4x256xf32>, 262 %arg1: tensor<4x256xf32>, %arg2: tensor<256xf32>, 263 %i: f32) -> (tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 264 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 265 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>) { 266 %c1_f32 = constant 1.0 : f32 267 %r:10 = linalg.generic { 268 indexing_maps = [ 269 affine_map<(d0, d1) -> (d0, d1)>, 270 affine_map<(d0, d1) -> (d1)>, 271 affine_map<(d0, d1) -> (d0, d1)>, 272 affine_map<(d0, d1) -> (d0, d1)>, 273 affine_map<(d0, d1) -> (d0, d1)>, 274 affine_map<(d0, d1) -> (d0, d1)>, 275 affine_map<(d0, d1) -> (d0, d1)>, 276 affine_map<(d0, d1) -> (d0, d1)>, 277 affine_map<(d0, d1) -> (d0, d1)>, 278 affine_map<(d0, d1) -> (d0, d1)>, 279 affine_map<(d0, d1) -> (d0, d1)>, 280 affine_map<(d0, d1) -> (d0, d1)>], 281 iterator_types = ["parallel", "parallel"]} 282 ins(%arg1, %arg2: tensor<4x256xf32>, tensor<256xf32>) 283 outs( 284 %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0, %arg0 : 285 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 286 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 287 tensor<4x256xf32>, tensor<4x256xf32>) { 288 ^bb0(%arg3 : f32, %arg4 : f32, %arg5: f32, %arg6: f32, %arg7: f32, %arg8: f32, 289 %arg9 : f32, %arg10 : f32, %arg11 : f32, %arg12 : f32, %arg13 : f32, 290 %arg14 : f32): 291 // CHECK-DAG: %[[CST0:.*]] = constant dense<2.000000e+00> : vector<4x256xf32> 292 // CHECK-DAG: %[[CST1:.*]] = constant dense<1.000000e+00> : vector<4x256xf32> 293 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 294 // CHECK: %[[V2:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x256xf32>, vector<4x256xf32> 295 // CHECK: %[[V0:.*]] = vector.transfer_read %[[ARG2]][%[[C0]]], {{.*}} : tensor<256xf32>, vector<256xf32> 296 // CHECK: %[[V3:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x256xf32>, vector<4x256xf32> 297 // CHECK: %[[V1:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x256xf32>, vector<4x256xf32> 298 // CHECK: %[[V0B:.*]] = vector.broadcast %[[V0]] : vector<256xf32> to vector<4x256xf32> 299 // CHECK: %[[ADD:.*]] = addf %[[V0B]], %[[V1]] : vector<4x256xf32> 300 %6 = addf %arg4, %arg6 : f32 301 // CHECK: %[[CMP:.*]] = cmpf ogt, %[[V2]], %[[V1]] : vector<4x256xf32> 302 %7 = cmpf ogt, %arg3, %arg6 : f32 303 // CHECK: %[[ARG3B:.*]] = vector.broadcast %[[ARG3]] : f32 to vector<4x256xf32> 304 %8 = constant 2.0 : f32 305 // CHECK: %[[DIV:.*]] = divf %[[V3]], %[[ARG3B]] : vector<4x256xf32> 306 %9 = divf %arg5, %i : f32 307 // CHECK: %[[EXP:.*]] = math.exp2 %[[V3]] : vector<4x256xf32> 308 %10 = math.exp2 %arg5 : f32 309 // CHECK: %[[MUL:.*]] = mulf %[[V3]], %[[CST0]] : vector<4x256xf32> 310 %11 = mulf %arg5, %8 : f32 311 // CHECK: %[[RSQRT:.*]] = math.rsqrt %[[V3]] : vector<4x256xf32> 312 %12 = math.rsqrt %arg5 : f32 313 // CHECK: %[[SEL:.*]] = select %[[CMP]], %[[V3]], %[[V1]] : vector<4x256xi1>, vector<4x256xf32> 314 %13 = select %7, %arg5, %arg6 : f32 315 // CHECK: %[[V0B:.*]] = vector.broadcast %[[V0]] : vector<256xf32> to vector<4x256xf32> 316 // CHECK: %[[SUB:.*]] = subf %[[V3]], %[[V0B]] : vector<4x256xf32> 317 %14 = subf %arg5, %arg4 : f32 318 // CHECK: %[[TAN:.*]] = math.tanh %[[V3]] : vector<4x256xf32> 319 %15 = math.tanh %arg5 : f32 320 // CHECK: %[[R0:.*]] = vector.transfer_write %[[ADD]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 321 // CHECK: %[[R1:.*]] = vector.transfer_write %[[CST0]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 322 // CHECK: %[[R2:.*]] = vector.transfer_write %[[CST1]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 323 // CHECK: %[[R3:.*]] = vector.transfer_write %[[DIV]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 324 // CHECK: %[[R4:.*]] = vector.transfer_write %[[EXP]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 325 // CHECK: %[[R5:.*]] = vector.transfer_write %[[MUL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 326 // CHECK: %[[R6:.*]] = vector.transfer_write %[[RSQRT]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 327 // CHECK: %[[R7:.*]] = vector.transfer_write %[[SEL]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 328 // CHECK: %[[R8:.*]] = vector.transfer_write %[[SUB]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 329 // CHECK: %[[R9:.*]] = vector.transfer_write %[[TAN]], %[[ARG0]][%[[C0]], %[[C0]]] {{.*}} : vector<4x256xf32>, tensor<4x256xf32> 330 linalg.yield %6, %8, %c1_f32, %9, %10, %11, %12, %13, %14, %15 : f32, f32, 331 f32, f32, f32, f32, f32, f32, f32, f32 332 } -> tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 333 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 334 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32> 335 // 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> 336 return %r#0, %r#1, %r#2, %r#3, %r#4, %r#5, %r#6, %r#7, %r#8, %r#9: 337 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 338 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32>, 339 tensor<4x256xf32>, tensor<4x256xf32>, tensor<4x256xf32> 340} 341 342// ----- 343 344// CHECK-LABEL: func @matmul_tensors 345// CHECK-SAME: (%[[ARG0:.*]]: tensor<8x4xf32>, %[[ARG1:.*]]: tensor<4x12xf32>, 346// CHECK-SAME: %[[ARG2:.*]]: tensor<8x12xf32>) -> tensor<8x12xf32> 347func @matmul_tensors( 348 %arg0: tensor<8x4xf32>, %arg1: tensor<4x12xf32>, %arg2: tensor<8x12xf32>) 349 -> tensor<8x12xf32> { 350 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 351 // CHECK-DAG: %[[VEC_C0:.*]] = constant dense<0.000000e+00> : vector<8x12xf32> 352 // CHECK-DAG: %[[V0:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : tensor<8x4xf32>, vector<8x4xf32> 353 // CHECK-DAG: %[[V1:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : tensor<4x12xf32>, vector<4x12xf32> 354 // CHECK-DAG: %[[V2:.*]] = vector.transfer_read %[[ARG2]][%[[C0]], %[[C0]]], {{.*}} : tensor<8x12xf32>, vector<8x12xf32> 355 // 356 // linalg contraction lowers to %tmp = vector.contract %a, %b, %c0 followed by addf %c, %tmp. 357 // a later canonicalization fuses the add into vector.contract. 358 // CHECK: %[[C:.*]] = vector.contract {{.*}} iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[V0]], %[[V1]], %[[VEC_C0]] : vector<8x4xf32>, vector<4x12xf32> into vector<8x12xf32> 359 // CHECK: %[[C2:.*]] = addf %[[V2]], %[[C]] : vector<8x12xf32> 360 // CHECK: %[[W:.*]] = vector.transfer_write %[[C2]], %[[ARG2]][%[[C0]], %[[C0]]] {masked = [false, false]} : vector<8x12xf32>, tensor<8x12xf32> 361 %0 = linalg.matmul ins(%arg0, %arg1: tensor<8x4xf32>, tensor<4x12xf32>) 362 outs(%arg2: tensor<8x12xf32>) 363 -> tensor<8x12xf32> 364 // CHECK: return %[[W]] : tensor<8x12xf32> 365 return %0 : tensor<8x12xf32> 366} 367 368// ----- 369 370// CHECK-LABEL: func @matmul_i8_i8_i32 371// CHECK-SAME: %[[ARG0:[a-z0-9]+]]: memref<4x6xi8> 372// CHECK-SAME: %[[ARG1:[a-z0-9]+]]: memref<6x12xi8> 373// CHECK-SAME: %[[ARG2:[a-z0-9]+]]: memref<4x12xi32> 374func @matmul_i8_i8_i32(%a: memref<4x6xi8>, %b: memref<6x12xi8>, %c: memref<4x12xi32>) { 375 // CHECK-DAG: %[[C0:.*]] = constant 0 : index 376 // CHECK-DAG: %[[VEC_C0:.*]] = constant dense<0> : vector<4x12xi8> 377 // CHECK-DAG: %[[V0:.*]] = vector.transfer_read %[[ARG0]][%[[C0]], %[[C0]]], {{.*}} : memref<4x6xi8>, vector<4x6xi8> 378 // CHECK-DAG: %[[V1:.*]] = vector.transfer_read %[[ARG1]][%[[C0]], %[[C0]]], {{.*}} : memref<6x12xi8>, vector<6x12xi8> 379 // CHECK-DAG: %[[V2:.*]] = vector.transfer_read %[[ARG2]][%[[C0]], %[[C0]]], {{.*}} : memref<4x12xi32>, vector<4x12xi32> 380 // 381 // linalg contraction lowers to %tmp = vector.contract %a, %b, %c0 followed by addf %c, %tmp. 382 // a later canonicalization fuses the add into vector.contract. 383 // CHECK: %[[C:.*]] = vector.contract {{.*}} iterator_types = ["parallel", "parallel", "reduction"], kind = #vector.kind<add>} %[[V0]], %[[V1]], %[[VEC_C0]] 384 // CHECK-SAME: vector<4x6xi8>, vector<6x12xi8> into vector<4x12xi8> 385 // CHECK: %[[C32:.*]] = sexti %[[C]] : vector<4x12xi8> to vector<4x12xi32> 386 // CHECK: %[[RES:.*]] = addi %[[V2]], %[[C32]] : vector<4x12xi32> 387 // CHECK: vector.transfer_write %[[RES]], %[[ARG2]][%[[C0]], %[[C0]]] {masked = [false, false]} 388 // CHECK-SAME: vector<4x12xi32>, memref<4x12xi32> 389 linalg.matmul_i8_i8_i32 ins(%a, %b : memref<4x6xi8>, memref<6x12xi8>) 390 outs(%c: memref<4x12xi32>) 391 return 392} 393 394// ----- 395 396// CHECK-LABEL: func @pad_static 397// CHECK-NOT: linalg.pad_tensor 398func @pad_static(%arg0: tensor<?x?x?xf32>, %pad_value: f32) -> tensor<2x3x4xf32> { 399 // CHECK: %[[C0:.*]] = constant 0 : index 400 // CHECK: %[[READ:.*]] = vector.transfer_read %{{.*}}[%[[C0]], %[[C0]], %[[C0]]] 401 // CHECK-SAME: : tensor<?x?x?xf32>, vector<2x3x4xf32> 402 // CHECK: %[[INIT:.*]] = linalg.init_tensor [2, 3, 4] : tensor<2x3x4xf32> 403 // CHECK: %[[WRITTEN:.*]] = vector.transfer_write %[[READ]], %[[INIT]][%[[C0]], %[[C0]], %[[C0]]] 404 // CHECK-SAME: {masked = [false, false, false]} : vector<2x3x4xf32>, tensor<2x3x4xf32> 405 %c0 = constant 0 : index 406 %0 = linalg.pad_tensor %arg0 low[0, %c0, 0] high[0, 0, %c0] { 407 ^bb0(%arg1: index, %arg2: index, %arg3: index): 408 linalg.yield %pad_value : f32 409 } : tensor<?x?x?xf32> to tensor<2x3x4xf32> 410 411 // CHECK: return %[[WRITTEN]] : tensor<2x3x4xf32> 412 return %0 : tensor<2x3x4xf32> 413} 414 415// CHECK-LABEL: func @pad_static_high_padding 416// CHECK: linalg.pad_tensor 417func @pad_static_high_padding(%arg0: tensor<?x?x?xf32>, %pad_value: f32) -> tensor<2x3x4xf32> { 418 %0 = linalg.pad_tensor %arg0 low[0, 0, 0] high[0, 1, 0] { 419 ^bb0(%arg1: index, %arg2: index, %arg3: index): 420 linalg.yield %pad_value : f32 421 } : tensor<?x?x?xf32> to tensor<2x3x4xf32> 422 return %0 : tensor<2x3x4xf32> 423} 424 425// CHECK-LABEL: func @pad_dynamic 426// CHECK: linalg.pad_tensor 427func @pad_dynamic(%arg0: tensor<1x2x2x?xf32>, %low: index, %high: index, 428 %pad_value: f32) -> tensor<6x?x?x?xf32> { 429 %0 = linalg.pad_tensor %arg0 low[2, %low, 3, 3] high[3, 3, %high, 2] { 430 ^bb0(%arg1: index, %arg2: index, %arg3: index, %arg4: index): 431 linalg.yield %pad_value : f32 432 } : tensor<1x2x2x?xf32> to tensor<6x?x?x?xf32> 433 return %0 : tensor<6x?x?x?xf32> 434} 435