1// RUN: mlir-opt %s -test-vector-multi-reduction-lowering-patterns | FileCheck %s 2 3func.func @vector_multi_reduction(%arg0: vector<2x4xf32>, %acc: vector<2xf32>) -> vector<2xf32> { 4 %0 = vector.multi_reduction <mul>, %arg0, %acc [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>, %[[ACC:.*]]: vector<2xf32>) 9// CHECK: %[[RESULT_VEC_0:.+]] = arith.constant dense<{{.*}}> : vector<2xf32> 10// CHECK: %[[C0:.+]] = arith.constant 0 : index 11// CHECK: %[[C1:.+]] = arith.constant 1 : index 12// CHECK: %[[V0:.+]] = vector.extract %[[INPUT]][0] 13// CHECK: %[[ACC0:.+]] = vector.extract %[[ACC]][0] 14// CHECK: %[[RV0:.+]] = vector.reduction <mul>, %[[V0]], %[[ACC0]] : vector<4xf32> into f32 15// CHECK: %[[RESULT_VEC_1:.+]] = vector.insertelement %[[RV0:.+]], %[[RESULT_VEC_0]][%[[C0]] : index] : vector<2xf32> 16// CHECK: %[[V1:.+]] = vector.extract %[[INPUT]][1] 17// CHECK: %[[ACC1:.+]] = vector.extract %[[ACC]][1] 18// CHECK: %[[RV1:.+]] = vector.reduction <mul>, %[[V1]], %[[ACC1]] : vector<4xf32> into f32 19// CHECK: %[[RESULT_VEC:.+]] = vector.insertelement %[[RV1:.+]], %[[RESULT_VEC_1]][%[[C1]] : index] : vector<2xf32> 20// CHECK: return %[[RESULT_VEC]] 21 22func.func @vector_multi_reduction_to_scalar(%arg0: vector<2x4xf32>, %acc: f32) -> f32 { 23 %0 = vector.multi_reduction <mul>, %arg0, %acc [0, 1] : vector<2x4xf32> to f32 24 return %0 : f32 25} 26// CHECK-LABEL: func @vector_multi_reduction_to_scalar 27// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xf32>, %[[ACC:.*]]: f32) 28// CHECK: %[[CASTED:.*]] = vector.shape_cast %[[INPUT]] : vector<2x4xf32> to vector<8xf32> 29// CHECK: %[[REDUCED:.*]] = vector.reduction <mul>, %[[CASTED]], %[[ACC]] : vector<8xf32> into f32 30// CHECK: %[[INSERTED:.*]] = vector.insertelement %[[REDUCED]], {{.*}} : vector<1xf32> 31// CHECK: %[[RES:.*]] = vector.extract %[[INSERTED]][0] : vector<1xf32> 32// CHECK: return %[[RES]] 33 34func.func @vector_reduction_inner(%arg0: vector<2x3x4x5xi32>, %acc: vector<2x3xi32>) -> vector<2x3xi32> { 35 %0 = vector.multi_reduction <add>, %arg0, %acc [2, 3] : vector<2x3x4x5xi32> to vector<2x3xi32> 36 return %0 : vector<2x3xi32> 37} 38// CHECK-LABEL: func @vector_reduction_inner 39// CHECK-SAME: %[[INPUT:.+]]: vector<2x3x4x5xi32>, %[[ACC:.*]]: vector<2x3xi32> 40// CHECK: %[[FLAT_RESULT_VEC_0:.+]] = arith.constant dense<0> : vector<6xi32> 41// CHECK-DAG: %[[C0:.+]] = arith.constant 0 : index 42// CHECK-DAG: %[[C1:.+]] = arith.constant 1 : index 43// CHECK-DAG: %[[C2:.+]] = arith.constant 2 : index 44// CHECK-DAG: %[[C3:.+]] = arith.constant 3 : index 45// CHECK-DAG: %[[C4:.+]] = arith.constant 4 : index 46// CHECK-DAG: %[[C5:.+]] = arith.constant 5 : index 47// CHECK: %[[RESHAPED_INPUT:.+]] = vector.shape_cast %[[INPUT]] : vector<2x3x4x5xi32> to vector<6x20xi32> 48// CHECK: %[[V0:.+]] = vector.extract %[[RESHAPED_INPUT]][0] : vector<6x20xi32> 49// CHECK: %[[ACC0:.+]] = vector.extract %[[ACC]][0, 0] : vector<2x3xi32> 50// CHECK: %[[V0R:.+]] = vector.reduction <add>, %[[V0]], %[[ACC0]] : vector<20xi32> into i32 51// CHECK: %[[FLAT_RESULT_VEC_1:.+]] = vector.insertelement %[[V0R]], %[[FLAT_RESULT_VEC_0]][%[[C0]] : index] : vector<6xi32> 52// CHECK: %[[V1:.+]] = vector.extract %[[RESHAPED_INPUT]][1] : vector<6x20xi32> 53// CHECK: %[[ACC1:.+]] = vector.extract %[[ACC]][0, 1] : vector<2x3xi32> 54// CHECK: %[[V1R:.+]] = vector.reduction <add>, %[[V1]], %[[ACC1]] : vector<20xi32> into i32 55// CHECK: %[[FLAT_RESULT_VEC_2:.+]] = vector.insertelement %[[V1R]], %[[FLAT_RESULT_VEC_1]][%[[C1]] : index] : vector<6xi32> 56// CHECK: %[[V2:.+]] = vector.extract %[[RESHAPED_INPUT]][2] : vector<6x20xi32> 57// CHECK: %[[ACC2:.+]] = vector.extract %[[ACC]][0, 2] : vector<2x3xi32> 58// CHECK: %[[V2R:.+]] = vector.reduction <add>, %[[V2]], %[[ACC2]] : vector<20xi32> into i32 59// CHECK: %[[FLAT_RESULT_VEC_3:.+]] = vector.insertelement %[[V2R]], %[[FLAT_RESULT_VEC_2]][%[[C2]] : index] : vector<6xi32> 60// CHECK: %[[V3:.+]] = vector.extract %[[RESHAPED_INPUT]][3] : vector<6x20xi32> 61// CHECK: %[[ACC3:.+]] = vector.extract %[[ACC]][1, 0] : vector<2x3xi32> 62// CHECK: %[[V3R:.+]] = vector.reduction <add>, %[[V3]], %[[ACC3]] : vector<20xi32> into i32 63// CHECK: %[[FLAT_RESULT_VEC_4:.+]] = vector.insertelement %[[V3R]], %[[FLAT_RESULT_VEC_3]][%[[C3]] : index] : vector<6xi32> 64// CHECK: %[[V4:.+]] = vector.extract %[[RESHAPED_INPUT]][4] : vector<6x20xi32> 65// CHECK: %[[ACC4:.+]] = vector.extract %[[ACC]][1, 1] : vector<2x3xi32> 66// CHECK: %[[V4R:.+]] = vector.reduction <add>, %[[V4]], %[[ACC4]] : vector<20xi32> into i32 67// CHECK: %[[FLAT_RESULT_VEC_5:.+]] = vector.insertelement %[[V4R]], %[[FLAT_RESULT_VEC_4]][%[[C4]] : index] : vector<6xi32> 68/// CHECK: %[[V5:.+]] = vector.extract %[[RESHAPED_INPUT]][5] : vector<6x20xi32> 69// CHECK: %[[ACC5:.+]] = vector.extract %[[ACC]][1, 2] : vector<2x3xi32> 70// CHECK: %[[V5R:.+]] = vector.reduction <add>, %[[V5]], %[[ACC5]] : vector<20xi32> into i32 71// CHECK: %[[FLAT_RESULT_VEC:.+]] = vector.insertelement %[[V5R]], %[[FLAT_RESULT_VEC_5]][%[[C5]] : index] : vector<6xi32> 72// CHECK: %[[RESULT:.+]] = vector.shape_cast %[[FLAT_RESULT_VEC]] : vector<6xi32> to vector<2x3xi32> 73// CHECK: return %[[RESULT]] 74 75 76func.func @vector_multi_reduction_transposed(%arg0: vector<2x3x4x5xf32>, %acc: vector<2x5xf32>) -> vector<2x5xf32> { 77 %0 = vector.multi_reduction <add>, %arg0, %acc [1, 2] : vector<2x3x4x5xf32> to vector<2x5xf32> 78 return %0 : vector<2x5xf32> 79} 80 81// CHECK-LABEL: func @vector_multi_reduction_transposed 82// CHECK-SAME: %[[INPUT:.+]]: vector<2x3x4x5xf32> 83// CHECK: %[[TRANSPOSED_INPUT:.+]] = vector.transpose %[[INPUT]], [0, 3, 1, 2] : vector<2x3x4x5xf32> to vector<2x5x3x4xf32> 84// CHECK: vector.shape_cast %[[TRANSPOSED_INPUT]] : vector<2x5x3x4xf32> to vector<10x12xf32> 85// CHECK: %[[RESULT:.+]] = vector.shape_cast %{{.*}} : vector<10xf32> to vector<2x5xf32> 86// CHECK: return %[[RESULT]] 87 88func.func @vector_multi_reduction_ordering(%arg0: vector<3x2x4xf32>, %acc: vector<2x4xf32>) -> vector<2x4xf32> { 89 %0 = vector.multi_reduction <mul>, %arg0, %acc [0] : vector<3x2x4xf32> to vector<2x4xf32> 90 return %0 : vector<2x4xf32> 91} 92// CHECK-LABEL: func @vector_multi_reduction_ordering 93// CHECK-SAME: %[[INPUT:.+]]: vector<3x2x4xf32>, %[[ACC:.*]]: vector<2x4xf32>) 94// CHECK: %[[RESULT_VEC_0:.+]] = arith.constant dense<{{.*}}> : vector<8xf32> 95// CHECK: %[[C0:.+]] = arith.constant 0 : index 96// CHECK: %[[C1:.+]] = arith.constant 1 : index 97// CHECK: %[[C2:.+]] = arith.constant 2 : index 98// CHECK: %[[C3:.+]] = arith.constant 3 : index 99// CHECK: %[[C4:.+]] = arith.constant 4 : index 100// CHECK: %[[C5:.+]] = arith.constant 5 : index 101// CHECK: %[[C6:.+]] = arith.constant 6 : index 102// CHECK: %[[C7:.+]] = arith.constant 7 : index 103// CHECK: %[[TRANSPOSED_INPUT:.+]] = vector.transpose %[[INPUT]], [1, 2, 0] : vector<3x2x4xf32> to vector<2x4x3xf32> 104// CHECK: %[[V0:.+]] = vector.extract %[[TRANSPOSED_INPUT]][0, 0] 105// CHECK: %[[ACC0:.+]] = vector.extract %[[ACC]][0, 0] : vector<2x4xf32> 106// CHECK: %[[RV0:.+]] = vector.reduction <mul>, %[[V0]], %[[ACC0]] : vector<3xf32> into f32 107// CHECK: %[[RESULT_VEC_1:.+]] = vector.insertelement %[[RV0:.+]], %[[RESULT_VEC_0]][%[[C0]] : index] : vector<8xf32> 108// CHECK: %[[V1:.+]] = vector.extract %[[TRANSPOSED_INPUT]][0, 1] 109// CHECK: %[[ACC1:.+]] = vector.extract %[[ACC]][0, 1] : vector<2x4xf32> 110// CHECK: %[[RV1:.+]] = vector.reduction <mul>, %[[V1]], %[[ACC1]] : vector<3xf32> into f32 111// CHECK: %[[RESULT_VEC_2:.+]] = vector.insertelement %[[RV1:.+]], %[[RESULT_VEC_1]][%[[C1]] : index] : vector<8xf32> 112// CHECK: %[[V2:.+]] = vector.extract %[[TRANSPOSED_INPUT]][0, 2] 113// CHECK: %[[ACC2:.+]] = vector.extract %[[ACC]][0, 2] : vector<2x4xf32> 114// CHECK: %[[RV2:.+]] = vector.reduction <mul>, %[[V2]], %[[ACC2]] : vector<3xf32> into f32 115// CHECK: %[[RESULT_VEC_3:.+]] = vector.insertelement %[[RV2:.+]], %[[RESULT_VEC_2]][%[[C2]] : index] : vector<8xf32> 116// CHECK: %[[V3:.+]] = vector.extract %[[TRANSPOSED_INPUT]][0, 3] 117// CHECK: %[[ACC3:.+]] = vector.extract %[[ACC]][0, 3] : vector<2x4xf32> 118// CHECK: %[[RV3:.+]] = vector.reduction <mul>, %[[V3]], %[[ACC3]] : vector<3xf32> into f32 119// CHECK: %[[RESULT_VEC_4:.+]] = vector.insertelement %[[RV3:.+]], %[[RESULT_VEC_3]][%[[C3]] : index] : vector<8xf32> 120// CHECK: %[[V4:.+]] = vector.extract %[[TRANSPOSED_INPUT]][1, 0] 121// CHECK: %[[ACC4:.+]] = vector.extract %[[ACC]][1, 0] : vector<2x4xf32> 122// CHECK: %[[RV4:.+]] = vector.reduction <mul>, %[[V4]], %[[ACC4]] : vector<3xf32> into f32 123// CHECK: %[[RESULT_VEC_5:.+]] = vector.insertelement %[[RV4:.+]], %[[RESULT_VEC_4]][%[[C4]] : index] : vector<8xf32> 124// CHECK: %[[V5:.+]] = vector.extract %[[TRANSPOSED_INPUT]][1, 1] 125// CHECK: %[[ACC5:.+]] = vector.extract %[[ACC]][1, 1] : vector<2x4xf32> 126// CHECK: %[[RV5:.+]] = vector.reduction <mul>, %[[V5]], %[[ACC5]] : vector<3xf32> into f32 127// CHECK: %[[RESULT_VEC_6:.+]] = vector.insertelement %[[RV5:.+]], %[[RESULT_VEC_5]][%[[C5]] : index] : vector<8xf32> 128// CHECK: %[[V6:.+]] = vector.extract %[[TRANSPOSED_INPUT]][1, 2] 129// CHECK: %[[ACC6:.+]] = vector.extract %[[ACC]][1, 2] : vector<2x4xf32> 130// CHECK: %[[RV6:.+]] = vector.reduction <mul>, %[[V6]], %[[ACC6]] : vector<3xf32> into f32 131// CHECK: %[[RESULT_VEC_7:.+]] = vector.insertelement %[[RV6:.+]], %[[RESULT_VEC_6]][%[[C6]] : index] : vector<8xf32> 132// CHECK: %[[V7:.+]] = vector.extract %[[TRANSPOSED_INPUT]][1, 3] 133// CHECK: %[[ACC7:.+]] = vector.extract %[[ACC]][1, 3] : vector<2x4xf32> 134// CHECK: %[[RV7:.+]] = vector.reduction <mul>, %[[V7]], %[[ACC7]] : vector<3xf32> into f32 135// CHECK: %[[RESULT_VEC:.+]] = vector.insertelement %[[RV7:.+]], %[[RESULT_VEC_7]][%[[C7]] : index] : vector<8xf32> 136// CHECK: %[[RESHAPED_VEC:.+]] = vector.shape_cast %[[RESULT_VEC]] : vector<8xf32> to vector<2x4xf32> 137// CHECK: return %[[RESHAPED_VEC]] 138