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