From df218110bee3f67740159811f433df162497d5fd Mon Sep 17 00:00:00 2001 From: nihui Date: Sat, 20 Jan 2018 15:49:33 +0800 Subject: [PATCH] unroll num_output for innerproduct, about 60% speed gain --- src/layer/arm/innerproduct_arm.cpp | 112 ++++++++++++++++++++++++++++- 1 file changed, 110 insertions(+), 2 deletions(-) diff --git a/src/layer/arm/innerproduct_arm.cpp b/src/layer/arm/innerproduct_arm.cpp index 59941d294..47e03d665 100644 --- a/src/layer/arm/innerproduct_arm.cpp +++ b/src/layer/arm/innerproduct_arm.cpp @@ -33,10 +33,118 @@ int InnerProduct_arm::forward(const Mat& bottom_blob, Mat& top_blob) const if (top_blob.empty()) return -100; - // num_output const float* weight_data_ptr = weight_data; + + int nn_num_output = num_output >> 2; + int remain_num_output_start = nn_num_output << 2; + + #pragma omp parallel for + for (int pp=0; pp> 2; + int remain = size & 3; +#else + int remain = size; +#endif // __ARM_NEON + +#if __ARM_NEON + for (; nn>0; nn--) + { + float32x4_t _m = vld1q_f32(m); + + float32x4_t _w0 = vld1q_f32(w0); + _sum0 = vmlaq_f32(_sum0, _m, _w0); + + float32x4_t _w1 = vld1q_f32(w1); + _sum1 = vmlaq_f32(_sum1, _m, _w1); + + float32x4_t _w2 = vld1q_f32(w2); + _sum2 = vmlaq_f32(_sum2, _m, _w2); + + float32x4_t _w3 = vld1q_f32(w3); + _sum3 = vmlaq_f32(_sum3, _m, _w3); + + m += 4; + w0 += 4; + w1 += 4; + w2 += 4; + w3 += 4; + } +#endif // __ARM_NEON + for (; remain>0; remain--) + { + sum0 += *m * *w0; + sum1 += *m * *w1; + sum2 += *m * *w2; + sum3 += *m * *w3; + + m++; + w0++; + w1++; + w2++; + w3++; + } + + } + +#if __ARM_NEON + float32x2_t _sum0ss = vadd_f32(vget_low_f32(_sum0), vget_high_f32(_sum0)); + float32x2_t _sum1ss = vadd_f32(vget_low_f32(_sum1), vget_high_f32(_sum1)); + float32x2_t _sum2ss = vadd_f32(vget_low_f32(_sum2), vget_high_f32(_sum2)); + float32x2_t _sum3ss = vadd_f32(vget_low_f32(_sum3), vget_high_f32(_sum3)); + + float32x2_t _sum01ss = vpadd_f32(_sum0ss, _sum1ss); + float32x2_t _sum23ss = vpadd_f32(_sum2ss, _sum3ss); + + sum0 += vget_lane_f32(_sum01ss, 0); + sum1 += vget_lane_f32(_sum01ss, 1); + sum2 += vget_lane_f32(_sum23ss, 0); + sum3 += vget_lane_f32(_sum23ss, 1); + +#endif // __ARM_NEON + + top_blob[p] = sum0; + top_blob[p+1] = sum1; + top_blob[p+2] = sum2; + top_blob[p+3] = sum3; + } + + // num_output #pragma omp parallel for - for (int p=0; p