Commit af806c9c8c for aom
commit af806c9c8cb616a29db052f2138cba46afe9da4f
Author: Jingning Han <jingning@google.com>
Date: Fri Jul 17 09:37:48 2026 -0700
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 346282f02c..0abd0137c6 100755
--- a/aom_dsp/aom_dsp_rtcd_defs.pl
+++ b/aom_dsp/aom_dsp_rtcd_defs.pl
@@ -1388,6 +1388,9 @@ if (aom_config("CONFIG_AV1_ENCODER") eq "yes") {
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 639e83200c..04096527a3 100644
--- a/aom_dsp/variance.c
+++ b/aom_dsp/variance.c
@@ -1145,6 +1145,53 @@ int64_t aom_calc_variance_stat_c(const uint8_t *src, int stride, int bw,
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 d431c850e5..9b264c83aa 100644
--- a/aom_dsp/x86/variance_avx2.c
+++ b/aom_dsp/x86/variance_avx2.c
@@ -1131,6 +1131,186 @@ int64_t aom_calc_variance_stat_avx2(const uint8_t *src, int stride, int bw,
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 3b75c340cf..5d00742bce 100644
--- a/av1/encoder/rdopt.c
+++ b/av1/encoder/rdopt.c
@@ -632,79 +632,10 @@ static void get_variance_stats_hbd(const MACROBLOCK *x, int64_t *src_var,
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 55b1b99eb5..43ff6f3b8d 100644
--- a/test/variance_test.cc
+++ b/test/variance_test.cc
@@ -3852,6 +3852,49 @@ TEST_P(CalcVarianceStatTest, CompareWithC) {
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