1// RUN: mlir-opt %s -test-vector-contraction-conversion=vector-outerproduct=1 | FileCheck %s
2
3#matvec_accesses = [
4  affine_map<(i, j) -> (i, j)>,
5  affine_map<(i, j) -> (j)>,
6  affine_map<(i, j) -> (i)>
7]
8#matvec_trait = {
9  indexing_maps = #matvec_accesses,
10  iterator_types = ["parallel", "reduction"]
11}
12#matvecmax_trait = {
13  indexing_maps = #matvec_accesses,
14  iterator_types = ["parallel", "reduction"],
15  kind = #vector.kind<max>
16}
17
18#mattransvec_accesses = [
19  affine_map<(i, j) -> (j, i)>,
20  affine_map<(i, j) -> (j)>,
21  affine_map<(i, j) -> (i)>
22]
23#mattransvec_trait = {
24  indexing_maps = #mattransvec_accesses,
25  iterator_types = ["parallel", "reduction"]
26}
27
28#vecmat_accesses = [
29  affine_map<(i, j) -> (j)>,
30  affine_map<(i, j) -> (i, j)>,
31  affine_map<(i, j) -> (i)>
32]
33#vecmat_trait = {
34  indexing_maps = #vecmat_accesses,
35  iterator_types = ["parallel", "reduction"]
36}
37
38#vecmattrans_accesses = [
39  affine_map<(i, j) -> (j)>,
40  affine_map<(i, j) -> (j, i)>,
41  affine_map<(i, j) -> (i)>
42]
43#vecmattrans_trait = {
44  indexing_maps = #vecmattrans_accesses,
45  iterator_types = ["parallel", "reduction"]
46}
47
48// CHECK-LABEL: func @matvec2x2
49// CHECK-SAME: %[[A:.*0]]: memref<vector<2x2xf32>>
50// CHECK-SAME: %[[B:.*1]]: memref<vector<2xf32>>
51// CHECK-SAME: %[[C:.*2]]: memref<vector<2xf32>>
52// CHECK: %[[T0:.*]] = memref.load %[[A]][] : memref<vector<2x2xf32>>
53// CHECK: %[[T1:.*]] = memref.load %[[B]][] : memref<vector<2xf32>>
54// CHECK: %[[T2:.*]] = memref.load %[[C]][] : memref<vector<2xf32>>
55// CHECK: %[[T3:.*]] = vector.transpose %[[T0]], [1, 0] : vector<2x2xf32> to vector<2x2xf32>
56// CHECK: %[[T4:.*]] = vector.extract %[[T3]][0] : vector<2x2xf32>
57// CHECK: %[[T5:.*]] = vector.extract %[[T1]][0] : vector<2xf32>
58// CHECK: %[[T6:.*]] = vector.outerproduct %[[T4]], %[[T5]], %[[T2]] {kind = #vector.kind<add>} : vector<2xf32>, f32
59// CHECK: %[[T7:.*]] = vector.extract %[[T3]][1] : vector<2x2xf32>
60// CHECK: %[[T8:.*]] = vector.extract %[[T1]][1] : vector<2xf32>
61// CHECK: %[[T9:.*]] = vector.outerproduct %[[T7]], %[[T8]], %[[T6]] {kind = #vector.kind<add>} : vector<2xf32>, f32
62// CHECK: memref.store %[[T9]], %[[C]][] : memref<vector<2xf32>>
63// CHECK: return
64func @matvec2x2(%arg0: memref<vector<2x2xf32>>, %arg1: memref<vector<2xf32>>,
65                                                %arg2: memref<vector<2xf32>>) {
66  %A = memref.load %arg0[] : memref<vector<2x2xf32>>
67  %x = memref.load %arg1[] : memref<vector<2xf32>>
68  %b = memref.load %arg2[] : memref<vector<2xf32>>
69  %0 = vector.contract #matvec_trait %A, %x, %b : vector<2x2xf32>, vector<2xf32> into vector<2xf32>
70  memref.store %0, %arg2[] : memref<vector<2xf32>>
71  return
72}
73
74// CHECK-LABEL: func @matvecmax2x2
75// CHECK-SAME: %[[A:.*0]]: memref<vector<2x2xf32>>
76// CHECK-SAME: %[[B:.*1]]: memref<vector<2xf32>>
77// CHECK-SAME: %[[C:.*2]]: memref<vector<2xf32>>
78// CHECK: %[[T0:.*]] = memref.load %[[A]][] : memref<vector<2x2xf32>>
79// CHECK: %[[T1:.*]] = memref.load %[[B]][] : memref<vector<2xf32>>
80// CHECK: %[[T2:.*]] = memref.load %[[C]][] : memref<vector<2xf32>>
81// CHECK: %[[T3:.*]] = vector.transpose %[[T0]], [1, 0] : vector<2x2xf32> to vector<2x2xf32>
82// CHECK: %[[T4:.*]] = vector.extract %[[T3]][0] : vector<2x2xf32>
83// CHECK: %[[T5:.*]] = vector.extract %[[T1]][0] : vector<2xf32>
84// CHECK: %[[T6:.*]] = vector.outerproduct %[[T4]], %[[T5]], %[[T2]] {kind = #vector.kind<max>} : vector<2xf32>, f32
85// CHECK: %[[T7:.*]] = vector.extract %[[T3]][1] : vector<2x2xf32>
86// CHECK: %[[T8:.*]] = vector.extract %[[T1]][1] : vector<2xf32>
87// CHECK: %[[T9:.*]] = vector.outerproduct %[[T7]], %[[T8]], %[[T6]] {kind = #vector.kind<max>} : vector<2xf32>, f32
88// CHECK: memref.store %[[T9]], %[[C]][] : memref<vector<2xf32>>
89// CHECK: return
90func @matvecmax2x2(%arg0: memref<vector<2x2xf32>>, %arg1: memref<vector<2xf32>>,
91                                                   %arg2: memref<vector<2xf32>>) {
92  %A = memref.load %arg0[] : memref<vector<2x2xf32>>
93  %x = memref.load %arg1[] : memref<vector<2xf32>>
94  %b = memref.load %arg2[] : memref<vector<2xf32>>
95  %0 = vector.contract #matvecmax_trait %A, %x, %b : vector<2x2xf32>, vector<2xf32> into vector<2xf32>
96  memref.store %0, %arg2[] : memref<vector<2xf32>>
97  return
98}
99
100// CHECK-LABEL: func @mattransvec2x2
101// CHECK-SAME: %[[A:.*0]]: memref<vector<2x2xf32>>
102// CHECK-SAME: %[[B:.*1]]: memref<vector<2xf32>>
103// CHECK-SAME: %[[C:.*2]]: memref<vector<2xf32>>
104// CHECK: %[[T0:.*]] = memref.load %[[A]][] : memref<vector<2x2xf32>>
105// CHECK: %[[T1:.*]] = memref.load %[[B]][] : memref<vector<2xf32>>
106// CHECK: %[[T2:.*]] = memref.load %[[C]][] : memref<vector<2xf32>>
107// CHECK: %[[T3:.*]] = vector.extract %[[T0]][0] : vector<2x2xf32>
108// CHECK: %[[T4:.*]] = vector.extract %[[T1]][0] : vector<2xf32>
109// CHECK: %[[T5:.*]] = vector.outerproduct %[[T3]], %[[T4]], %[[T2]] {kind = #vector.kind<add>} : vector<2xf32>, f32
110// CHECK: %[[T6:.*]] = vector.extract %[[T0]][1] : vector<2x2xf32>
111// CHECK: %[[T7:.*]] = vector.extract %[[T1]][1] : vector<2xf32>
112// CHECK: %[[T8:.*]] = vector.outerproduct %[[T6]], %[[T7]], %[[T5]] {kind = #vector.kind<add>} : vector<2xf32>, f32
113// CHECK: memref.store %[[T8]], %[[C]][] : memref<vector<2xf32>>
114// CHECK: return
115func @mattransvec2x2(%arg0: memref<vector<2x2xf32>>, %arg1: memref<vector<2xf32>>,
116                                                     %arg2: memref<vector<2xf32>>) {
117  %A = memref.load %arg0[] : memref<vector<2x2xf32>>
118  %x = memref.load %arg1[] : memref<vector<2xf32>>
119  %b = memref.load %arg2[] : memref<vector<2xf32>>
120  %0 = vector.contract #mattransvec_trait %A, %x, %b : vector<2x2xf32>, vector<2xf32> into vector<2xf32>
121  memref.store %0, %arg2[] : memref<vector<2xf32>>
122  return
123}
124
125// CHECK-LABEL: func @vecmat2x2
126// CHECK-SAME: %[[A:.*0]]: memref<vector<2x2xf32>>
127// CHECK-SAME: %[[B:.*1]]: memref<vector<2xf32>>
128// CHECK-SAME: %[[C:.*2]]: memref<vector<2xf32>>
129// CHECK: %[[T0:.*]] = memref.load %[[A]][] : memref<vector<2x2xf32>>
130// CHECK: %[[T1:.*]] = memref.load %[[B]][] : memref<vector<2xf32>>
131// CHECK: %[[T2:.*]] = memref.load %[[C]][] : memref<vector<2xf32>>
132// CHECK: %[[T3:.*]] = vector.transpose %[[T0]], [1, 0] : vector<2x2xf32> to vector<2x2xf32>
133// CHECK: %[[T4:.*]] = vector.extract %[[T3]][0] : vector<2x2xf32>
134// CHECK: %[[T5:.*]] = vector.extract %[[T1]][0] : vector<2xf32>
135// CHECK: %[[T6:.*]] = vector.outerproduct %[[T4]], %[[T5]], %[[T2]] {kind = #vector.kind<add>} : vector<2xf32>, f32
136// CHECK: %[[T7:.*]] = vector.extract %[[T3]][1] : vector<2x2xf32>
137// CHECK: %[[T8:.*]] = vector.extract %[[T1]][1] : vector<2xf32>
138// CHECK: %[[T9:.*]] = vector.outerproduct %[[T7]], %[[T8]], %[[T6]] {kind = #vector.kind<add>} : vector<2xf32>, f32
139// CHECK: memref.store %[[T9]], %[[C]][] : memref<vector<2xf32>>
140// CHECK: return
141func @vecmat2x2(%arg0: memref<vector<2x2xf32>>, %arg1: memref<vector<2xf32>>,
142                                                %arg2: memref<vector<2xf32>>) {
143  %A = memref.load %arg0[] : memref<vector<2x2xf32>>
144  %x = memref.load %arg1[] : memref<vector<2xf32>>
145  %b = memref.load %arg2[] : memref<vector<2xf32>>
146  %0 = vector.contract #vecmat_trait %x, %A, %b : vector<2xf32>, vector<2x2xf32> into vector<2xf32>
147  memref.store %0, %arg2[] : memref<vector<2xf32>>
148  return
149}
150
151// CHECK-LABEL: func @vecmattrans2x2
152// CHECK-SAME: %[[A:.*0]]: memref<vector<2x2xf32>>
153// CHECK-SAME: %[[B:.*1]]: memref<vector<2xf32>>
154// CHECK-SAME: %[[C:.*2]]: memref<vector<2xf32>>
155// CHECK: %[[T0:.*]] = memref.load %[[A]][] : memref<vector<2x2xf32>>
156// CHECK: %[[T1:.*]] = memref.load %[[B]][] : memref<vector<2xf32>>
157// CHECK: %[[T2:.*]] = memref.load %[[C]][] : memref<vector<2xf32>>
158// CHECK: %[[T3:.*]] = vector.extract %[[T0]][0] : vector<2x2xf32>
159// CHECK: %[[T4:.*]] = vector.extract %[[T1]][0] : vector<2xf32>
160// CHECK: %[[T5:.*]] = vector.outerproduct %[[T3]], %[[T4]], %[[T2]] {kind = #vector.kind<add>} : vector<2xf32>, f32
161// CHECK: %[[T6:.*]] = vector.extract %[[T0]][1] : vector<2x2xf32>
162// CHECK: %[[T7:.*]] = vector.extract %[[T1]][1] : vector<2xf32>
163// CHECK: %[[T8:.*]] = vector.outerproduct %[[T6]], %[[T7]], %[[T5]] {kind = #vector.kind<add>} : vector<2xf32>, f32
164// CHECK: memref.store %[[T8]], %[[C]][] : memref<vector<2xf32>>
165// CHECK: return
166func @vecmattrans2x2(%arg0: memref<vector<2x2xf32>>, %arg1: memref<vector<2xf32>>,
167                                                     %arg2: memref<vector<2xf32>>) {
168  %A = memref.load %arg0[] : memref<vector<2x2xf32>>
169  %x = memref.load %arg1[] : memref<vector<2xf32>>
170  %b = memref.load %arg2[] : memref<vector<2xf32>>
171  %0 = vector.contract #vecmattrans_trait %x, %A, %b : vector<2xf32>, vector<2x2xf32> into vector<2xf32>
172  memref.store %0, %arg2[] : memref<vector<2xf32>>
173  return
174}
175