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,