Refactor and optimize highbd_dist_wtd_convolve_y_neon Simplify the final rounding step and tidy up the naming so that we can reuse the helper functions defined for highbd_dist_wtd_convolve_x_neon. Change-Id: I77e2da4bc3e8e656ee56f02a89b7f2ddd1e76249
diff --git a/av1/common/arm/highbd_compound_convolve_neon.c b/av1/common/arm/highbd_compound_convolve_neon.c index d4e9325..d8b823f 100644 --- a/av1/common/arm/highbd_compound_convolve_neon.c +++ b/av1/common/arm/highbd_compound_convolve_neon.c
@@ -22,50 +22,52 @@ #include "av1/common/filter.h" #include "av1/common/arm/highbd_convolve_neon.h" -static INLINE uint16x4_t highbd_convolve8_4_x(const int16x4_t s[8], - const int16x8_t x_filter, - const int32x4_t shift, - const int32x4_t offset) { - const int16x4_t x_filter_0_3 = vget_low_s16(x_filter); - const int16x4_t x_filter_4_7 = vget_high_s16(x_filter); +static INLINE uint16x4_t highbd_convolve8_4( + const int16x4_t s0, const int16x4_t s1, const int16x4_t s2, + const int16x4_t s3, const int16x4_t s4, const int16x4_t s5, + const int16x4_t s6, const int16x4_t s7, const int16x8_t filter, + const int32x4_t shift, const int32x4_t offset) { + const int16x4_t filter_0_3 = vget_low_s16(filter); + const int16x4_t filter_4_7 = vget_high_s16(filter); - int32x4_t sum = vmlal_lane_s16(offset, s[0], x_filter_0_3, 0); - sum = vmlal_lane_s16(sum, s[1], x_filter_0_3, 1); - sum = vmlal_lane_s16(sum, s[2], x_filter_0_3, 2); - sum = vmlal_lane_s16(sum, s[3], x_filter_0_3, 3); - sum = vmlal_lane_s16(sum, s[4], x_filter_4_7, 0); - sum = vmlal_lane_s16(sum, s[5], x_filter_4_7, 1); - sum = vmlal_lane_s16(sum, s[6], x_filter_4_7, 2); - sum = vmlal_lane_s16(sum, s[7], x_filter_4_7, 3); + int32x4_t sum = vmlal_lane_s16(offset, s0, filter_0_3, 0); + sum = vmlal_lane_s16(sum, s1, filter_0_3, 1); + sum = vmlal_lane_s16(sum, s2, filter_0_3, 2); + sum = vmlal_lane_s16(sum, s3, filter_0_3, 3); + sum = vmlal_lane_s16(sum, s4, filter_4_7, 0); + sum = vmlal_lane_s16(sum, s5, filter_4_7, 1); + sum = vmlal_lane_s16(sum, s6, filter_4_7, 2); + sum = vmlal_lane_s16(sum, s7, filter_4_7, 3); sum = vshlq_s32(sum, shift); return vqmovun_s32(sum); } -static INLINE uint16x8_t highbd_convolve8_8_x(const int16x8_t s[8], - const int16x8_t x_filter, - const int32x4_t shift, - const int32x4_t offset) { - const int16x4_t x_filter_0_3 = vget_low_s16(x_filter); - const int16x4_t x_filter_4_7 = vget_high_s16(x_filter); +static INLINE uint16x8_t highbd_convolve8_8( + const int16x8_t s0, const int16x8_t s1, const int16x8_t s2, + const int16x8_t s3, const int16x8_t s4, const int16x8_t s5, + const int16x8_t s6, const int16x8_t s7, const int16x8_t filter, + const int32x4_t shift, const int32x4_t offset) { + const int16x4_t filter_0_3 = vget_low_s16(filter); + const int16x4_t filter_4_7 = vget_high_s16(filter); - int32x4_t sum0 = vmlal_lane_s16(offset, vget_low_s16(s[0]), x_filter_0_3, 0); - sum0 = vmlal_lane_s16(sum0, vget_low_s16(s[1]), x_filter_0_3, 1); - sum0 = vmlal_lane_s16(sum0, vget_low_s16(s[2]), x_filter_0_3, 2); - sum0 = vmlal_lane_s16(sum0, vget_low_s16(s[3]), x_filter_0_3, 3); - sum0 = vmlal_lane_s16(sum0, vget_low_s16(s[4]), x_filter_4_7, 0); - sum0 = vmlal_lane_s16(sum0, vget_low_s16(s[5]), x_filter_4_7, 1); - sum0 = vmlal_lane_s16(sum0, vget_low_s16(s[6]), x_filter_4_7, 2); - sum0 = vmlal_lane_s16(sum0, vget_low_s16(s[7]), x_filter_4_7, 3); + int32x4_t sum0 = vmlal_lane_s16(offset, vget_low_s16(s0), filter_0_3, 0); + sum0 = vmlal_lane_s16(sum0, vget_low_s16(s1), filter_0_3, 1); + sum0 = vmlal_lane_s16(sum0, vget_low_s16(s2), filter_0_3, 2); + sum0 = vmlal_lane_s16(sum0, vget_low_s16(s3), filter_0_3, 3); + sum0 = vmlal_lane_s16(sum0, vget_low_s16(s4), filter_4_7, 0); + sum0 = vmlal_lane_s16(sum0, vget_low_s16(s5), filter_4_7, 1); + sum0 = vmlal_lane_s16(sum0, vget_low_s16(s6), filter_4_7, 2); + sum0 = vmlal_lane_s16(sum0, vget_low_s16(s7), filter_4_7, 3); - int32x4_t sum1 = vmlal_lane_s16(offset, vget_high_s16(s[0]), x_filter_0_3, 0); - sum1 = vmlal_lane_s16(sum1, vget_high_s16(s[1]), x_filter_0_3, 1); - sum1 = vmlal_lane_s16(sum1, vget_high_s16(s[2]), x_filter_0_3, 2); - sum1 = vmlal_lane_s16(sum1, vget_high_s16(s[3]), x_filter_0_3, 3); - sum1 = vmlal_lane_s16(sum1, vget_high_s16(s[4]), x_filter_4_7, 0); - sum1 = vmlal_lane_s16(sum1, vget_high_s16(s[5]), x_filter_4_7, 1); - sum1 = vmlal_lane_s16(sum1, vget_high_s16(s[6]), x_filter_4_7, 2); - sum1 = vmlal_lane_s16(sum1, vget_high_s16(s[7]), x_filter_4_7, 3); + int32x4_t sum1 = vmlal_lane_s16(offset, vget_high_s16(s0), filter_0_3, 0); + sum1 = vmlal_lane_s16(sum1, vget_high_s16(s1), filter_0_3, 1); + sum1 = vmlal_lane_s16(sum1, vget_high_s16(s2), filter_0_3, 2); + sum1 = vmlal_lane_s16(sum1, vget_high_s16(s3), filter_0_3, 3); + sum1 = vmlal_lane_s16(sum1, vget_high_s16(s4), filter_4_7, 0); + sum1 = vmlal_lane_s16(sum1, vget_high_s16(s5), filter_4_7, 1); + sum1 = vmlal_lane_s16(sum1, vget_high_s16(s6), filter_4_7, 2); + sum1 = vmlal_lane_s16(sum1, vget_high_s16(s7), filter_4_7, 3); sum0 = vshlq_s32(sum0, shift); sum1 = vshlq_s32(sum1, shift); @@ -96,10 +98,18 @@ load_s16_4x8(s + 3 * src_stride, 1, &s3[0], &s3[1], &s3[2], &s3[3], &s3[4], &s3[5], &s3[6], &s3[7]); - uint16x4_t d0 = highbd_convolve8_4_x(s0, x_filter, shift, offset_vec); - uint16x4_t d1 = highbd_convolve8_4_x(s1, x_filter, shift, offset_vec); - uint16x4_t d2 = highbd_convolve8_4_x(s2, x_filter, shift, offset_vec); - uint16x4_t d3 = highbd_convolve8_4_x(s3, x_filter, shift, offset_vec); + uint16x4_t d0 = + highbd_convolve8_4(s0[0], s0[1], s0[2], s0[3], s0[4], s0[5], s0[6], + s0[7], x_filter, shift, offset_vec); + uint16x4_t d1 = + highbd_convolve8_4(s1[0], s1[1], s1[2], s1[3], s1[4], s1[5], s1[6], + s1[7], x_filter, shift, offset_vec); + uint16x4_t d2 = + highbd_convolve8_4(s2[0], s2[1], s2[2], s2[3], s2[4], s2[5], s2[6], + s2[7], x_filter, shift, offset_vec); + uint16x4_t d3 = + highbd_convolve8_4(s3[0], s3[1], s3[2], s3[3], s3[4], s3[5], s3[6], + s3[7], x_filter, shift, offset_vec); store_u16_4x4(d, dst_stride, d0, d1, d2, d3); @@ -126,10 +136,18 @@ load_s16_8x8(s + 3 * src_stride, 1, &s3[0], &s3[1], &s3[2], &s3[3], &s3[4], &s3[5], &s3[6], &s3[7]); - uint16x8_t d0 = highbd_convolve8_8_x(s0, x_filter, shift, offset_vec); - uint16x8_t d1 = highbd_convolve8_8_x(s1, x_filter, shift, offset_vec); - uint16x8_t d2 = highbd_convolve8_8_x(s2, x_filter, shift, offset_vec); - uint16x8_t d3 = highbd_convolve8_8_x(s3, x_filter, shift, offset_vec); + uint16x8_t d0 = + highbd_convolve8_8(s0[0], s0[1], s0[2], s0[3], s0[4], s0[5], s0[6], + s0[7], x_filter, shift, offset_vec); + uint16x8_t d1 = + highbd_convolve8_8(s1[0], s1[1], s1[2], s1[3], s1[4], s1[5], s1[6], + s1[7], x_filter, shift, offset_vec); + uint16x8_t d2 = + highbd_convolve8_8(s2[0], s2[1], s2[2], s2[3], s2[4], s2[5], s2[6], + s2[7], x_filter, shift, offset_vec); + uint16x8_t d3 = + highbd_convolve8_8(s3[0], s3[1], s3[2], s3[3], s3[4], s3[5], s3[6], + s3[7], x_filter, shift, offset_vec); store_u16_8x4(d, dst_stride, d0, d1, d2, d3); @@ -192,16 +210,13 @@ } } -static INLINE void highbd_convolve_dist_wtd_y_8tap_neon( +static INLINE void highbd_dist_wtd_convolve_y_8tap_neon( const uint16_t *src_ptr, int src_stride, uint16_t *dst_ptr, int dst_stride, int w, int h, const int16_t *y_filter_ptr, ConvolveParams *conv_params, const int offset) { const int16x8_t y_filter = vld1q_s16(y_filter_ptr); - const int32x4_t shift_s32 = vdupq_n_s32(-conv_params->round_0); - const int weight_bits = FILTER_BITS - conv_params->round_1; - const int32x4_t zero_s32 = vdupq_n_s32(0); - const int32x4_t weight_s32 = vdupq_n_s32(1 << weight_bits); - const int32x4_t offset_s32 = vdupq_n_s32(offset); + const int32x4_t shift = vdupq_n_s32(-conv_params->round_0); + const int32x4_t offset_vec = vdupq_n_s32(offset); if (w <= 4) { const int16_t *s = (const int16_t *)src_ptr; @@ -215,33 +230,28 @@ int16x4_t s7, s8, s9, s10; load_s16_4x4(s, src_stride, &s7, &s8, &s9, &s10); - uint16x4_t d0 = highbd_convolve8_wtd_4_s32_s16( - s0, s1, s2, s3, s4, s5, s6, s7, y_filter, shift_s32, zero_s32, - weight_s32, offset_s32); - uint16x4_t d1 = highbd_convolve8_wtd_4_s32_s16( - s1, s2, s3, s4, s5, s6, s7, s8, y_filter, shift_s32, zero_s32, - weight_s32, offset_s32); - uint16x8_t d01 = vcombine_u16(d0, d1); + uint16x4_t d0 = highbd_convolve8_4(s0, s1, s2, s3, s4, s5, s6, s7, + y_filter, shift, offset_vec); + uint16x4_t d1 = highbd_convolve8_4(s1, s2, s3, s4, s5, s6, s7, s8, + y_filter, shift, offset_vec); + uint16x4_t d2 = highbd_convolve8_4(s2, s3, s4, s5, s6, s7, s8, s9, + y_filter, shift, offset_vec); + uint16x4_t d3 = highbd_convolve8_4(s3, s4, s5, s6, s7, s8, s9, s10, + y_filter, shift, offset_vec); - if (w == 2) { - store_u16q_2x1(d + 0 * dst_stride, d01, 0); - store_u16q_2x1(d + 1 * dst_stride, d01, 2); - } else { - vst1_u16(d + 0 * dst_stride, vget_low_u16(d01)); - vst1_u16(d + 1 * dst_stride, vget_high_u16(d01)); - } + store_u16_4x4(d, dst_stride, d0, d1, d2, d3); - s0 = s2; - s1 = s3; - s2 = s4; - s3 = s5; - s4 = s6; - s5 = s7; - s6 = s8; - s += 2 * src_stride; - d += 2 * dst_stride; - h -= 2; - } while (h > 0); + s0 = s4; + s1 = s5; + s2 = s6; + s3 = s7; + s4 = s8; + s5 = s9; + s6 = s10; + s += 4 * src_stride; + d += 4 * dst_stride; + h -= 4; + } while (h != 0); } else { do { int height = h; @@ -253,33 +263,35 @@ s += 7 * src_stride; do { - int16x8_t s7, s8; - load_s16_8x2(s, src_stride, &s7, &s8); + int16x8_t s7, s8, s9, s10; + load_s16_8x4(s, src_stride, &s7, &s8, &s9, &s10); - uint16x8_t d0 = highbd_convolve8_wtd_8_s32_s16( - s0, s1, s2, s3, s4, s5, s6, s7, y_filter, shift_s32, zero_s32, - weight_s32, offset_s32); - uint16x8_t d1 = highbd_convolve8_wtd_8_s32_s16( - s1, s2, s3, s4, s5, s6, s7, s8, y_filter, shift_s32, zero_s32, - weight_s32, offset_s32); + uint16x8_t d0 = highbd_convolve8_8(s0, s1, s2, s3, s4, s5, s6, s7, + y_filter, shift, offset_vec); + uint16x8_t d1 = highbd_convolve8_8(s1, s2, s3, s4, s5, s6, s7, s8, + y_filter, shift, offset_vec); + uint16x8_t d2 = highbd_convolve8_8(s2, s3, s4, s5, s6, s7, s8, s9, + y_filter, shift, offset_vec); + uint16x8_t d3 = highbd_convolve8_8(s3, s4, s5, s6, s7, s8, s9, s10, + y_filter, shift, offset_vec); - store_u16_8x2(d, dst_stride, d0, d1); + store_u16_8x4(d, dst_stride, d0, d1, d2, d3); - s0 = s2; - s1 = s3; - s2 = s4; - s3 = s5; - s4 = s6; - s5 = s7; - s6 = s8; - s += 2 * src_stride; - d += 2 * dst_stride; - height -= 2; - } while (height > 0); + s0 = s4; + s1 = s5; + s2 = s6; + s3 = s7; + s4 = s8; + s5 = s9; + s6 = s10; + s += 4 * src_stride; + d += 4 * dst_stride; + height -= 4; + } while (height != 0); src_ptr += 8; dst_ptr += 8; w -= 8; - } while (w > 0); + } while (w != 0); } } @@ -293,9 +305,13 @@ int dst16_stride = conv_params->dst_stride; const int im_stride = MAX_SB_SIZE; const int vert_offset = filter_params_y->taps / 2 - 1; + assert(FILTER_BITS == COMPOUND_ROUND1_BITS); const int offset_bits = bd + 2 * FILTER_BITS - conv_params->round_0; - const int round_offset = (1 << (offset_bits - conv_params->round_1)) + - (1 << (offset_bits - conv_params->round_1 - 1)); + const int round_offset_avg = (1 << (offset_bits - conv_params->round_1)) + + (1 << (offset_bits - conv_params->round_1 - 1)); + const int round_offset_conv = (1 << (conv_params->round_0 - 1)) + + (1 << (bd + FILTER_BITS)) + + (1 << (bd + FILTER_BITS - 1)); const int round_bits = 2 * FILTER_BITS - conv_params->round_0 - conv_params->round_1; assert(round_bits >= 0); @@ -305,24 +321,24 @@ src -= vert_offset * src_stride; - // vertical filter if (conv_params->do_average) { - highbd_convolve_dist_wtd_y_8tap_neon(src, src_stride, im_block, im_stride, + highbd_dist_wtd_convolve_y_8tap_neon(src, src_stride, im_block, im_stride, w, h, y_filter_ptr, conv_params, - round_offset); + round_offset_conv); } else { - highbd_convolve_dist_wtd_y_8tap_neon(src, src_stride, dst16, dst16_stride, + highbd_dist_wtd_convolve_y_8tap_neon(src, src_stride, dst16, dst16_stride, w, h, y_filter_ptr, conv_params, - round_offset); + round_offset_conv); } if (conv_params->do_average) { if (conv_params->use_dist_wtd_comp_avg) { highbd_dist_wtd_comp_avg_neon(im_block, im_stride, dst, dst_stride, w, h, - conv_params, round_bits, round_offset, bd); + conv_params, round_bits, round_offset_avg, + bd); } else { highbd_comp_avg_neon(im_block, im_stride, dst, dst_stride, w, h, - conv_params, round_bits, round_offset, bd); + conv_params, round_bits, round_offset_avg, bd); } } }
diff --git a/av1/common/arm/highbd_convolve_neon.h b/av1/common/arm/highbd_convolve_neon.h index cfbe173..e7aebd2 100644 --- a/av1/common/arm/highbd_convolve_neon.h +++ b/av1/common/arm/highbd_convolve_neon.h
@@ -50,21 +50,6 @@ return vqmovun_s32(sum); } -static INLINE uint16x4_t highbd_convolve8_wtd_4_s32_s16( - const int16x4_t s0, const int16x4_t s1, const int16x4_t s2, - const int16x4_t s3, const int16x4_t s4, const int16x4_t s5, - const int16x4_t s6, const int16x4_t s7, const int16x8_t y_filter, - const int32x4_t shift_s32, const int32x4_t offset, const int32x4_t weight, - const int32x4_t offset2) { - int32x4_t sum = - highbd_convolve8_4_s32(s0, s1, s2, s3, s4, s5, s6, s7, y_filter, offset); - - sum = vqrshlq_s32(sum, shift_s32); - sum = vmlaq_s32(offset2, sum, weight); - - return vqmovun_s32(sum); -} - // Like above but also perform round shifting and subtract correction term static INLINE uint16x4_t highbd_convolve8_4_srsub_s32_s16( const int16x4_t s0, const int16x4_t s1, const int16x4_t s2, @@ -120,25 +105,6 @@ vqrshrun_n_s32(sum1, COMPOUND_ROUND1_BITS)); } -static INLINE uint16x8_t highbd_convolve8_wtd_8_s32_s16( - const int16x8_t s0, const int16x8_t s1, const int16x8_t s2, - const int16x8_t s3, const int16x8_t s4, const int16x8_t s5, - const int16x8_t s6, const int16x8_t s7, const int16x8_t y_filter, - const int32x4_t shift_s32, const int32x4_t offset, const int32x4_t weight, - const int32x4_t offset2) { - int32x4_t sum0; - int32x4_t sum1; - highbd_convolve8_8_s32(s0, s1, s2, s3, s4, s5, s6, s7, y_filter, offset, - &sum0, &sum1); - - sum0 = vqrshlq_s32(sum0, shift_s32); - sum1 = vqrshlq_s32(sum1, shift_s32); - sum0 = vmlaq_s32(offset2, sum0, weight); - sum1 = vmlaq_s32(offset2, sum1, weight); - - return vcombine_u16(vqmovun_s32(sum0), vqmovun_s32(sum1)); -} - // Like above but also perform round shifting and subtract correction term static INLINE uint16x8_t highbd_convolve8_8_srsub_s32_s16( const int16x8_t s0, const int16x8_t s1, const int16x8_t s2,