1// RUN: mlir-opt %s -test-vector-multi-reduction-lowering-patterns | FileCheck %s
2
3func @vector_multi_reduction(%arg0: vector<2x4xf32>) -> vector<2xf32> {
4    %0 = vector.multi_reduction #vector.kind<mul>, %arg0 [1] : vector<2x4xf32> to vector<2xf32>
5    return %0 : vector<2xf32>
6}
7// CHECK-LABEL: func @vector_multi_reduction
8//  CHECK-SAME:   %[[INPUT:.+]]: vector<2x4xf32>
9//       CHECK:       %[[RESULT_VEC_0:.+]] = constant dense<{{.*}}> : vector<2xf32>
10//       CHECK:       %[[C0:.+]] = constant 0 : i32
11//       CHECK:       %[[C1:.+]] = constant 1 : i32
12//       CHECK:       %[[V0:.+]] = vector.extract %[[INPUT]][0]
13//       CHECK:       %[[RV0:.+]] = vector.reduction "mul", %[[V0]] : vector<4xf32> into f32
14//       CHECK:       %[[RESULT_VEC_1:.+]] = vector.insertelement %[[RV0:.+]], %[[RESULT_VEC_0]][%[[C0]] : i32] : vector<2xf32>
15//       CHECK:       %[[V1:.+]] = vector.extract %[[INPUT]][1]
16//       CHECK:       %[[RV1:.+]] = vector.reduction "mul", %[[V1]] : vector<4xf32> into f32
17//       CHECK:       %[[RESULT_VEC:.+]] = vector.insertelement %[[RV1:.+]], %[[RESULT_VEC_1]][%[[C1]] : i32] : vector<2xf32>
18//       CHECK:       return %[[RESULT_VEC]]
19
20func @vector_reduction_inner(%arg0: vector<2x3x4x5xi32>) -> vector<2x3xi32> {
21    %0 = vector.multi_reduction #vector.kind<add>, %arg0 [2, 3] : vector<2x3x4x5xi32> to vector<2x3xi32>
22    return %0 : vector<2x3xi32>
23}
24// CHECK-LABEL: func @vector_reduction_inner
25//  CHECK-SAME:   %[[INPUT:.+]]: vector<2x3x4x5xi32>
26//       CHECK:       %[[FLAT_RESULT_VEC_0:.+]] = constant dense<0> : vector<6xi32>
27//   CHECK-DAG:       %[[C0:.+]] = constant 0 : i32
28//   CHECK-DAG:       %[[C1:.+]] = constant 1 : i32
29//   CHECK-DAG:       %[[C2:.+]] = constant 2 : i32
30//   CHECK-DAG:       %[[C3:.+]] = constant 3 : i32
31//   CHECK-DAG:       %[[C4:.+]] = constant 4 : i32
32//   CHECK-DAG:       %[[C5:.+]] = constant 5 : i32
33//       CHECK:       %[[RESHAPED_INPUT:.+]] = vector.shape_cast %[[INPUT]] : vector<2x3x4x5xi32> to vector<6x20xi32>
34//       CHECK:       %[[V0:.+]] = vector.extract %[[RESHAPED_INPUT]][0] : vector<6x20xi32>
35//       CHECK:       %[[V0R:.+]] = vector.reduction "add", %[[V0]] : vector<20xi32> into i32
36//       CHECK:       %[[FLAT_RESULT_VEC_1:.+]] = vector.insertelement %[[V0R]], %[[FLAT_RESULT_VEC_0]][%[[C0]] : i32] : vector<6xi32>
37//       CHECK:       %[[V1:.+]] = vector.extract %[[RESHAPED_INPUT]][1] : vector<6x20xi32>
38//       CHECK:       %[[V1R:.+]] = vector.reduction "add", %[[V1]] : vector<20xi32> into i32
39//       CHECK:       %[[FLAT_RESULT_VEC_2:.+]] = vector.insertelement %[[V1R]], %[[FLAT_RESULT_VEC_1]][%[[C1]] : i32] : vector<6xi32>
40//       CHECK:       %[[V2:.+]] = vector.extract %[[RESHAPED_INPUT]][2] : vector<6x20xi32>
41//       CHECK:       %[[V2R:.+]] = vector.reduction "add", %[[V2]] : vector<20xi32> into i32
42//       CHECK:       %[[FLAT_RESULT_VEC_3:.+]] = vector.insertelement %[[V2R]], %[[FLAT_RESULT_VEC_2]][%[[C2]] : i32] : vector<6xi32>
43//       CHECK:       %[[V3:.+]] = vector.extract %[[RESHAPED_INPUT]][3] : vector<6x20xi32>
44//       CHECK:       %[[V3R:.+]] = vector.reduction "add", %[[V3]] : vector<20xi32> into i32
45//       CHECK:       %[[FLAT_RESULT_VEC_4:.+]] = vector.insertelement %[[V3R]], %[[FLAT_RESULT_VEC_3]][%[[C3]] : i32] : vector<6xi32>
46//       CHECK:       %[[V4:.+]] = vector.extract %[[RESHAPED_INPUT]][4] : vector<6x20xi32>
47//       CHECK:       %[[V4R:.+]] = vector.reduction "add", %[[V4]] : vector<20xi32> into i32
48//       CHECK:       %[[FLAT_RESULT_VEC_5:.+]] = vector.insertelement %[[V4R]], %[[FLAT_RESULT_VEC_4]][%[[C4]] : i32] : vector<6xi32>
49///       CHECK:      %[[V5:.+]] = vector.extract %[[RESHAPED_INPUT]][5] : vector<6x20xi32>
50//       CHECK:       %[[V5R:.+]] = vector.reduction "add", %[[V5]] : vector<20xi32> into i32
51//       CHECK:       %[[FLAT_RESULT_VEC:.+]] = vector.insertelement %[[V5R]], %[[FLAT_RESULT_VEC_5]][%[[C5]] : i32] : vector<6xi32>
52//       CHECK:       %[[RESULT:.+]] = vector.shape_cast %[[FLAT_RESULT_VEC]] : vector<6xi32> to vector<2x3xi32>
53//       CHECK:       return %[[RESULT]]
54
55
56func @vector_multi_reduction_transposed(%arg0: vector<2x3x4x5xf32>) -> vector<2x5xf32> {
57    %0 = vector.multi_reduction #vector.kind<add>, %arg0 [1, 2] : vector<2x3x4x5xf32> to vector<2x5xf32>
58    return %0 : vector<2x5xf32>
59}
60
61// CHECK-LABEL: func @vector_multi_reduction_transposed
62//  CHECK-SAME:    %[[INPUT:.+]]: vector<2x3x4x5xf32>
63//       CHECK:     %[[TRANSPOSED_INPUT:.+]] = vector.transpose %[[INPUT]], [0, 3, 1, 2] : vector<2x3x4x5xf32> to vector<2x5x3x4xf32>
64//       CHEKC:     vector.shape_cast %[[TRANSPOSED_INPUT]] : vector<2x5x3x4xf32> to vector<10x12xf32>
65//       CHECK:     %[[RESULT:.+]] = vector.shape_cast %{{.*}} : vector<10xf32> to vector<2x5xf32>
66//       CHECK:       return %[[RESULT]]
67