Browse Source

unroll num_output for innerproduct, about 60% speed gain

tags/20180129
nihui 8 years ago
parent
commit
df218110be
1 changed files with 110 additions and 2 deletions
  1. +110
    -2
      src/layer/arm/innerproduct_arm.cpp

+ 110
- 2
src/layer/arm/innerproduct_arm.cpp View File

@@ -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;



Loading…
Cancel
Save