Add AVX2 functions to high bit-depth variance_stats Change-Id: Ib66ac33bf65aa3347cfd0155a067b283f9a3bac7
diff --git a/aom_dsp/aom_dsp_rtcd_defs.pl b/aom_dsp/aom_dsp_rtcd_defs.pl index 346282f..0abd013 100755 --- a/aom_dsp/aom_dsp_rtcd_defs.pl +++ b/aom_dsp/aom_dsp_rtcd_defs.pl
@@ -1388,6 +1388,9 @@ add_proto qw/int64_t/, "aom_calc_variance_stat", "const uint8_t *src, int stride, int bw, int bh"; specialize qw/aom_calc_variance_stat avx2/; + add_proto qw/int64_t/, "aom_highbd_calc_variance_stat", "const uint16_t *src, int stride, int bw, int bh"; + specialize qw/aom_highbd_calc_variance_stat avx2/; + add_proto qw/uint64_t/, "aom_mse_16xh_16bit", "uint8_t *dst, int dstride,uint16_t *src, int w, int h"; specialize qw/aom_mse_16xh_16bit sse2 avx2 neon/;
diff --git a/aom_dsp/variance.c b/aom_dsp/variance.c index 639e832..0409652 100644 --- a/aom_dsp/variance.c +++ b/aom_dsp/variance.c
@@ -1145,6 +1145,53 @@ return var_stats; } +int64_t aom_highbd_calc_variance_stat_c(const uint16_t *src, int stride, int bw, + int bh) { + DECLARE_ALIGNED(16, uint16_t, dclevel[(MAX_SB_SIZE + 2) * (MAX_SB_SIZE + 2)]); + int pstride = bw + 2; + uint16_t *pred_ptr = &dclevel[pstride + 1]; + + static const int gau_filter[3][3] = { + { 1, 2, 1 }, + { 2, 4, 2 }, + { 1, 2, 1 }, + }; + + for (int idy = -1; idy < bh + 1; ++idy) { + for (int idx = -1; idx < bw + 1; ++idx) { + int offset_idy = idy; + int offset_idx = idx; + if (idy == -1) offset_idy = 0; + if (idy == bh) offset_idy = bh - 1; + if (idx == -1) offset_idx = 0; + if (idx == bw) offset_idx = bw - 1; + + int offset = offset_idy * stride + offset_idx; + pred_ptr[idy * pstride + idx] = src[offset]; + } + } + + int64_t var_stats = 0; + + for (int idy = 0; idy < bh; ++idy) { + for (int idx = 0; idx < bw; ++idx) { + int sum = 0; + for (int iy = 0; iy < 3; ++iy) + for (int ix = 0; ix < 3; ++ix) + sum += pred_ptr[(idy + iy - 1) * pstride + (idx + ix - 1)] * + gau_filter[iy][ix]; + + sum = sum >> 4; + + int64_t diff = pred_ptr[idy * pstride + idx] - sum; + var_stats += diff * diff; + } + } + var_stats <<= 4; + + return var_stats; +} + #if CONFIG_AV1_HIGHBITDEPTH uint64_t aom_mse_wxh_16bit_highbd_c(uint16_t *dst, int dstride, uint16_t *src, int sstride, int w, int h) {
diff --git a/aom_dsp/x86/variance_avx2.c b/aom_dsp/x86/variance_avx2.c index d431c85..9b264c8 100644 --- a/aom_dsp/x86/variance_avx2.c +++ b/aom_dsp/x86/variance_avx2.c
@@ -1131,6 +1131,186 @@ return total_var << 4; } +static inline int64_t yy_hsum_epi64_si64(__m256i v) { + __m128i v128 = + _mm_add_epi64(_mm256_castsi256_si128(v), _mm256_extracti128_si256(v, 1)); + __m128i tmp = _mm_srli_si128(v128, 8); + v128 = _mm_add_epi64(v128, tmp); + +#if AOM_ARCH_X86_64 + return _mm_cvtsi128_si64(v128); +#else + int64_t tmp32; + _mm_storel_epi64((__m128i *)&tmp32, v128); + return tmp32; +#endif +} + +static inline int64_t xx_hsum_epi64_si64(__m128i v) { + __m128i tmp = _mm_srli_si128(v, 8); + v = _mm_add_epi64(v, tmp); + +#if AOM_ARCH_X86_64 + return _mm_cvtsi128_si64(v); +#else + int64_t tmp32; + _mm_storel_epi64((__m128i *)&tmp32, v); + return tmp32; +#endif +} + +int64_t aom_highbd_calc_variance_stat_avx2(const uint16_t *src, int stride, + int bw, int bh) { + // Temporary buffer to store horizontal filter results H[y][x] + DECLARE_ALIGNED(32, uint16_t, H_buf[128 * 128]); + + // Step 1: Compute Horizontal 1D Filter H[y][x] = P(y, x-1) + 2*P(y, x) + P(y, + // x + 1) + for (int y = 0; y < bh; ++y) { + const uint16_t *src_row = src + y * stride; + uint16_t *H_row = H_buf + y * bw; + + if (bw >= 8) { + for (int x = 0; x < bw; x += 8) { + __m128i v_curr = _mm_loadu_si128((const __m128i *)(src_row + x)); + __m128i v_left, v_right; + + if (x == 0) { + v_left = _mm_insert_epi16(_mm_slli_si128(v_curr, 2), src_row[0], 0); + } else { + v_left = _mm_loadu_si128((const __m128i *)(src_row + x - 1)); + } + + if (x + 8 < bw) { + v_right = _mm_loadu_si128((const __m128i *)(src_row + x + 1)); + } else { + v_right = + _mm_insert_epi16(_mm_srli_si128(v_curr, 2), src_row[bw - 1], 7); + } + + __m128i u16_H = _mm_add_epi16(_mm_add_epi16(v_left, v_right), + _mm_slli_epi16(v_curr, 1)); + + _mm_storeu_si128((__m128i *)(H_row + x), u16_H); + } + } else { // bw == 4 + __m128i v_curr = _mm_loadl_epi64((const __m128i *)src_row); + __m128i v_left = + _mm_insert_epi16(_mm_slli_si128(v_curr, 2), src_row[0], 0); + __m128i v_right = + _mm_insert_epi16(_mm_srli_si128(v_curr, 2), src_row[3], 3); + + __m128i u16_H = _mm_add_epi16(_mm_add_epi16(v_left, v_right), + _mm_slli_epi16(v_curr, 1)); + + _mm_storel_epi64((__m128i *)H_row, u16_H); + } + } + + // Step 2: Compute Vertical Filter V[y][x] = H(y-1, x) + 2*H(y, x) + H(y + 1, + // x), smooth = V >> 4, diff = P - smooth, and accum (diff^2) + int64_t total_var = 0; + + if (bw >= 16) { + __m256i acc_var_64 = _mm256_setzero_si256(); + + for (int y = 0; y < bh; ++y) { + const uint16_t *src_row = src + y * stride; + const uint16_t *H_curr_row = H_buf + y * bw; + const uint16_t *H_top_row = (y == 0) ? H_curr_row : H_buf + (y - 1) * bw; + const uint16_t *H_bot_row = + (y == bh - 1) ? H_curr_row : H_buf + (y + 1) * bw; + + for (int x = 0; x < bw; x += 16) { + __m256i H_top = _mm256_loadu_si256((const __m256i *)(H_top_row + x)); + __m256i H_curr = _mm256_loadu_si256((const __m256i *)(H_curr_row + x)); + __m256i H_bot = _mm256_loadu_si256((const __m256i *)(H_bot_row + x)); + + __m256i u16_V = _mm256_add_epi16(_mm256_add_epi16(H_top, H_bot), + _mm256_slli_epi16(H_curr, 1)); + + __m256i u16_sum = _mm256_srli_epi16(u16_V, 4); + + __m256i v_p_curr = _mm256_loadu_si256((const __m256i *)(src_row + x)); + + __m256i diff = _mm256_sub_epi16(v_p_curr, u16_sum); + __m256i diff_sq = _mm256_madd_epi16(diff, diff); + + __m256i diff_sq_lo = + _mm256_cvtepi32_epi64(_mm256_castsi256_si128(diff_sq)); + __m256i diff_sq_hi = + _mm256_cvtepi32_epi64(_mm256_extracti128_si256(diff_sq, 1)); + acc_var_64 = _mm256_add_epi64(acc_var_64, diff_sq_lo); + acc_var_64 = _mm256_add_epi64(acc_var_64, diff_sq_hi); + } + } + + total_var = yy_hsum_epi64_si64(acc_var_64); + } else if (bw == 8) { + __m128i acc_var_64 = _mm_setzero_si128(); + + for (int y = 0; y < bh; ++y) { + const uint16_t *src_row = src + y * stride; + const uint16_t *H_curr_row = H_buf + y * 8; + const uint16_t *H_top_row = (y == 0) ? H_curr_row : H_buf + (y - 1) * 8; + const uint16_t *H_bot_row = + (y == bh - 1) ? H_curr_row : H_buf + (y + 1) * 8; + + __m128i H_top = _mm_loadu_si128((const __m128i *)H_top_row); + __m128i H_curr = _mm_loadu_si128((const __m128i *)H_curr_row); + __m128i H_bot = _mm_loadu_si128((const __m128i *)H_bot_row); + + __m128i u16_V = + _mm_add_epi16(_mm_add_epi16(H_top, H_bot), _mm_slli_epi16(H_curr, 1)); + + __m128i u16_sum = _mm_srli_epi16(u16_V, 4); + + __m128i v_p_curr = _mm_loadu_si128((const __m128i *)src_row); + + __m128i diff = _mm_sub_epi16(v_p_curr, u16_sum); + __m128i diff_sq = _mm_madd_epi16(diff, diff); + + __m128i diff_sq_lo = _mm_cvtepi32_epi64(diff_sq); + __m128i diff_sq_hi = _mm_cvtepi32_epi64(_mm_srli_si128(diff_sq, 8)); + acc_var_64 = _mm_add_epi64(acc_var_64, diff_sq_lo); + acc_var_64 = _mm_add_epi64(acc_var_64, diff_sq_hi); + } + + total_var = xx_hsum_epi64_si64(acc_var_64); + } else { // bw == 4 + __m128i acc_var_64 = _mm_setzero_si128(); + + for (int y = 0; y < bh; ++y) { + const uint16_t *src_row = src + y * stride; + const uint16_t *H_curr_row = H_buf + y * 4; + const uint16_t *H_top_row = (y == 0) ? H_curr_row : H_buf + (y - 1) * 4; + const uint16_t *H_bot_row = + (y == bh - 1) ? H_curr_row : H_buf + (y + 1) * 4; + + __m128i H_top = _mm_loadl_epi64((const __m128i *)H_top_row); + __m128i H_curr = _mm_loadl_epi64((const __m128i *)H_curr_row); + __m128i H_bot = _mm_loadl_epi64((const __m128i *)H_bot_row); + + __m128i u16_V = + _mm_add_epi16(_mm_add_epi16(H_top, H_bot), _mm_slli_epi16(H_curr, 1)); + + __m128i u16_sum = _mm_srli_epi16(u16_V, 4); + + __m128i v_p_curr = _mm_loadl_epi64((const __m128i *)src_row); + + __m128i diff = _mm_sub_epi16(v_p_curr, u16_sum); + __m128i diff_sq = _mm_madd_epi16(diff, diff); + + __m128i diff_sq_lo = _mm_cvtepi32_epi64(diff_sq); + acc_var_64 = _mm_add_epi64(acc_var_64, diff_sq_lo); + } + + total_var = xx_hsum_epi64_si64(acc_var_64); + } + + return total_var << 4; +} + void aom_get_var_sse_sum_8x8_quad_avx2(const uint8_t *src_ptr, int source_stride, const uint8_t *ref_ptr, int ref_stride,
diff --git a/av1/encoder/rdopt.c b/av1/encoder/rdopt.c index 3b75c34..5d00742 100644 --- a/av1/encoder/rdopt.c +++ b/av1/encoder/rdopt.c
@@ -632,79 +632,10 @@ int bw = block_size_wide[bsize]; int bh = block_size_high[bsize]; - static const int gau_filter[3][3] = { - { 1, 2, 1 }, - { 2, 4, 2 }, - { 1, 2, 1 }, - }; - - DECLARE_ALIGNED(16, uint16_t, dclevel[(MAX_SB_SIZE + 2) * (MAX_SB_SIZE + 2)]); - - uint16_t *pred_ptr = &dclevel[bw + 1]; - int pred_stride = xd->plane[0].dst.stride; - - for (int idy = -1; idy < bh + 1; ++idy) { - for (int idx = -1; idx < bw + 1; ++idx) { - int offset_idy = idy; - int offset_idx = idx; - if (idy == -1) offset_idy = 0; - if (idy == bh) offset_idy = bh - 1; - if (idx == -1) offset_idx = 0; - if (idx == bw) offset_idx = bw - 1; - - int offset = offset_idy * pred_stride + offset_idx; - pred_ptr[idy * bw + idx] = CONVERT_TO_SHORTPTR(pd->dst.buf)[offset]; - } - } - - *rec_var = 0; - for (int idy = 0; idy < bh; ++idy) { - for (int idx = 0; idx < bw; ++idx) { - int sum = 0; - for (int iy = 0; iy < 3; ++iy) - for (int ix = 0; ix < 3; ++ix) - sum += pred_ptr[(idy + iy - 1) * bw + (idx + ix - 1)] * - gau_filter[iy][ix]; - - sum = sum >> 4; - - int64_t diff = pred_ptr[idy * bw + idx] - sum; - *rec_var += diff * diff; - } - } - *rec_var <<= 4; - - int src_stride = p->src.stride; - for (int idy = -1; idy < bh + 1; ++idy) { - for (int idx = -1; idx < bw + 1; ++idx) { - int offset_idy = idy; - int offset_idx = idx; - if (idy == -1) offset_idy = 0; - if (idy == bh) offset_idy = bh - 1; - if (idx == -1) offset_idx = 0; - if (idx == bw) offset_idx = bw - 1; - - int offset = offset_idy * src_stride + offset_idx; - pred_ptr[idy * bw + idx] = CONVERT_TO_SHORTPTR(p->src.buf)[offset]; - } - } - - *src_var = 0; - for (int idy = 0; idy < bh; ++idy) { - for (int idx = 0; idx < bw; ++idx) { - int sum = 0; - for (int iy = 0; iy < 3; ++iy) - for (int ix = 0; ix < 3; ++ix) - sum += pred_ptr[(idy + iy - 1) * bw + (idx + ix - 1)] * - gau_filter[iy][ix]; - - sum = sum >> 4; - - int64_t diff = pred_ptr[idy * bw + idx] - sum; - *src_var += diff * diff; - } - } - *src_var <<= 4; + *rec_var = aom_highbd_calc_variance_stat(CONVERT_TO_SHORTPTR(pd->dst.buf), + pd->dst.stride, bw, bh); + *src_var = aom_highbd_calc_variance_stat(CONVERT_TO_SHORTPTR(p->src.buf), + p->src.stride, bw, bh); } static void get_variance_stats(const MACROBLOCK *x, int64_t *src_var,
diff --git a/test/variance_test.cc b/test/variance_test.cc index 55b1b99..43ff6f3 100644 --- a/test/variance_test.cc +++ b/test/variance_test.cc
@@ -3852,6 +3852,49 @@ INSTANTIATE_TEST_SUITE_P(AVX2, CalcVarianceStatTest, ::testing::Values(&aom_calc_variance_stat_avx2)); + +#if CONFIG_AV1_HIGHBITDEPTH +using CalcVarianceStatHbdFunc = int64_t (*)(const uint16_t *src, int stride, + int bw, int bh); + +class CalcVarianceStatHbdTest + : public ::testing::TestWithParam<CalcVarianceStatHbdFunc> { + protected: + void SetUp() override { + target_func_ = GetParam(); + rnd_.Reset(ACMRandom::DeterministicSeed()); + } + + CalcVarianceStatHbdFunc target_func_; + ACMRandom rnd_; +}; + +TEST_P(CalcVarianceStatHbdTest, CompareWithC) { + static const int kSizes[] = { 4, 8, 16, 32, 64, 128 }; + DECLARE_ALIGNED(32, uint16_t, src[128 * 128]); + + for (int w : kSizes) { + for (int h : kSizes) { + SCOPED_TRACE(::testing::Message() << "bw=" << w << " bh=" << h); + int stride = 128; + for (int iter = 0; iter < 500; ++iter) { + for (int r = 0; r < h; ++r) { + for (int c = 0; c < w; ++c) { + src[r * stride + c] = rnd_.Rand16() & 0xfff; // 12-bit + } + } + int64_t res_c = aom_highbd_calc_variance_stat_c(src, stride, w, h); + int64_t res_target = target_func_(src, stride, w, h); + EXPECT_EQ(res_c, res_target) << "iter=" << iter; + } + } + } +} + +INSTANTIATE_TEST_SUITE_P( + AVX2, CalcVarianceStatHbdTest, + ::testing::Values(&aom_highbd_calc_variance_stat_avx2)); +#endif // CONFIG_AV1_HIGHBITDEPTH #endif // HAVE_AVX2 } // namespace