From 4cd1a5c0e35cd3e38507bf6cd4ac30a2507dba74 Mon Sep 17 00:00:00 2001 From: nihui Date: Sat, 23 Jan 2021 21:59:48 +0800 Subject: [PATCH] simplify innerproduct x86 arm packing class --- src/layer/arm/innerproduct_arm.cpp | 703 ++++++++++------------------- src/layer/x86/innerproduct_x86.cpp | 390 +++++----------- 2 files changed, 340 insertions(+), 753 deletions(-) diff --git a/src/layer/arm/innerproduct_arm.cpp b/src/layer/arm/innerproduct_arm.cpp index 8f5948c1d..ac39fb63a 100644 --- a/src/layer/arm/innerproduct_arm.cpp +++ b/src/layer/arm/innerproduct_arm.cpp @@ -625,34 +625,29 @@ int InnerProduct_arm::create_pipeline_fp16s(const Option& opt) { const int num_input = weight_data_size / num_output; - int elempack = 1; int out_elempack = 1; if (opt.use_packing_layout) { - elempack = opt.use_fp16_arithmetic && num_input % 8 == 0 ? 8 : num_input % 4 == 0 ? 4 : 1; out_elempack = opt.use_fp16_arithmetic && num_output % 8 == 0 ? 8 : num_output % 4 == 0 ? 4 : 1; } // src = inch-outch - // dst = pb-pa-inch/pa-outch/pb + // dst = pb-inch-outch/pb { Mat weight_data_r2 = weight_data.reshape(num_input, num_output); - weight_data_fp16.create(num_input / elempack, num_output / out_elempack, (size_t)2u * elempack * out_elempack, elempack * out_elempack); + weight_data_fp16.create(num_input, num_output / out_elempack, (size_t)2u * out_elempack, out_elempack); for (int q = 0; q + (out_elempack - 1) < num_output; q += out_elempack) { __fp16* g0 = weight_data_fp16.row<__fp16>(q / out_elempack); - for (int p = 0; p + (elempack - 1) < num_input; p += elempack) + for (int p = 0; p < num_input; p++) { - for (int i = 0; i < elempack; i++) + for (int j = 0; j < out_elempack; j++) { - for (int j = 0; j < out_elempack; j++) - { - *g0++ = (__fp16)(weight_data_r2.row(q + j)[p + i]); - } + *g0++ = (__fp16)(weight_data_r2.row(q + j)[p]); } } } @@ -866,7 +861,6 @@ int InnerProduct_arm::forward_fp16s(const Mat& bottom_blob, Mat& top_blob, const flatten->forward(bottom_blob, bottom_blob_flattened, opt_flatten); } - int size = bottom_blob_flattened.w; size_t elemsize = bottom_blob_flattened.elemsize; int elempack = bottom_blob_flattened.elempack; @@ -877,7 +871,7 @@ int InnerProduct_arm::forward_fp16s(const Mat& bottom_blob, Mat& top_blob, const if (top_blob.empty()) return -100; - if (elempack == 4 && out_elempack == 4) + if (out_elempack == 4) { // num_output #pragma omp parallel for num_threads(opt.num_threads) @@ -894,7 +888,8 @@ int InnerProduct_arm::forward_fp16s(const Mat& bottom_blob, Mat& top_blob, const const __fp16* sptr = bottom_blob_flattened; - for (int i = 0; i < size; i++) + int i = 0; + for (; i + 3 < num_input; i += 4) { float32x4_t _val = vcvt_f32_f16(vld1_f16(sptr)); @@ -911,32 +906,7 @@ int InnerProduct_arm::forward_fp16s(const Mat& bottom_blob, Mat& top_blob, const sptr += 4; kptr += 16; } - - _sum = activation_ps(_sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - vst1_f16(outptr + p * 4, vcvt_f16_f32(_sum)); - } - } - - if (elempack == 1 && out_elempack == 4) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float32x4_t _sum = vdupq_n_f32(0.f); - - if (bias_term) - { - _sum = vld1q_f32(((const float*)bias_data) + p * 4); - } - - const __fp16* kptr = weight_data_fp16.row(p); - - const __fp16* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) + for (; i < num_input; i++) { float32x4_t _val = vdupq_n_f32((float)sptr[0]); @@ -955,47 +925,7 @@ int InnerProduct_arm::forward_fp16s(const Mat& bottom_blob, Mat& top_blob, const } } - if (elempack == 4 && out_elempack == 1) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float sum = 0.f; - - if (bias_term) - { - sum = bias_data[p]; - } - - const __fp16* kptr = weight_data_fp16.row(p); - - const __fp16* sptr = bottom_blob_flattened; - - float32x4_t _sum = vdupq_n_f32(0.f); - - for (int i = 0; i < size; i++) - { - float32x4_t _val = vcvt_f32_f16(vld1_f16(sptr)); - - float32x4_t _w = vcvt_f32_f16(vld1_f16(kptr)); - - _sum = vfmaq_f32(_sum, _val, _w); - - sptr += 4; - kptr += 4; - } - - sum += vaddvq_f32(_sum); // dot - - sum = activation_ss(sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - outptr[p] = (__fp16)sum; - } - } - - if (elempack == 1 && out_elempack == 1) + if (out_elempack == 1) { // num_output #pragma omp parallel for num_threads(opt.num_threads) @@ -1012,7 +942,7 @@ int InnerProduct_arm::forward_fp16s(const Mat& bottom_blob, Mat& top_blob, const float32x4_t _sum = vdupq_n_f32(0.f); int i = 0; - for (; i + 3 < size; i += 4) + for (; i + 3 < num_input; i += 4) { float32x4_t _m = vcvt_f32_f16(vld1_f16(sptr)); float32x4_t _w = vcvt_f32_f16(vld1_f16(kptr)); @@ -1022,7 +952,7 @@ int InnerProduct_arm::forward_fp16s(const Mat& bottom_blob, Mat& top_blob, const sptr += 4; kptr += 4; } - for (; i < size; i++) + for (; i < num_input; i++) { float v = (float)(*sptr); float k = (float)(*kptr); @@ -1498,7 +1428,6 @@ int InnerProduct_arm::forward_fp16sa(const Mat& bottom_blob, Mat& top_blob, cons flatten->forward(bottom_blob, bottom_blob_flattened, opt_flatten); } - int size = bottom_blob_flattened.w; size_t elemsize = bottom_blob_flattened.elemsize; int elempack = bottom_blob_flattened.elempack; @@ -1513,365 +1442,253 @@ int InnerProduct_arm::forward_fp16sa(const Mat& bottom_blob, Mat& top_blob, cons if (top_blob.empty()) return -100; - if (elempack == 8 && out_elempack == 8) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float16x8_t _sum = vdupq_n_f16(0.f); - - if (bias_term) - { - _sum = vld1q_f16((const __fp16*)bias_data_fp16 + p * 8); - } - - const __fp16* kptr = weight_data_fp16.row(p); - - const __fp16* sptr = bottom_blob_flattened; - - int nn = size; // size always > 0 - - asm volatile( - "eor v1.16b, v1.16b, v1.16b \n" - "eor v2.16b, v2.16b, v2.16b \n" - "eor v3.16b, v3.16b, v3.16b \n" - - "0: \n" - - "prfm pldl1keep, [%2, #128] \n" - "ld1 {v0.8h}, [%2], #16 \n" // _val - - "prfm pldl1keep, [%3, #512] \n" - "ld1 {v8.8h, v9.8h, v10.8h, v11.8h}, [%3], #64 \n" // w0123 - - "fmla %1.8h, v8.8h, v0.h[0] \n" - "fmla v1.8h, v9.8h, v0.h[1] \n" - - "prfm pldl1keep, [%3, #512] \n" - "ld1 {v12.8h, v13.8h, v14.8h, v15.8h}, [%3], #64 \n" // w4567 - - "fmla v2.8h, v10.8h, v0.h[2] \n" - "fmla v3.8h, v11.8h, v0.h[3] \n" - "fmla %1.8h, v12.8h, v0.h[4] \n" - "fmla v1.8h, v13.8h, v0.h[5] \n" - - "subs %w0, %w0, #1 \n" - - "fmla v2.8h, v14.8h, v0.h[6] \n" - "fmla v3.8h, v15.8h, v0.h[7] \n" - - "bne 0b \n" - - "fadd %1.8h, %1.8h, v1.8h \n" - "fadd v2.8h, v2.8h, v3.8h \n" - "fadd %1.8h, %1.8h, v2.8h \n" - - : "=r"(nn), // %0 - "=w"(_sum), // %1 - "=r"(sptr), // %2 - "=r"(kptr) // %3 - : "0"(nn), - "1"(_sum), - "2"(sptr), - "3"(kptr) - : "cc", "memory", "v0", "v1", "v2", "v3", "v8", "v9", "v10", "v11", "v12", "v13", "v14", "v15", "v16"); - - _sum = activation_ps(_sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - vst1q_f16(outptr + p * 8, _sum); - } - } - - if (elempack == 1 && out_elempack == 8) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float16x8_t _sum = vdupq_n_f16(0.f); - - if (bias_term) - { - _sum = vld1q_f16((const __fp16*)bias_data_fp16 + p * 8); - } - - const __fp16* kptr = weight_data_fp16.row(p); - - const __fp16* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) - { - float16x8_t _val = vdupq_n_f16(sptr[0]); - - float16x8_t _w = vld1q_f16(kptr); - - _sum = vfmaq_f16(_sum, _val, _w); - - sptr += 1; - kptr += 8; - } - - _sum = activation_ps(_sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - vst1q_f16(outptr + p * 8, _sum); - } - } - - if (elempack == 4 && out_elempack == 8) + if (out_elempack == 8) { // num_output #pragma omp parallel for num_threads(opt.num_threads) for (int p = 0; p < num_output / out_elempack; p++) { - float16x8_t _sum = vdupq_n_f16(0.f); + float16x8_t _sum0 = vdupq_n_f16(0.f); + float16x8_t _sum1 = vdupq_n_f16(0.f); + float16x8_t _sum2 = vdupq_n_f16(0.f); + float16x8_t _sum3 = vdupq_n_f16(0.f); + float16x8_t _sum4 = vdupq_n_f16(0.f); + float16x8_t _sum5 = vdupq_n_f16(0.f); + float16x8_t _sum6 = vdupq_n_f16(0.f); + float16x8_t _sum7 = vdupq_n_f16(0.f); if (bias_term) { - _sum = vld1q_f16((const __fp16*)bias_data_fp16 + p * 8); + _sum0 = vld1q_f16((const __fp16*)bias_data_fp16 + p * 8); } const __fp16* kptr = weight_data_fp16.row(p); const __fp16* sptr = bottom_blob_flattened; - for (int i = 0; i < size; i++) + int i = 0; + for (; i + 7 < num_input; i += 8) { - float16x4_t _val = vld1_f16(sptr); - - float16x8_t _w0 = vld1q_f16(kptr); - float16x8_t _w1 = vld1q_f16(kptr + 8); - float16x8_t _w2 = vld1q_f16(kptr + 16); - float16x8_t _w3 = vld1q_f16(kptr + 24); - - _sum = vfmaq_lane_f16(_sum, _w0, _val, 0); - _sum = vfmaq_lane_f16(_sum, _w1, _val, 1); - _sum = vfmaq_lane_f16(_sum, _w2, _val, 2); - _sum = vfmaq_lane_f16(_sum, _w3, _val, 3); - - sptr += 4; - kptr += 32; + asm volatile( + "prfm pldl1keep, [%8, #128] \n" + "ld1 {v0.8h}, [%8], #16 \n" // _val + + "prfm pldl1keep, [%9, #512] \n" + "ld1 {v8.8h, v9.8h, v10.8h, v11.8h}, [%9], #64 \n" // w0123 + + "prfm pldl1keep, [%9, #512] \n" + "ld1 {v12.8h, v13.8h, v14.8h, v15.8h}, [%9], #64 \n" // w4567 + + "fmla %0.8h, v8.8h, v0.h[0] \n" + "fmla %1.8h, v9.8h, v0.h[1] \n" + "fmla %2.8h, v10.8h, v0.h[2] \n" + "fmla %3.8h, v11.8h, v0.h[3] \n" + "fmla %4.8h, v12.8h, v0.h[4] \n" + "fmla %5.8h, v13.8h, v0.h[5] \n" + "fmla %6.8h, v14.8h, v0.h[6] \n" + "fmla %7.8h, v15.8h, v0.h[7] \n" + + : "=w"(_sum0), // %0 + "=w"(_sum1), // %1 + "=w"(_sum2), // %2 + "=w"(_sum3), // %3 + "=w"(_sum4), // %4 + "=w"(_sum5), // %5 + "=w"(_sum6), // %6 + "=w"(_sum7), // %7 + "=r"(sptr), // %8 + "=r"(kptr) // %9 + : "0"(_sum0), + "1"(_sum1), + "2"(_sum2), + "3"(_sum3), + "4"(_sum4), + "5"(_sum5), + "6"(_sum6), + "7"(_sum7), + "8"(sptr), + "9"(kptr) + : "cc", "memory", "v0", "v8", "v9", "v10", "v11", "v12", "v13", "v14", "v15"); } - - _sum = activation_ps(_sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - vst1q_f16(outptr + p * 8, _sum); - } - } - - if (elempack == 8 && out_elempack == 1) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float sum = 0.f; - - if (bias_term) + for (; i + 3 < num_input; i += 4) { - sum = bias_data[p]; + asm volatile( + "prfm pldl1keep, [%4, #128] \n" + "ld1 {v0.4h}, [%4], #8 \n" // _val + + "prfm pldl1keep, [%5, #512] \n" + "ld1 {v8.8h, v9.8h, v10.8h, v11.8h}, [%5], #64 \n" // w0123 + + "fmla %0.8h, v8.8h, v0.h[0] \n" + "fmla %1.8h, v9.8h, v0.h[1] \n" + "fmla %2.8h, v10.8h, v0.h[2] \n" + "fmla %3.8h, v11.8h, v0.h[3] \n" + + : "=w"(_sum0), // %0 + "=w"(_sum1), // %1 + "=w"(_sum2), // %2 + "=w"(_sum3), // %3 + "=r"(sptr), // %4 + "=r"(kptr) // %5 + : "0"(_sum0), + "1"(_sum1), + "2"(_sum2), + "3"(_sum3), + "4"(sptr), + "5"(kptr) + : "cc", "memory", "v0", "v8", "v9", "v10", "v11"); } - - const __fp16* kptr = weight_data_fp16.row(p); - - const __fp16* sptr = bottom_blob_flattened; - - float16x8_t _sum = vdupq_n_f16(0.f); - - for (int i = 0; i < size; i++) + for (; i < num_input; i++) { - float16x8_t _val = vld1q_f16(sptr); + float16x8_t _val = vdupq_n_f16(sptr[0]); float16x8_t _w = vld1q_f16(kptr); - _sum = vfmaq_f16(_sum, _val, _w); + _sum0 = vfmaq_f16(_sum0, _val, _w); - sptr += 8; + sptr += 1; kptr += 8; } - float16x4_t _s4 = vadd_f16(vget_low_f16(_sum), vget_high_f16(_sum)); - sum += vaddvq_f32(vcvt_f32_f16(_s4)); // dot - - sum = activation_ss(sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - outptr[p] = sum; - } - } - - if (elempack == 8 && out_elempack == 4) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float16x4_t _sum = vdup_n_f16(0.f); - - if (bias_term) - { - _sum = vld1_f16((const __fp16*)bias_data_fp16 + p * 4); - } - - const __fp16* kptr = weight_data_fp16.row(p); - - const __fp16* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) - { - float16x8_t _val = vld1q_f16(sptr); - - float16x4_t _w0 = vld1_f16(kptr); - float16x4_t _w1 = vld1_f16(kptr + 4); - float16x4_t _w2 = vld1_f16(kptr + 8); - float16x4_t _w3 = vld1_f16(kptr + 12); - float16x4_t _w4 = vld1_f16(kptr + 16); - float16x4_t _w5 = vld1_f16(kptr + 20); - float16x4_t _w6 = vld1_f16(kptr + 24); - float16x4_t _w7 = vld1_f16(kptr + 28); - - _sum = vfma_laneq_f16(_sum, _w0, _val, 0); - _sum = vfma_laneq_f16(_sum, _w1, _val, 1); - _sum = vfma_laneq_f16(_sum, _w2, _val, 2); - _sum = vfma_laneq_f16(_sum, _w3, _val, 3); - _sum = vfma_laneq_f16(_sum, _w4, _val, 4); - _sum = vfma_laneq_f16(_sum, _w5, _val, 5); - _sum = vfma_laneq_f16(_sum, _w6, _val, 6); - _sum = vfma_laneq_f16(_sum, _w7, _val, 7); - - sptr += 8; - kptr += 32; - } + _sum0 = vaddq_f16(_sum0, _sum1); + _sum2 = vaddq_f16(_sum2, _sum3); + _sum4 = vaddq_f16(_sum4, _sum5); + _sum6 = vaddq_f16(_sum6, _sum7); + _sum0 = vaddq_f16(_sum0, _sum2); + _sum4 = vaddq_f16(_sum4, _sum6); + _sum0 = vaddq_f16(_sum0, _sum4); - _sum = activation_ps(_sum, activation_type, activation_params); + _sum0 = activation_ps(_sum0, activation_type, activation_params); __fp16* outptr = (__fp16*)top_blob; - vst1_f16(outptr + p * 4, _sum); + vst1q_f16(outptr + p * 8, _sum0); } } - if (elempack == 4 && out_elempack == 4) + if (out_elempack == 4) { // num_output #pragma omp parallel for num_threads(opt.num_threads) for (int p = 0; p < num_output / out_elempack; p++) { - float16x4_t _sum = vdup_n_f16(0.f); + float16x4_t _sum0 = vdup_n_f16(0.f); + float16x4_t _sum1 = vdup_n_f16(0.f); + float16x4_t _sum2 = vdup_n_f16(0.f); + float16x4_t _sum3 = vdup_n_f16(0.f); + float16x4_t _sum4 = vdup_n_f16(0.f); + float16x4_t _sum5 = vdup_n_f16(0.f); + float16x4_t _sum6 = vdup_n_f16(0.f); + float16x4_t _sum7 = vdup_n_f16(0.f); if (bias_term) { - _sum = vld1_f16((const __fp16*)bias_data_fp16 + p * 4); + _sum0 = vld1_f16((const __fp16*)bias_data_fp16 + p * 4); } const __fp16* kptr = weight_data_fp16.row(p); const __fp16* sptr = bottom_blob_flattened; - for (int i = 0; i < size; i++) + int i = 0; + for (; i + 7 < num_input; i += 8) { - float16x4_t _val = vld1_f16(sptr); - - float16x4_t _w0 = vld1_f16(kptr); - float16x4_t _w1 = vld1_f16(kptr + 4); - float16x4_t _w2 = vld1_f16(kptr + 8); - float16x4_t _w3 = vld1_f16(kptr + 12); - - _sum = vfma_lane_f16(_sum, _w0, _val, 0); - _sum = vfma_lane_f16(_sum, _w1, _val, 1); - _sum = vfma_lane_f16(_sum, _w2, _val, 2); - _sum = vfma_lane_f16(_sum, _w3, _val, 3); - - sptr += 4; - kptr += 16; + asm volatile( + "prfm pldl1keep, [%8, #128] \n" + "ld1 {v0.8h}, [%8], #16 \n" // _val + + "prfm pldl1keep, [%9, #256] \n" + "ld1 {v8.4h, v9.4h, v10.4h, v11.4h}, [%9], #32 \n" // w0123 + + "prfm pldl1keep, [%9, #256] \n" + "ld1 {v12.4h, v13.4h, v14.4h, v15.4h}, [%9], #32 \n" // w4567 + + "fmla %0.4h, v8.4h, v0.h[0] \n" + "fmla %1.4h, v9.4h, v0.h[1] \n" + "fmla %2.4h, v10.4h, v0.h[2] \n" + "fmla %3.4h, v11.4h, v0.h[3] \n" + "fmla %4.4h, v12.4h, v0.h[4] \n" + "fmla %5.4h, v13.4h, v0.h[5] \n" + "fmla %6.4h, v14.4h, v0.h[6] \n" + "fmla %7.4h, v15.4h, v0.h[7] \n" + + : "=w"(_sum0), // %0 + "=w"(_sum1), // %1 + "=w"(_sum2), // %2 + "=w"(_sum3), // %3 + "=w"(_sum4), // %4 + "=w"(_sum5), // %5 + "=w"(_sum6), // %6 + "=w"(_sum7), // %7 + "=r"(sptr), // %8 + "=r"(kptr) // %9 + : "0"(_sum0), + "1"(_sum1), + "2"(_sum2), + "3"(_sum3), + "4"(_sum4), + "5"(_sum5), + "6"(_sum6), + "7"(_sum7), + "8"(sptr), + "9"(kptr) + : "cc", "memory", "v0", "v8", "v9", "v10", "v11", "v12", "v13", "v14", "v15"); } - - _sum = activation_ps(_sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - vst1_f16(outptr + p * 4, _sum); - } - } - - if (elempack == 1 && out_elempack == 4) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float16x4_t _sum = vdup_n_f16(0.f); - - if (bias_term) + for (; i + 3 < num_input; i += 4) { - _sum = vld1_f16((const __fp16*)bias_data_fp16 + p * 4); + asm volatile( + "prfm pldl1keep, [%4, #128] \n" + "ld1 {v0.4h}, [%4], #8 \n" // _val + + "prfm pldl1keep, [%5, #256] \n" + "ld1 {v8.4h, v9.4h, v10.4h, v11.4h}, [%5], #32 \n" // w0123 + + "fmla %0.4h, v8.4h, v0.h[0] \n" + "fmla %1.4h, v9.4h, v0.h[1] \n" + "fmla %2.4h, v10.4h, v0.h[2] \n" + "fmla %3.4h, v11.4h, v0.h[3] \n" + + : "=w"(_sum0), // %0 + "=w"(_sum1), // %1 + "=w"(_sum2), // %2 + "=w"(_sum3), // %3 + "=r"(sptr), // %4 + "=r"(kptr) // %5 + : "0"(_sum0), + "1"(_sum1), + "2"(_sum2), + "3"(_sum3), + "4"(sptr), + "5"(kptr) + : "cc", "memory", "v0", "v8", "v9", "v10", "v11"); } - - const __fp16* kptr = weight_data_fp16.row(p); - - const __fp16* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) + for (; i < num_input; i++) { float16x4_t _val = vdup_n_f16(sptr[0]); float16x4_t _w = vld1_f16(kptr); - _sum = vfma_f16(_sum, _val, _w); + _sum0 = vfma_f16(_sum0, _val, _w); sptr += 1; kptr += 4; } - _sum = activation_ps(_sum, activation_type, activation_params); - - __fp16* outptr = (__fp16*)top_blob; - vst1_f16(outptr + p * 4, _sum); - } - } - - if (elempack == 4 && out_elempack == 1) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float sum = 0.f; - - if (bias_term) - { - sum = bias_data[p]; - } - - const __fp16* kptr = weight_data_fp16.row(p); + _sum0 = vadd_f16(_sum0, _sum1); + _sum2 = vadd_f16(_sum2, _sum3); + _sum4 = vadd_f16(_sum4, _sum5); + _sum6 = vadd_f16(_sum6, _sum7); + _sum0 = vadd_f16(_sum0, _sum2); + _sum4 = vadd_f16(_sum4, _sum6); + _sum0 = vadd_f16(_sum0, _sum4); - const __fp16* sptr = bottom_blob_flattened; - - float16x4_t _sum = vdup_n_f16(0.f); - - for (int i = 0; i < size; i++) - { - float16x4_t _val = vld1_f16(sptr); - - float16x4_t _w = vld1_f16(kptr); - - _sum = vfma_f16(_sum, _val, _w); - - sptr += 4; - kptr += 4; - } - - sum += vaddvq_f32(vcvt_f32_f16(_sum)); // dot - - sum = activation_ss(sum, activation_type, activation_params); + _sum0 = activation_ps(_sum0, activation_type, activation_params); __fp16* outptr = (__fp16*)top_blob; - outptr[p] = (__fp16)sum; + vst1_f16(outptr + p * 4, _sum0); } } - if (elempack == 1 && out_elempack == 1) + if (out_elempack == 1) { // num_output #pragma omp parallel for num_threads(opt.num_threads) @@ -1888,7 +1705,7 @@ int InnerProduct_arm::forward_fp16sa(const Mat& bottom_blob, Mat& top_blob, cons float16x8_t _sum = vdupq_n_f16(0.f); int i = 0; - for (; i + 7 < size; i += 8) + for (; i + 7 < num_input; i += 8) { float16x8_t _m = vld1q_f16(sptr); float16x8_t _w = vld1q_f16(kptr); @@ -1898,7 +1715,7 @@ int InnerProduct_arm::forward_fp16sa(const Mat& bottom_blob, Mat& top_blob, cons sptr += 8; kptr += 8; } - for (; i < size; i++) + for (; i < num_input; i++) { __fp16 v = *sptr; __fp16 k = *kptr; @@ -1927,28 +1744,24 @@ int InnerProduct_arm::create_pipeline_bf16s(const Option& opt) { const int num_input = weight_data_size / num_output; - int elempack = opt.use_packing_layout && num_input % 4 == 0 ? 4 : 1; int out_elempack = opt.use_packing_layout && num_output % 4 == 0 ? 4 : 1; // src = inch-outch - // dst = pb-pa-inch/pa-outch/pb + // dst = pb-inch-outch/pb { Mat weight_data_r2 = weight_data.reshape(num_input, num_output); - weight_data_bf16.create(num_input / elempack, num_output / out_elempack, (size_t)2u * elempack * out_elempack, elempack * out_elempack); + weight_data_bf16.create(num_input, num_output / out_elempack, (size_t)2u * out_elempack, out_elempack); for (int q = 0; q + (out_elempack - 1) < num_output; q += out_elempack) { unsigned short* g0 = weight_data_bf16.row(q / out_elempack); - for (int p = 0; p + (elempack - 1) < num_input; p += elempack) + for (int p = 0; p < num_input; p++) { - for (int i = 0; i < elempack; i++) + for (int j = 0; j < out_elempack; j++) { - for (int j = 0; j < out_elempack; j++) - { - *g0++ = float32_to_bfloat16(weight_data_r2.row(q + j)[p + i]); - } + *g0++ = float32_to_bfloat16(weight_data_r2.row(q + j)[p]); } } } @@ -2171,7 +1984,6 @@ int InnerProduct_arm::forward_bf16s(const Mat& bottom_blob, Mat& top_blob, const flatten->forward(bottom_blob, bottom_blob_flattened, opt_flatten); } - int size = bottom_blob_flattened.w; size_t elemsize = bottom_blob_flattened.elemsize; int elempack = bottom_blob_flattened.elempack; @@ -2183,24 +1995,28 @@ int InnerProduct_arm::forward_bf16s(const Mat& bottom_blob, Mat& top_blob, const return -100; #if __ARM_NEON - if (elempack == 4 && out_elempack == 4) + if (out_elempack == 4) { // num_output #pragma omp parallel for num_threads(opt.num_threads) for (int p = 0; p < num_output / out_elempack; p++) { - float32x4_t _sum = vdupq_n_f32(0.f); + 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); if (bias_term) { - _sum = vld1q_f32(((const float*)bias_data) + p * 4); + _sum0 = vld1q_f32(((const float*)bias_data) + p * 4); } const unsigned short* kptr = weight_data_bf16.row(p); const unsigned short* sptr = bottom_blob_flattened; - for (int i = 0; i < size; i++) + int i = 0; + for (; i + 3 < num_input; i += 4) { float32x4_t _val = vcvt_f32_bf16(vld1_u16(sptr)); @@ -2210,112 +2026,45 @@ int InnerProduct_arm::forward_bf16s(const Mat& bottom_blob, Mat& top_blob, const float32x4_t _w3 = vcvt_f32_bf16(vld1_u16(kptr + 12)); #if __aarch64__ - _sum = vmlaq_laneq_f32(_sum, _w0, _val, 0); - _sum = vmlaq_laneq_f32(_sum, _w1, _val, 1); - _sum = vmlaq_laneq_f32(_sum, _w2, _val, 2); - _sum = vmlaq_laneq_f32(_sum, _w3, _val, 3); + _sum0 = vmlaq_laneq_f32(_sum0, _w0, _val, 0); + _sum1 = vmlaq_laneq_f32(_sum1, _w1, _val, 1); + _sum2 = vmlaq_laneq_f32(_sum2, _w2, _val, 2); + _sum3 = vmlaq_laneq_f32(_sum3, _w3, _val, 3); #else - _sum = vmlaq_lane_f32(_sum, _w0, vget_low_f32(_val), 0); - _sum = vmlaq_lane_f32(_sum, _w1, vget_low_f32(_val), 1); - _sum = vmlaq_lane_f32(_sum, _w2, vget_high_f32(_val), 0); - _sum = vmlaq_lane_f32(_sum, _w3, vget_high_f32(_val), 1); + _sum0 = vmlaq_lane_f32(_sum0, _w0, vget_low_f32(_val), 0); + _sum1 = vmlaq_lane_f32(_sum1, _w1, vget_low_f32(_val), 1); + _sum2 = vmlaq_lane_f32(_sum2, _w2, vget_high_f32(_val), 0); + _sum3 = vmlaq_lane_f32(_sum3, _w3, vget_high_f32(_val), 1); #endif sptr += 4; kptr += 16; } - - _sum = activation_ps(_sum, activation_type, activation_params); - - unsigned short* outptr = (unsigned short*)top_blob; - vst1_u16(outptr + p * 4, vcvt_bf16_f32(_sum)); - } - } - - if (elempack == 1 && out_elempack == 4) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float32x4_t _sum = vdupq_n_f32(0.f); - - if (bias_term) - { - _sum = vld1q_f32(((const float*)bias_data) + p * 4); - } - - const unsigned short* kptr = weight_data_bf16.row(p); - - const unsigned short* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) + for (; i < num_input; i++) { float32x4_t _val = vdupq_n_f32(bfloat16_to_float32(sptr[0])); float32x4_t _w = vcvt_f32_bf16(vld1_u16(kptr)); - _sum = vmlaq_f32(_sum, _val, _w); + _sum0 = vmlaq_f32(_sum0, _val, _w); sptr += 1; kptr += 4; } - _sum = activation_ps(_sum, activation_type, activation_params); - - unsigned short* outptr = (unsigned short*)top_blob; - vst1_u16(outptr + p * 4, vcvt_bf16_f32(_sum)); - } - } - - if (elempack == 4 && out_elempack == 1) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float sum = 0.f; - - if (bias_term) - { - sum = bias_data[p]; - } + _sum0 = vaddq_f32(_sum0, _sum1); + _sum2 = vaddq_f32(_sum2, _sum3); + _sum0 = vaddq_f32(_sum0, _sum2); - const unsigned short* kptr = weight_data_bf16.row(p); - - const unsigned short* sptr = bottom_blob_flattened; - - float32x4_t _sum = vdupq_n_f32(0.f); - - for (int i = 0; i < size; i++) - { - float32x4_t _val = vcvt_f32_bf16(vld1_u16(sptr)); - - float32x4_t _w = vcvt_f32_bf16(vld1_u16(kptr)); - - _sum = vmlaq_f32(_sum, _val, _w); - - sptr += 4; - kptr += 4; - } - -#if __aarch64__ - sum += vaddvq_f32(_sum); // dot -#else - float32x2_t _ss = vadd_f32(vget_low_f32(_sum), vget_high_f32(_sum)); - _ss = vpadd_f32(_ss, _ss); - sum += vget_lane_f32(_ss, 0); -#endif - - sum = activation_ss(sum, activation_type, activation_params); + _sum0 = activation_ps(_sum0, activation_type, activation_params); unsigned short* outptr = (unsigned short*)top_blob; - outptr[p] = float32_to_bfloat16(sum); + vst1_u16(outptr + p * 4, vcvt_bf16_f32(_sum0)); } } #endif // __ARM_NEON - if (elempack == 1 && out_elempack == 1) + if (out_elempack == 1) { // num_output #pragma omp parallel for num_threads(opt.num_threads) @@ -2333,7 +2082,7 @@ int InnerProduct_arm::forward_bf16s(const Mat& bottom_blob, Mat& top_blob, const int i = 0; #if __ARM_NEON float32x4_t _sum = vdupq_n_f32(0.f); - for (; i + 3 < size; i += 4) + for (; i + 3 < num_input; i += 4) { float32x4_t _m = vcvt_f32_bf16(vld1_u16(sptr)); float32x4_t _w = vcvt_f32_bf16(vld1_u16(kptr)); @@ -2344,7 +2093,7 @@ int InnerProduct_arm::forward_bf16s(const Mat& bottom_blob, Mat& top_blob, const kptr += 4; } #endif // __ARM_NEON - for (; i < size; i++) + for (; i < num_input; i++) { float v = bfloat16_to_float32(*sptr); float k = bfloat16_to_float32(*kptr); diff --git a/src/layer/x86/innerproduct_x86.cpp b/src/layer/x86/innerproduct_x86.cpp index 05d6f1a73..c737a5f22 100644 --- a/src/layer/x86/innerproduct_x86.cpp +++ b/src/layer/x86/innerproduct_x86.cpp @@ -57,47 +57,37 @@ int InnerProduct_x86::create_pipeline(const Option& opt) const int num_input = weight_data_size / num_output; - int elempack = 1; int out_elempack = 1; #if __SSE2__ if (opt.use_packing_layout) { #if __AVX__ - elempack = num_input % 8 == 0 ? 8 : num_input % 4 == 0 ? 4 : 1; out_elempack = num_output % 8 == 0 ? 8 : num_output % 4 == 0 ? 4 : 1; #else - elempack = num_input % 4 == 0 ? 4 : 1; out_elempack = num_output % 4 == 0 ? 4 : 1; #endif } #endif // __SSE2__ - if (elempack == 1 && out_elempack == 1) - { - weight_data_packed = weight_data; - } - else + if (out_elempack != 1) { // src = inch-outch - // dst = pb-pa-inch/pa-outch/pb + // dst = pb-inch-outch/pb { Mat weight_data_r2 = weight_data.reshape(num_input, num_output); - weight_data_packed.create(num_input / elempack, num_output / out_elempack, (size_t)4u * elempack * out_elempack, elempack * out_elempack); + weight_data_packed.create(num_input, num_output / out_elempack, (size_t)4u * out_elempack, out_elempack); for (int q = 0; q + (out_elempack - 1) < num_output; q += out_elempack) { float* g0 = weight_data_packed.row(q / out_elempack); - for (int p = 0; p + (elempack - 1) < num_input; p += elempack) + for (int p = 0; p < num_input; p++) { - for (int i = 0; i < elempack; i++) + for (int j = 0; j < out_elempack; j++) { - for (int j = 0; j < out_elempack; j++) - { - *g0++ = weight_data_r2.row(q + j)[p + i]; - } + *g0++ = weight_data_r2.row(q + j)[p]; } } } @@ -399,7 +389,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio for (int p = 0; p < num_output; p++) { - const float* kptr = (const float*)weight_data_packed + num_input * p; + const float* kptr = (const float*)weight_data + num_input * p; const float* m = bottom_blob.row(j); __m256 _sum0 = _mm256_set1_ps(0.f); @@ -713,7 +703,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio for (int p = 0; p < num_output; p++) { - const float* kptr = (const float*)weight_data_packed + num_input * p; + const float* kptr = (const float*)weight_data + num_input * p; const float* m = bottom_blob.row(j); __m128 _sum0 = _mm_set1_ps(0.f); @@ -791,7 +781,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio for (int p = 0; p < num_output; p++) { - const float* kptr = (const float*)weight_data_packed + num_input * p; + const float* kptr = (const float*)weight_data + num_input * p; const float* m = bottom_blob.row(j); float sum = 0.f; @@ -889,7 +879,6 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio flatten->forward(bottom_blob, bottom_blob_flattened, opt_flatten); } - int size = bottom_blob_flattened.w; size_t elemsize = bottom_blob_flattened.elemsize; int elempack = bottom_blob_flattened.elempack; @@ -912,24 +901,32 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio #if __SSE2__ #if __AVX__ - if (elempack == 8 && out_elempack == 8) + if (out_elempack == 8) { // num_output #pragma omp parallel for num_threads(opt.num_threads) for (int p = 0; p < num_output / out_elempack; p++) { - __m256 _sum = _mm256_set1_ps(0.f); + __m256 _sum0 = _mm256_set1_ps(0.f); + __m256 _sum1 = _mm256_set1_ps(0.f); + __m256 _sum2 = _mm256_set1_ps(0.f); + __m256 _sum3 = _mm256_set1_ps(0.f); + __m256 _sum4 = _mm256_set1_ps(0.f); + __m256 _sum5 = _mm256_set1_ps(0.f); + __m256 _sum6 = _mm256_set1_ps(0.f); + __m256 _sum7 = _mm256_set1_ps(0.f); if (bias_term) { - _sum = _mm256_loadu_ps((const float*)bias_data + p * 8); + _sum0 = _mm256_loadu_ps((const float*)bias_data + p * 8); } const float* kptr = weight_data_packed.row(p); const float* sptr = bottom_blob_flattened; - for (int i = 0; i < size; i++) + int i = 0; + for (; i + 7 < num_input; i += 8) { __m256 _val0 = _mm256_broadcast_ss(sptr); __m256 _val1 = _mm256_broadcast_ss(sptr + 1); @@ -941,85 +938,26 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio __m256 _val7 = _mm256_broadcast_ss(sptr + 7); __m256 _w0 = _mm256_loadu_ps(kptr); - _sum = _mm256_fmadd_ps(_val0, _w0, _sum); + _sum0 = _mm256_fmadd_ps(_val0, _w0, _sum0); __m256 _w1 = _mm256_loadu_ps(kptr + 8); - _sum = _mm256_fmadd_ps(_val1, _w1, _sum); + _sum1 = _mm256_fmadd_ps(_val1, _w1, _sum1); __m256 _w2 = _mm256_loadu_ps(kptr + 16); - _sum = _mm256_fmadd_ps(_val2, _w2, _sum); + _sum2 = _mm256_fmadd_ps(_val2, _w2, _sum2); __m256 _w3 = _mm256_loadu_ps(kptr + 24); - _sum = _mm256_fmadd_ps(_val3, _w3, _sum); + _sum3 = _mm256_fmadd_ps(_val3, _w3, _sum3); __m256 _w4 = _mm256_loadu_ps(kptr + 32); - _sum = _mm256_fmadd_ps(_val4, _w4, _sum); + _sum4 = _mm256_fmadd_ps(_val4, _w4, _sum4); __m256 _w5 = _mm256_loadu_ps(kptr + 40); - _sum = _mm256_fmadd_ps(_val5, _w5, _sum); + _sum5 = _mm256_fmadd_ps(_val5, _w5, _sum5); __m256 _w6 = _mm256_loadu_ps(kptr + 48); - _sum = _mm256_fmadd_ps(_val6, _w6, _sum); + _sum6 = _mm256_fmadd_ps(_val6, _w6, _sum6); __m256 _w7 = _mm256_loadu_ps(kptr + 56); - _sum = _mm256_fmadd_ps(_val7, _w7, _sum); + _sum7 = _mm256_fmadd_ps(_val7, _w7, _sum7); sptr += 8; kptr += 64; } - - _sum = activation_avx(_sum, activation_type, activation_params); - - float* outptr = top_blob; - _mm256_storeu_ps(outptr + p * 8, _sum); - } - } - - if (elempack == 1 && out_elempack == 8) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - __m256 _sum = _mm256_set1_ps(0.f); - - if (bias_term) - { - _sum = _mm256_loadu_ps((const float*)bias_data + p * 8); - } - - const float* kptr = weight_data_packed.row(p); - - const float* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) - { - __m256 _val = _mm256_set1_ps(sptr[0]); - __m256 _w = _mm256_loadu_ps(kptr); - _sum = _mm256_fmadd_ps(_val, _w, _sum); - - sptr += 1; - kptr += 8; - } - - _sum = activation_avx(_sum, activation_type, activation_params); - - float* outptr = top_blob; - _mm256_storeu_ps(outptr + p * 8, _sum); - } - } - - if (elempack == 4 && out_elempack == 8) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - __m256 _sum = _mm256_set1_ps(0.f); - - if (bias_term) - { - _sum = _mm256_loadu_ps((const float*)bias_data + p * 8); - } - - const float* kptr = weight_data_packed.row(p); - - const float* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) + for (; i + 3 < num_input; i += 4) { __m256 _val0 = _mm256_broadcast_ss(sptr); __m256 _val1 = _mm256_broadcast_ss(sptr + 1); @@ -1027,81 +965,72 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio __m256 _val3 = _mm256_broadcast_ss(sptr + 3); __m256 _w0 = _mm256_loadu_ps(kptr); - _sum = _mm256_fmadd_ps(_val0, _w0, _sum); + _sum0 = _mm256_fmadd_ps(_val0, _w0, _sum0); __m256 _w1 = _mm256_loadu_ps(kptr + 8); - _sum = _mm256_fmadd_ps(_val1, _w1, _sum); + _sum1 = _mm256_fmadd_ps(_val1, _w1, _sum1); __m256 _w2 = _mm256_loadu_ps(kptr + 16); - _sum = _mm256_fmadd_ps(_val2, _w2, _sum); + _sum2 = _mm256_fmadd_ps(_val2, _w2, _sum2); __m256 _w3 = _mm256_loadu_ps(kptr + 24); - _sum = _mm256_fmadd_ps(_val3, _w3, _sum); + _sum3 = _mm256_fmadd_ps(_val3, _w3, _sum3); sptr += 4; kptr += 32; } - - _sum = activation_avx(_sum, activation_type, activation_params); - - float* outptr = top_blob; - _mm256_storeu_ps(outptr + p * 8, _sum); - } - } - - if (elempack == 8 && out_elempack == 1) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float sum = 0.f; - - if (bias_term) - { - sum = bias_data[p]; - } - - const float* kptr = weight_data_packed.row(p); - - const float* sptr = bottom_blob_flattened; - - __m256 _sum = _mm256_set1_ps(0.f); - - for (int i = 0; i < size; i++) + for (; i < num_input; i++) { - __m256 _val = _mm256_loadu_ps(sptr); + __m256 _val = _mm256_set1_ps(sptr[0]); __m256 _w = _mm256_loadu_ps(kptr); - _sum = _mm256_fmadd_ps(_val, _w, _sum); + _sum0 = _mm256_fmadd_ps(_val, _w, _sum0); - sptr += 8; + sptr += 1; kptr += 8; } - sum += _mm256_reduce_add_ps(_sum); // dot + _sum0 = _mm256_add_ps(_sum0, _sum1); + _sum2 = _mm256_add_ps(_sum2, _sum3); + _sum4 = _mm256_add_ps(_sum4, _sum5); + _sum6 = _mm256_add_ps(_sum6, _sum7); + _sum0 = _mm256_add_ps(_sum0, _sum2); + _sum4 = _mm256_add_ps(_sum4, _sum6); + _sum0 = _mm256_add_ps(_sum0, _sum4); - sum = activation_ss(sum, activation_type, activation_params); + _sum0 = activation_avx(_sum0, activation_type, activation_params); float* outptr = top_blob; - outptr[p] = sum; + _mm256_storeu_ps(outptr + p * 8, _sum0); } } +#endif // __AVX__ - if (elempack == 8 && out_elempack == 4) + if (out_elempack == 4) { // num_output #pragma omp parallel for num_threads(opt.num_threads) for (int p = 0; p < num_output / out_elempack; p++) { - __m128 _sum = _mm_set1_ps(0.f); + __m128 _sum0 = _mm_set1_ps(0.f); + __m128 _sum1 = _mm_set1_ps(0.f); + __m128 _sum2 = _mm_set1_ps(0.f); + __m128 _sum3 = _mm_set1_ps(0.f); +#if __AVX__ + __m128 _sum4 = _mm_set1_ps(0.f); + __m128 _sum5 = _mm_set1_ps(0.f); + __m128 _sum6 = _mm_set1_ps(0.f); + __m128 _sum7 = _mm_set1_ps(0.f); +#endif if (bias_term) { - _sum = _mm_loadu_ps((const float*)bias_data + p * 4); + _sum0 = _mm_loadu_ps((const float*)bias_data + p * 4); } const float* kptr = weight_data_packed.row(p); const float* sptr = bottom_blob_flattened; - for (int i = 0; i < size; i++) + int i = 0; +#if __AVX__ + for (; i + 7 < num_input; i += 8) { __m128 _val0 = _mm_broadcast_ss(sptr); __m128 _val1 = _mm_broadcast_ss(sptr + 1); @@ -1113,52 +1042,27 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio __m128 _val7 = _mm_broadcast_ss(sptr + 7); __m128 _w0 = _mm_loadu_ps(kptr); - _sum = _mm_fmadd_ps(_val0, _w0, _sum); + _sum0 = _mm_fmadd_ps(_val0, _w0, _sum0); __m128 _w1 = _mm_loadu_ps(kptr + 4); - _sum = _mm_fmadd_ps(_val1, _w1, _sum); + _sum1 = _mm_fmadd_ps(_val1, _w1, _sum1); __m128 _w2 = _mm_loadu_ps(kptr + 8); - _sum = _mm_fmadd_ps(_val2, _w2, _sum); + _sum2 = _mm_fmadd_ps(_val2, _w2, _sum2); __m128 _w3 = _mm_loadu_ps(kptr + 12); - _sum = _mm_fmadd_ps(_val3, _w3, _sum); + _sum3 = _mm_fmadd_ps(_val3, _w3, _sum3); __m128 _w4 = _mm_loadu_ps(kptr + 16); - _sum = _mm_fmadd_ps(_val4, _w4, _sum); + _sum4 = _mm_fmadd_ps(_val4, _w4, _sum4); __m128 _w5 = _mm_loadu_ps(kptr + 20); - _sum = _mm_fmadd_ps(_val5, _w5, _sum); + _sum5 = _mm_fmadd_ps(_val5, _w5, _sum5); __m128 _w6 = _mm_loadu_ps(kptr + 24); - _sum = _mm_fmadd_ps(_val6, _w6, _sum); + _sum6 = _mm_fmadd_ps(_val6, _w6, _sum6); __m128 _w7 = _mm_loadu_ps(kptr + 28); - _sum = _mm_fmadd_ps(_val7, _w7, _sum); + _sum7 = _mm_fmadd_ps(_val7, _w7, _sum7); sptr += 8; kptr += 32; } - - _sum = activation_sse(_sum, activation_type, activation_params); - - float* outptr = top_blob; - _mm_storeu_ps(outptr + p * 4, _sum); - } - } -#endif // __AVX__ - - if (elempack == 4 && out_elempack == 4) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - __m128 _sum = _mm_set1_ps(0.f); - - if (bias_term) - { - _sum = _mm_loadu_ps((const float*)bias_data + p * 4); - } - - const float* kptr = weight_data_packed.row(p); - - const float* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) +#endif + for (; i + 3 < num_input; i += 4) { __m128 _val0 = _mm_set1_ps(sptr[0]); __m128 _val1 = _mm_set1_ps(sptr[1]); @@ -1166,102 +1070,49 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio __m128 _val3 = _mm_set1_ps(sptr[3]); __m128 _w0 = _mm_loadu_ps(kptr); - _sum = _mm_add_ps(_mm_mul_ps(_val0, _w0), _sum); + _sum0 = _mm_add_ps(_mm_mul_ps(_val0, _w0), _sum0); __m128 _w1 = _mm_loadu_ps(kptr + 4); - _sum = _mm_add_ps(_mm_mul_ps(_val1, _w1), _sum); + _sum1 = _mm_add_ps(_mm_mul_ps(_val1, _w1), _sum1); __m128 _w2 = _mm_loadu_ps(kptr + 8); - _sum = _mm_add_ps(_mm_mul_ps(_val2, _w2), _sum); + _sum2 = _mm_add_ps(_mm_mul_ps(_val2, _w2), _sum2); __m128 _w3 = _mm_loadu_ps(kptr + 12); - _sum = _mm_add_ps(_mm_mul_ps(_val3, _w3), _sum); + _sum3 = _mm_add_ps(_mm_mul_ps(_val3, _w3), _sum3); sptr += 4; kptr += 16; } - - _sum = activation_sse(_sum, activation_type, activation_params); - - float* outptr = top_blob; - _mm_storeu_ps(outptr + p * 4, _sum); - } - } - - if (elempack == 1 && out_elempack == 4) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - __m128 _sum = _mm_set1_ps(0.f); - - if (bias_term) - { - _sum = _mm_loadu_ps((const float*)bias_data + p * 4); - } - - const float* kptr = weight_data_packed.row(p); - - const float* sptr = bottom_blob_flattened; - - for (int i = 0; i < size; i++) + for (; i < num_input; i++) { __m128 _val = _mm_set1_ps(sptr[0]); __m128 _w = _mm_loadu_ps(kptr); - _sum = _mm_add_ps(_mm_mul_ps(_val, _w), _sum); + _sum0 = _mm_add_ps(_mm_mul_ps(_val, _w), _sum0); sptr += 1; kptr += 4; } - _sum = activation_sse(_sum, activation_type, activation_params); - - float* outptr = top_blob; - _mm_storeu_ps(outptr + p * 4, _sum); - } - } - - if (elempack == 4 && out_elempack == 1) - { - // num_output - #pragma omp parallel for num_threads(opt.num_threads) - for (int p = 0; p < num_output / out_elempack; p++) - { - float sum = 0.f; - - if (bias_term) - { - sum = bias_data[p]; - } - - const float* kptr = weight_data_packed.row(p); - - const float* sptr = bottom_blob_flattened; - - __m128 _sum = _mm_set1_ps(0.f); - - for (int i = 0; i < size; i++) - { - __m128 _val = _mm_loadu_ps(sptr); - __m128 _w = _mm_loadu_ps(kptr); - _sum = _mm_add_ps(_mm_mul_ps(_val, _w), _sum); - - sptr += 4; - kptr += 4; - } - - sum += _mm_reduce_add_ps(_sum); // dot + _sum0 = _mm_add_ps(_sum0, _sum1); + _sum2 = _mm_add_ps(_sum2, _sum3); +#if __AVX__ + _sum4 = _mm_add_ps(_sum4, _sum5); + _sum6 = _mm_add_ps(_sum6, _sum7); +#endif + _sum0 = _mm_add_ps(_sum0, _sum2); +#if __AVX__ + _sum4 = _mm_add_ps(_sum4, _sum6); + _sum0 = _mm_add_ps(_sum0, _sum4); +#endif - sum = activation_ss(sum, activation_type, activation_params); + _sum0 = activation_sse(_sum0, activation_type, activation_params); float* outptr = top_blob; - outptr[p] = sum; + _mm_storeu_ps(outptr + p * 4, _sum0); } } #endif // __SSE2__ - if (elempack == 1 && out_elempack == 1) + if (out_elempack == 1) { - const float* weight_data_ptr = weight_data; - #if __SSE2__ #if __AVX__ int remain_num_output_start = 0; @@ -1285,14 +1136,14 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio sums[7] = bias_data[p + 7]; } - const float* w0 = weight_data_ptr + size * p; - const float* w1 = weight_data_ptr + size * (p + 1); - const float* w2 = weight_data_ptr + size * (p + 2); - const float* w3 = weight_data_ptr + size * (p + 3); - const float* w4 = weight_data_ptr + size * (p + 4); - const float* w5 = weight_data_ptr + size * (p + 5); - const float* w6 = weight_data_ptr + size * (p + 6); - const float* w7 = weight_data_ptr + size * (p + 7); + const float* w0 = (const float*)weight_data + num_input * p; + const float* w1 = (const float*)weight_data + num_input * (p + 1); + const float* w2 = (const float*)weight_data + num_input * (p + 2); + const float* w3 = (const float*)weight_data + num_input * (p + 3); + const float* w4 = (const float*)weight_data + num_input * (p + 4); + const float* w5 = (const float*)weight_data + num_input * (p + 5); + const float* w6 = (const float*)weight_data + num_input * (p + 6); + const float* w7 = (const float*)weight_data + num_input * (p + 7); const float* m = bottom_blob_flattened; @@ -1306,31 +1157,24 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio __m256 _sum7 = _mm256_set1_ps(0.f); int i = 0; - for (; i + 7 < size; i += 8) + for (; i + 7 < num_input; i += 8) { __m256 _m = _mm256_loadu_ps(m); __m256 _w0 = _mm256_loadu_ps(w0); _sum0 = _mm256_fmadd_ps(_m, _w0, _sum0); - __m256 _w1 = _mm256_loadu_ps(w1); _sum1 = _mm256_fmadd_ps(_m, _w1, _sum1); - __m256 _w2 = _mm256_loadu_ps(w2); _sum2 = _mm256_fmadd_ps(_m, _w2, _sum2); - __m256 _w3 = _mm256_loadu_ps(w3); _sum3 = _mm256_fmadd_ps(_m, _w3, _sum3); - __m256 _w4 = _mm256_loadu_ps(w4); _sum4 = _mm256_fmadd_ps(_m, _w4, _sum4); - __m256 _w5 = _mm256_loadu_ps(w5); _sum5 = _mm256_fmadd_ps(_m, _w5, _sum5); - __m256 _w6 = _mm256_loadu_ps(w6); _sum6 = _mm256_fmadd_ps(_m, _w6, _sum6); - __m256 _w7 = _mm256_loadu_ps(w7); _sum7 = _mm256_fmadd_ps(_m, _w7, _sum7); @@ -1344,7 +1188,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio w6 += 8; w7 += 8; } - for (; i < size; i++) + for (; i < num_input; i++) { sums[0] += *m * *w0; sums[1] += *m * *w1; @@ -1396,10 +1240,10 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio sums[3] = bias_data[p + 3]; } - const float* w0 = weight_data_ptr + size * p; - const float* w1 = weight_data_ptr + size * (p + 1); - const float* w2 = weight_data_ptr + size * (p + 2); - const float* w3 = weight_data_ptr + size * (p + 3); + const float* w0 = (const float*)weight_data + num_input * p; + const float* w1 = (const float*)weight_data + num_input * (p + 1); + const float* w2 = (const float*)weight_data + num_input * (p + 2); + const float* w3 = (const float*)weight_data + num_input * (p + 3); const float* m = bottom_blob_flattened; @@ -1409,19 +1253,16 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio __m256 _sum1 = _mm256_set1_ps(0.f); __m256 _sum2 = _mm256_set1_ps(0.f); __m256 _sum3 = _mm256_set1_ps(0.f); - for (; i + 7 < size; i += 8) + for (; i + 7 < num_input; i += 8) { __m256 _m = _mm256_loadu_ps(m); __m256 _w0 = _mm256_loadu_ps(w0); _sum0 = _mm256_fmadd_ps(_m, _w0, _sum0); - __m256 _w1 = _mm256_loadu_ps(w1); _sum1 = _mm256_fmadd_ps(_m, _w1, _sum1); - __m256 _w2 = _mm256_loadu_ps(w2); _sum2 = _mm256_fmadd_ps(_m, _w2, _sum2); - __m256 _w3 = _mm256_loadu_ps(w3); _sum3 = _mm256_fmadd_ps(_m, _w3, _sum3); @@ -1436,19 +1277,16 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio __m128 _sum1l = _mm_set1_ps(0.f); __m128 _sum2l = _mm_set1_ps(0.f); __m128 _sum3l = _mm_set1_ps(0.f); - for (; i + 3 < size; i += 4) + for (; i + 3 < num_input; i += 4) { __m128 _m = _mm_loadu_ps(m); __m128 _w0 = _mm_loadu_ps(w0); _sum0l = _mm_add_ps(_mm_mul_ps(_m, _w0), _sum0l); - __m128 _w1 = _mm_loadu_ps(w1); _sum1l = _mm_add_ps(_mm_mul_ps(_m, _w1), _sum1l); - __m128 _w2 = _mm_loadu_ps(w2); _sum2l = _mm_add_ps(_mm_mul_ps(_m, _w2), _sum2l); - __m128 _w3 = _mm_loadu_ps(w3); _sum3l = _mm_add_ps(_mm_mul_ps(_m, _w3), _sum3l); @@ -1458,7 +1296,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio w2 += 4; w3 += 4; } - for (; i < size; i++) + for (; i < num_input; i++) { sums[0] += *m * *w0; sums[1] += *m * *w1; @@ -1501,7 +1339,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio if (bias_term) sum = bias_data[p]; - const float* w = weight_data_ptr + size * p; + const float* w = (const float*)weight_data + num_input * p; const float* m = bottom_blob_flattened; @@ -1509,7 +1347,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio #if __SSE2__ #if __AVX__ __m256 _sum = _mm256_set1_ps(0.f); - for (; i + 7 < size; i += 8) + for (; i + 7 < num_input; i += 8) { __m256 _m = _mm256_loadu_ps(m); @@ -1521,7 +1359,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio } #endif // __AVX__ __m128 _suml = _mm_set1_ps(0.f); - for (; i + 3 < size; i += 4) + for (; i + 3 < num_input; i += 4) { __m128 _m = _mm_loadu_ps(m); @@ -1532,7 +1370,7 @@ int InnerProduct_x86::forward(const Mat& bottom_blob, Mat& top_blob, const Optio w += 4; } #endif // __SSE2__ - for (; i < size; i++) + for (; i < num_input; i++) { sum += *m * *w; m++;