1// RUN: mlir-opt %s -test-vector-contraction-conversion | FileCheck %s 2// RUN: mlir-opt %s -test-vector-contraction-conversion=vector-lower-matrix-intrinsics=1 | FileCheck %s --check-prefix=MATRIX 3// RUN: mlir-opt %s -test-vector-contraction-conversion=vector-outerproduct=1 | FileCheck %s --check-prefix=OUTERPRODUCT 4// RUN: mlir-opt %s -test-vector-contraction-conversion=vector-filter-outerproduct=1 | FileCheck %s --check-prefix=FILTEROUTERPRODUCT 5 6#dotp_accesses = [ 7 affine_map<(i) -> (i)>, 8 affine_map<(i) -> (i)>, 9 affine_map<(i) -> ()> 10] 11#dotp_trait = { 12 indexing_maps = #dotp_accesses, 13 iterator_types = ["reduction"] 14} 15 16// CHECK-LABEL: func @extract_contract1 17// CHECK-SAME: %[[A:.*0]]: vector<4xf32>, 18// CHECK-SAME: %[[B:.*1]]: vector<4xf32>, 19// CHECK-SAME: %[[C:.*2]]: f32 20// CHECK: %[[F:.*]] = mulf %[[A]], %[[B]] : vector<4xf32> 21// CHECK: %[[R:.*]] = vector.reduction "add", %[[F]] : vector<4xf32> into f32 22// CHECK: %[[ACC:.*]] = addf %[[R]], %[[C]] : f32 23// CHECK: return %[[ACC]] : f32 24 25func @extract_contract1(%arg0: vector<4xf32>, %arg1: vector<4xf32>, %arg2: f32) -> f32 { 26 %0 = vector.contract #dotp_trait %arg0, %arg1, %arg2 27 : vector<4xf32>, vector<4xf32> into f32 28 return %0 : f32 29} 30 31// CHECK-LABEL: func @extract_contract1_int 32// CHECK-SAME: %[[A:.*0]]: vector<4xi32>, 33// CHECK-SAME: %[[B:.*1]]: vector<4xi32>, 34// CHECK-SAME: %[[C:.*2]]: i32 35// CHECK: %[[F:.*]] = muli %[[A]], %[[B]] : vector<4xi32> 36// CHECK: %[[R:.*]] = vector.reduction "add", %[[F]] : vector<4xi32> into i32 37// CHECK: %[[ACC:.*]] = addi %[[R]], %[[C]] : i32 38// CHECK: return %[[ACC]] : i32 39 40func @extract_contract1_int(%arg0: vector<4xi32>, %arg1: vector<4xi32>, %arg2: i32) -> i32 { 41 %0 = vector.contract #dotp_trait %arg0, %arg1, %arg2 42 : vector<4xi32>, vector<4xi32> into i32 43 return %0 : i32 44} 45 46#matvec_accesses = [ 47 affine_map<(i, j) -> (i, j)>, 48 affine_map<(i, j) -> (j)>, 49 affine_map<(i, j) -> (i)> 50] 51#matvec_trait = { 52 indexing_maps = #matvec_accesses, 53 iterator_types = ["parallel", "reduction"] 54} 55 56// CHECK-LABEL: func @extract_contract2 57// CHECK-SAME: %[[A:.*0]]: vector<2x3xf32>, 58// CHECK-SAME: %[[B:.*1]]: vector<3xf32>, 59// CHECK-SAME: %[[C:.*2]]: vector<2xf32> 60// CHECK: %[[R:.*]] = constant dense<0.000000e+00> : vector<2xf32> 61// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x3xf32> 62// CHECK: %[[T2:.*]] = mulf %[[T0]], %[[B]] : vector<3xf32> 63// CHECK: %[[T3:.*]] = vector.reduction "add", %[[T2]] : vector<3xf32> into f32 64// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[R]] [0] : f32 into vector<2xf32> 65// CHECK: %[[T5:.*]] = vector.extract %[[A]][1] : vector<2x3xf32> 66// CHECK: %[[T7:.*]] = mulf %[[T5]], %[[B]] : vector<3xf32> 67// CHECK: %[[T8:.*]] = vector.reduction "add", %[[T7]] : vector<3xf32> into f32 68// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T4]] [1] : f32 into vector<2xf32> 69// CHECK: %[[T10:.*]] = addf %[[T9]], %[[C]] : vector<2xf32> 70// CHECK: return %[[T10]] : vector<2xf32> 71 72func @extract_contract2(%arg0: vector<2x3xf32>, 73 %arg1: vector<3xf32>, 74 %arg2: vector<2xf32>) -> vector<2xf32> { 75 %0 = vector.contract #matvec_trait %arg0, %arg1, %arg2 76 : vector<2x3xf32>, vector<3xf32> into vector<2xf32> 77 return %0 : vector<2xf32> 78} 79 80// CHECK-LABEL: func @extract_contract2_int 81// CHECK-SAME: %[[A:.*0]]: vector<2x3xi32>, 82// CHECK-SAME: %[[B:.*1]]: vector<3xi32>, 83// CHECK-SAME: %[[C:.*2]]: vector<2xi32> 84// CHECK: %[[R:.*]] = constant dense<0> : vector<2xi32> 85// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x3xi32> 86// CHECK: %[[T2:.*]] = muli %[[T0]], %[[B]] : vector<3xi32> 87// CHECK: %[[T3:.*]] = vector.reduction "add", %[[T2]] : vector<3xi32> into i32 88// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[R]] [0] : i32 into vector<2xi32> 89// CHECK: %[[T5:.*]] = vector.extract %[[A]][1] : vector<2x3xi32> 90// CHECK: %[[T7:.*]] = muli %[[T5]], %[[B]] : vector<3xi32> 91// CHECK: %[[T8:.*]] = vector.reduction "add", %[[T7]] : vector<3xi32> into i32 92// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T4]] [1] : i32 into vector<2xi32> 93// CHECK: %[[T10:.*]] = addi %[[T9]], %[[C]] : vector<2xi32> 94// CHECK: return %[[T10]] : vector<2xi32> 95func @extract_contract2_int(%arg0: vector<2x3xi32>, 96 %arg1: vector<3xi32>, 97 %arg2: vector<2xi32>) -> vector<2xi32> { 98 %0 = vector.contract #matvec_trait %arg0, %arg1, %arg2 99 : vector<2x3xi32>, vector<3xi32> into vector<2xi32> 100 return %0 : vector<2xi32> 101} 102 103#vecmat_accesses = [ 104 affine_map<(i, j) -> (j)>, 105 affine_map<(i, j) -> (i, j)>, 106 affine_map<(i, j) -> (i)> 107] 108#vecmat_trait = { 109 indexing_maps = #vecmat_accesses, 110 iterator_types = ["parallel", "reduction"] 111} 112 113// CHECK-LABEL: func @extract_contract3 114// CHECK-SAME: %[[A:.*0]]: vector<3xf32>, 115// CHECK-SAME: %[[B:.*1]]: vector<2x3xf32>, 116// CHECK-SAME: %[[C:.*2]]: vector<2xf32> 117// CHECK: %[[R:.*]] = constant dense<0.000000e+00> : vector<2xf32> 118// CHECK: %[[T0:.*]] = vector.extract %[[B]][0] : vector<2x3xf32> 119// CHECK: %[[T2:.*]] = mulf %[[T0]], %[[A]] : vector<3xf32> 120// CHECK: %[[T3:.*]] = vector.reduction "add", %[[T2]] : vector<3xf32> into f32 121// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[R]] [0] : f32 into vector<2xf32> 122// CHECK: %[[T5:.*]] = vector.extract %[[B]][1] : vector<2x3xf32> 123// CHECK: %[[T7:.*]] = mulf %[[T5]], %[[A]] : vector<3xf32> 124// CHECK: %[[T8:.*]] = vector.reduction "add", %[[T7]] : vector<3xf32> into f32 125// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T4]] [1] : f32 into vector<2xf32> 126// CHECK: %[[T10:.*]] = addf %[[T9]], %[[C]] : vector<2xf32> 127// CHECK: return %[[T10]] : vector<2xf32> 128 129func @extract_contract3(%arg0: vector<3xf32>, 130 %arg1: vector<2x3xf32>, 131 %arg2: vector<2xf32>) -> vector<2xf32> { 132 %0 = vector.contract #vecmat_trait %arg0, %arg1, %arg2 133 : vector<3xf32>, vector<2x3xf32> into vector<2xf32> 134 return %0 : vector<2xf32> 135} 136 137#matmat_accesses = [ 138 affine_map<(i, j, k) -> (i, k)>, 139 affine_map<(i, j, k) -> (k, j)>, 140 affine_map<(i, j, k) -> (i, j)> 141] 142#matmat_trait = { 143 indexing_maps = #matmat_accesses, 144 iterator_types = ["parallel", "parallel", "reduction"] 145} 146 147// CHECK-LABEL: func @extract_contract4 148// CHECK-SAME: %[[A:.*0]]: vector<2x2xf32>, 149// CHECK-SAME: %[[B:.*1]]: vector<2x2xf32>, 150// CHECK-SAME: %[[C:.*2]]: vector<2x2xf32> 151// CHECK: %[[R:.*]] = constant dense<0.000000e+00> : vector<2x2xf32> 152// ... bunch of extract insert to transpose B into Bt 153// CHECK: %[[Bt:.*]] = vector.insert %{{.*}}, %{{.*}} [1, 1] : f32 into vector<2x2xf32> 154// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x2xf32> 155// CHECK: %[[T2:.*]] = vector.extract %[[Bt]][0] : vector<2x2xf32> 156// CHECK: %[[T9:.*]] = mulf %[[T0]], %[[T2]] : vector<2xf32> 157// CHECK: %[[T10:.*]] = vector.reduction "add", %[[T9]] : vector<2xf32> into f32 158// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[R]] [0, 0] : f32 into vector<2x2xf32> 159// 160// CHECK: %[[T12:.*]] = vector.extract %[[Bt]][1] : vector<2x2xf32> 161// CHECK: %[[T19:.*]] = mulf %[[T0]], %[[T12]] : vector<2xf32> 162// CHECK: %[[T20:.*]] = vector.reduction "add", %[[T19]] : vector<2xf32> into f32 163// CHECK: %[[T21:.*]] = vector.insert %[[T20]], %[[T11]] [0, 1] : f32 into vector<2x2xf32> 164// 165// CHECK: %[[T23:.*]] = vector.extract %[[A]][1] : vector<2x2xf32> 166// CHECK: %[[T24:.*]] = vector.extract %[[Bt]][0] : vector<2x2xf32> 167// CHECK: %[[T32:.*]] = mulf %[[T23]], %[[T24]] : vector<2xf32> 168// CHECK: %[[T33:.*]] = vector.reduction "add", %[[T32]] : vector<2xf32> into f32 169// CHECK: %[[T34:.*]] = vector.insert %[[T33]], %[[T21]] [1, 0] : f32 into vector<2x2xf32> 170// 171// CHECK: %[[T40:.*]] = vector.extract %[[Bt]][1] : vector<2x2xf32> 172// CHECK: %[[T41:.*]] = mulf %[[T23]], %[[T40]] : vector<2xf32> 173// CHECK: %[[T42:.*]] = vector.reduction "add", %[[T41]] : vector<2xf32> into f32 174// CHECK: %[[T43:.*]] = vector.insert %[[T42]], %[[T34]] [1, 1] : f32 into vector<2x2xf32> 175// 176// CHECK: %[[T52:.*]] = addf %[[T43]], %[[C]] : vector<2x2xf32> 177// CHECK: return %[[T52]] : vector<2x2xf32> 178 179func @extract_contract4(%arg0: vector<2x2xf32>, 180 %arg1: vector<2x2xf32>, 181 %arg2: vector<2x2xf32>) -> vector<2x2xf32> { 182 %0 = vector.contract #matmat_trait %arg0, %arg1, %arg2 183 : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 184 return %0 : vector<2x2xf32> 185} 186 187#contraction2d_accesses = [ 188 affine_map<(i, j) -> (i, j)>, 189 affine_map<(i, j) -> (i, j)>, 190 affine_map<(i, j) -> ()> 191] 192#contraction2d_trait = { 193 indexing_maps = #contraction2d_accesses, 194 iterator_types = ["reduction", "reduction"] 195} 196 197// CHECK-LABEL: func @full_contract1 198// CHECK-SAME: %[[A:.*0]]: vector<2x3xf32>, 199// CHECK-SAME: %[[B:.*1]]: vector<2x3xf32>, 200// CHECK-SAME: %[[C:.*2]]: f32 201// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x3xf32> 202// CHECK: %[[T1:.*]] = vector.extract %[[B]][0] : vector<2x3xf32> 203// CHECK: %[[T2:.*]] = mulf %[[T0]], %[[T1]] : vector<3xf32> 204// CHECK: %[[T3:.*]] = vector.reduction "add", %[[T2]] : vector<3xf32> into f32 205// CHECK: %[[T4:.*]] = addf %[[T3]], %[[C]] : f32 206// CHECK: %[[T5:.*]] = vector.extract %[[A]][1] : vector<2x3xf32> 207// CHECK: %[[T6:.*]] = vector.extract %[[B]][1] : vector<2x3xf32> 208// CHECK: %[[T7:.*]] = mulf %[[T5]], %[[T6]] : vector<3xf32> 209// CHECK: %[[T8:.*]] = vector.reduction "add", %[[T7]] : vector<3xf32> into f32 210// CHECK: %[[T9:.*]] = addf %[[T8]], %[[T4]] : f32 211// CHECK: return %[[T9]] : f32 212 213func @full_contract1(%arg0: vector<2x3xf32>, 214 %arg1: vector<2x3xf32>, 215 %arg2: f32) -> f32 { 216 %0 = vector.contract #contraction2d_trait %arg0, %arg1, %arg2 217 : vector<2x3xf32>, vector<2x3xf32> into f32 218 return %0 : f32 219} 220 221#contraction2d_trans_accesses = [ 222 affine_map<(i, j) -> (i, j)>, 223 affine_map<(i, j) -> (j, i)>, 224 affine_map<(i, j) -> ()> 225] 226#contraction2d_trans_trait = { 227 indexing_maps = #contraction2d_trans_accesses, 228 iterator_types = ["reduction", "reduction"] 229} 230 231// CHECK-LABEL: func @full_contract2 232// CHECK-SAME: %[[A:.*0]]: vector<2x3xf32>, 233// CHECK-SAME: %[[B:.*1]]: vector<3x2xf32>, 234// CHECK-SAME: %[[C:.*2]]: f32 235// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<3xf32> 236// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x3xf32> 237// CHECK: %[[T1:.*]] = vector.extract %[[B]][0, 0] : vector<3x2xf32> 238// CHECK: %[[T3:.*]] = vector.insert %[[T1]], %[[Z]] [0] : f32 into vector<3xf32> 239// CHECK: %[[T4:.*]] = vector.extract %[[B]][1, 0] : vector<3x2xf32> 240// CHECK: %[[T6:.*]] = vector.insert %[[T4]], %[[T3]] [1] : f32 into vector<3xf32> 241// CHECK: %[[T7:.*]] = vector.extract %[[B]][2, 0] : vector<3x2xf32> 242// CHECK: %[[T9:.*]] = vector.insert %[[T7]], %[[T6]] [2] : f32 into vector<3xf32> 243// CHECK: %[[T10:.*]] = mulf %[[T0]], %[[T9]] : vector<3xf32> 244// CHECK: %[[T11:.*]] = vector.reduction "add", %[[T10]] : vector<3xf32> into f32 245// CHECK: %[[ACC0:.*]] = addf %[[T11]], %[[C]] : f32 246// 247// CHECK: %[[T12:.*]] = vector.extract %[[A]][1] : vector<2x3xf32> 248// CHECK: %[[T13:.*]] = vector.extract %[[B]][0, 1] : vector<3x2xf 249// CHECK: %[[T15:.*]] = vector.insert %[[T13]], %[[Z]] [0] : f32 into vector<3xf32> 250// CHECK: %[[T16:.*]] = vector.extract %[[B]][1, 1] : vector<3x2xf32> 251// CHECK: %[[T18:.*]] = vector.insert %[[T16]], %[[T15]] [1] : f32 into vector<3xf32> 252// CHECK: %[[T19:.*]] = vector.extract %[[B]][2, 1] : vector<3x2xf32> 253// CHECK: %[[T21:.*]] = vector.insert %[[T19]], %[[T18]] [2] : f32 into vector<3xf32> 254// CHECK: %[[T22:.*]] = mulf %[[T12]], %[[T21]] : vector<3xf32> 255// CHECK: %[[T23:.*]] = vector.reduction "add", %[[T22]] : vector<3xf32> into f32 256// CHECK: %[[ACC1:.*]] = addf %[[T23]], %[[ACC0]] : f32 257// CHECK: return %[[ACC1]] : f32 258 259func @full_contract2(%arg0: vector<2x3xf32>, 260 %arg1: vector<3x2xf32>, 261 %arg2: f32) -> f32 { 262 %0 = vector.contract #contraction2d_trans_trait %arg0, %arg1, %arg2 263 : vector<2x3xf32>, vector<3x2xf32> into f32 264 return %0 : f32 265} 266 267// CHECK-LABEL: func @outerproduct_noacc 268// CHECK-SAME: %[[A:.*0]]: vector<2xf32>, 269// CHECK-SAME: %[[B:.*1]]: vector<3xf32> 270// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 271// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2xf32> 272// CHECK: %[[T1:.*]] = splat %[[T0]] : vector<3xf32> 273// CHECK: %[[T2:.*]] = mulf %[[T1]], %[[B]] : vector<3xf32> 274// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[C0]] [0] : vector<3xf32> into vector<2x3xf32> 275// CHECK: %[[T4:.*]] = vector.extract %[[A]][1] : vector<2xf32> 276// CHECK: %[[T5:.*]] = splat %[[T4]] : vector<3xf32> 277// CHECK: %[[T6:.*]] = mulf %[[T5]], %[[B]] : vector<3xf32> 278// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T3]] [1] : vector<3xf32> into vector<2x3xf32> 279// CHECK: return %[[T7]] : vector<2x3xf32> 280 281func @outerproduct_noacc(%arg0: vector<2xf32>, 282 %arg1: vector<3xf32>) -> vector<2x3xf32> { 283 %0 = vector.outerproduct %arg0, %arg1 : vector<2xf32>, vector<3xf32> 284 return %0: vector<2x3xf32> 285} 286 287// CHECK-LABEL: func @outerproduct_acc 288// CHECK-SAME: %[[A:.*0]]: vector<2xf32>, 289// CHECK-SAME: %[[B:.*1]]: vector<3xf32>, 290// CHECK-SAME: %[[C:.*2]]: vector<2x3xf32> 291// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 292// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2xf32> 293// CHECK: %[[T1:.*]] = splat %[[T0]] : vector<3xf32> 294// CHECK: %[[T2:.*]] = vector.extract %[[C]][0] : vector<2x3xf32> 295// CHECK: %[[T3:.*]] = vector.fma %[[T1]], %[[B]], %[[T2]] : vector<3xf32> 296// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[C0]] [0] : vector<3xf32> into vector<2x3xf32> 297// CHECK: %[[T5:.*]] = vector.extract %[[A]][1] : vector<2xf32> 298// CHECK: %[[T6:.*]] = splat %[[T5]] : vector<3xf32> 299// CHECK: %[[T7:.*]] = vector.extract %[[C]][1] : vector<2x3xf32> 300// CHECK: %[[T8:.*]] = vector.fma %[[T6]], %[[B]], %[[T7]] : vector<3xf32> 301// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T4]] [1] : vector<3xf32> into vector<2x3xf32> 302// CHECK: return %[[T9]] : vector<2x3xf32> 303 304func @outerproduct_acc(%arg0: vector<2xf32>, 305 %arg1: vector<3xf32>, 306 %arg2: vector<2x3xf32>) -> vector<2x3xf32> { 307 %0 = vector.outerproduct %arg0, %arg1, %arg2 : vector<2xf32>, vector<3xf32> 308 return %0: vector<2x3xf32> 309} 310 311// CHECK-LABEL: func @outerproduct_noacc_int 312// CHECK-SAME: %[[A:.*0]]: vector<2xi32>, 313// CHECK-SAME: %[[B:.*1]]: vector<3xi32> 314// CHECK: %[[C0:.*]] = constant dense<0> : vector<2x3xi32> 315// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2xi32> 316// CHECK: %[[T1:.*]] = splat %[[T0]] : vector<3xi32> 317// CHECK: %[[T2:.*]] = muli %[[T1]], %[[B]] : vector<3xi32> 318// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[C0]] [0] : vector<3xi32> into vector<2x3xi32> 319// CHECK: %[[T4:.*]] = vector.extract %[[A]][1] : vector<2xi32> 320// CHECK: %[[T5:.*]] = splat %[[T4]] : vector<3xi32> 321// CHECK: %[[T6:.*]] = muli %[[T5]], %[[B]] : vector<3xi32> 322// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T3]] [1] : vector<3xi32> into vector<2x3xi32> 323// CHECK: return %[[T7]] : vector<2x3xi32> 324func @outerproduct_noacc_int(%arg0: vector<2xi32>, 325 %arg1: vector<3xi32>) -> vector<2x3xi32> { 326 %0 = vector.outerproduct %arg0, %arg1 : vector<2xi32>, vector<3xi32> 327 return %0: vector<2x3xi32> 328} 329 330// CHECK-LABEL: func @outerproduct_acc_int 331// CHECK-SAME: %[[A:.*0]]: vector<2xi32>, 332// CHECK-SAME: %[[B:.*1]]: vector<3xi32>, 333// CHECK-SAME: %[[C:.*2]]: vector<2x3xi32> 334// CHECK: %[[C0:.*]] = constant dense<0> : vector<2x3xi32> 335// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2xi32> 336// CHECK: %[[T1:.*]] = splat %[[T0]] : vector<3xi32> 337// CHECK: %[[T2:.*]] = vector.extract %[[C]][0] : vector<2x3xi32> 338// CHECK: %[[T3:.*]] = muli %[[T1]], %[[B]] : vector<3xi32> 339// CHECK: %[[T4:.*]] = addi %[[T3]], %[[T2]] : vector<3xi32> 340// CHECK: %[[T5:.*]] = vector.insert %[[T4]], %[[C0]] [0] : vector<3xi32> into vector<2x3xi32> 341// CHECK: %[[T6:.*]] = vector.extract %[[A]][1] : vector<2xi32> 342// CHECK: %[[T7:.*]] = splat %[[T6]] : vector<3xi32> 343// CHECK: %[[T8:.*]] = vector.extract %[[C]][1] : vector<2x3xi32> 344// CHECK: %[[T9:.*]] = muli %[[T7]], %[[B]] : vector<3xi32> 345// CHECK: %[[T10:.*]] = addi %[[T9]], %[[T8]] : vector<3xi32> 346// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T5]] [1] : vector<3xi32> into vector<2x3xi32> 347// CHECK: return %[[T11]] : vector<2x3xi32> 348func @outerproduct_acc_int(%arg0: vector<2xi32>, 349 %arg1: vector<3xi32>, 350 %arg2: vector<2x3xi32>) -> vector<2x3xi32> { 351 %0 = vector.outerproduct %arg0, %arg1, %arg2 : vector<2xi32>, vector<3xi32> 352 return %0: vector<2x3xi32> 353} 354 355// CHECK-LABEL: func @axpy_fp( 356// CHECK-SAME: %[[A:.*0]]: vector<16xf32>, 357// CHECK-SAME: %[[B:.*1]]: f32) 358// CHECK: %[[T0:.*]] = splat %[[B]] : vector<16xf32> 359// CHECK: %[[T1:.*]] = mulf %[[A]], %[[T0]] : vector<16xf32> 360// CHECK: return %[[T1]] : vector<16xf32> 361func @axpy_fp(%arg0: vector<16xf32>, %arg1: f32) -> vector<16xf32> { 362 %0 = vector.outerproduct %arg0, %arg1: vector<16xf32>, f32 363 return %0: vector<16xf32> 364} 365 366// CHECK-LABEL: func @axpy_fp_add( 367// CHECK-SAME: %[[A:.*0]]: vector<16xf32>, 368// CHECK-SAME: %[[B:.*1]]: f32, 369// CHECK-SAME: %[[C:.*2]]: vector<16xf32>) 370// CHECK: %[[T0:.*]] = splat %[[B]] : vector<16xf32> 371// CHECK: %[[T1:.*]] = vector.fma %[[A]], %[[T0]], %[[C]] : vector<16xf32> 372// CHECK: return %[[T1]] : vector<16xf32> 373func @axpy_fp_add(%arg0: vector<16xf32>, %arg1: f32, %arg2 : vector<16xf32>) -> vector<16xf32> { 374 %0 = vector.outerproduct %arg0, %arg1, %arg2: vector<16xf32>, f32 375 return %0: vector<16xf32> 376} 377 378// CHECK-LABEL: func @axpy_int( 379// CHECK-SAME: %[[A:.*0]]: vector<16xi32>, 380// CHECK-SAME: %[[B:.*1]]: i32) 381// CHECK: %[[T0:.*]] = splat %[[B]] : vector<16xi32> 382// CHECK: %[[T1:.*]] = muli %[[A]], %[[T0]] : vector<16xi32> 383// CHECK: return %[[T1]] : vector<16xi32> 384func @axpy_int(%arg0: vector<16xi32>, %arg1: i32) -> vector<16xi32> { 385 %0 = vector.outerproduct %arg0, %arg1: vector<16xi32>, i32 386 return %0: vector<16xi32> 387} 388 389// CHECK-LABEL: func @axpy_int_add( 390// CHECK-SAME: %[[A:.*0]]: vector<16xi32>, 391// CHECK-SAME: %[[B:.*1]]: i32, 392// CHECK-SAME: %[[C:.*2]]: vector<16xi32>) 393// CHECK: %[[T0:.*]] = splat %[[B]] : vector<16xi32> 394// CHECK: %[[T1:.*]] = muli %[[A]], %[[T0]] : vector<16xi32> 395// CHECK: %[[T2:.*]] = addi %[[T1]], %[[C]] : vector<16xi32> 396// CHECK: return %[[T2]] : vector<16xi32> 397func @axpy_int_add(%arg0: vector<16xi32>, %arg1: i32, %arg2: vector<16xi32>) -> vector<16xi32> { 398 %0 = vector.outerproduct %arg0, %arg1, %arg2: vector<16xi32>, i32 399 return %0: vector<16xi32> 400} 401 402// CHECK-LABEL: func @transpose23 403// CHECK-SAME: %[[A:.*]]: vector<2x3xf32> 404// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<3x2xf32> 405// CHECK: %[[T0:.*]] = vector.extract %[[A]][0, 0] : vector<2x3xf32> 406// CHECK: %[[T1:.*]] = vector.insert %[[T0]], %[[Z]] [0, 0] : f32 into vector<3x2xf32> 407// CHECK: %[[T2:.*]] = vector.extract %[[A]][1, 0] : vector<2x3xf32> 408// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[T1]] [0, 1] : f32 into vector<3x2xf32> 409// CHECK: %[[T4:.*]] = vector.extract %[[A]][0, 1] : vector<2x3xf32> 410// CHECK: %[[T5:.*]] = vector.insert %[[T4]], %[[T3]] [1, 0] : f32 into vector<3x2xf32> 411// CHECK: %[[T6:.*]] = vector.extract %[[A]][1, 1] : vector<2x3xf32> 412// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T5]] [1, 1] : f32 into vector<3x2xf32> 413// CHECK: %[[T8:.*]] = vector.extract %[[A]][0, 2] : vector<2x3xf32> 414// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T7]] [2, 0] : f32 into vector<3x2xf32> 415// CHECK: %[[T10:.*]] = vector.extract %[[A]][1, 2] : vector<2x3xf32> 416// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T9]] [2, 1] : f32 into vector<3x2xf32> 417// CHECK: return %[[T11]] : vector<3x2xf32> 418 419func @transpose23(%arg0: vector<2x3xf32>) -> vector<3x2xf32> { 420 %0 = vector.transpose %arg0, [1, 0] : vector<2x3xf32> to vector<3x2xf32> 421 return %0 : vector<3x2xf32> 422} 423 424// CHECK-LABEL: func @nop_shape_cast 425// CHECK-SAME: %[[A:.*]]: vector<16xf32> 426// CHECK: return %[[A]] : vector<16xf32> 427 428func @nop_shape_cast(%arg0: vector<16xf32>) -> vector<16xf32> { 429 %0 = vector.shape_cast %arg0 : vector<16xf32> to vector<16xf32> 430 return %0 : vector<16xf32> 431} 432 433// CHECK-LABEL: func @cancel_shape_cast 434// FIXME: PR49590 435// HECK-SAME: %[[A:.*]]: vector<16xf32> 436// HECK: return %[[A]] : vector<16xf32> 437 438func @cancel_shape_cast(%arg0: vector<16xf32>) -> vector<16xf32> { 439 %0 = vector.shape_cast %arg0 : vector<16xf32> to vector<4x4xf32> 440 %1 = vector.shape_cast %0 : vector<4x4xf32> to vector<16xf32> 441 return %1 : vector<16xf32> 442} 443 444// Shape up and downcasts for 2-D vectors, for supporting conversion to 445// llvm.matrix operations 446// CHECK-LABEL: func @shape_casts 447func @shape_casts(%a: vector<2x2xf32>) -> (vector<4xf32>, vector<2x2xf32>) { 448 // CHECK-DAG: %[[cst22:.*]] = constant dense<0.000000e+00> : vector<2x2xf32> 449 // CHECK-DAG: %[[cst:.*]] = constant dense<0.000000e+00> : vector<4xf32> 450 // CHECK: %[[ex0:.*]] = vector.extract %{{.*}}[0] : vector<2x2xf32> 451 // 452 // CHECK: %[[in0:.*]] = vector.insert_strided_slice %[[ex0]], %[[cst]] 453 // CHECK-SAME: {offsets = [0], strides = [1]} : vector<2xf32> into vector<4xf32> 454 // 455 // CHECK: %[[ex1:.*]] = vector.extract %{{.*}}[1] : vector<2x2xf32> 456 // 457 // CHECK: %[[in2:.*]] = vector.insert_strided_slice %[[ex1]], %[[in0]] 458 // CHECK-SAME: {offsets = [2], strides = [1]} : vector<2xf32> into vector<4xf32> 459 // 460 %0 = vector.shape_cast %a : vector<2x2xf32> to vector<4xf32> 461 // CHECK: %[[add:.*]] = addf %[[in2]], %[[in2]] : vector<4xf32> 462 %r0 = addf %0, %0: vector<4xf32> 463 // 464 // CHECK: %[[ss0:.*]] = vector.extract_strided_slice %[[add]] 465 // CHECK-SAME: {offsets = [0], sizes = [2], strides = [1]} : 466 // CHECK-SAME: vector<4xf32> to vector<2xf32> 467 // 468 // CHECK: %[[res0:.*]] = vector.insert %[[ss0]], %[[cst22]] [0] : 469 // CHECK-SAME: vector<2xf32> into vector<2x2xf32> 470 // 471 // CHECK: %[[s2:.*]] = vector.extract_strided_slice %[[add]] 472 // CHECK-SAME: {offsets = [2], sizes = [2], strides = [1]} : 473 // CHECK-SAME: vector<4xf32> to vector<2xf32> 474 // 475 // CHECK: %[[res1:.*]] = vector.insert %[[s2]], %[[res0]] [1] : 476 // CHECK-SAME: vector<2xf32> into vector<2x2xf32> 477 // 478 %1 = vector.shape_cast %r0 : vector<4xf32> to vector<2x2xf32> 479 // CHECK: return %[[add]], %[[res1]] : vector<4xf32>, vector<2x2xf32> 480 return %r0, %1 : vector<4xf32>, vector<2x2xf32> 481} 482 483// CHECK-LABEL: func @shape_cast_2d2d 484// CHECK-SAME: %[[A:.*]]: vector<3x2xf32> 485// CHECK: %[[C:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 486// CHECK: %[[T0:.*]] = vector.extract %[[A]][0, 0] : vector<3x2xf32> 487// CHECK: %[[T1:.*]] = vector.insert %[[T0]], %[[C]] [0, 0] : f32 into vector<2x3xf32> 488// CHECK: %[[T2:.*]] = vector.extract %[[A]][0, 1] : vector<3x2xf32> 489// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[T1]] [0, 1] : f32 into vector<2x3xf32> 490// CHECK: %[[T4:.*]] = vector.extract %[[A]][1, 0] : vector<3x2xf32> 491// CHECK: %[[T5:.*]] = vector.insert %[[T4]], %[[T3]] [0, 2] : f32 into vector<2x3xf32> 492// CHECK: %[[T6:.*]] = vector.extract %[[A]][1, 1] : vector<3x2xf32> 493// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T5]] [1, 0] : f32 into vector<2x3xf32> 494// CHECK: %[[T8:.*]] = vector.extract %[[A]][2, 0] : vector<3x2xf32> 495// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T7]] [1, 1] : f32 into vector<2x3xf32> 496// CHECK: %[[T10:.*]] = vector.extract %[[A]][2, 1] : vector<3x2xf32> 497// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T9]] [1, 2] : f32 into vector<2x3xf32> 498// CHECK: return %[[T11]] : vector<2x3xf32> 499 500func @shape_cast_2d2d(%arg0 : vector<3x2xf32>) -> vector<2x3xf32> { 501 %s = vector.shape_cast %arg0: vector<3x2xf32> to vector<2x3xf32> 502 return %s : vector<2x3xf32> 503} 504 505// CHECK-LABEL: func @shape_cast_3d1d 506// CHECK-SAME: %[[A:.*]]: vector<1x3x2xf32> 507// CHECK: %[[C:.*]] = constant dense<0.000000e+00> : vector<6xf32> 508// CHECK: %[[T0:.*]] = vector.extract %[[A]][0, 0, 0] : vector<1x3x2xf32> 509// CHECK: %[[T1:.*]] = vector.insert %[[T0]], %[[C]] [0] : f32 into vector<6xf32> 510// CHECK: %[[T2:.*]] = vector.extract %[[A]][0, 0, 1] : vector<1x3x2xf32> 511// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[T1]] [1] : f32 into vector<6xf32> 512// CHECK: %[[T4:.*]] = vector.extract %[[A]][0, 1, 0] : vector<1x3x2xf32> 513// CHECK: %[[T5:.*]] = vector.insert %[[T4]], %[[T3]] [2] : f32 into vector<6xf32> 514// CHECK: %[[T6:.*]] = vector.extract %[[A]][0, 1, 1] : vector<1x3x2xf32> 515// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T5]] [3] : f32 into vector<6xf32> 516// CHECK: %[[T8:.*]] = vector.extract %[[A]][0, 2, 0] : vector<1x3x2xf32> 517// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T7]] [4] : f32 into vector<6xf32> 518// CHECK: %[[T10:.*]] = vector.extract %[[A]][0, 2, 1] : vector<1x3x2xf32> 519// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T9]] [5] : f32 into vector<6xf32> 520// CHECK: return %[[T11]] : vector<6xf32> 521 522func @shape_cast_3d1d(%arg0 : vector<1x3x2xf32>) -> vector<6xf32> { 523 %s = vector.shape_cast %arg0 : vector<1x3x2xf32> to vector<6xf32> 524 return %s : vector<6xf32> 525} 526 527// CHECK-LABEL: func @shape_cast_1d3d 528// CHECK-SAME: %[[A:.*]]: vector<6xf32> 529// CHECK: %[[C:.*]] = constant dense<0.000000e+00> : vector<2x1x3xf32> 530// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<6xf32> 531// CHECK: %[[T1:.*]] = vector.insert %[[T0]], %[[C]] [0, 0, 0] : f32 into vector<2x1x3xf32> 532// CHECK: %[[T2:.*]] = vector.extract %[[A]][1] : vector<6xf32> 533// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[T1]] [0, 0, 1] : f32 into vector<2x1x3xf32> 534// CHECK: %[[T4:.*]] = vector.extract %[[A]][2] : vector<6xf32> 535// CHECK: %[[T5:.*]] = vector.insert %[[T4]], %[[T3]] [0, 0, 2] : f32 into vector<2x1x3xf32> 536// CHECK: %[[T6:.*]] = vector.extract %[[A]][3] : vector<6xf32> 537// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T5]] [1, 0, 0] : f32 into vector<2x1x3xf32> 538// CHECK: %[[T8:.*]] = vector.extract %[[A]][4] : vector<6xf32> 539// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T7]] [1, 0, 1] : f32 into vector<2x1x3xf32> 540// CHECK: %[[T10:.*]] = vector.extract %[[A]][5] : vector<6xf32> 541// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T9]] [1, 0, 2] : f32 into vector<2x1x3xf32> 542// CHECK: return %[[T11]] : vector<2x1x3xf32> 543 544func @shape_cast_1d3d(%arg0 : vector<6xf32>) -> vector<2x1x3xf32> { 545 %s = vector.shape_cast %arg0 : vector<6xf32> to vector<2x1x3xf32> 546 return %s : vector<2x1x3xf32> 547} 548 549// MATRIX-LABEL: func @matmul 550// MATRIX-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x4xf32>, 551// MATRIX-SAME: %[[B:[a-zA-Z0-9]*]]: vector<4x3xf32>, 552// MATRIX-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 553// MATRIX: %[[vcst:.*]] = constant dense<0.000000e+00> : vector<8xf32> 554// MATRIX: %[[vcst_0:.*]] = constant dense<0.000000e+00> : vector<12xf32> 555// MATRIX: %[[vcst_1:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 556// MATRIX: %[[a0:.*]] = vector.extract %[[A]][0] : vector<2x4xf32> 557// MATRIX: %[[a1:.*]] = vector.insert_strided_slice %[[a0]], %[[vcst]] {offsets = [0], strides = [1]} : vector<4xf32> into vector<8xf32> 558// MATRIX: %[[a2:.*]] = vector.extract %[[A]][1] : vector<2x4xf32> 559// MATRIX: %[[a3:.*]] = vector.insert_strided_slice %[[a2]], %[[a1]] {offsets = [4], strides = [1]} : vector<4xf32> into vector<8xf32> 560// MATRIX: %[[b0:.*]] = vector.extract %[[B]][0] : vector<4x3xf32> 561// MATRIX: %[[b1:.*]] = vector.insert_strided_slice %[[b0]], %[[vcst_0]] {offsets = [0], strides = [1]} : vector<3xf32> into vector<12xf32> 562// MATRIX: %[[b2:.*]] = vector.extract %[[B]][1] : vector<4x3xf32> 563// MATRIX: %[[b3:.*]] = vector.insert_strided_slice %[[b2]], %[[b1]] {offsets = [3], strides = [1]} : vector<3xf32> into vector<12xf32> 564// MATRIX: %[[b4:.*]] = vector.extract %[[B]][2] : vector<4x3xf32> 565// MATRIX: %[[b5:.*]] = vector.insert_strided_slice %[[b4]], %[[b3]] {offsets = [6], strides = [1]} : vector<3xf32> into vector<12xf32> 566// MATRIX: %[[b6:.*]] = vector.extract %[[B]][3] : vector<4x3xf32> 567// MATRIX: %[[b7:.*]] = vector.insert_strided_slice %[[b6]], %[[b5]] {offsets = [9], strides = [1]} : vector<3xf32> into vector<12xf32> 568// MATRIX: %[[mm1:.*]] = vector.matrix_multiply %[[a3]], %[[b7]] {lhs_columns = 4 : i32, lhs_rows = 2 : i32, rhs_columns = 3 : i32} : (vector<8xf32>, vector<12xf32>) -> vector<6xf32> 569// MATRIX: %[[mm2:.*]] = vector.extract_strided_slice %[[mm1]] {offsets = [0], sizes = [3], strides = [1]} : vector<6xf32> to vector<3xf32> 570// MATRIX: %[[mm3:.*]] = vector.insert %[[mm2]], %[[vcst_1]] [0] : vector<3xf32> into vector<2x3xf32> 571// MATRIX: %[[mm4:.*]] = vector.extract_strided_slice %[[mm1]] {offsets = [3], sizes = [3], strides = [1]} : vector<6xf32> to vector<3xf32> 572// MATRIX: %[[mm5:.*]] = vector.insert %[[mm4]], %[[mm3]] [1] : vector<3xf32> into vector<2x3xf32> 573// MATRIX: %[[mm6:.*]] = addf %[[C]], %[[mm5]] : vector<2x3xf32> 574 575// OUTERPRODUCT-LABEL: func @matmul 576// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x4xf32>, 577// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<4x3xf32>, 578// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 579// OUTERPRODUCT: %[[At:.*]] = vector.transpose %[[A]], [1, 0] 580// OUTERPRODUCT-SAME: : vector<2x4xf32> to vector<4x2xf32> 581// 582// OUTERPRODUCT: %[[a0:.*]] = vector.extract %[[At]][0] : vector<4x2xf32> 583// OUTERPRODUCT: %[[b0:.*]] = vector.extract %[[B]][0] : vector<4x3xf32> 584// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[a0]], %[[b0]], %[[C]] 585// OUTERPRODUCT-SAME: : vector<2xf32>, vector<3xf32> 586// 587// OUTERPRODUCT: %[[a1:.*]] = vector.extract %[[At]][1] : vector<4x2xf32> 588// OUTERPRODUCT: %[[b1:.*]] = vector.extract %[[B]][1] : vector<4x3xf32> 589// OUTERPRODUCT: %[[c1:.*]] = vector.outerproduct %[[a1]], %[[b1]], %[[c0]] 590// OUTERPRODUCT-SAME: : vector<2xf32>, vector<3xf32> 591// 592// OUTERPRODUCT: %[[a2:.*]] = vector.extract %[[At]][2] : vector<4x2xf32> 593// OUTERPRODUCT: %[[b2:.*]] = vector.extract %[[B]][2] : vector<4x3xf32> 594// OUTERPRODUCT: %[[c2:.*]] = vector.outerproduct %[[a2]], %[[b2]], %[[c1]] 595// OUTERPRODUCT-SAME: : vector<2xf32>, vector<3xf32> 596// 597// OUTERPRODUCT: %[[a3:.*]] = vector.extract %[[At]][3] : vector<4x2xf32> 598// OUTERPRODUCT: %[[b3:.*]] = vector.extract %[[B]][3] : vector<4x3xf32> 599// OUTERPRODUCT: %[[c3:.*]] = vector.outerproduct %[[a3]], %[[b3]], %[[c2]] 600// OUTERPRODUCT-SAME: : vector<2xf32>, vector<3xf32> 601// 602// OUTERPRODUCT: return %[[c3]] : vector<2x3xf32> 603 604// REDUCE-LABEL: func @matmul 605// REDUCE-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x4xf32>, 606// REDUCE-SAME: %[[B:[a-zA-Z0-9]*]]: vector<4x3xf32>, 607// REDUCE-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 608// 609// REDUCE: %[[RES:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 610// REDUCE: %[[Bt:.*]] = vector.transpose %[[B]], [1, 0] 611// REDUCE-SAME: : vector<4x3f32> to vector<3x4xf32> 612// 613// REDUCE: %[[a0:.*]] = vector.extract %[[A]][0] : vector<2x4xf32> 614// REDUCE-NEXT: %[[b0:.*]] = vector.extract %[[Bt]][0] : vector<3x4xf32> 615// REDUCE-NEXT: %[[ab00:.*]] = mul %[[a0]], %[[b0]] : vector<4xf32> 616// REDUCE-NEXT: %[[s00:.*]] = vector.reduction "add", %[[ab00]] : vector<4xf32> into f32 617// REDUCE-NEXT: %[[r00:.*]] = vector.insert %[[s00]], %[[RES]] [0, 0] : f32 into vector<2x3xf32> 618// 619// ... 620// 621// REDUCE: %[[a1:.*]] = vector.extract %[[A]][1] : vector<2x4xf32> 622// REDUCE-NEXT: %[[b2:.*]] = vector.extract %[[Bt]][2] : vector<3x4xf32> 623// REDUCE-NEXT: %[[ab12:.*]] = mul %[[a1]], %[[b02]] : vector<4xf32> 624// REDUCE-NEXT: %[[s12:.*]] = vector.reduction "add", %[[ab12]] : vector<4xf32> into f32 625// REDUCE-NEXT: %[[r12:.*]] = vector.insert %[[s12]], %{{.*}} [1, 2] : f32 into vector<2x3xf32> 626// 627// REDUCE: return %[[c3]] : vector<2x3xf32> 628func @matmul(%arg0: vector<2x4xf32>, 629 %arg1: vector<4x3xf32>, 630 %arg2: vector<2x3xf32>) -> vector<2x3xf32> { 631 %0 = vector.contract #matmat_trait %arg0, %arg1, %arg2 632 : vector<2x4xf32>, vector<4x3xf32> into vector<2x3xf32> 633 return %0 : vector<2x3xf32> 634} 635 636// CHECK-LABEL: func @broadcast_vec1d_from_scalar 637// CHECK-SAME: %[[A:.*0]]: f32 638// CHECK: %[[T0:.*]] = splat %[[A]] : vector<2xf32> 639// CHECK: return %[[T0]] : vector<2xf32> 640 641func @broadcast_vec1d_from_scalar(%arg0: f32) -> vector<2xf32> { 642 %0 = vector.broadcast %arg0 : f32 to vector<2xf32> 643 return %0 : vector<2xf32> 644} 645 646// CHECK-LABEL: func @broadcast_vec2d_from_scalar 647// CHECK-SAME: %[[A:.*0]]: f32 648// CHECK: %[[T0:.*]] = splat %[[A]] : vector<2x3xf32> 649// CHECK: return %[[T0]] : vector<2x3xf32> 650 651func @broadcast_vec2d_from_scalar(%arg0: f32) -> vector<2x3xf32> { 652 %0 = vector.broadcast %arg0 : f32 to vector<2x3xf32> 653 return %0 : vector<2x3xf32> 654} 655 656// CHECK-LABEL: func @broadcast_vec3d_from_scalar 657// CHECK-SAME: %[[A:.*0]]: f32 658// CHECK: %[[T0:.*]] = splat %[[A]] : vector<2x3x4xf32> 659// CHECK: return %[[T0]] : vector<2x3x4xf32> 660 661func @broadcast_vec3d_from_scalar(%arg0: f32) -> vector<2x3x4xf32> { 662 %0 = vector.broadcast %arg0 : f32 to vector<2x3x4xf32> 663 return %0 : vector<2x3x4xf32> 664} 665 666// CHECK-LABEL: func @broadcast_vec1d_from_vec1d 667// CHECK-SAME: %[[A:.*0]]: vector<2xf32> 668// CHECK: return %[[A]] : vector<2xf32> 669 670func @broadcast_vec1d_from_vec1d(%arg0: vector<2xf32>) -> vector<2xf32> { 671 %0 = vector.broadcast %arg0 : vector<2xf32> to vector<2xf32> 672 return %0 : vector<2xf32> 673} 674 675// CHECK-LABEL: func @broadcast_vec2d_from_vec1d 676// CHECK-SAME: %[[A:.*0]]: vector<2xf32> 677// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<3x2xf32> 678// CHECK: %[[T0:.*]] = vector.insert %[[A]], %[[C0]] [0] : vector<2xf32> into vector<3x2xf32> 679// CHECK: %[[T1:.*]] = vector.insert %[[A]], %[[T0]] [1] : vector<2xf32> into vector<3x2xf32> 680// CHECK: %[[T2:.*]] = vector.insert %[[A]], %[[T1]] [2] : vector<2xf32> into vector<3x2xf32> 681// CHECK: return %[[T2]] : vector<3x2xf32> 682 683func @broadcast_vec2d_from_vec1d(%arg0: vector<2xf32>) -> vector<3x2xf32> { 684 %0 = vector.broadcast %arg0 : vector<2xf32> to vector<3x2xf32> 685 return %0 : vector<3x2xf32> 686} 687 688// CHECK-LABEL: func @broadcast_vec3d_from_vec1d 689// CHECK-SAME: %[[A:.*0]]: vector<2xf32> 690// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<3x2xf32> 691// CHECK: %[[C1:.*]] = constant dense<0.000000e+00> : vector<4x3x2xf32> 692// CHECK: %[[T0:.*]] = vector.insert %[[A]], %[[C0]] [0] : vector<2xf32> into vector<3x2xf32> 693// CHECK: %[[T1:.*]] = vector.insert %[[A]], %[[T0]] [1] : vector<2xf32> into vector<3x2xf32> 694// CHECK: %[[T2:.*]] = vector.insert %[[A]], %[[T1]] [2] : vector<2xf32> into vector<3x2xf32> 695// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[C1]] [0] : vector<3x2xf32> into vector<4x3x2xf32> 696// CHECK: %[[T4:.*]] = vector.insert %[[T2]], %[[T3]] [1] : vector<3x2xf32> into vector<4x3x2xf32> 697// CHECK: %[[T5:.*]] = vector.insert %[[T2]], %[[T4]] [2] : vector<3x2xf32> into vector<4x3x2xf32> 698// CHECK: %[[T6:.*]] = vector.insert %[[T2]], %[[T5]] [3] : vector<3x2xf32> into vector<4x3x2xf32> 699// CHECK: return %[[T6]] : vector<4x3x2xf32> 700 701func @broadcast_vec3d_from_vec1d(%arg0: vector<2xf32>) -> vector<4x3x2xf32> { 702 %0 = vector.broadcast %arg0 : vector<2xf32> to vector<4x3x2xf32> 703 return %0 : vector<4x3x2xf32> 704} 705 706// CHECK-LABEL: func @broadcast_vec3d_from_vec2d 707// CHECK-SAME: %[[A:.*0]]: vector<3x2xf32> 708// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<4x3x2xf32> 709// CHECK: %[[T0:.*]] = vector.insert %[[A]], %[[C0]] [0] : vector<3x2xf32> into vector<4x3x2xf32> 710// CHECK: %[[T1:.*]] = vector.insert %[[A]], %[[T0]] [1] : vector<3x2xf32> into vector<4x3x2xf32> 711// CHECK: %[[T2:.*]] = vector.insert %[[A]], %[[T1]] [2] : vector<3x2xf32> into vector<4x3x2xf32> 712// CHECK: %[[T3:.*]] = vector.insert %[[A]], %[[T2]] [3] : vector<3x2xf32> into vector<4x3x2xf32> 713// CHECK: return %[[T3]] : vector<4x3x2xf32> 714 715func @broadcast_vec3d_from_vec2d(%arg0: vector<3x2xf32>) -> vector<4x3x2xf32> { 716 %0 = vector.broadcast %arg0 : vector<3x2xf32> to vector<4x3x2xf32> 717 return %0 : vector<4x3x2xf32> 718} 719 720// CHECK-LABEL: func @broadcast_stretch 721// CHECK-SAME: %[[A:.*0]]: vector<1xf32> 722// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<1xf32> 723// CHECK: %[[T1:.*]] = splat %[[T0]] : vector<4xf32> 724// CHECK: return %[[T1]] : vector<4xf32> 725 726func @broadcast_stretch(%arg0: vector<1xf32>) -> vector<4xf32> { 727 %0 = vector.broadcast %arg0 : vector<1xf32> to vector<4xf32> 728 return %0 : vector<4xf32> 729} 730 731// CHECK-LABEL: func @broadcast_stretch_at_start 732// CHECK-SAME: %[[A:.*0]]: vector<1x4xf32> 733// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<3x4xf32> 734// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<1x4xf32> 735// CHECK: %[[T1:.*]] = vector.insert %[[T0]], %[[C0]] [0] : vector<4xf32> into vector<3x4xf32> 736// CHECK: %[[T2:.*]] = vector.insert %[[T0]], %[[T1]] [1] : vector<4xf32> into vector<3x4xf32> 737// CHECK: %[[T3:.*]] = vector.insert %[[T0]], %[[T2]] [2] : vector<4xf32> into vector<3x4xf32> 738// CHECK: return %[[T3]] : vector<3x4xf32> 739 740func @broadcast_stretch_at_start(%arg0: vector<1x4xf32>) -> vector<3x4xf32> { 741 %0 = vector.broadcast %arg0 : vector<1x4xf32> to vector<3x4xf32> 742 return %0 : vector<3x4xf32> 743} 744 745// CHECK-LABEL: func @broadcast_stretch_at_end 746// CHECK-SAME: %[[A:.*0]]: vector<4x1xf32> 747// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<4x3xf32> 748// CHECK: %[[T0:.*]] = vector.extract %[[A]][0, 0] : vector<4x1xf32> 749// CHECK: %[[T2:.*]] = splat %[[T0]] : vector<3xf32> 750// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[C0]] [0] : vector<3xf32> into vector<4x3xf32> 751// CHECK: %[[T4:.*]] = vector.extract %[[A]][1, 0] : vector<4x1xf32> 752// CHECK: %[[T6:.*]] = splat %[[T4]] : vector<3xf32> 753// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T3]] [1] : vector<3xf32> into vector<4x3xf32> 754// CHECK: %[[T8:.*]] = vector.extract %[[A]][2, 0] : vector<4x1xf32> 755// CHECK: %[[T10:.*]] = splat %[[T8]] : vector<3xf32> 756// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T7]] [2] : vector<3xf32> into vector<4x3xf32> 757// CHECK: %[[T12:.*]] = vector.extract %[[A]][3, 0] : vector<4x1xf32> 758// CHECK: %[[T14:.*]] = splat %[[T12]] : vector<3xf32> 759// CHECK: %[[T15:.*]] = vector.insert %[[T14]], %[[T11]] [3] : vector<3xf32> into vector<4x3xf32> 760// CHECK: return %[[T15]] : vector<4x3xf32> 761 762func @broadcast_stretch_at_end(%arg0: vector<4x1xf32>) -> vector<4x3xf32> { 763 %0 = vector.broadcast %arg0 : vector<4x1xf32> to vector<4x3xf32> 764 return %0 : vector<4x3xf32> 765} 766 767// CHECK-LABEL: func @broadcast_stretch_in_middle 768// CHECK-SAME: %[[A:.*0]]: vector<4x1x2xf32> 769// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<4x3x2xf32> 770// CHECK: %[[C1:.*]] = constant dense<0.000000e+00> : vector<3x2xf32> 771// CHECK: %[[T0:.*]] = vector.extract %[[A]][0, 0] : vector<4x1x2xf32> 772// CHECK: %[[T2:.*]] = vector.insert %[[T0]], %[[C1]] [0] : vector<2xf32> into vector<3x2xf32> 773// CHECK: %[[T3:.*]] = vector.insert %[[T0]], %[[T2]] [1] : vector<2xf32> into vector<3x2xf32> 774// CHECK: %[[T4:.*]] = vector.insert %[[T0]], %[[T3]] [2] : vector<2xf32> into vector<3x2xf32> 775// CHECK: %[[T5:.*]] = vector.insert %[[T4]], %[[C0]] [0] : vector<3x2xf32> into vector<4x3x2xf32> 776// CHECK: %[[T6:.*]] = vector.extract %[[A]][1, 0] : vector<4x1x2xf32> 777// CHECK: %[[T8:.*]] = vector.insert %[[T6]], %[[C1]] [0] : vector<2xf32> into vector<3x2xf32> 778// CHECK: %[[T9:.*]] = vector.insert %[[T6]], %[[T8]] [1] : vector<2xf32> into vector<3x2xf32> 779// CHECK: %[[T10:.*]] = vector.insert %[[T6]], %[[T9]] [2] : vector<2xf32> into vector<3x2xf32> 780// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T5]] [1] : vector<3x2xf32> into vector<4x3x2xf32> 781// CHECK: %[[T12:.*]] = vector.extract %[[A]][2, 0] : vector<4x1x2xf32> 782// CHECK: %[[T14:.*]] = vector.insert %[[T12]], %[[C1]] [0] : vector<2xf32> into vector<3x2xf32> 783// CHECK: %[[T15:.*]] = vector.insert %[[T12]], %[[T14]] [1] : vector<2xf32> into vector<3x2xf32> 784// CHECK: %[[T16:.*]] = vector.insert %[[T12]], %[[T15]] [2] : vector<2xf32> into vector<3x2xf32> 785// CHECK: %[[T17:.*]] = vector.insert %[[T16]], %[[T11]] [2] : vector<3x2xf32> into vector<4x3x2xf32> 786// CHECK: %[[T18:.*]] = vector.extract %[[A]][3, 0] : vector<4x1x2xf32> 787// CHECK: %[[T20:.*]] = vector.insert %[[T18]], %[[C1]] [0] : vector<2xf32> into vector<3x2xf32> 788// CHECK: %[[T21:.*]] = vector.insert %[[T18]], %[[T20]] [1] : vector<2xf32> into vector<3x2xf32> 789// CHECK: %[[T22:.*]] = vector.insert %[[T18]], %[[T21]] [2] : vector<2xf32> into vector<3x2xf32> 790// CHECK: %[[T23:.*]] = vector.insert %[[T22]], %[[T17]] [3] : vector<3x2xf32> into vector<4x3x2xf32> 791// CHECK: return %[[T23]] : vector<4x3x2xf32> 792 793func @broadcast_stretch_in_middle(%arg0: vector<4x1x2xf32>) -> vector<4x3x2xf32> { 794 %0 = vector.broadcast %arg0 : vector<4x1x2xf32> to vector<4x3x2xf32> 795 return %0 : vector<4x3x2xf32> 796} 797 798// CHECK-LABEL: func @genbool_1d 799// CHECK: %[[T0:.*]] = constant dense<[true, true, true, true, false, false, false, false]> : vector<8xi1> 800// CHECK: return %[[T0]] : vector<8xi1> 801 802func @genbool_1d() -> vector<8xi1> { 803 %0 = vector.constant_mask [4] : vector<8xi1> 804 return %0 : vector<8xi1> 805} 806 807// CHECK-LABEL: func @genbool_2d 808// CHECK: %[[C1:.*]] = constant dense<[true, true, false, false]> : vector<4xi1> 809// CHECK: %[[C2:.*]] = constant dense<false> : vector<4x4xi1> 810// CHECK: %[[T0:.*]] = vector.insert %[[C1]], %[[C2]] [0] : vector<4xi1> into vector<4x4xi1> 811// CHECK: %[[T1:.*]] = vector.insert %[[C1]], %[[T0]] [1] : vector<4xi1> into vector<4x4xi1> 812// CHECK: return %[[T1]] : vector<4x4xi1> 813 814func @genbool_2d() -> vector<4x4xi1> { 815 %v = vector.constant_mask [2, 2] : vector<4x4xi1> 816 return %v: vector<4x4xi1> 817} 818 819// CHECK-LABEL: func @genbool_3d 820// CHECK: %[[C1:.*]] = constant dense<[true, true, true, false]> : vector<4xi1> 821// CHECK: %[[C2:.*]] = constant dense<false> : vector<3x4xi1> 822// CHECK: %[[C3:.*]] = constant dense<false> : vector<2x3x4xi1> 823// CHECK: %[[T0:.*]] = vector.insert %[[C1]], %[[C2]] [0] : vector<4xi1> into vector<3x4xi1> 824// CHECK: %[[T1:.*]] = vector.insert %[[T0]], %[[C3]] [0] : vector<3x4xi1> into vector<2x3x4xi1> 825// CHECK: return %[[T1]] : vector<2x3x4xi1> 826 827func @genbool_3d() -> vector<2x3x4xi1> { 828 %v = vector.constant_mask [1, 1, 3] : vector<2x3x4xi1> 829 return %v: vector<2x3x4xi1> 830} 831 832// CHECK-LABEL: func @genbool_var_1d( 833// CHECK-SAME: %[[A:.*]]: index) 834// CHECK: %[[T0:.*]] = vector.create_mask %[[A]] : vector<3xi1> 835// CHECK: return %[[T0]] : vector<3xi1> 836 837func @genbool_var_1d(%arg0: index) -> vector<3xi1> { 838 %0 = vector.create_mask %arg0 : vector<3xi1> 839 return %0 : vector<3xi1> 840} 841 842// CHECK-LABEL: func @genbool_var_2d( 843// CHECK-SAME: %[[A:.*0]]: index, 844// CHECK-SAME: %[[B:.*1]]: index) 845// CHECK: %[[C1:.*]] = constant dense<false> : vector<3xi1> 846// CHECK: %[[C2:.*]] = constant dense<false> : vector<2x3xi1> 847// CHECK: %[[c0:.*]] = constant 0 : index 848// CHECK: %[[c1:.*]] = constant 1 : index 849// CHECK: %[[T0:.*]] = vector.create_mask %[[B]] : vector<3xi1> 850// CHECK: %[[T1:.*]] = cmpi slt, %[[c0]], %[[A]] : index 851// CHECK: %[[T2:.*]] = select %[[T1]], %[[T0]], %[[C1]] : vector<3xi1> 852// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[C2]] [0] : vector<3xi1> into vector<2x3xi1> 853// CHECK: %[[T4:.*]] = cmpi slt, %[[c1]], %[[A]] : index 854// CHECK: %[[T5:.*]] = select %[[T4]], %[[T0]], %[[C1]] : vector<3xi1> 855// CHECK: %[[T6:.*]] = vector.insert %[[T5]], %[[T3]] [1] : vector<3xi1> into vector<2x3xi1> 856// CHECK: return %[[T6]] : vector<2x3xi1> 857 858func @genbool_var_2d(%arg0: index, %arg1: index) -> vector<2x3xi1> { 859 %0 = vector.create_mask %arg0, %arg1 : vector<2x3xi1> 860 return %0 : vector<2x3xi1> 861} 862 863// CHECK-LABEL: func @genbool_var_3d( 864// CHECK-SAME: %[[A:.*0]]: index, 865// CHECK-SAME: %[[B:.*1]]: index, 866// CHECK-SAME: %[[C:.*2]]: index) 867// CHECK: %[[C1:.*]] = constant dense<false> : vector<7xi1> 868// CHECK: %[[C2:.*]] = constant dense<false> : vector<1x7xi1> 869// CHECK: %[[C3:.*]] = constant dense<false> : vector<2x1x7xi1> 870// CHECK: %[[c0:.*]] = constant 0 : index 871// CHECK: %[[c1:.*]] = constant 1 : index 872// CHECK: %[[T0:.*]] = vector.create_mask %[[C]] : vector<7xi1> 873// CHECK: %[[T1:.*]] = cmpi slt, %[[c0]], %[[B]] : index 874// CHECK: %[[T2:.*]] = select %[[T1]], %[[T0]], %[[C1]] : vector<7xi1> 875// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[C2]] [0] : vector<7xi1> into vector<1x7xi1> 876// CHECK: %[[T4:.*]] = cmpi slt, %[[c0]], %[[A]] : index 877// CHECK: %[[T5:.*]] = select %[[T4]], %[[T3]], %[[C2]] : vector<1x7xi1> 878// CHECK: %[[T6:.*]] = vector.insert %[[T5]], %[[C3]] [0] : vector<1x7xi1> into vector<2x1x7xi1> 879// CHECK: %[[T7:.*]] = cmpi slt, %[[c1]], %[[A]] : index 880// CHECK: %[[T8:.*]] = select %[[T7]], %[[T3]], %[[C2]] : vector<1x7xi1> 881// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T6]] [1] : vector<1x7xi1> into vector<2x1x7xi1> 882// CHECK: return %[[T9]] : vector<2x1x7xi1> 883 884func @genbool_var_3d(%arg0: index, %arg1: index, %arg2: index) -> vector<2x1x7xi1> { 885 %0 = vector.create_mask %arg0, %arg1, %arg2 : vector<2x1x7xi1> 886 return %0 : vector<2x1x7xi1> 887} 888 889#matmat_accesses_0 = [ 890 affine_map<(m, n, k) -> (m, k)>, 891 affine_map<(m, n, k) -> (k, n)>, 892 affine_map<(m, n, k) -> (m, n)> 893] 894#matmat_trait_0 = { 895 indexing_maps = #matmat_accesses_0, 896 iterator_types = ["parallel", "parallel", "reduction"] 897} 898 899// OUTERPRODUCT-LABEL: func @matmul_0 900// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x1xf32>, 901// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<1x3xf32>, 902// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 903// OUTERPRODUCT: %[[At:.*]] = vector.transpose %[[A]], [1, 0] 904// OUTERPRODUCT: %[[a0:.*]] = vector.extract %[[At]][0] : vector<1x2xf32> 905// OUTERPRODUCT: %[[b0:.*]] = vector.extract %[[B]][0] : vector<1x3xf32> 906// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[a0]], %[[b0]], %[[C]] 907// OUTERPRODUCT: return %[[c0]] : vector<2x3xf32> 908func @matmul_0(%arg0: vector<2x1xf32>, %arg1: vector<1x3xf32>, %arg2: vector<2x3xf32>) 909-> vector<2x3xf32> 910{ 911 %0 = vector.contract #matmat_trait_0 %arg0, %arg1, %arg2 912 : vector<2x1xf32>, vector<1x3xf32> into vector<2x3xf32> 913 return %0 : vector<2x3xf32> 914} 915 916#matmat_accesses_1 = [ 917 affine_map<(m, n, k) -> (m, k)>, 918 affine_map<(m, n, k) -> (n, k)>, 919 affine_map<(m, n, k) -> (m, n)> 920] 921#matmat_trait_1 = { 922 indexing_maps = #matmat_accesses_1, 923 iterator_types = ["parallel", "parallel", "reduction"] 924} 925 926// OUTERPRODUCT-LABEL: func @matmul_1 927// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x1xf32>, 928// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<3x1xf32>, 929// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 930// OUTERPRODUCT: %[[At:.*]] = vector.transpose %[[A]], [1, 0] 931// OUTERPRODUCT: %[[Bt:.*]] = vector.transpose %[[B]], [1, 0] 932// OUTERPRODUCT: %[[a0:.*]] = vector.extract %[[At]][0] : vector<1x2xf32> 933// OUTERPRODUCT: %[[b0:.*]] = vector.extract %[[Bt]][0] : vector<1x3xf32> 934// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[a0]], %[[b0]], %[[C]] 935// OUTERPRODUCT: return %[[c0]] : vector<2x3xf32> 936func @matmul_1(%arg0: vector<2x1xf32>, %arg1: vector<3x1xf32>, %arg2: vector<2x3xf32>) 937-> vector<2x3xf32> 938{ 939 %0 = vector.contract #matmat_trait_1 %arg0, %arg1, %arg2 940 : vector<2x1xf32>, vector<3x1xf32> into vector<2x3xf32> 941 return %0 : vector<2x3xf32> 942} 943 944#matmat_accesses_2 = [ 945 affine_map<(m, n, k) -> (k, m)>, 946 affine_map<(m, n, k) -> (k, n)>, 947 affine_map<(m, n, k) -> (m, n)> 948] 949#matmat_trait_2 = { 950 indexing_maps = #matmat_accesses_2, 951 iterator_types = ["parallel", "parallel", "reduction"] 952} 953 954// OUTERPRODUCT-LABEL: func @matmul_2 955// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<1x2xf32>, 956// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<1x3xf32>, 957// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 958// OUTERPRODUCT: %[[a0:.*]] = vector.extract %[[A]][0] : vector<1x2xf32> 959// OUTERPRODUCT: %[[b0:.*]] = vector.extract %[[B]][0] : vector<1x3xf32> 960// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[a0]], %[[b0]], %[[C]] 961// OUTERPRODUCT: return %[[c0]] : vector<2x3xf32> 962func @matmul_2(%arg0: vector<1x2xf32>, %arg1: vector<1x3xf32>, %arg2: vector<2x3xf32>) 963-> vector<2x3xf32> 964{ 965 %0 = vector.contract #matmat_trait_2 %arg0, %arg1, %arg2 966 : vector<1x2xf32>, vector<1x3xf32> into vector<2x3xf32> 967 return %0 : vector<2x3xf32> 968} 969 970#matmat_accesses_3 = [ 971 affine_map<(m, n, k) -> (k, m)>, 972 affine_map<(m, n, k) -> (n, k)>, 973 affine_map<(m, n, k) -> (m, n)> 974] 975#matmat_trait_3 = { 976 indexing_maps = #matmat_accesses_3, 977 iterator_types = ["parallel", "parallel", "reduction"] 978} 979 980// OUTERPRODUCT-LABEL: func @matmul_3 981// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<1x2xf32>, 982// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<3x1xf32>, 983// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 984// OUTERPRODUCT: %[[Bt:.*]] = vector.transpose %[[B]], [1, 0] 985// OUTERPRODUCT: %[[a0:.*]] = vector.extract %[[A]][0] : vector<1x2xf32> 986// OUTERPRODUCT: %[[b0:.*]] = vector.extract %[[Bt]][0] : vector<1x3xf32> 987// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[a0]], %[[b0]], %[[C]] 988// OUTERPRODUCT: return %[[c0]] : vector<2x3xf32> 989func @matmul_3(%arg0: vector<1x2xf32>, %arg1: vector<3x1xf32>, %arg2: vector<2x3xf32>) 990-> vector<2x3xf32> 991{ 992 %0 = vector.contract #matmat_trait_3 %arg0, %arg1, %arg2 993 : vector<1x2xf32>, vector<3x1xf32> into vector<2x3xf32> 994 return %0 : vector<2x3xf32> 995} 996 997#matmat_accesses_4 = [ 998 affine_map<(m, n, k) -> (m, k)>, 999 affine_map<(m, n, k) -> (k, n)>, 1000 affine_map<(m, n, k) -> (n, m)> 1001] 1002#matmat_trait_4 = { 1003 indexing_maps = #matmat_accesses_4, 1004 iterator_types = ["parallel", "parallel", "reduction"] 1005} 1006 1007// OUTERPRODUCT-LABEL: func @matmul_4 1008// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x1xf32>, 1009// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<1x3xf32>, 1010// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<3x2xf32> 1011// OUTERPRODUCT: %[[At:.*]] = vector.transpose %[[A]], [1, 0] 1012// OUTERPRODUCT: %[[b0:.*]] = vector.extract %[[B]][0] : vector<1x3xf32> 1013// OUTERPRODUCT: %[[a0:.*]] = vector.extract %[[At]][0] : vector<1x2xf32> 1014// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[b0]], %[[a0]], %[[C]] 1015// OUTERPRODUCT: return %[[c0]] : vector<3x2xf32> 1016func @matmul_4(%arg0: vector<2x1xf32>, %arg1: vector<1x3xf32>, %arg2: vector<3x2xf32>) 1017-> vector<3x2xf32> 1018{ 1019 %0 = vector.contract #matmat_trait_4 %arg0, %arg1, %arg2 1020 : vector<2x1xf32>, vector<1x3xf32> into vector<3x2xf32> 1021 return %0 : vector<3x2xf32> 1022} 1023 1024#matmat_accesses_5 = [ 1025 affine_map<(m, n, k) -> (m, k)>, 1026 affine_map<(m, n, k) -> (k, n)>, 1027 affine_map<(m, n, k) -> (n, m)> 1028] 1029#matmat_trait_5 = { 1030 indexing_maps = #matmat_accesses_5, 1031 iterator_types = ["parallel", "parallel", "reduction"] 1032} 1033 1034// OUTERPRODUCT-LABEL: func @matmul_5 1035// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x1xf32>, 1036// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<1x3xf32>, 1037// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<3x2xf32> 1038// OUTERPRODUCT: %[[At:.*]] = vector.transpose %[[A]], [1, 0] 1039// OUTERPRODUCT-DAG: %[[a0:.*]] = vector.extract %[[At]][0] : vector<1x2xf32> 1040// OUTERPRODUCT-DAG: %[[b0:.*]] = vector.extract %[[B]][0] : vector<1x3xf32> 1041// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[b0]], %[[a0]], %[[C]] 1042// OUTERPRODUCT: return %[[c0]] : vector<3x2xf32> 1043func @matmul_5(%arg0: vector<2x1xf32>, %arg1: vector<1x3xf32>, %arg2: vector<3x2xf32>) 1044-> vector<3x2xf32> 1045{ 1046 %0 = vector.contract #matmat_trait_5 %arg0, %arg1, %arg2 1047 : vector<2x1xf32>, vector<1x3xf32> into vector<3x2xf32> 1048 return %0 : vector<3x2xf32> 1049} 1050 1051#matmat_accesses_6 = [ 1052 affine_map<(m, n, k) -> (m, k)>, 1053 affine_map<(m, n, k) -> (k, n)>, 1054 affine_map<(m, n, k) -> (n, m)> 1055] 1056#matmat_trait_6 = { 1057 indexing_maps = #matmat_accesses_6, 1058 iterator_types = ["parallel", "parallel", "reduction"] 1059} 1060 1061// OUTERPRODUCT-LABEL: func @matmul_6 1062// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x1xf32>, 1063// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<1x3xf32>, 1064// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<3x2xf32> 1065// OUTERPRODUCT: %[[At:.*]] = vector.transpose %[[A]], [1, 0] 1066// OUTERPRODUCT-DAG: %[[a0:.*]] = vector.extract %[[At]][0] : vector<1x2xf32> 1067// OUTERPRODUCT-DAG: %[[b0:.*]] = vector.extract %[[B]][0] : vector<1x3xf32> 1068// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[b0]], %[[a0]], %[[C]] 1069// OUTERPRODUCT: return %[[c0]] : vector<3x2xf32> 1070func @matmul_6(%arg0: vector<2x1xf32>, %arg1: vector<1x3xf32>, %arg2: vector<3x2xf32>) 1071-> vector<3x2xf32> 1072{ 1073 %0 = vector.contract #matmat_trait_6 %arg0, %arg1, %arg2 1074 : vector<2x1xf32>, vector<1x3xf32> into vector<3x2xf32> 1075 return %0 : vector<3x2xf32> 1076} 1077 1078#matmat_accesses_7 = [ 1079 affine_map<(m, n, k) -> (m, k)>, 1080 affine_map<(m, n, k) -> (k, n)>, 1081 affine_map<(m, n, k) -> (n, m)> 1082] 1083#matmat_trait_7 = { 1084 indexing_maps = #matmat_accesses_7, 1085 iterator_types = ["parallel", "parallel", "reduction"] 1086} 1087 1088// OUTERPRODUCT-LABEL: func @matmul_7 1089// OUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x1xf32>, 1090// OUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<1x3xf32>, 1091// OUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<3x2xf32> 1092// OUTERPRODUCT: %[[At:.*]] = vector.transpose %[[A]], [1, 0] 1093// OUTERPRODUCT-DAG: %[[a0:.*]] = vector.extract %[[At]][0] : vector<1x2xf32> 1094// OUTERPRODUCT-DAG: %[[b0:.*]] = vector.extract %[[B]][0] : vector<1x3xf32> 1095// OUTERPRODUCT: %[[c0:.*]] = vector.outerproduct %[[b0]], %[[a0]], %[[C]] 1096// OUTERPRODUCT: return %[[c0]] : vector<3x2xf32> 1097func @matmul_7(%arg0: vector<2x1xf32>, %arg1: vector<1x3xf32>, %arg2: vector<3x2xf32>) 1098-> vector<3x2xf32> 1099{ 1100 %0 = vector.contract #matmat_trait_7 %arg0, %arg1, %arg2 1101 : vector<2x1xf32>, vector<1x3xf32> into vector<3x2xf32> 1102 return %0 : vector<3x2xf32> 1103} 1104 1105// FILTEROUTERPRODUCT-LABEL: func @matmul_4_filtered 1106// FILTEROUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<4x4xf32>, 1107// FILTEROUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<4x4xf32>, 1108// FILTEROUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<4x4xf32> 1109// FILTEROUTERPRODUCT: %[[c0:.*]] = vector.contract {{{.*}}} %[[A]], %[[B]], %[[C]] 1110func @matmul_4_filtered(%arg0: vector<4x4xf32>, %arg1: vector<4x4xf32>, %arg2: vector<4x4xf32>) 1111-> vector<4x4xf32> 1112{ 1113 %0 = vector.contract #matmat_trait_0 %arg0, %arg1, %arg2 1114 : vector<4x4xf32>, vector<4x4xf32> into vector<4x4xf32> 1115 return %0 : vector<4x4xf32> 1116} 1117 1118// FILTEROUTERPRODUCT-LABEL: func @matmul_4_not_filtered 1119// FILTEROUTERPRODUCT-SAME: %[[A:[a-zA-Z0-9]*]]: vector<3x4xf32>, 1120// FILTEROUTERPRODUCT-SAME: %[[B:[a-zA-Z0-9]*]]: vector<4x4xf32>, 1121// FILTEROUTERPRODUCT-SAME: %[[C:[a-zA-Z0-9]*]]: vector<3x4xf32> 1122// FILTEROUTERPRODUCT: %[[c0:.*]] = vector.contract {{{.*}}} %[[A]], %[[B]], %[[C]] 1123func @matmul_4_not_filtered(%arg0: vector<3x4xf32>, %arg1: vector<4x4xf32>, %arg2: vector<3x4xf32>) 1124-> vector<3x4xf32> 1125{ 1126 %0 = vector.contract #matmat_trait_0 %arg0, %arg1, %arg2 1127 : vector<3x4xf32>, vector<4x4xf32> into vector<3x4xf32> 1128 return %0 : vector<3x4xf32> 1129} 1130