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 4#dotp_accesses = [ 5 affine_map<(i) -> (i)>, 6 affine_map<(i) -> (i)>, 7 affine_map<(i) -> ()> 8] 9#dotp_trait = { 10 indexing_maps = #dotp_accesses, 11 iterator_types = ["reduction"] 12} 13 14// CHECK-LABEL: func @extract_contract1 15// CHECK-SAME: %[[A:.*0]]: vector<4xf32>, 16// CHECK-SAME: %[[B:.*1]]: vector<4xf32>, 17// CHECK-SAME: %[[C:.*2]]: f32 18// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<4xf32> 19// CHECK: %[[F:.*]] = vector.fma %[[A]], %[[B]], %[[Z]] : vector<4xf32> 20// CHECK: %[[R:.*]] = vector.reduction "add", %[[F]], %[[C]] : vector<4xf32> into f32 21// CHECK: return %[[R]] : f32 22 23func @extract_contract1(%arg0: vector<4xf32>, %arg1: vector<4xf32>, %arg2: f32) -> f32 { 24 %0 = vector.contract #dotp_trait %arg0, %arg1, %arg2 25 : vector<4xf32>, vector<4xf32> into f32 26 return %0 : f32 27} 28 29#matvec_accesses = [ 30 affine_map<(i, j) -> (i, j)>, 31 affine_map<(i, j) -> (j)>, 32 affine_map<(i, j) -> (i)> 33] 34#matvec_trait = { 35 indexing_maps = #matvec_accesses, 36 iterator_types = ["parallel", "reduction"] 37} 38 39// CHECK-LABEL: func @extract_contract2 40// CHECK-SAME: %[[A:.*0]]: vector<2x3xf32>, 41// CHECK-SAME: %[[B:.*1]]: vector<3xf32>, 42// CHECK-SAME: %[[C:.*2]]: vector<2xf32> 43// CHECK: %[[R:.*]] = constant dense<0.000000e+00> : vector<2xf32> 44// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<3xf32> 45// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x3xf32> 46// CHECK: %[[T1:.*]] = vector.extract %[[C]][0] : vector<2xf32> 47// CHECK: %[[T2:.*]] = vector.fma %[[T0]], %[[B]], %[[Z]] : vector<3xf32> 48// CHECK: %[[T3:.*]] = vector.reduction "add", %[[T2]], %[[T1]] : vector<3xf32> into f32 49// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[R]] [0] : f32 into vector<2xf32> 50// CHECK: %[[T5:.*]] = vector.extract %[[A]][1] : vector<2x3xf32> 51// CHECK: %[[T6:.*]] = vector.extract %[[C]][1] : vector<2xf32> 52// CHECK: %[[T7:.*]] = vector.fma %[[T5]], %[[B]], %[[Z]] : vector<3xf32> 53// CHECK: %[[T8:.*]] = vector.reduction "add", %[[T7]], %[[T6]] : vector<3xf32> into f32 54// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T4]] [1] : f32 into vector<2xf32> 55// CHECK: return %[[T9]] : vector<2xf32> 56 57func @extract_contract2(%arg0: vector<2x3xf32>, 58 %arg1: vector<3xf32>, 59 %arg2: vector<2xf32>) -> vector<2xf32> { 60 %0 = vector.contract #matvec_trait %arg0, %arg1, %arg2 61 : vector<2x3xf32>, vector<3xf32> into vector<2xf32> 62 return %0 : vector<2xf32> 63} 64 65#vecmat_accesses = [ 66 affine_map<(i, j) -> (j)>, 67 affine_map<(i, j) -> (i, j)>, 68 affine_map<(i, j) -> (i)> 69] 70#vecmat_trait = { 71 indexing_maps = #vecmat_accesses, 72 iterator_types = ["parallel", "reduction"] 73} 74 75// CHECK-LABEL: func @extract_contract3 76// CHECK-SAME: %[[A:.*0]]: vector<3xf32>, 77// CHECK-SAME: %[[B:.*1]]: vector<2x3xf32>, 78// CHECK-SAME: %[[C:.*2]]: vector<2xf32> 79// CHECK: %[[R:.*]] = constant dense<0.000000e+00> : vector<2xf32> 80// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<3xf32> 81// CHECK: %[[T0:.*]] = vector.extract %[[B]][0] : vector<2x3xf32> 82// CHECK: %[[T1:.*]] = vector.extract %[[C]][0] : vector<2xf32> 83// CHECK: %[[T2:.*]] = vector.fma %[[A]], %[[T0]], %[[Z]] : vector<3xf32> 84// CHECK: %[[T3:.*]] = vector.reduction "add", %[[T2]], %[[T1]] : vector<3xf32> into f32 85// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[R]] [0] : f32 into vector<2xf32> 86// CHECK: %[[T5:.*]] = vector.extract %[[B]][1] : vector<2x3xf32> 87// CHECK: %[[T6:.*]] = vector.extract %[[C]][1] : vector<2xf32> 88// CHECK: %[[T7:.*]] = vector.fma %[[A]], %[[T5]], %[[Z]] : vector<3xf32> 89// CHECK: %[[T8:.*]] = vector.reduction "add", %[[T7]], %[[T6]] : vector<3xf32> into f32 90// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T4]] [1] : f32 into vector<2xf32> 91// CHECK: return %[[T9]] : vector<2xf32> 92 93func @extract_contract3(%arg0: vector<3xf32>, 94 %arg1: vector<2x3xf32>, 95 %arg2: vector<2xf32>) -> vector<2xf32> { 96 %0 = vector.contract #vecmat_trait %arg0, %arg1, %arg2 97 : vector<3xf32>, vector<2x3xf32> into vector<2xf32> 98 return %0 : vector<2xf32> 99} 100 101#matmat_accesses = [ 102 affine_map<(i, j, k) -> (i, k)>, 103 affine_map<(i, j, k) -> (k, j)>, 104 affine_map<(i, j, k) -> (i, j)> 105] 106#matmat_trait = { 107 indexing_maps = #matmat_accesses, 108 iterator_types = ["parallel", "parallel", "reduction"] 109} 110 111// CHECK-LABEL: func @extract_contract4 112// CHECK-SAME: %[[A:.*0]]: vector<2x2xf32>, 113// CHECK-SAME: %[[B:.*1]]: vector<2x2xf32>, 114// CHECK-SAME: %[[C:.*2]]: vector<2x2xf32> 115// CHECK: %[[R:.*]] = constant dense<0.000000e+00> : vector<2x2xf32> 116// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<2xf32> 117// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x2xf32> 118// CHECK: %[[T1:.*]] = vector.extract %[[C]][0] : vector<2x2xf32> 119// CHECK: %[[T2:.*]] = vector.extract %[[B]][0] : vector<2x2xf32> 120// CHECK: %[[T3:.*]] = vector.extract %[[T2]][0] : vector<2xf32> 121// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[Z]] [0] : f32 into vector<2xf32> 122// CHECK: %[[T5:.*]] = vector.extract %[[B]][1] : vector<2x2xf32> 123// CHECK: %[[T6:.*]] = vector.extract %[[T5]][0] : vector<2xf32> 124// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T4]] [1] : f32 into vector<2xf32> 125// CHECK: %[[T8:.*]] = vector.extract %[[T1]][0] : vector<2xf32> 126// CHECK: %[[T9:.*]] = vector.fma %[[T0]], %[[T7]], %[[Z]] : vector<2xf32> 127// CHECK: %[[T10:.*]] = vector.reduction "add", %[[T9]], %[[T8]] : vector<2xf32> into f32 128// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[Z]] [0] : f32 into vector<2xf32> 129// CHECK: %[[T12:.*]] = vector.extract %[[B]][0] : vector<2x2xf32> 130// CHECK: %[[T13:.*]] = vector.extract %[[T12]][1] : vector<2xf32> 131// CHECK: %[[T14:.*]] = vector.insert %[[T13]], %[[Z]] [0] : f32 into vector<2xf32> 132// CHECK: %[[T15:.*]] = vector.extract %[[B]][1] : vector<2x2xf32> 133// CHECK: %[[T16:.*]] = vector.extract %[[T15]][1] : vector<2xf32> 134// CHECK: %[[T17:.*]] = vector.insert %[[T16]], %[[T14]] [1] : f32 into vector<2xf32> 135// CHECK: %[[T18:.*]] = vector.extract %[[T1]][1] : vector<2xf32> 136// CHECK: %[[T19:.*]] = vector.fma %[[T0]], %[[T17]], %[[Z]] : vector<2xf32> 137// CHECK: %[[T20:.*]] = vector.reduction "add", %[[T19]], %[[T18]] : vector<2xf32> into f32 138// CHECK: %[[T21:.*]] = vector.insert %[[T20]], %[[T11]] [1] : f32 into vector<2xf32> 139// CHECK: %[[T22:.*]] = vector.insert %[[T21]], %[[R]] [0] : vector<2xf32> into vector<2x2xf32> 140// CHECK: %[[T23:.*]] = vector.extract %[[A]][1] : vector<2x2xf32> 141// CHECK: %[[T24:.*]] = vector.extract %[[C]][1] : vector<2x2xf32> 142// CHECK: %[[T25:.*]] = vector.extract %[[B]][0] : vector<2x2xf32> 143// CHECK: %[[T26:.*]] = vector.extract %[[T25]][0] : vector<2xf32> 144// CHECK: %[[T27:.*]] = vector.insert %[[T26]], %[[Z]] [0] : f32 into vector<2xf32> 145// CHECK: %[[T28:.*]] = vector.extract %[[B]][1] : vector<2x2xf32> 146// CHECK: %[[T29:.*]] = vector.extract %[[T28]][0] : vector<2xf32> 147// CHECK: %[[T30:.*]] = vector.insert %[[T29]], %[[T27]] [1] : f32 into vector<2xf32> 148// CHECK: %[[T31:.*]] = vector.extract %[[T24]][0] : vector<2xf32> 149// CHECK: %[[T32:.*]] = vector.fma %[[T23]], %[[T30]], %[[Z]] : vector<2xf32> 150// CHECK: %[[T33:.*]] = vector.reduction "add", %[[T32]], %[[T31]] : vector<2xf32> into f32 151// CHECK: %[[T34:.*]] = vector.insert %[[T33]], %[[Z]] [0] : f32 into vector<2xf32> 152// CHECK: %[[T35:.*]] = vector.extract %[[B]][0] : vector<2x2xf32> 153// CHECK: %[[T36:.*]] = vector.extract %[[T35]][1] : vector<2xf32> 154// CHECK: %[[T37:.*]] = vector.insert %[[T36]], %[[Z]] [0] : f32 into vector<2xf32> 155// CHECK: %[[T38:.*]] = vector.extract %[[B]][1] : vector<2x2xf32> 156// CHECK: %[[T39:.*]] = vector.extract %[[T38]][1] : vector<2xf32> 157// CHECK: %[[T40:.*]] = vector.insert %[[T39]], %[[T37]] [1] : f32 into vector<2xf32> 158// CHECK: %[[T41:.*]] = vector.extract %[[T24]][1] : vector<2xf32> 159// CHECK: %[[T42:.*]] = vector.fma %[[T23]], %[[T40]], %[[Z]] : vector<2xf32> 160// CHECK: %[[T43:.*]] = vector.reduction "add", %[[T42]], %[[T41]] : vector<2xf32> into f32 161// CHECK: %[[T44:.*]] = vector.insert %[[T43]], %[[T34]] [1] : f32 into vector<2xf32> 162// CHECK: %[[T45:.*]] = vector.insert %[[T44]], %[[T22]] [1] : vector<2xf32> into vector<2x2xf32> 163// CHECK: return %[[T45]] : vector<2x2xf32> 164 165func @extract_contract4(%arg0: vector<2x2xf32>, 166 %arg1: vector<2x2xf32>, 167 %arg2: vector<2x2xf32>) -> vector<2x2xf32> { 168 %0 = vector.contract #matmat_trait %arg0, %arg1, %arg2 169 : vector<2x2xf32>, vector<2x2xf32> into vector<2x2xf32> 170 return %0 : vector<2x2xf32> 171} 172 173#contraction2d_accesses = [ 174 affine_map<(i, j) -> (i, j)>, 175 affine_map<(i, j) -> (i, j)>, 176 affine_map<(i, j) -> ()> 177] 178#contraction2d_trait = { 179 indexing_maps = #contraction2d_accesses, 180 iterator_types = ["reduction", "reduction"] 181} 182 183// CHECK-LABEL: func @full_contract1 184// CHECK-SAME: %[[A:.*0]]: vector<2x3xf32>, 185// CHECK-SAME: %[[B:.*1]]: vector<2x3xf32>, 186// CHECK-SAME: %[[C:.*2]]: f32 187// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<3xf32> 188// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x3xf32> 189// CHECK: %[[T1:.*]] = vector.extract %[[B]][0] : vector<2x3xf32> 190// CHECK: %[[T2:.*]] = vector.fma %[[T0]], %[[T1]], %[[Z]] : vector<3xf32> 191// CHECK: %[[T3:.*]] = vector.reduction "add", %[[T2]], %[[C]] : vector<3xf32> into f32 192// CHECK: %[[T4:.*]] = vector.extract %[[A]][1] : vector<2x3xf32> 193// CHECK: %[[T5:.*]] = vector.extract %[[B]][1] : vector<2x3xf32> 194// CHECK: %[[T6:.*]] = vector.fma %[[T4]], %[[T5]], %[[Z]] : vector<3xf32> 195// CHECK: %[[T7:.*]] = vector.reduction "add", %[[T6]], %[[T3]] : vector<3xf32> into f32 196// CHECK: return %[[T7]] : f32 197 198func @full_contract1(%arg0: vector<2x3xf32>, 199 %arg1: vector<2x3xf32>, 200 %arg2: f32) -> f32 { 201 %0 = vector.contract #contraction2d_trait %arg0, %arg1, %arg2 202 : vector<2x3xf32>, vector<2x3xf32> into f32 203 return %0 : f32 204} 205 206#contraction2d_trans_accesses = [ 207 affine_map<(i, j) -> (i, j)>, 208 affine_map<(i, j) -> (j, i)>, 209 affine_map<(i, j) -> ()> 210] 211#contraction2d_trans_trait = { 212 indexing_maps = #contraction2d_trans_accesses, 213 iterator_types = ["reduction", "reduction"] 214} 215 216// CHECK-LABEL: func @full_contract2 217// CHECK-SAME: %[[A:.*0]]: vector<2x3xf32>, 218// CHECK-SAME: %[[B:.*1]]: vector<3x2xf32>, 219// CHECK-SAME: %[[C:.*2]]: f32 220// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<3xf32> 221// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2x3xf32> 222// CHECK: %[[T1:.*]] = vector.extract %[[B]][0] : vector<3x2xf32> 223// CHECK: %[[T2:.*]] = vector.extract %[[T1]][0] : vector<2xf32> 224// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[Z]] [0] : f32 into vector<3xf32> 225// CHECK: %[[T4:.*]] = vector.extract %[[B]][1] : vector<3x2xf32> 226// CHECK: %[[T5:.*]] = vector.extract %[[T4]][0] : vector<2xf32> 227// CHECK: %[[T6:.*]] = vector.insert %[[T5]], %[[T3]] [1] : f32 into vector<3xf32> 228// CHECK: %[[T7:.*]] = vector.extract %[[B]][2] : vector<3x2xf32> 229// CHECK: %[[T8:.*]] = vector.extract %[[T7]][0] : vector<2xf32> 230// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T6]] [2] : f32 into vector<3xf32> 231// CHECK: %[[T10:.*]] = vector.fma %[[T0]], %[[T9]], %[[Z]] : vector<3xf32> 232// CHECK: %[[T11:.*]] = vector.reduction "add", %[[T10]], %[[C]] : vector<3xf32> into f32 233// CHECK: %[[T12:.*]] = vector.extract %[[A]][1] : vector<2x3xf32> 234// CHECK: %[[T13:.*]] = vector.extract %[[B]][0] : vector<3x2xf32> 235// CHECK: %[[T14:.*]] = vector.extract %[[T13]][1] : vector<2xf32> 236// CHECK: %[[T15:.*]] = vector.insert %[[T14]], %[[Z]] [0] : f32 into vector<3xf32> 237// CHECK: %[[T16:.*]] = vector.extract %[[B]][1] : vector<3x2xf32> 238// CHECK: %[[T17:.*]] = vector.extract %[[T16]][1] : vector<2xf32> 239// CHECK: %[[T18:.*]] = vector.insert %[[T17]], %[[T15]] [1] : f32 into vector<3xf32> 240// CHECK: %[[T19:.*]] = vector.extract %[[B]][2] : vector<3x2xf32> 241// CHECK: %[[T20:.*]] = vector.extract %[[T19]][1] : vector<2xf32> 242// CHECK: %[[T21:.*]] = vector.insert %[[T20]], %[[T18]] [2] : f32 into vector<3xf32> 243// CHECK: %[[T22:.*]] = vector.fma %[[T12]], %[[T21]], %[[Z]] : vector<3xf32> 244// CHECK: %[[T23:.*]] = vector.reduction "add", %[[T22]], %[[T11]] : vector<3xf32> into f32 245// CHECK: return %[[T23]] : f32 246 247func @full_contract2(%arg0: vector<2x3xf32>, 248 %arg1: vector<3x2xf32>, 249 %arg2: f32) -> f32 { 250 %0 = vector.contract #contraction2d_trans_trait %arg0, %arg1, %arg2 251 : vector<2x3xf32>, vector<3x2xf32> into f32 252 return %0 : f32 253} 254 255// CHECK-LABEL: func @outerproduct_noacc 256// CHECK-SAME: %[[A:.*0]]: vector<2xf32>, 257// CHECK-SAME: %[[B:.*1]]: vector<3xf32> 258// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 259// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2xf32> 260// CHECK: %[[T1:.*]] = vector.broadcast %[[T0]] : f32 to vector<3xf32> 261// CHECK: %[[T2:.*]] = mulf %[[T1]], %[[B]] : vector<3xf32> 262// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[C0]] [0] : vector<3xf32> into vector<2x3xf32> 263// CHECK: %[[T4:.*]] = vector.extract %[[A]][1] : vector<2xf32> 264// CHECK: %[[T5:.*]] = vector.broadcast %[[T4]] : f32 to vector<3xf32> 265// CHECK: %[[T6:.*]] = mulf %[[T5]], %[[B]] : vector<3xf32> 266// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T3]] [1] : vector<3xf32> into vector<2x3xf32> 267// CHECK: return %[[T7]] : vector<2x3xf32> 268 269func @outerproduct_noacc(%arg0: vector<2xf32>, 270 %arg1: vector<3xf32>) -> vector<2x3xf32> { 271 %0 = vector.outerproduct %arg0, %arg1 : vector<2xf32>, vector<3xf32> 272 return %0: vector<2x3xf32> 273} 274 275// CHECK-LABEL: func @outerproduct_acc 276// CHECK-SAME: %[[A:.*0]]: vector<2xf32>, 277// CHECK-SAME: %[[B:.*1]]: vector<3xf32>, 278// CHECK-SAME: %[[C:.*2]]: vector<2x3xf32> 279// CHECK: %[[C0:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 280// CHECK: %[[T0:.*]] = vector.extract %[[A]][0] : vector<2xf32> 281// CHECK: %[[T1:.*]] = vector.broadcast %[[T0]] : f32 to vector<3xf32> 282// CHECK: %[[T2:.*]] = vector.extract %[[C]][0] : vector<2x3xf32> 283// CHECK: %[[T3:.*]] = vector.fma %[[T1]], %[[B]], %[[T2]] : vector<3xf32> 284// CHECK: %[[T4:.*]] = vector.insert %[[T3]], %[[C0]] [0] : vector<3xf32> into vector<2x3xf32> 285// CHECK: %[[T5:.*]] = vector.extract %[[A]][1] : vector<2xf32> 286// CHECK: %[[T6:.*]] = vector.broadcast %[[T5]] : f32 to vector<3xf32> 287// CHECK: %[[T7:.*]] = vector.extract %[[C]][1] : vector<2x3xf32> 288// CHECK: %[[T8:.*]] = vector.fma %[[T6]], %[[B]], %[[T7]] : vector<3xf32> 289// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T4]] [1] : vector<3xf32> into vector<2x3xf32> 290// CHECK: return %[[T9]] : vector<2x3xf32> 291 292func @outerproduct_acc(%arg0: vector<2xf32>, 293 %arg1: vector<3xf32>, 294 %arg2: vector<2x3xf32>) -> vector<2x3xf32> { 295 %0 = vector.outerproduct %arg0, %arg1, %arg2 : vector<2xf32>, vector<3xf32> 296 return %0: vector<2x3xf32> 297} 298 299// CHECK-LABEL: func @transpose23 300// CHECK-SAME: %[[A:.*]]: vector<2x3xf32> 301// CHECK: %[[Z:.*]] = constant dense<0.000000e+00> : vector<3x2xf32> 302// CHECK: %[[T0:.*]] = vector.extract %[[A]][0, 0] : vector<2x3xf32> 303// CHECK: %[[T1:.*]] = vector.insert %[[T0]], %[[Z]] [0, 0] : f32 into vector<3x2xf32> 304// CHECK: %[[T2:.*]] = vector.extract %[[A]][1, 0] : vector<2x3xf32> 305// CHECK: %[[T3:.*]] = vector.insert %[[T2]], %[[T1]] [0, 1] : f32 into vector<3x2xf32> 306// CHECK: %[[T4:.*]] = vector.extract %[[A]][0, 1] : vector<2x3xf32> 307// CHECK: %[[T5:.*]] = vector.insert %[[T4]], %[[T3]] [1, 0] : f32 into vector<3x2xf32> 308// CHECK: %[[T6:.*]] = vector.extract %[[A]][1, 1] : vector<2x3xf32> 309// CHECK: %[[T7:.*]] = vector.insert %[[T6]], %[[T5]] [1, 1] : f32 into vector<3x2xf32> 310// CHECK: %[[T8:.*]] = vector.extract %[[A]][0, 2] : vector<2x3xf32> 311// CHECK: %[[T9:.*]] = vector.insert %[[T8]], %[[T7]] [2, 0] : f32 into vector<3x2xf32> 312// CHECK: %[[T10:.*]] = vector.extract %[[A]][1, 2] : vector<2x3xf32> 313// CHECK: %[[T11:.*]] = vector.insert %[[T10]], %[[T9]] [2, 1] : f32 into vector<3x2xf32> 314// CHECK: return %[[T11]] : vector<3x2xf32> 315 316func @transpose23(%arg0: vector<2x3xf32>) -> vector<3x2xf32> { 317 %0 = vector.transpose %arg0, [1, 0] : vector<2x3xf32> to vector<3x2xf32> 318 return %0 : vector<3x2xf32> 319} 320 321// Shape up and downcasts for 2-D vectors, for supporting conversion to 322// llvm.matrix operations 323// CHECK-LABEL: func @shape_casts 324func @shape_casts(%a: vector<2x2xf32>) -> (vector<4xf32>, vector<2x2xf32>) { 325 // CHECK: %[[cst:.*]] = constant dense<0.000000e+00> : vector<4xf32> 326 // CHECK: %[[cst22:.*]] = constant dense<0.000000e+00> : vector<2x2xf32> 327 // CHECK: %[[ex0:.*]] = vector.extract %{{.*}}[0] : vector<2x2xf32> 328 // 329 // CHECK: %[[in0:.*]] = vector.insert_strided_slice %[[ex0]], %[[cst]] 330 // CHECK-SAME: {offsets = [0], strides = [1]} : vector<2xf32> into vector<4xf32> 331 // 332 // CHECK: %[[ex1:.*]] = vector.extract %{{.*}}[1] : vector<2x2xf32> 333 // 334 // CHECK: %[[in2:.*]] = vector.insert_strided_slice %[[ex1]], %[[in0]] 335 // CHECK-SAME: {offsets = [2], strides = [1]} : vector<2xf32> into vector<4xf32> 336 // 337 %0 = vector.shape_cast %a : vector<2x2xf32> to vector<4xf32> 338 // CHECK: %[[add:.*]] = addf %[[in2]], %[[in2]] : vector<4xf32> 339 %r0 = addf %0, %0: vector<4xf32> 340 // 341 // CHECK: %[[ss0:.*]] = vector.strided_slice %[[add]] 342 // CHECK-SAME: {offsets = [0], sizes = [2], strides = [1]} : 343 // CHECK-SAME: vector<4xf32> to vector<2xf32> 344 // 345 // CHECK: %[[res0:.*]] = vector.insert %[[ss0]], %[[cst22]] [0] : 346 // CHECK-SAME: vector<2xf32> into vector<2x2xf32> 347 // 348 // CHECK: %[[s2:.*]] = vector.strided_slice %[[add]] 349 // CHECK-SAME: {offsets = [2], sizes = [2], strides = [1]} : 350 // CHECK-SAME: vector<4xf32> to vector<2xf32> 351 // 352 // CHECK: %[[res1:.*]] = vector.insert %[[s2]], %[[res0]] [1] : 353 // CHECK-SAME: vector<2xf32> into vector<2x2xf32> 354 // 355 %1 = vector.shape_cast %r0 : vector<4xf32> to vector<2x2xf32> 356 // CHECK: return %[[add]], %[[res1]] : vector<4xf32>, vector<2x2xf32> 357 return %r0, %1 : vector<4xf32>, vector<2x2xf32> 358} 359 360// MATRIX-LABEL: func @matmul 361// MATRIX-SAME: %[[A:[a-zA-Z0-9]*]]: vector<2x4xf32>, 362// MATRIX-SAME: %[[B:[a-zA-Z0-9]*]]: vector<4x3xf32>, 363// MATRIX-SAME: %[[C:[a-zA-Z0-9]*]]: vector<2x3xf32> 364// MATRIX: %[[vcst:.*]] = constant dense<0.000000e+00> : vector<8xf32> 365// MATRIX: %[[vcst_0:.*]] = constant dense<0.000000e+00> : vector<12xf32> 366// MATRIX: %[[vcst_1:.*]] = constant dense<0.000000e+00> : vector<2x3xf32> 367// MATRIX: %[[a0:.*]] = vector.extract %[[A]][0] : vector<2x4xf32> 368// MATRIX: %[[a1:.*]] = vector.insert_strided_slice %[[a0]], %[[vcst]] {offsets = [0], strides = [1]} : vector<4xf32> into vector<8xf32> 369// MATRIX: %[[a2:.*]] = vector.extract %[[A]][1] : vector<2x4xf32> 370// MATRIX: %[[a3:.*]] = vector.insert_strided_slice %[[a2]], %[[a1]] {offsets = [4], strides = [1]} : vector<4xf32> into vector<8xf32> 371// MATRIX: %[[b0:.*]] = vector.extract %[[B]][0] : vector<4x3xf32> 372// MATRIX: %[[b1:.*]] = vector.insert_strided_slice %[[b0]], %[[vcst_0]] {offsets = [0], strides = [1]} : vector<3xf32> into vector<12xf32> 373// MATRIX: %[[b2:.*]] = vector.extract %[[B]][1] : vector<4x3xf32> 374// MATRIX: %[[b3:.*]] = vector.insert_strided_slice %[[b2]], %[[b1]] {offsets = [3], strides = [1]} : vector<3xf32> into vector<12xf32> 375// MATRIX: %[[b4:.*]] = vector.extract %[[B]][2] : vector<4x3xf32> 376// MATRIX: %[[b5:.*]] = vector.insert_strided_slice %[[b4]], %[[b3]] {offsets = [6], strides = [1]} : vector<3xf32> into vector<12xf32> 377// MATRIX: %[[b6:.*]] = vector.extract %[[B]][3] : vector<4x3xf32> 378// MATRIX: %[[b7:.*]] = vector.insert_strided_slice %[[b6]], %[[b5]] {offsets = [9], strides = [1]} : vector<3xf32> into vector<12xf32> 379// 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> 380// MATRIX: %[[mm2:.*]] = vector.strided_slice %[[mm1]] {offsets = [0], sizes = [3], strides = [1]} : vector<6xf32> to vector<3xf32> 381// MATRIX: %[[mm3:.*]] = vector.insert %[[mm2]], %[[vcst_1]] [0] : vector<3xf32> into vector<2x3xf32> 382// MATRIX: %[[mm4:.*]] = vector.strided_slice %[[mm1]] {offsets = [3], sizes = [3], strides = [1]} : vector<6xf32> to vector<3xf32> 383// MATRIX: %[[mm5:.*]] = vector.insert %[[mm4]], %[[mm3]] [1] : vector<3xf32> into vector<2x3xf32> 384// MATRIX: %[[mm6:.*]] = addf %[[C]], %[[mm5]] : vector<2x3xf32> 385func @matmul(%arg0: vector<2x4xf32>, 386 %arg1: vector<4x3xf32>, 387 %arg2: vector<2x3xf32>) -> vector<2x3xf32> { 388 %0 = vector.contract #matmat_trait %arg0, %arg1, %arg2 389 : vector<2x4xf32>, vector<4x3xf32> into vector<2x3xf32> 390 return %0 : vector<2x3xf32> 391} 392