Browse Source

simplify innerproduct x86 arm packing class

tags/20210124
nihui 5 years ago
parent
commit
4cd1a5c0e3
2 changed files with 340 additions and 753 deletions
  1. +226
    -477
      src/layer/arm/innerproduct_arm.cpp
  2. +114
    -276
      src/layer/x86/innerproduct_x86.cpp

+ 226
- 477
src/layer/arm/innerproduct_arm.cpp View File

@@ -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<const __fp16>(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<const __fp16>(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<const __fp16>(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<const __fp16>(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<const __fp16>(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<const __fp16>(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<const __fp16>(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<const __fp16>(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<const __fp16>(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<const __fp16>(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<unsigned short>(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<const unsigned short>(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<const unsigned short>(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<const unsigned short>(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);


+ 114
- 276
src/layer/x86/innerproduct_x86.cpp View File

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


Loading…
Cancel
Save