From 790829bc62f487cccc36d7040eeb0ed55fc831d7 Mon Sep 17 00:00:00 2001 From: nihui Date: Sun, 12 Nov 2017 13:56:46 +0800 Subject: [PATCH] partition dot tiles and reuse kernel register, over 20% improvement for tiny image --- src/layer/arm/convolution_3x3.h | 950 ++++++++++++++++++++++++++++++ src/layer/arm/convolution_arm.cpp | 22 +- 2 files changed, 971 insertions(+), 1 deletion(-) diff --git a/src/layer/arm/convolution_3x3.h b/src/layer/arm/convolution_3x3.h index 52fba9bca..8bfaddda1 100644 --- a/src/layer/arm/convolution_3x3.h +++ b/src/layer/arm/convolution_3x3.h @@ -1303,6 +1303,956 @@ static void conv3x3s1_winograd64_neon(const Mat& bottom_blob, Mat& top_blob, con copy_cut_border(top_blob_bordered, top_blob, 0, top_blob_bordered.h - top_blob.h, 0, top_blob_bordered.w - top_blob.w); } +static void conv3x3s1_winograd64_neon2(const Mat& bottom_blob, Mat& top_blob, const Mat& kernel_tm, const Mat& _bias) +{ + int w = bottom_blob.w; + int h = bottom_blob.h; + int inch = bottom_blob.c; + + int outw = top_blob.w; + int outh = top_blob.h; + int outch = top_blob.c; + + // pad to 6n+2 + Mat bottom_blob_bordered = bottom_blob; + + outw = (outw + 5) / 6 * 6; + outh = (outh + 5) / 6 * 6; + + w = outw + 2; + h = outh + 2; + copy_make_border(bottom_blob, bottom_blob_bordered, 0, h - bottom_blob.h, 0, w - bottom_blob.w, 0, 0.f); + + const float* bias = _bias; + + // BEGIN transform input + Mat bottom_blob_tm; + { + int w_tm = outw / 6 * 8; + int h_tm = outh / 6 * 8; + bottom_blob_tm.create(2*8, 4 * w_tm/8 * h_tm/8, inch); + const int tiles = w_tm/8 * h_tm/8; + +// const float itm[8][8] = { +// {1.0f, 0.0f, -5.25f, 0.00f, 5.25f, 0.00f, -1.0f, 0.0f}, +// +// {0.0f, 1.0f, 1.00f, -4.25f, -4.25f, 1.00f, 1.0f, 0.0f}, +// {0.0f, -1.0f, 1.00f, 4.25f, -4.25f, -1.00f, 1.0f, 0.0f}, +// +// {0.0f, 0.5f, 0.25f, -2.50f, -1.25f, 2.00f, 1.0f, 0.0f}, +// {0.0f, -0.5f, 0.25f, 2.50f, -1.25f, -2.00f, 1.0f, 0.0f}, +// +// {0.0f, 2.0f, 4.00f, -2.50f, -5.00f, 0.50f, 1.0f, 0.0f}, +// {0.0f, -2.0f, 4.00f, 2.50f, -5.00f, -0.50f, 1.0f, 0.0f}, +// +// {0.0f, -1.0f, 0.00f, 5.25f, 0.00f, -5.25f, 0.0f, 1.0f} +// }; + + // 0 = r00 - r06 + (r04 - r02) * 5.25 + // 7 = r07 - r01 + (r03 - r05) * 5.25 + + // 1 = (r02 + r06 - r04 * 4.25) + (r01 - r03 * 4.25 + r05) + // 2 = (r02 + r06 - r04 * 4.25) - (r01 - r03 * 4.25 + r05) + + // 3 = (r06 + r02 * 0.25 - r04 * 1.25) + (r01 * 0.5 - r03 * 2.5 + r05 * 2) + // 4 = (r06 + r02 * 0.25 - r04 * 1.25) - (r01 * 0.5 - r03 * 2.5 + r05 * 2) + + // reuse r04 * 1.25 + // reuse r03 * 2.5 + // 5 = (r06 + (r02 - r04 * 1.25) * 4) + (r01 * 2 - r03 * 2.5 + r05 * 0.5) + // 6 = (r06 + (r02 - r04 * 1.25) * 4) - (r01 * 2 - r03 * 2.5 + r05 * 0.5) + + #pragma omp parallel for + for (int q = 0; q> 2; + int remain = tiles & 3; +#else + int remain = tiles; +#endif // __ARM_NEON + +#if __ARM_NEON +#if __aarch64__ + for (; nn>0; nn--) + { + float32x4_t _output0_tm = vld1q_f32(output0_tm); + float32x4_t _output0_tmn = vld1q_f32(output0_tm+4); + + float32x4_t _r0 = vld1q_f32(r0); + float32x4_t _r0n = vld1q_f32(r0+4); + float32x4_t _r1 = vld1q_f32(r1); + float32x4_t _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0n); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1n); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + } +#else + if (nn > 0) + { +#if 1 + asm volatile( + "mov r4, %1 \n" + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "pld [%1, #256] \n" + "vld1.f32 {d16-d19}, [%1 :128]! \n"// q8 q9 = _output0_tm + + "vmla.f32 q8, q12, %q8 \n" + "vmla.f32 q9, q13, %q9 \n" + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "0: \n" + + "pld [%1, #256] \n" + "vld1.f32 {d20-d23}, [%1 :128]! \n"// q10 q11 = _output0_tm + + "vmla.f32 q10, q12, %q12 \n" + "vmla.f32 q11, q13, %q13 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "vmla.f32 q8, q14, %q10 \n" + "vmla.f32 q9, q15, %q11 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "vmla.f32 q10, q14, %q14 \n" + "vmla.f32 q11, q15, %q15 \n" + + "vst1.f32 {d16-d19}, [r4 :128]! \n" + + "pld [%1, #256] \n" + "vld1.f32 {d16-d19}, [%1 :128]! \n"// q8 q9 = _output0_tm + + "vmla.f32 q8, q12, %q8 \n" + "vmla.f32 q9, q13, %q9 \n" + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "vst1.f32 {d20-d23}, [r4 :128]! \n" + + "pld [%1, #256] \n" + "vld1.f32 {d20-d23}, [%1 :128]! \n"// q10 q11 = _output0_tm + + "vmla.f32 q10, q12, %q12 \n" + "vmla.f32 q11, q13, %q13 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "vmla.f32 q8, q14, %q10 \n" + "vmla.f32 q9, q15, %q11 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "vmla.f32 q10, q14, %q14 \n" + "vmla.f32 q11, q15, %q15 \n" + + "vst1.f32 {d16-d19}, [r4 :128]! \n" + + "pld [%1, #256] \n" + "vld1.f32 {d16-d19}, [%1 :128]! \n"// q8 q9 = _output0_tm + + "vmla.f32 q8, q12, %q8 \n" + "vmla.f32 q9, q13, %q9 \n" + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "vst1.f32 {d20-d23}, [r4 :128]! \n" + + "pld [%1, #256] \n" + "vld1.f32 {d20-d23}, [%1 :128]! \n"// q10 q11 = _output0_tm + + "vmla.f32 q10, q12, %q12 \n" + "vmla.f32 q11, q13, %q13 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "vmla.f32 q8, q14, %q10 \n" + "vmla.f32 q9, q15, %q11 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "vmla.f32 q10, q14, %q14 \n" + "vmla.f32 q11, q15, %q15 \n" + + "vst1.f32 {d16-d19}, [r4 :128]! \n" + + "pld [%1, #256] \n" + "vld1.f32 {d16-d19}, [%1 :128]! \n"// q8 q9 = _output0_tm + + "vmla.f32 q8, q12, %q8 \n" + "vmla.f32 q9, q13, %q9 \n" + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "vst1.f32 {d20-d23}, [r4 :128]! \n" + + "pld [%1, #256] \n" + "vld1.f32 {d20-d23}, [%1 :128]! \n"// q10 q11 = _output0_tm + + "vmla.f32 q10, q12, %q12 \n" + "vmla.f32 q11, q13, %q13 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "vmla.f32 q8, q14, %q10 \n" + "vmla.f32 q9, q15, %q11 \n" + + "pld [%3, #256] \n" + "vld1.f32 {d28-d31}, [%3 :128]! \n"// q14 q15 = _r1 + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "vmla.f32 q10, q14, %q14 \n" + "vmla.f32 q11, q15, %q15 \n" + + "vst1.f32 {d16-d19}, [r4 :128]! \n" + + "pld [%1, #256] \n" + "vld1.f32 {d16-d19}, [%1 :128]! \n"// q8 q9 = _output0_tm + + "vmla.f32 q8, q12, %q8 \n" + "vmla.f32 q9, q13, %q9 \n" + + "pld [%2, #256] \n" + "vld1.f32 {d24-d27}, [%2 :128]! \n"// q12 q13 = _r0 + + "subs %0, #1 \n" + + "vst1.f32 {d20-d23}, [r4 :128]! \n" + + "bne 0b \n" + + "sub %1, #32 \n" + "sub %2, #64 \n" + : "=r"(nn), // %0 + "=r"(output0_tm), // %1 + "=r"(r0), // %2 + "=r"(r1) // %3 + : "0"(nn), + "1"(output0_tm), + "2"(r0), + "3"(r1), + "w"(_k0), // %8 + "w"(_k0n), // %9 + "w"(_k1), // %10 + "w"(_k1n), // %11 + "w"(_k0nn), // %12 + "w"(_k0nnn), // %13 + "w"(_k1nn), // %14 + "w"(_k1nnn) // %15 + : "cc", "memory", "r4", "q8", "q9", "q10", "q11", "q12", "q13", "q14", "q15" + ); + } +#endif + +#endif // __aarch64__ +#endif // __ARM_NEON + for (; remain>0; remain--) + { +#if __ARM_NEON +#if __aarch64__ + float32x4_t _output0_tm = vld1q_f32(output0_tm); + float32x4_t _output0_tmn = vld1q_f32(output0_tm+4); + + float32x4_t _r0 = vld1q_f32(r0); + float32x4_t _r0n = vld1q_f32(r0+4); + float32x4_t _r1 = vld1q_f32(r1); + float32x4_t _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0n); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1n); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; + + _output0_tm = vld1q_f32(output0_tm); + _output0_tmn = vld1q_f32(output0_tm+4); + + _r0 = vld1q_f32(r0); + _r0n = vld1q_f32(r0+4); + _r1 = vld1q_f32(r1); + _r1n = vld1q_f32(r1+4); + + r0 += 8; + r1 += 8; + + _output0_tm = vmlaq_f32(_output0_tm, _r0, _k0nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r0n, _k0nnn); + _output0_tm = vmlaq_f32(_output0_tm, _r1, _k1nn); + _output0_tmn = vmlaq_f32(_output0_tmn, _r1n, _k1nnn); + + vst1q_f32(output0_tm, _output0_tm); + vst1q_f32(output0_tm+4, _output0_tmn); + + output0_tm += 8; +#else + asm volatile( + "mov r4, %0 \n" + + "pld [%1, #256] \n" + "vld1.f32 {d24-d27}, [%1 :128]! \n"// q12 q13 = _r0 + + "pld [%0, #256] \n" + "vld1.f32 {d16-d19}, [%0 :128]! \n"// q8 q9 = _output0_tm + + "vmla.f32 q8, q12, %q6 \n" + + "pld [%2, #256] \n" + "vld1.f32 {d28-d31}, [%2 :128]! \n"// q14 q15 = _r1 + "vmla.f32 q9, q13, %q7 \n" + + "pld [%1, #256] \n" + "vld1.f32 {d24-d27}, [%1 :128]! \n"// q12 q13 = _r0 + + "vmla.f32 q8, q14, %q8 \n" + + "pld [%0, #256] \n" + "vld1.f32 {d20-d23}, [%0 :128] \n"// q10 q11 = _output0_tm + "vmla.f32 q9, q15, %q9 \n" + + "vmla.f32 q10, q12, %q10 \n" + "vmla.f32 q11, q13, %q11 \n" + + "vst1.f32 {d16-d19}, [r4 :128] \n" + + "pld [%2, #256] \n" + "vld1.f32 {d28-d31}, [%2 :128]! \n"// q14 q15 = _r1 + + "vmla.f32 q10, q14, %q12 \n" + "vmla.f32 q11, q15, %q13 \n" + + "vst1.f32 {d20-d23}, [%0 :128]! \n" + : "=r"(output0_tm), // %0 + "=r"(r0), // %1 + "=r"(r1) // %2 + : "0"(output0_tm), + "1"(r0), + "2"(r1), + "w"(_k0), // %6 + "w"(_k0n), // %7 + "w"(_k1), // %8 + "w"(_k1n), // %9 + "w"(_k0nn), // %10 + "w"(_k0nnn), // %11 + "w"(_k1nn), // %12 + "w"(_k1nnn) // %13 + : "cc", "memory", "r4", "q8", "q9", "q10", "q11", "q12", "q13", "q14", "q15" + ); +#endif // __aarch64__ +#else + for (int m=0; m<16; m++) + { + output0_tm[m] += r0[m] * k0[m]; + output0_tm[m] += r1[m] * k1[m]; + } + + r0 += 16; + r1 += 16; + output0_tm += 16; +#endif // __ARM_NEON + } + +#if __ARM_NEON +#if __aarch64__ + k0 += 16; + k1 += 16; +#endif // __aarch64__ +#else + k0 += 16; + k1 += 16; +#endif // __ARM_NEON + } + } + + for (; q