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