| @@ -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<nn_num_output; pp++) | |||
| { | |||
| int p = pp * 4; | |||
| float sum0 = 0.f; | |||
| float sum1 = 0.f; | |||
| float sum2 = 0.f; | |||
| float sum3 = 0.f; | |||
| if (bias_term) | |||
| { | |||
| sum0 = bias_data[p]; | |||
| sum1 = bias_data[p+1]; | |||
| sum2 = bias_data[p+2]; | |||
| sum3 = bias_data[p+3]; | |||
| } | |||
| const float* w0 = weight_data_ptr + size * channels * p; | |||
| const float* w1 = weight_data_ptr + size * channels * (p+1); | |||
| const float* w2 = weight_data_ptr + size * channels * (p+2); | |||
| const float* w3 = weight_data_ptr + size * channels * (p+3); | |||
| #if __ARM_NEON | |||
| float32x4_t _sum0 = vdupq_n_f32(0.f); | |||
| float32x4_t _sum1 = vdupq_n_f32(0.f); | |||
| float32x4_t _sum2 = vdupq_n_f32(0.f); | |||
| float32x4_t _sum3 = vdupq_n_f32(0.f); | |||
| #endif // __ARM_NEON | |||
| // channels | |||
| for (int q=0; q<channels; q++) | |||
| { | |||
| const float* m = bottom_blob.channel(q); | |||
| #if __ARM_NEON | |||
| int nn = size >> 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<num_output; p++) | |||
| for (int p=remain_num_output_start; p<num_output; p++) | |||
| { | |||
| float sum = 0.f; | |||