diff --git a/av2/common/av2_rtcd_defs.pl b/av2/common/av2_rtcd_defs.pl index 9be6c38a5d..90e311299f 100644 --- a/av2/common/av2_rtcd_defs.pl +++ b/av2/common/av2_rtcd_defs.pl @@ -410,7 +410,7 @@ () add_proto qw/uint64_t compute_distortion_block/, "const uint16_t *org, const int org_stride, const uint16_t *rec16, const int rec_stride, const int x, const int y, const int log2_filter_unit_size_y, const int log2_filter_unit_size_x, const int height, - const int width"; + const int width, const int bd"; specialize qw/compute_distortion_block avx2/; add_proto qw/void ccso_derive_src_block/, "const uint16_t *src_y, uint8_t *const src_cls0, diff --git a/av2/common/x86/highbd_ccso_avx2.c b/av2/common/x86/highbd_ccso_avx2.c index bb937fc088..66068a5e17 100644 --- a/av2/common/x86/highbd_ccso_avx2.c +++ b/av2/common/x86/highbd_ccso_avx2.c @@ -15,26 +15,125 @@ #include "av2/common/ccso.h" +static const uint8_t shuf_even_mask_8bit[32] = { 0, 2, 4, 6, 8, 10, 12, 14, + 0, 0, 0, 0, 0, 0, 0, 0, + 0, 2, 4, 6, 8, 10, 12, 14, + 0, 0, 0, 0, 0, 0, 0, 0 }; + +static const uint8_t shuf_even_mask_16bit[32] = { 0, 1, 4, 5, 8, 9, 12, 13, + 0, 0, 0, 0, 0, 0, 0, 0, + 0, 1, 4, 5, 8, 9, 12, 13, + 0, 0, 0, 0, 0, 0, 0, 0 }; + +// The logic for d > qstep -> 2, d < -qstep -> 0, otherwise 1 is implemented as +// 1 + (d > qstep) - (d < -qstep) __m256i cal_filter_support_edge0_avx2(__m256i d, __m256i cmp_thr1, - __m256i cmp_thr2, __m256i all1, - __m256i cmp_idxa, __m256i cmp_idxc) { - __m256i idx_mask1a = _mm256_cmpgt_epi16(d, cmp_thr1); - __m256i idx_mask1b = _mm256_cmpgt_epi16(cmp_thr2, d); - __m256i idx_mask1c = _mm256_xor_si256(idx_mask1a, idx_mask1b); - idx_mask1c = _mm256_xor_si256(idx_mask1c, all1); - idx_mask1a = _mm256_and_si256(idx_mask1a, cmp_idxa); - idx_mask1c = _mm256_and_si256(idx_mask1c, cmp_idxc); - - __m256i idx = _mm256_add_epi16(idx_mask1a, idx_mask1c); - return idx; + __m256i cmp_thr2, __m256i cmp_idxc) { + const __m256i gt = _mm256_cmpgt_epi16(d, cmp_thr1); + const __m256i lt = _mm256_cmpgt_epi16(cmp_thr2, d); + return _mm256_add_epi16(_mm256_sub_epi16(cmp_idxc, gt), lt); } -__m256i cal_filter_support_edge1_avx2(__m256i d, __m256i cmp_thr2, __m256i all1, +// The logic for d < -qstep -> 0, otherwise 1 is implemented as +// 1 - (d < -qstep) +__m256i cal_filter_support_edge1_avx2(__m256i d, __m256i cmp_thr2, __m256i cmp_idxc) { - __m256i idx_mask1a = _mm256_cmpgt_epi16(cmp_thr2, d); - __m256i idx_mask1b = _mm256_xor_si256(idx_mask1a, all1); - __m256i idx = _mm256_and_si256(idx_mask1b, cmp_idxc); - return idx; + const __m256i lt = _mm256_cmpgt_epi16(cmp_thr2, d); + return _mm256_add_epi16(cmp_idxc, lt); +} + +static AVM_FORCE_INLINE __m256i extract_even_16bit_avx2(const uint16_t *src) { + const __m256i even_mask = + _mm256_loadu_si256((const __m256i *)(shuf_even_mask_16bit)); + const __m256i src_even_0 = + _mm256_shuffle_epi8(_mm256_loadu_si256((const __m256i *)src), even_mask); + const __m256i src_even_1 = _mm256_shuffle_epi8( + _mm256_loadu_si256((const __m256i *)(src + 16)), even_mask); + return _mm256_permute4x64_epi64(_mm256_unpacklo_epi64(src_even_0, src_even_1), + 0xD8); +} + +static AVM_FORCE_INLINE __m256i +extract_even_8bit_w32_avx2(const uint8_t *src_cls) { + const __m256i even_mask = + _mm256_loadu_si256((const __m256i *)(shuf_even_mask_8bit)); + const __m256i cls_even_0 = _mm256_shuffle_epi8( + _mm256_loadu_si256((const __m256i *)src_cls), even_mask); + const __m256i cls_even_1 = _mm256_shuffle_epi8( + _mm256_loadu_si256((const __m256i *)(src_cls + 32)), even_mask); + return _mm256_permute4x64_epi64(_mm256_unpacklo_epi64(cls_even_0, cls_even_1), + 0xD8); +} + +static AVM_FORCE_INLINE __m128i +extract_even_8bit_w16_avx2(const uint8_t *src_cls) { + const __m128i even_mask = + _mm_loadu_si128((const __m128i *)(shuf_even_mask_8bit)); + const __m128i cls_even_0 = + _mm_shuffle_epi8(_mm_loadu_si128((const __m128i *)src_cls), even_mask); + const __m128i cls_even_1 = _mm_shuffle_epi8( + _mm_loadu_si128((const __m128i *)(src_cls + 16)), even_mask); + return _mm_unpacklo_epi64(cls_even_0, cls_even_1); +} + +static AVM_FORCE_INLINE void add_offset_avx2(uint16_t *dst_rec2, int xOff, + __m256i offset, __m256i allmax) { + __m256i res = _mm256_add_epi16( + offset, _mm256_loadu_si256((const __m256i *)(dst_rec2 + xOff))); + res = _mm256_max_epi16(_mm256_min_epi16(res, allmax), _mm256_setzero_si256()); + _mm256_storeu_si256((__m256i *)(dst_rec2 + xOff), res); +} + +/*! + * This function gets the filter offsets for 32 pixels at once + * from the look-up table populated in filter_offset_lut[8]. + * The look-up table index (lut_idx) is a 8-bit value of the form + * + * bit 7 | 6 5 4 | 3 2 1 0 + * 0 | band | (cls0 << 2) + cls1 + * + * Bits 0 - 4 of the lut_idx are used to select an offset from + * each of the 8 filter_offset_lut[i] via _mm256_shuffle_epi8 + * instruction assuming it belongs to that band. Then, using + * a sequence of _mm256_blendv_epi8() operations the correct + * filter offset corresponding to the pixel's actual band is + * selected, i.e., + * bit 4: selects between bands 0/1, 2/3, 4/5, and 6/7 + * bit 5: selects between bands 0-1 and 2-3, and between bands 4-5 and 6-7 + * bit 6: selects between bands 0-3 and bands 4-7 + */ +static AVM_FORCE_INLINE __m256i get_offset_from_index_avx2( + const __m256i *filter_offset_lut, __m256i lut_idx, int max_band) { + if (max_band == 1) return _mm256_shuffle_epi8(filter_offset_lut[0], lut_idx); + + const __m256i mask_bit_4 = _mm256_slli_epi16(lut_idx, 3); + const __m256i res_band_01 = _mm256_blendv_epi8( + _mm256_shuffle_epi8(filter_offset_lut[0], lut_idx), + _mm256_shuffle_epi8(filter_offset_lut[1], lut_idx), mask_bit_4); + + if (max_band == 2) return res_band_01; + + const __m256i mask_bit_5 = _mm256_slli_epi16(lut_idx, 2); + const __m256i res_band_23 = _mm256_blendv_epi8( + _mm256_shuffle_epi8(filter_offset_lut[2], lut_idx), + _mm256_shuffle_epi8(filter_offset_lut[3], lut_idx), mask_bit_4); + const __m256i res_band_0123 = + _mm256_blendv_epi8(res_band_01, res_band_23, mask_bit_5); + + if (max_band <= 4) return res_band_0123; + + const __m256i res_band_45 = _mm256_blendv_epi8( + _mm256_shuffle_epi8(filter_offset_lut[4], lut_idx), + _mm256_shuffle_epi8(filter_offset_lut[5], lut_idx), mask_bit_4); + const __m256i res_band_67 = _mm256_blendv_epi8( + _mm256_shuffle_epi8(filter_offset_lut[6], lut_idx), + _mm256_shuffle_epi8(filter_offset_lut[7], lut_idx), mask_bit_4); + const __m256i res_band_4567 = + _mm256_blendv_epi8(res_band_45, res_band_67, mask_bit_5); + + const __m256i mask_bit_6 = _mm256_slli_epi16(lut_idx, 1); + + return _mm256_blendv_epi8(res_band_0123, res_band_4567, mask_bit_6); } // AVX2 implementation for ccso band offset only case. @@ -145,45 +244,150 @@ void ccso_filter_block_hbd_wo_buf_bo_only_avx2( } } -void ccso_filter_block_hbd_wo_buf_avx2( - const uint16_t *src_y, uint16_t *dts_yuv, const int x, const int y, - const int pic_width, const int pic_height, int *rec_luma_idx, - const int8_t *offset_buf, - // const int* src_y_stride, const int* dst_stride, - const int src_y_stride, const int dst_stride, const int y_uv_hscale, - const int y_uv_vscale, - // const int pad_stride, no pad size anymore - const int quant_step_size, const int inv_quant_step, const int *rec_idx, - const int max_val, const int blk_size_x, const int blk_size_y, - const bool isSingleBand, const uint8_t shift_bits, const int edge_clf, - const uint8_t ccso_bo_only) { - assert(ccso_bo_only == 0); - (void)ccso_bo_only; - __m256i cmp_thr1 = _mm256_set1_epi16(quant_step_size); - __m256i cmp_thr2 = _mm256_set1_epi16(inv_quant_step); - __m256i cmp_idxa = _mm256_set1_epi16(2); // d > quant_step_size - __m256i cmp_idxc = - _mm256_set1_epi16(1); // -quant_step_size <= d <= quant_step_size +static AVM_FORCE_INLINE void ccso_filter_wo_buf_row_width_32_avx2( + const uint16_t *src_rec, uint16_t *dst_rec2, int xOff, int tap1_pos, + int tap2_pos, int y_uv_hscale, uint8_t shift_bits, int max_band, + int edge_clf, __m256i cmp_thr1, __m256i cmp_thr2, __m256i cmp_idxc, + const __m256i *filter_offset_lut, __m256i allmax) { + __m256i rec_cur_low, rec_cur_high; + __m256i rec_tap1_low, rec_tap1_high, rec_tap2_low, rec_tap2_high; + + if (y_uv_hscale == 0) { + const uint16_t *src = src_rec + xOff; + rec_cur_low = _mm256_loadu_si256((const __m256i *)src); + rec_cur_high = _mm256_loadu_si256((const __m256i *)(src + 16)); + rec_tap1_low = _mm256_loadu_si256((const __m256i *)(src + tap1_pos)); + rec_tap1_high = _mm256_loadu_si256((const __m256i *)(src + tap1_pos + 16)); + rec_tap2_low = _mm256_loadu_si256((const __m256i *)(src + tap2_pos)); + rec_tap2_high = _mm256_loadu_si256((const __m256i *)(src + tap2_pos + 16)); + } else { + const uint16_t *src = src_rec + (xOff << 1); + rec_cur_low = extract_even_16bit_avx2(src); + rec_cur_high = extract_even_16bit_avx2(src + 32); + rec_tap1_low = extract_even_16bit_avx2(src + tap1_pos); + rec_tap1_high = extract_even_16bit_avx2(src + tap1_pos + 32); + rec_tap2_low = extract_even_16bit_avx2(src + tap2_pos); + rec_tap2_high = extract_even_16bit_avx2(src + tap2_pos + 32); + } - __m128i tmp = _mm_lddqu_si128((const __m128i *)offset_buf); - //__m256i ccso_lut = _mm256_setr_m128i(tmp, tmp); - __m256i ccso_lut = - _mm256_insertf128_si256(_mm256_castsi128_si256(tmp), (tmp), 0x1); - __m256i all0 = _mm256_set1_epi16(0); - __m256i all1 = _mm256_set1_epi16(-1); - __m256i allmax = _mm256_set1_epi16(max_val); - __m128i shufsub = - _mm_set_epi8(0, 0, 0, 0, 0, 0, 0, 0, 13, 12, 9, 8, 5, 4, 1, 0); - //__m256i masksub1 = _mm256_set_m128i(shufsub, shufsub); - __m256i masksub1 = - _mm256_insertf128_si256(_mm256_castsi128_si256(shufsub), (shufsub), 0x1); - __m256i d1, d2; + const __m256i d1_low = _mm256_sub_epi16(rec_tap1_low, rec_cur_low); + const __m256i d1_high = _mm256_sub_epi16(rec_tap1_high, rec_cur_high); + const __m256i d2_low = _mm256_sub_epi16(rec_tap2_low, rec_cur_low); + const __m256i d2_high = _mm256_sub_epi16(rec_tap2_high, rec_cur_high); + + __m256i eo_idx0_low, eo_idx0_high, eo_idx1_low, eo_idx1_high; + if (edge_clf == 0) { + eo_idx0_low = + cal_filter_support_edge0_avx2(d1_low, cmp_thr1, cmp_thr2, cmp_idxc); + eo_idx0_high = + cal_filter_support_edge0_avx2(d1_high, cmp_thr1, cmp_thr2, cmp_idxc); + eo_idx1_low = + cal_filter_support_edge0_avx2(d2_low, cmp_thr1, cmp_thr2, cmp_idxc); + eo_idx1_high = + cal_filter_support_edge0_avx2(d2_high, cmp_thr1, cmp_thr2, cmp_idxc); + } else { + eo_idx0_low = cal_filter_support_edge1_avx2(d1_low, cmp_thr2, cmp_idxc); + eo_idx0_high = cal_filter_support_edge1_avx2(d1_high, cmp_thr2, cmp_idxc); + eo_idx1_low = cal_filter_support_edge1_avx2(d2_low, cmp_thr2, cmp_idxc); + eo_idx1_high = cal_filter_support_edge1_avx2(d2_high, cmp_thr2, cmp_idxc); + } + + // lut_idx = (band_num << 4) + (rec_luma_idx[0] << 2) + rec_luma_idx[1] + const __m256i bo_idx_low = _mm256_srli_epi16(rec_cur_low, shift_bits); + const __m256i bo_idx_high = _mm256_srli_epi16(rec_cur_high, shift_bits); + + __m256i idx_low = + _mm256_add_epi16(_mm256_slli_epi16(eo_idx0_low, 2), eo_idx1_low); + idx_low = _mm256_add_epi16(idx_low, _mm256_slli_epi16(bo_idx_low, 4)); + __m256i idx_high = + _mm256_add_epi16(_mm256_slli_epi16(eo_idx0_high, 2), eo_idx1_high); + idx_high = _mm256_add_epi16(idx_high, _mm256_slli_epi16(bo_idx_high, 4)); + + const __m256i lut_idx = + _mm256_permute4x64_epi64(_mm256_packus_epi16(idx_low, idx_high), 0xD8); + + const __m256i offset = + get_offset_from_index_avx2(filter_offset_lut, lut_idx, max_band); + add_offset_avx2(dst_rec2, xOff, + _mm256_cvtepi8_epi16(_mm256_castsi256_si128(offset)), allmax); + add_offset_avx2(dst_rec2, xOff + 16, + _mm256_cvtepi8_epi16(_mm256_extracti128_si256(offset, 1)), + allmax); +} + +static AVM_FORCE_INLINE void ccso_filter_wo_buf_row_width_16_avx2( + const uint16_t *src_rec, uint16_t *dst_rec2, int xOff, int tap1_pos, + int tap2_pos, int y_uv_hscale, uint8_t shift_bits, int max_band, + int edge_clf, __m256i cmp_thr1, __m256i cmp_thr2, __m256i cmp_idxc, + const __m256i *filter_offset_lut, __m256i allmax) { + __m256i rec_cur, rec_tap1, rec_tap2; + + if (y_uv_hscale == 0) { + const uint16_t *src = src_rec + xOff; + rec_cur = _mm256_loadu_si256((const __m256i *)src); + rec_tap1 = _mm256_loadu_si256((const __m256i *)(src + tap1_pos)); + rec_tap2 = _mm256_loadu_si256((const __m256i *)(src + tap2_pos)); + } else { + const uint16_t *src = src_rec + (xOff << 1); + rec_cur = extract_even_16bit_avx2(src); + rec_tap1 = extract_even_16bit_avx2(src + tap1_pos); + rec_tap2 = extract_even_16bit_avx2(src + tap2_pos); + } + + // d1 = rec[tap1_pos] - rec[0], d2 = rec[tap2_pos] - rec[0] + const __m256i d1 = _mm256_sub_epi16(rec_tap1, rec_cur); + const __m256i d2 = _mm256_sub_epi16(rec_tap2, rec_cur); + + __m256i eo_idx0, eo_idx1; + if (edge_clf == 0) { + eo_idx0 = cal_filter_support_edge0_avx2(d1, cmp_thr1, cmp_thr2, cmp_idxc); + eo_idx1 = cal_filter_support_edge0_avx2(d2, cmp_thr1, cmp_thr2, cmp_idxc); + } else { + eo_idx0 = cal_filter_support_edge1_avx2(d1, cmp_thr2, cmp_idxc); + eo_idx1 = cal_filter_support_edge1_avx2(d2, cmp_thr2, cmp_idxc); + } + + // lut_idx = (band_num << 4) + (rec_luma_idx[0] << 2) + rec_luma_idx[1] + const __m256i bo_idx = _mm256_srli_epi16(rec_cur, shift_bits); + __m256i lut_idx_16 = _mm256_add_epi16(_mm256_slli_epi16(eo_idx0, 2), eo_idx1); + lut_idx_16 = _mm256_add_epi16(lut_idx_16, _mm256_slli_epi16(bo_idx, 4)); - int tap1_pos = rec_idx[0]; - int tap2_pos = rec_idx[1]; + const __m128i lut_idx = + _mm_packus_epi16(_mm256_castsi256_si128(lut_idx_16), + _mm256_extracti128_si256(lut_idx_16, 1)); + const __m256i offset_256 = get_offset_from_index_avx2( + filter_offset_lut, _mm256_castsi128_si256(lut_idx), max_band); + const __m128i offset = _mm256_castsi256_si128(offset_256); + add_offset_avx2(dst_rec2, xOff, _mm256_cvtepi8_epi16(offset), allmax); +} + +static AVM_FORCE_INLINE void ccso_filter_wo_buf_row_width_remainder_avx2( + const uint16_t *src_rec, uint16_t *dst_rec2, int xOff, int x_remainder, + int *rec_luma_idx, const int8_t *offset_buf, int y_uv_hscale, + int quant_step_size, int inv_quant_step, const int *rec_idx, int max_val, + bool isSingleBand, uint8_t shift_bits, int edge_clf) { + for (int i = xOff; i < xOff + x_remainder; i++) { + const uint16_t *src = &src_rec[i << y_uv_hscale]; + cal_filter_support(rec_luma_idx, src, quant_step_size, inv_quant_step, + rec_idx, edge_clf); + const int band_num = isSingleBand ? 0 : src[0] >> shift_bits; + const int lut_idx = + (band_num << 4) + (rec_luma_idx[0] << 2) + rec_luma_idx[1]; + dst_rec2[i] = clamp(offset_buf[lut_idx] + dst_rec2[i], 0, max_val); + } +} + +static AVM_FORCE_INLINE void ccso_filter_wo_buf_block( + const uint16_t *src_y, uint16_t *dts_yuv, int x, int y, int pic_width, + int pic_height, int blk_size_x, int blk_size_y, int *rec_luma_idx, + const int8_t *offset_buf, int src_y_stride, int dst_stride, int y_uv_vscale, + int quant_step_size, int inv_quant_step, const int *rec_idx, int max_val, + bool isSingleBand, uint8_t shift_bits, int edge_clf, int y_uv_hscale, + int max_band) { int y_offset; int x_offset, x_remainder; + if (y + blk_size_y >= pic_height) y_offset = pic_height - y; else @@ -196,148 +400,100 @@ void ccso_filter_block_hbd_wo_buf_avx2( x_offset = blk_size_x; x_remainder = 0; } - for (int yOff = 0; yOff < y_offset; yOff++) { - // uint16_t* dst_rec2 = dts_yuv + x + dst_stride[yOff]; - uint16_t *dst_rec2 = dts_yuv + x + yOff * dst_stride; - // const uint16_t* src_rec2 = src_y + ((src_y_stride[yOff] << y_uv_vscale) + - // (x << y_uv_hscale)) + pad_stride; - const uint16_t *src_rec2 = - src_y + ((yOff << y_uv_vscale) * src_y_stride + (x << y_uv_hscale)); - // int stride = src_y_stride[yOff] << y_uv_vscale; - for (int xOff = 0; xOff < x_offset; xOff += 16) { - // uint16_t* rec_tmp = &src_rec2[xOff << y_uv_hscale]; - __m256i rec_curlo = _mm256_lddqu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale))); - __m256i rec_cur_final; + const __m256i cmp_thr1 = _mm256_set1_epi16(quant_step_size); + const __m256i cmp_thr2 = _mm256_set1_epi16(inv_quant_step); + const __m256i cmp_idxc = _mm256_set1_epi16(1); + const __m256i allmax = _mm256_set1_epi16((short)max_val); + + __m256i filter_offset_lut[8]; + for (int band_num = 0; band_num < 8; band_num++) { + if (band_num < max_band) { + filter_offset_lut[band_num] = _mm256_broadcastsi128_si256( + _mm_loadu_si128((const __m128i *)(offset_buf + (band_num << 4)))); + } else { + filter_offset_lut[band_num] = _mm256_setzero_si256(); + } + } - //__m256i rec_tap1 = _mm256_loadu_si256((const __m256i*)(src_rec2 + (xOff - //<< y_uv_hscale) + tap1_pos)); - __m256i rec_tap1lo = _mm256_lddqu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap1_pos)); - //__m256i rec_tap2 = _mm256_loadu_si256((const __m256i*)(src_rec2 + (xOff - //<< y_uv_hscale) + tap2_pos)); - __m256i rec_tap2lo = _mm256_lddqu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap2_pos)); + const int tap1_pos = rec_idx[0]; + const int tap2_pos = rec_idx[1]; - if (y_uv_hscale > 0) { - __m256i rec_curhi = _mm256_lddqu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + 16)); - rec_curlo = _mm256_shuffle_epi8(rec_curlo, masksub1); - rec_curhi = _mm256_shuffle_epi8(rec_curhi, masksub1); - __m256i rec_cur = _mm256_unpacklo_epi64(rec_curlo, rec_curhi); - rec_cur = _mm256_permute4x64_epi64(rec_cur, 0xD8); - //__m256i rec_cur = _mm256_setr_m128i(_mm256_castsi256_si128(rec_curlo), - //_mm256_castsi256_si128(rec_curhi)); - - __m256i rec_tap1hi = _mm256_lddqu_si256(( - const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap1_pos + 16)); - rec_tap1lo = _mm256_shuffle_epi8(rec_tap1lo, masksub1); - rec_tap1hi = _mm256_shuffle_epi8(rec_tap1hi, masksub1); - __m256i rec_tap1 = _mm256_unpacklo_epi64(rec_tap1lo, rec_tap1hi); - rec_tap1 = _mm256_permute4x64_epi64(rec_tap1, 0xD8); - //__m256i rec_tap1 = - //_mm256_setr_m128i(_mm256_castsi256_si128(rec_tap1lo), - //_mm256_castsi256_si128(rec_tap1hi)); - - __m256i rec_tap2hi = _mm256_lddqu_si256(( - const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap2_pos + 16)); - rec_tap2lo = _mm256_shuffle_epi8(rec_tap2lo, masksub1); - rec_tap2hi = _mm256_shuffle_epi8(rec_tap2hi, masksub1); - __m256i rec_tap2 = _mm256_unpacklo_epi64(rec_tap2lo, rec_tap2hi); - rec_tap2 = _mm256_permute4x64_epi64(rec_tap2, 0xD8); - //__m256i rec_tap2 = - //_mm256_setr_m128i(_mm256_castsi256_si128(rec_tap2lo), - //_mm256_castsi256_si128(rec_tap2hi)); - - // int d1 = rec_tmp[tap1_pos] - rec_tmp[0]; - // int d2 = rec_tmp[tap2_pos] - rec_tmp[0]; - d1 = _mm256_sub_epi16(rec_tap1, rec_cur); - d2 = _mm256_sub_epi16(rec_tap2, rec_cur); - rec_cur_final = rec_cur; - } else { - d1 = _mm256_sub_epi16(rec_tap1lo, rec_curlo); - d2 = _mm256_sub_epi16(rec_tap2lo, rec_curlo); - rec_cur_final = rec_curlo; - } - __m256i dst_rec = _mm256_lddqu_si256((const __m256i *)(dst_rec2 + xOff)); - __m256i idx1, idx2; - if (edge_clf == 0) { - idx1 = cal_filter_support_edge0_avx2(d1, cmp_thr1, cmp_thr2, all1, - cmp_idxa, cmp_idxc); - idx2 = cal_filter_support_edge0_avx2(d2, cmp_thr1, cmp_thr2, all1, - cmp_idxa, cmp_idxc); - } else { // if (edge_clf == 1) - idx1 = cal_filter_support_edge1_avx2(d1, cmp_thr2, all1, cmp_idxc); - idx2 = cal_filter_support_edge1_avx2(d2, cmp_thr2, all1, cmp_idxc); - } + uint16_t *dst_rec2 = dts_yuv + x; + const uint16_t *src_rec2 = src_y + (x << y_uv_hscale); + const int src_rec2_stride = src_y_stride << y_uv_vscale; - __m256i offset; - // const int band_num = src_y[x_pos] >> shift_bits; - __m256i num_band = - isSingleBand ? all0 : _mm256_srli_epi16(rec_cur_final, shift_bits); + for (int yOff = 0; yOff < y_offset; yOff++) { + int xOff = 0; - // const int lut_idx_ext = (band_num << 4) + (src_cls[0] << 2) + - // src_cls[1]; - num_band = _mm256_slli_epi16(num_band, 4); - idx1 = _mm256_slli_epi16(idx1, 2); - idx2 = _mm256_add_epi16(idx1, idx2); - idx2 = _mm256_add_epi16(num_band, idx2); - - if (isSingleBand) { - idx2 = _mm256_packus_epi16(idx2, idx2); - idx2 = _mm256_permute4x64_epi64(idx2, 0x08); - offset = _mm256_shuffle_epi8(ccso_lut, idx2); - offset = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(offset, 0)); - } else { - // multiple band offset implementation - __m256i idx2_lo = _mm256_unpacklo_epi16(idx2, all0); - __m256i idx2_hi = _mm256_unpackhi_epi16(idx2, all0); + for (; xOff + 32 <= x_offset; xOff += 32) { + ccso_filter_wo_buf_row_width_32_avx2( + src_rec2, dst_rec2, xOff, tap1_pos, tap2_pos, y_uv_hscale, shift_bits, + max_band, edge_clf, cmp_thr1, cmp_thr2, cmp_idxc, filter_offset_lut, + allmax); + } - __m256i offset_val_lo = - _mm256_i32gather_epi32((int *)offset_buf, idx2_lo, 1); - __m256i offset_val_hi = - _mm256_i32gather_epi32((int *)offset_buf, idx2_hi, 1); - offset_val_lo = - _mm256_shuffle_epi8(offset_val_lo, _mm256_set1_epi32(0x0c080400u)); - offset_val_hi = - _mm256_shuffle_epi8(offset_val_hi, _mm256_set1_epi32(0x0c080400u)); - __m256i offset_val = - _mm256_unpacklo_epi32(offset_val_lo, offset_val_hi); - __m256i sign_bits = _mm256_cmpgt_epi8(all0, offset_val); - offset = _mm256_unpacklo_epi8(offset_val, sign_bits); - } + for (; xOff < x_offset; xOff += 16) { + ccso_filter_wo_buf_row_width_16_avx2( + src_rec2, dst_rec2, xOff, tap1_pos, tap2_pos, y_uv_hscale, shift_bits, + max_band, edge_clf, cmp_thr1, cmp_thr2, cmp_idxc, filter_offset_lut, + allmax); + } - // uint16_t val = clamp(offset_val + dst_rec2[xOff], 0, (1 << - // cm->seq_params.bit_depth) - 1); - __m256i recon = _mm256_add_epi16(offset, dst_rec); - recon = _mm256_min_epi16(recon, allmax); - recon = _mm256_max_epi16(recon, all0); + if (x_remainder) + ccso_filter_wo_buf_row_width_remainder_avx2( + src_rec2, dst_rec2, x_offset, x_remainder, rec_luma_idx, offset_buf, + y_uv_hscale, quant_step_size, inv_quant_step, rec_idx, max_val, + isSingleBand, shift_bits, edge_clf); - // dst_rec2[xOff] = val; - _mm256_storeu_si256((__m256i *)(dst_rec2 + xOff), recon); + dst_rec2 += dst_stride; + src_rec2 += src_rec2_stride; + } +} + +void ccso_filter_block_hbd_wo_buf_avx2( + const uint16_t *src_y, uint16_t *dts_yuv, const int x, const int y, + const int pic_width, const int pic_height, int *rec_luma_idx, + const int8_t *offset_buf, + // const int* src_y_stride, const int* dst_stride, + const int src_y_stride, const int dst_stride, const int y_uv_hscale, + const int y_uv_vscale, + // const int pad_stride, no pad size anymore + const int quant_step_size, const int inv_quant_step, const int *rec_idx, + const int max_val, const int blk_size_x, const int blk_size_y, + const bool isSingleBand, const uint8_t shift_bits, const int edge_clf, + const uint8_t ccso_bo_only) { + assert(ccso_bo_only == 0); + (void)ccso_bo_only; + + // Number of bands in use: 1, 2, 4 or 8. + const int max_band = isSingleBand ? 1 : (max_val >> shift_bits) + 1; + +#define CCSO_FILTER_WO_BUF_BLOCK(MAX_BAND, HORIZONTAL_SCALE) \ + ccso_filter_wo_buf_block( \ + src_y, dts_yuv, x, y, pic_width, pic_height, blk_size_x, blk_size_y, \ + rec_luma_idx, offset_buf, src_y_stride, dst_stride, y_uv_vscale, \ + quant_step_size, inv_quant_step, rec_idx, max_val, isSingleBand, \ + shift_bits, edge_clf, (HORIZONTAL_SCALE), (MAX_BAND)) + + if (y_uv_hscale == 0) { + switch (max_band) { + case 1: CCSO_FILTER_WO_BUF_BLOCK(1, 0); break; + case 2: CCSO_FILTER_WO_BUF_BLOCK(2, 0); break; + case 4: CCSO_FILTER_WO_BUF_BLOCK(4, 0); break; + case 8: CCSO_FILTER_WO_BUF_BLOCK(8, 0); break; + default: assert(0); break; } - for (int xOff = x_offset; xOff < x_offset + x_remainder; xOff++) { - // cal_filter_support(rec_luma_idx, &src_y[((src_y_stride[yOff] << - // y_uv_vscale) + ((x + xOff) << y_uv_hscale)) + pad_stride], - // quant_step_size, inv_quant_step, rec_idx); - cal_filter_support(rec_luma_idx, - &src_y[((yOff << y_uv_vscale) * src_y_stride + - ((x + xOff) << y_uv_hscale))], - quant_step_size, inv_quant_step, rec_idx, edge_clf); - const int band_num = isSingleBand - ? 0 - : src_y[((yOff << y_uv_vscale) * src_y_stride + - ((x + xOff) << y_uv_hscale))] >> - shift_bits; - int offset_val = offset_buf[(band_num << 4) + (rec_luma_idx[0] << 2) + - rec_luma_idx[1]]; - // dts_yuv[dst_stride[yOff] + x + xOff] = clamp(offset_val + - // dts_yuv[dst_stride[yOff] + x + xOff], 0, max_val); - dts_yuv[yOff * dst_stride + x + xOff] = - clamp(offset_val + dts_yuv[yOff * dst_stride + x + xOff], 0, max_val); + } else { + switch (max_band) { + case 1: CCSO_FILTER_WO_BUF_BLOCK(1, 1); break; + case 2: CCSO_FILTER_WO_BUF_BLOCK(2, 1); break; + case 4: CCSO_FILTER_WO_BUF_BLOCK(4, 1); break; + case 8: CCSO_FILTER_WO_BUF_BLOCK(8, 1); break; + default: assert(0); break; } } +#undef CCSO_FILTER_WO_BUF_BLOCK } void ccso_derive_src_block_avx2(const uint16_t *src_y, uint8_t *const src_cls0, uint8_t *const src_cls1, const int src_y_stride, @@ -351,7 +507,6 @@ void ccso_derive_src_block_avx2(const uint16_t *src_y, uint8_t *const src_cls0, const int inv_quant_step = neg_qstep; __m256i cmp_thr1 = _mm256_set1_epi16(quant_step_size); __m256i cmp_thr2 = _mm256_set1_epi16(inv_quant_step); - __m256i cmp_idxa = _mm256_set1_epi16(2); // d > quant_step_size __m256i cmp_idxc = _mm256_set1_epi16(1); // -quant_step_size <= d <= quant_step_size @@ -359,13 +514,7 @@ void ccso_derive_src_block_avx2(const uint16_t *src_y, uint8_t *const src_cls0, //__m256i ccso_lut = _mm256_setr_m128i(tmp, tmp); //__m256i all0 = _mm256_set1_epi16(0); __m128i all0_128 = _mm_setzero_si128(); - __m256i all1 = _mm256_set1_epi16(-1); //__m256i allmax = _mm256_set1_epi16(max_val); - __m128i shufsub = - _mm_set_epi8(0, 0, 0, 0, 0, 0, 0, 0, 13, 12, 9, 8, 5, 4, 1, 0); - //__m256i masksub1 = _mm256_set_m128i(shufsub, shufsub); - __m256i masksub1 = - _mm256_insertf128_si256(_mm256_castsi128_si256(shufsub), (shufsub), 0x1); __m256i masksub2 = _mm256_set_epi32(0, 0, 0, 0, 5, 4, 1, 0); __m256i d1, d2; @@ -399,74 +548,30 @@ void ccso_derive_src_block_avx2(const uint16_t *src_y, uint8_t *const src_cls0, // int stride = src_y_stride[yOff] << y_uv_vscale; for (int xOff = 0; xOff < x_offset; xOff += 16) { - // uint16_t* rec_tmp = &src_rec2[xOff << y_uv_hscale]; - __m256i rec_curlo = _mm256_loadu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale))); - - //__m256i rec_tap1 = _mm256_loadu_si256((const __m256i*)(src_rec2 + (xOff - //<< y_uv_hscale) + tap1_pos)); - __m256i rec_tap1lo = _mm256_loadu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap1_pos)); - - //__m256i rec_tap2 = _mm256_loadu_si256((const __m256i*)(src_rec2 + (xOff - //<< y_uv_hscale) + tap2_pos)); - __m256i rec_tap2lo = _mm256_loadu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap2_pos)); - + __m256i rec_cur, rec_tap1, rec_tap2; if (y_uv_hscale > 0) { - __m256i rec_curhi = _mm256_loadu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + 16)); - rec_curlo = _mm256_shuffle_epi8(rec_curlo, masksub1); - rec_curhi = _mm256_shuffle_epi8(rec_curhi, masksub1); - rec_curlo = _mm256_permutevar8x32_epi32(rec_curlo, masksub2); - rec_curhi = _mm256_permutevar8x32_epi32(rec_curhi, masksub2); - //__m256i rec_cur = _mm256_setr_m128i(_mm256_castsi256_si128(rec_curlo), - // _mm256_castsi256_si128(rec_curhi)); - __m256i rec_cur = _mm256_insertf128_si256( - rec_curlo, _mm256_castsi256_si128(rec_curhi), 0x1); - - __m256i rec_tap1hi = _mm256_loadu_si256(( - const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap1_pos + 16)); - rec_tap1lo = _mm256_shuffle_epi8(rec_tap1lo, masksub1); - rec_tap1hi = _mm256_shuffle_epi8(rec_tap1hi, masksub1); - rec_tap1lo = _mm256_permutevar8x32_epi32(rec_tap1lo, masksub2); - rec_tap1hi = _mm256_permutevar8x32_epi32(rec_tap1hi, masksub2); - //__m256i rec_tap1 - //=_mm256_setr_m128i(_mm256_castsi256_si128(rec_tap1lo), - // _mm256_castsi256_si128(rec_tap1hi)); - __m256i rec_tap1 = _mm256_insertf128_si256( - rec_tap1lo, _mm256_castsi256_si128(rec_tap1hi), 0x1); - - __m256i rec_tap2hi = _mm256_loadu_si256(( - const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + tap2_pos + 16)); - rec_tap2lo = _mm256_shuffle_epi8(rec_tap2lo, masksub1); - rec_tap2hi = _mm256_shuffle_epi8(rec_tap2hi, masksub1); - rec_tap2lo = _mm256_permutevar8x32_epi32(rec_tap2lo, masksub2); - rec_tap2hi = _mm256_permutevar8x32_epi32(rec_tap2hi, masksub2); - //__m256i rec_tap2 - //=_mm256_setr_m128i(_mm256_castsi256_si128(rec_tap2lo), - // _mm256_castsi256_si128(rec_tap2hi)); - __m256i rec_tap2 = _mm256_insertf128_si256( - rec_tap2lo, _mm256_castsi256_si128(rec_tap2hi), 0x1); - - // int d1 = rec_tmp[tap1_pos] - rec_tmp[0]; - // int d2 = rec_tmp[tap2_pos] - rec_tmp[0]; - d1 = _mm256_sub_epi16(rec_tap1, rec_cur); - d2 = _mm256_sub_epi16(rec_tap2, rec_cur); + const uint16_t *src = src_rec2 + (xOff << 1); + rec_cur = extract_even_16bit_avx2(src); + rec_tap1 = extract_even_16bit_avx2(src + tap1_pos); + rec_tap2 = extract_even_16bit_avx2(src + tap2_pos); } else { - d1 = _mm256_sub_epi16(rec_tap1lo, rec_curlo); - d2 = _mm256_sub_epi16(rec_tap2lo, rec_curlo); + const uint16_t *src = src_rec2 + xOff; + rec_cur = _mm256_loadu_si256((const __m256i *)src); + rec_tap1 = _mm256_loadu_si256((const __m256i *)(src + tap1_pos)); + rec_tap2 = _mm256_loadu_si256((const __m256i *)(src + tap2_pos)); } + // d1 = rec[tap1_pos] - rec[0], d2 = rec[tap2_pos] - rec[0] + d1 = _mm256_sub_epi16(rec_tap1, rec_cur); + d2 = _mm256_sub_epi16(rec_tap2, rec_cur); + __m256i idx1, idx2; if (edge_clf == 0) { - idx1 = cal_filter_support_edge0_avx2(d1, cmp_thr1, cmp_thr2, all1, - cmp_idxa, cmp_idxc); - idx2 = cal_filter_support_edge0_avx2(d2, cmp_thr1, cmp_thr2, all1, - cmp_idxa, cmp_idxc); + idx1 = cal_filter_support_edge0_avx2(d1, cmp_thr1, cmp_thr2, cmp_idxc); + idx2 = cal_filter_support_edge0_avx2(d2, cmp_thr1, cmp_thr2, cmp_idxc); } else { // if (edge_clf == 1) - idx1 = cal_filter_support_edge1_avx2(d1, cmp_thr2, all1, cmp_idxc); - idx2 = cal_filter_support_edge1_avx2(d2, cmp_thr2, all1, cmp_idxc); + idx1 = cal_filter_support_edge1_avx2(d1, cmp_thr2, cmp_idxc); + idx2 = cal_filter_support_edge1_avx2(d2, cmp_thr2, cmp_idxc); } idx1 = _mm256_packs_epi16(idx1, idx1); @@ -623,27 +728,92 @@ void ccso_filter_block_hbd_with_buf_bo_only_avx2( } } -void ccso_filter_block_hbd_with_buf_avx2( - const uint16_t *src_y, uint16_t *dts_yuv, const uint8_t *src_cls0, - const uint8_t *src_cls1, const int src_y_stride, const int dst_stride, - const int ccso_stride, const int x, const int y, const int pic_width, - const int pic_height, const int8_t *filter_offset, const int blk_size_x, - const int blk_size_y, const int y_uv_hscale, const int y_uv_vscale, - const int max_val, const uint8_t shift_bits, const uint8_t ccso_bo_only) { - (void)ccso_bo_only; - __m256i all0 = _mm256_set1_epi16(0); - __m256i allmax = _mm256_set1_epi16(((short)max_val)); - __m128i shufsub = - _mm_set_epi8(0, 0, 0, 0, 0, 0, 0, 0, 13, 12, 9, 8, 5, 4, 1, 0); - //__m256i masksub1 = _mm256_set_m128i(shufsub, shufsub); - __m256i masksub1 = - _mm256_insertf128_si256(_mm256_castsi128_si256(shufsub), (shufsub), 0x1); - __m256i masksub2 = _mm256_set_epi32(0, 0, 0, 0, 5, 4, 1, 0); - // because src_cls has a type of uint8_t, selected directly using 8bit - // elements - __m128i mask_cls_sub1 = - _mm_set_epi8(0, 0, 0, 0, 0, 0, 0, 0, 14, 12, 10, 8, 6, 4, 2, 0); +static AVM_FORCE_INLINE void ccso_filter_row_width_32_avx2( + const uint16_t *src_rec, const uint8_t *src_cls0, const uint8_t *src_cls1, + uint16_t *dst_rec2, int xOff, int y_uv_hscale, int shift_bits, + const __m256i *filter_offset_lut, __m256i allmax, int max_band) { + __m256i src_rec_low, src_rec_high, cls0, cls1; + + if (y_uv_hscale == 0) { + src_rec_low = _mm256_loadu_si256((const __m256i *)(src_rec + xOff)); + src_rec_high = _mm256_loadu_si256((const __m256i *)(src_rec + xOff + 16)); + cls0 = _mm256_loadu_si256((const __m256i *)(src_cls0 + xOff)); + cls1 = _mm256_loadu_si256((const __m256i *)(src_cls1 + xOff)); + } else { + const uint16_t *src = src_rec + (xOff << 1); + src_rec_low = extract_even_16bit_avx2(src); + src_rec_high = extract_even_16bit_avx2(src + 32); + cls0 = extract_even_8bit_w32_avx2(src_cls0 + (xOff << 1)); + cls1 = extract_even_8bit_w32_avx2(src_cls1 + (xOff << 1)); + } + + const __m256i eo_idx = _mm256_add_epi8(_mm256_slli_epi16(cls0, 2), cls1); + const __m256i bo_idx = _mm256_permute4x64_epi64( + _mm256_packus_epi16(_mm256_srli_epi16(src_rec_low, shift_bits), + _mm256_srli_epi16(src_rec_high, shift_bits)), + 0xD8); + const __m256i lut_idx = _mm256_add_epi8(eo_idx, _mm256_slli_epi16(bo_idx, 4)); + + const __m256i offset = + get_offset_from_index_avx2(filter_offset_lut, lut_idx, max_band); + const __m256i offset_low = + _mm256_cvtepi8_epi16(_mm256_castsi256_si128(offset)); + const __m256i offset_high = + _mm256_cvtepi8_epi16(_mm256_extracti128_si256(offset, 1)); + + add_offset_avx2(dst_rec2, xOff, offset_low, allmax); + add_offset_avx2(dst_rec2, xOff + 16, offset_high, allmax); +} + +static AVM_FORCE_INLINE void ccso_filter_row_width_16_avx2( + const uint16_t *src_rec, const uint8_t *src_cls0, const uint8_t *src_cls1, + uint16_t *dst_rec2, int xOff, int y_uv_hscale, int shift_bits, + const __m256i *filter_offset_lut, __m256i allmax, int max_band) { + __m256i src_rec_reg; + __m128i cls0, cls1; + + if (y_uv_hscale == 0) { + src_rec_reg = _mm256_loadu_si256((const __m256i *)(src_rec + xOff)); + cls0 = _mm_loadu_si128((const __m128i *)(src_cls0 + xOff)); + cls1 = _mm_loadu_si128((const __m128i *)(src_cls1 + xOff)); + } else { + const uint16_t *src = src_rec + (xOff << 1); + src_rec_reg = extract_even_16bit_avx2(src); + cls0 = extract_even_8bit_w16_avx2(src_cls0 + (xOff << 1)); + cls1 = extract_even_8bit_w16_avx2(src_cls1 + (xOff << 1)); + } + + const __m128i eo_idx = _mm_add_epi8(_mm_slli_epi16(cls0, 2), cls1); + const __m256i band = _mm256_srli_epi16(src_rec_reg, shift_bits); + const __m128i bo_idx = _mm_packus_epi16(_mm256_castsi256_si128(band), + _mm256_extracti128_si256(band, 1)); + const __m128i lut_idx = _mm_add_epi8(eo_idx, _mm_slli_epi16(bo_idx, 4)); + const __m256i offset_256 = get_offset_from_index_avx2( + filter_offset_lut, _mm256_castsi128_si256(lut_idx), max_band); + const __m128i offset = _mm256_castsi256_si128(offset_256); + + add_offset_avx2(dst_rec2, xOff, _mm256_cvtepi8_epi16(offset), allmax); +} + +static AVM_FORCE_INLINE void ccso_filter_row_width_remainder_avx2( + const uint16_t *src_rec, const uint8_t *src_cls0, const uint8_t *src_cls1, + uint16_t *dst_rec2, int xOff, int x_remainder, int y_uv_hscale, + int shift_bits, const int8_t *filter_offset, int max_val) { + for (int i = xOff; i < xOff + x_remainder; i++) { + const int sx = i << y_uv_hscale; + const int band_num = src_rec[sx] >> shift_bits; + const int lut_idx = (band_num << 4) + (src_cls0[sx] << 2) + src_cls1[sx]; + dst_rec2[i] = clamp(filter_offset[lut_idx] + dst_rec2[i], 0, max_val); + } +} + +static AVM_FORCE_INLINE void ccso_filter_block( + const uint16_t *src_y, uint16_t *dts_yuv, const uint8_t *src_cls0, + const uint8_t *src_cls1, int src_y_stride, int dst_stride, int ccso_stride, + int x, int y, int pic_width, int pic_height, int blk_size_x, int blk_size_y, + const int8_t *filter_offset, int y_uv_hscale, int y_uv_vscale, int max_val, + uint8_t shift_bits, int max_band) { int y_offset; int x_offset, x_remainder; @@ -659,161 +829,117 @@ void ccso_filter_block_hbd_with_buf_avx2( x_offset = blk_size_x; x_remainder = 0; } - for (int yOff = 0; yOff < y_offset; yOff++) { - uint16_t *dst_rec2 = dts_yuv + x + yOff * dst_stride; - const uint16_t *src_rec2 = - src_y + ((yOff << y_uv_vscale) * src_y_stride + (x << y_uv_hscale)); - const uint8_t *src_cls0_2 = - src_cls0 + ((yOff << y_uv_vscale) * ccso_stride + (x << y_uv_hscale)); - const uint8_t *src_cls1_2 = - src_cls1 + ((yOff << y_uv_vscale) * ccso_stride + (x << y_uv_hscale)); - - // int stride = src_y_stride[yOff] << y_uv_vscale; - for (int xOff = 0; xOff < x_offset; xOff += 16) { - // uint16_t* rec_tmp = &src_rec2[xOff << y_uv_hscale]; - __m256i rec_curlo = _mm256_loadu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale))); - __m128i cur_src_cls0lo = _mm_lddqu_si128( - (const __m128i *)(src_cls0_2 + (xOff << y_uv_hscale))); - __m128i cur_src_cls1lo = _mm_lddqu_si128( - (const __m128i *)(src_cls1_2 + (xOff << y_uv_hscale))); - __m256i rec_cur_final; - __m256i cur_src_cls0_final; - __m256i cur_src_cls1_final; - - if (y_uv_hscale > 0) { - __m256i rec_curhi = _mm256_loadu_si256( - (const __m256i *)(src_rec2 + (xOff << y_uv_hscale) + 16)); - rec_curlo = _mm256_shuffle_epi8(rec_curlo, masksub1); - rec_curhi = _mm256_shuffle_epi8(rec_curhi, masksub1); - rec_curlo = _mm256_permutevar8x32_epi32(rec_curlo, masksub2); - rec_curhi = _mm256_permutevar8x32_epi32(rec_curhi, masksub2); - //__m256i rec_cur = _mm256_setr_m128i(_mm256_castsi256_si128(rec_curlo), - // _mm256_castsi256_si128(rec_curhi)); - __m256i rec_cur = _mm256_insertf128_si256( - rec_curlo, _mm256_castsi256_si128(rec_curhi), 0x1); - - __m128i cur_src_cls0hi = _mm_lddqu_si128( - (const __m128i *)(src_cls0_2 + (xOff << y_uv_hscale) + 16)); - cur_src_cls0lo = _mm_shuffle_epi8(cur_src_cls0lo, mask_cls_sub1); - cur_src_cls0hi = _mm_shuffle_epi8(cur_src_cls0hi, mask_cls_sub1); - __m256i cur_scr_cls0_256 = - //_mm256_setr_m128i(cur_src_cls0lo, cur_src_cls0hi); - _mm256_insertf128_si256(_mm256_castsi128_si256(cur_src_cls0lo), - (cur_src_cls0hi), 0x1); - - cur_scr_cls0_256 = - _mm256_permutevar8x32_epi32(cur_scr_cls0_256, masksub2); - __m128i cur_scr_cls0_128 = _mm256_castsi256_si128(cur_scr_cls0_256); - cur_scr_cls0_256 = _mm256_cvtepi8_epi16(cur_scr_cls0_128); - - __m128i cur_src_cls1hi = _mm_lddqu_si128( - (const __m128i *)(src_cls1_2 + (xOff << y_uv_hscale) + 16)); - cur_src_cls1lo = _mm_shuffle_epi8(cur_src_cls1lo, mask_cls_sub1); - cur_src_cls1hi = _mm_shuffle_epi8(cur_src_cls1hi, mask_cls_sub1); - __m256i cur_scr_cls1_256 = - //_mm256_setr_m128i(cur_src_cls1lo, cur_src_cls1hi); - _mm256_insertf128_si256(_mm256_castsi128_si256(cur_src_cls1lo), - (cur_src_cls1hi), 0x1); - - cur_scr_cls1_256 = - _mm256_permutevar8x32_epi32(cur_scr_cls1_256, masksub2); - __m128i cur_scr_cls1_128 = _mm256_castsi256_si128(cur_scr_cls1_256); - cur_scr_cls1_256 = _mm256_cvtepi8_epi16(cur_scr_cls1_128); - - cur_src_cls0_final = cur_scr_cls0_256; - cur_src_cls1_final = cur_scr_cls1_256; - rec_cur_final = rec_cur; - } else { - cur_src_cls0_final = _mm256_cvtepi8_epi16(cur_src_cls0lo); - cur_src_cls1_final = _mm256_cvtepi8_epi16(cur_src_cls1lo); - rec_cur_final = rec_curlo; - } - __m256i dst_rec = _mm256_loadu_si256((const __m256i *)(dst_rec2 + xOff)); + const __m256i allmax = _mm256_set1_epi16((short)max_val); + + uint16_t *dst_rec2 = dts_yuv + x; + const uint16_t *src_rec2 = src_y + (x << y_uv_hscale); + const uint8_t *src_cls0_2 = src_cls0 + (x << y_uv_hscale); + const uint8_t *src_cls1_2 = src_cls1 + (x << y_uv_hscale); + const int src_rec2_stride = src_y_stride << y_uv_vscale; + const int src_cls_stride = ccso_stride << y_uv_vscale; + + __m256i filter_offset_lut[8]; + for (int band_num = 0; band_num < 8; band_num++) { + if (band_num < max_band) { + filter_offset_lut[band_num] = _mm256_broadcastsi128_si256( + _mm_loadu_si128((const __m128i *)(filter_offset + (band_num << 4)))); + } else { + filter_offset_lut[band_num] = _mm256_setzero_si256(); + } + } - // const int band_num = src_y[x_pos] >> shift_bits; - __m256i num_band = _mm256_srli_epi16(rec_cur_final, shift_bits); - __m256i lut_idx_ext = all0; + for (int yOff = 0; yOff < y_offset; yOff++) { + int xOff = 0; - // const int lut_idx_ext = (band_num << 4) + (src_cls[0] << 2) + - // src_cls[1]; - num_band = _mm256_slli_epi16(num_band, 4); - lut_idx_ext = _mm256_add_epi16(lut_idx_ext, num_band); + for (; xOff + 32 <= x_offset; xOff += 32) { + ccso_filter_row_width_32_avx2(src_rec2, src_cls0_2, src_cls1_2, dst_rec2, + xOff, y_uv_hscale, shift_bits, + filter_offset_lut, allmax, max_band); + } - cur_src_cls0_final = _mm256_slli_epi16(cur_src_cls0_final, 2); - lut_idx_ext = _mm256_add_epi16(lut_idx_ext, cur_src_cls0_final); + for (; xOff < x_offset; xOff += 16) { + ccso_filter_row_width_16_avx2(src_rec2, src_cls0_2, src_cls1_2, dst_rec2, + xOff, y_uv_hscale, shift_bits, + filter_offset_lut, allmax, max_band); + } - lut_idx_ext = _mm256_add_epi16(lut_idx_ext, cur_src_cls1_final); + if (x_remainder) + ccso_filter_row_width_remainder_avx2( + src_rec2, src_cls0_2, src_cls1_2, dst_rec2, x_offset, x_remainder, + y_uv_hscale, shift_bits, filter_offset, max_val); - DECLARE_ALIGNED(32, uint16_t, offset_idx[16]); - int16_t offset_array[16]; - _mm256_store_si256((__m256i *)offset_idx, lut_idx_ext); - for (int i = 0; i < 16; i++) { - offset_array[i] = (int16_t)(filter_offset[offset_idx[i]]); - } - __m256i offset = _mm256_loadu_si256((const __m256i *)offset_array); + dst_rec2 += dst_stride; + src_rec2 += src_rec2_stride; + src_cls0_2 += src_cls_stride; + src_cls1_2 += src_cls_stride; + } +} - // uint16_t val = clamp(offset_val + dst_rec2[xOff], 0, (1 << - // cm->seq_params.bit_depth) - 1); - __m256i recon = _mm256_add_epi16(offset, dst_rec); - recon = _mm256_min_epi16(recon, allmax); - recon = _mm256_max_epi16(recon, all0); +void ccso_filter_block_hbd_with_buf_avx2( + const uint16_t *src_y, uint16_t *dts_yuv, const uint8_t *src_cls0, + const uint8_t *src_cls1, const int src_y_stride, const int dst_stride, + const int ccso_stride, const int x, const int y, const int pic_width, + const int pic_height, const int8_t *filter_offset, const int blk_size_x, + const int blk_size_y, const int y_uv_hscale, const int y_uv_vscale, + const int max_val, const uint8_t shift_bits, const uint8_t ccso_bo_only) { + (void)ccso_bo_only; - // dst_rec2[xOff] = val; - _mm256_storeu_si256((__m256i *)(dst_rec2 + xOff), recon); + // Number of bands in use: 1, 2, 4 or 8. + const int max_band = (max_val >> shift_bits) + 1; + +#define CCSO_FILTER_BLOCK(MAX_BAND, HORIZONTAL_SCALE) \ + ccso_filter_block(src_y, dts_yuv, src_cls0, src_cls1, src_y_stride, \ + dst_stride, ccso_stride, x, y, pic_width, pic_height, \ + blk_size_x, blk_size_y, filter_offset, (HORIZONTAL_SCALE), \ + y_uv_vscale, max_val, shift_bits, (MAX_BAND)) + + if (y_uv_hscale == 0) { + switch (max_band) { + case 1: CCSO_FILTER_BLOCK(1, 0); break; + case 2: CCSO_FILTER_BLOCK(2, 0); break; + case 4: CCSO_FILTER_BLOCK(4, 0); break; + case 8: CCSO_FILTER_BLOCK(8, 0); break; + default: assert(0); break; } - for (int xOff = x_offset; xOff < x_offset + x_remainder; xOff++) { - // cal_filter_support(rec_luma_idx, &src_y[((src_y_stride[yOff] << - // y_uv_vscale) + ((x + xOff) << y_uv_hscale)) + pad_stride], - // quant_step_size, inv_quant_step, rec_idx); - int cur_src_cls0 = src_cls0[(yOff << y_uv_vscale) * ccso_stride + - ((x + xOff) << y_uv_hscale)]; - int cur_src_cls1 = src_cls1[(yOff << y_uv_vscale) * ccso_stride + - ((x + xOff) << y_uv_hscale)]; - const int band_num = src_y[((yOff << y_uv_vscale) * src_y_stride + - ((x + xOff) << y_uv_hscale))] >> - shift_bits; - int offset_val = - filter_offset[(band_num << 4) + (cur_src_cls0 << 2) + cur_src_cls1]; - // dts_yuv[dst_stride[yOff] + x + xOff] = clamp(offset_val + - // dts_yuv[dst_stride[yOff] + x + xOff], 0, max_val); - dts_yuv[yOff * dst_stride + x + xOff] = - clamp(offset_val + dts_yuv[yOff * dst_stride + x + xOff], 0, max_val); + } else { + switch (max_band) { + case 1: CCSO_FILTER_BLOCK(1, 1); break; + case 2: CCSO_FILTER_BLOCK(2, 1); break; + case 4: CCSO_FILTER_BLOCK(4, 1); break; + case 8: CCSO_FILTER_BLOCK(8, 1); break; + default: assert(0); break; } } +#undef CCSO_FILTER_BLOCK } -static INLINE int SquareDifference(__m256i a, __m256i b) { - const __m256i Z = _mm256_setzero_si256(); - - const __m256i alo = _mm256_unpacklo_epi16(a, Z); - const __m256i blo = _mm256_unpacklo_epi16(b, Z); - const __m256i dlo = _mm256_sub_epi32(alo, blo); - - const __m256i ahi = _mm256_unpackhi_epi16(a, Z); - const __m256i bhi = _mm256_unpackhi_epi16(b, Z); - const __m256i dhi = _mm256_sub_epi32(ahi, bhi); - - const __m256i dloSq = _mm256_mullo_epi32(dlo, dlo); - const __m256i dhiSq = _mm256_mullo_epi32(dhi, dhi); - - const __m256i dlhSq = _mm256_add_epi32(dloSq, dhiSq); - const __m256i masksub2 = _mm256_set_epi32(0, 0, 0, 0, 5, 4, 1, 0); - - __m256i dloSum = _mm256_hadd_epi32(dlhSq, dlhSq); - dloSum = _mm256_permutevar8x32_epi32(dloSum, masksub2); - dloSum = _mm256_hadd_epi32(dloSum, dloSum); - dloSum = _mm256_hadd_epi32(dloSum, dloSum); - - return (_mm_cvtsi128_si32(_mm256_castsi256_si128(dloSum))); +// Horizontal sum of eight non-negative 32 bit values, widened to 64 bit so +// the result will not overflow. +static INLINE uint64_t hsum_epi32_to_u64_avx2(__m256i v) { + const __m256i zero = _mm256_setzero_si256(); + const __m256i lo = _mm256_unpacklo_epi32(v, zero); + const __m256i hi = _mm256_unpackhi_epi32(v, zero); + const __m256i sum64 = _mm256_add_epi64(lo, hi); + __m128i s = _mm_add_epi64(_mm256_castsi256_si128(sum64), + _mm256_extracti128_si256(sum64, 1)); + s = _mm_add_epi64(s, _mm_srli_si128(s, 8)); +#if ARCH_X86_64 + return (uint64_t)_mm_cvtsi128_si64(s); +#else + { + uint64_t tmp; + _mm_storel_epi64((__m128i *)&tmp, s); + return tmp; + } +#endif } uint64_t compute_distortion_block_avx2( const uint16_t *org, const int org_stride, const uint16_t *rec16, const int rec_stride, const int x, const int y, const int log2_filter_unit_size_y, const int log2_filter_unit_size_x, - const int height, const int width) { + const int height, const int width, const int bd) { const int blk_size_y = 1 << log2_filter_unit_size_y; const int blk_size_x = 1 << log2_filter_unit_size_x; int y_offset; @@ -833,6 +959,21 @@ uint64_t compute_distortion_block_avx2( } uint64_t sum = 0; + // The look-up table below holds the mask which is used to decide the nth row + // upto which the values can be accumulated without overflowing 32-bit + // unsigned lane.The maximum possible sum accumulated per row is(blk_size_x / + // 16)* 2*(2 ^ bd - 1)^2. + // + // Columns of the table correspond to block sizes 32x32, 64x64, 128x128, + // and 256x256. + static const int max_rows_to_sum_lut[3][4] = { + { 16383, 8191, 4095, 2047 }, // bd 8 + { 1023, 511, 255, 127 }, // bd 10 + { 63, 31, 15, 7 }, // bd 12 + }; + const int max_rows_to_sum = max_rows_to_sum_lut[(bd - 8) >> 1][AVMMAX( + 0, log2_filter_unit_size_x - 5)]; + __m256i acc = _mm256_setzero_si256(); for (int yOff = 0; yOff < y_offset; yOff++) { const uint16_t *org2 = org + (yOff * org_stride + x); const uint16_t *rec2 = rec16 + (yOff * rec_stride + x); @@ -841,10 +982,15 @@ uint64_t compute_distortion_block_avx2( _mm256_loadu_si256((const __m256i *)(org2 + xOff)); const __m256i rec_cur = _mm256_loadu_si256((const __m256i *)(rec2 + xOff)); - int err = SquareDifference(org_cur, rec_cur); - sum += err; + const __m256i diff = _mm256_sub_epi16(org_cur, rec_cur); + acc = _mm256_add_epi32(acc, _mm256_madd_epi16(diff, diff)); + } + if ((yOff & max_rows_to_sum) == max_rows_to_sum) { + sum += hsum_epi32_to_u64_avx2(acc); + acc = _mm256_setzero_si256(); } } + sum += hsum_epi32_to_u64_avx2(acc); // process remaining irregular block to avoid scalar processing for every row for (int yOff = 0; yOff < y_offset; yOff++) { diff --git a/av2/encoder/pickccso.c b/av2/encoder/pickccso.c index 0e97ff4ca8..b5ddce7b32 100644 --- a/av2/encoder/pickccso.c +++ b/av2/encoder/pickccso.c @@ -617,7 +617,9 @@ uint64_t compute_distortion_block_c(const uint16_t *org, const int org_stride, const int x, const int y, const int log2_filter_unit_size_y, const int log2_filter_unit_size_x, - const int height, const int width) { + const int height, const int width, + const int bd) { + (void)bd; int err; uint64_t ssd = 0; int y_offset; @@ -680,28 +682,35 @@ static void compute_distortion(const AV2_COMMON *cm, const CcsoCtx *ctx, } // All unified into pixel size uint64_t sb_ssd = 0; - const uint16_t *org_unit = org; - const uint16_t *rec_unit = rec16; - const int y_end = AVMMIN(height - y, blk_size_y); - const int x_end = AVMMIN(width - x, blk_size_x); - for (int unit_y = 0; unit_y < y_end; unit_y += unit_size_y) { - for (int unit_x = 0; unit_x < x_end; unit_x += unit_size_x) { - // skip if unit skip - const int mbmi_idx = get_mi_grid_idx( - &cm->mi_params, (y + unit_y) >> v_scale, (x + unit_x) >> h_scale); - if (cm->bru.enabled && - cm->mi_params.mi_grid_base[mbmi_idx]->sb_active_mode != - BRU_ACTIVE_SB) { - continue; + if (cm->bru.enabled) { + const uint16_t *org_unit = org; + const uint16_t *rec_unit = rec16; + const int y_end = AVMMIN(height - y, blk_size_y); + const int x_end = AVMMIN(width - x, blk_size_x); + for (int unit_y = 0; unit_y < y_end; unit_y += unit_size_y) { + for (int unit_x = 0; unit_x < x_end; unit_x += unit_size_x) { + // skip if unit skip + const int mbmi_idx = + get_mi_grid_idx(&cm->mi_params, (y + unit_y) >> v_scale, + (x + unit_x) >> h_scale); + if (cm->mi_params.mi_grid_base[mbmi_idx]->sb_active_mode != + BRU_ACTIVE_SB) { + continue; + } + // skip if unit skip + sb_ssd += compute_distortion_block( + org_unit, org_stride, rec_unit, rec_stride, x + unit_x, + y + unit_y, unit_log2_y, unit_log2_x, height, width, + cm->seq_params.bit_depth); } - // skip if unit skip - sb_ssd += compute_distortion_block( - org_unit, org_stride, rec_unit, rec_stride, x + unit_x, - y + unit_y, unit_log2_y, unit_log2_x, height, width); + // offset org, rec16 here + org_unit += (org_stride << unit_log2_x); + rec_unit += (rec_stride << unit_log2_x); } - // offset org, rec16 here - org_unit += (org_stride << unit_log2_x); - rec_unit += (rec_stride << unit_log2_x); + } else { + sb_ssd = compute_distortion_block(org, org_stride, rec16, rec_stride, x, + y, blk_log2_y, blk_log2_x, height, + width, cm->seq_params.bit_depth); } distortion_buf[(y >> blk_log2_y) * distortion_buf_stride + (x >> blk_log2_x)] = sb_ssd; diff --git a/test/av2_ccso_simd_cmp.cc b/test/av2_ccso_simd_cmp.cc index e6bbe32d57..caea2b05d4 100644 --- a/test/av2_ccso_simd_cmp.cc +++ b/test/av2_ccso_simd_cmp.cc @@ -78,8 +78,8 @@ class CCSOFilterTest : public FunctionEquivalenceTest { filter_sup_ = this->rng_(7); derive_ccso_sample_pos(src_loc_, src_y_stride_, filter_sup_); - const uint8_t quant_sz[4] = { 16, 8, 32, 64 }; - thr_ = quant_sz[this->rng_(4)]; + const uint8_t quant_sz[5] = { 16, 8, 32, 64, 0 }; + thr_ = quant_sz[this->rng_(5)]; neg_thr_ = -1 * thr_; const uint8_t shift_bits_a[2] = { 8, 10 }; @@ -120,23 +120,32 @@ class CCSOFilterTest : public FunctionEquivalenceTest { class CCSOWOBUFTest : public CCSOFilterTest { protected: void Execute() { - const int max_band_log2 = 3; - shift_bits_ = isSingleBand_ ? shift_bits_ : shift_bits_ - max_band_log2; - params_.ref_func(src_y_, dst_ref_, 0, 0, pic_width_, pic_height_, src_cls_, - offset_buf_, src_y_stride_, dst_stride_, y_uv_hscale_, - y_uv_vscale_, thr_, neg_thr_, src_loc_, max_val_, - CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), isSingleBand_, - shift_bits_, edge_clf_, 0); - - ASM_REGISTER_STATE_CHECK(params_.tst_func( - src_y_, dst_tst_, 0, 0, pic_width_, pic_height_, src_cls_, offset_buf_, - src_y_stride_, dst_stride_, y_uv_hscale_, y_uv_vscale_, thr_, neg_thr_, - src_loc_, max_val_, CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), - isSingleBand_, shift_bits_, edge_clf_, 0)); - - for (int r = 0; r < blk_size_; ++r) { - for (int c = 0; c < blk_size_; ++c) { - ASSERT_EQ(dst_ref_[r * dst_stride_ + c], dst_tst_[r * dst_stride_ + c]); + const int cur_bit_depth = 10; + max_val_ = (1 << cur_bit_depth) - 1; + // Test for number of bands = 1, 2, 4, 8. + for (int num_bands_log2 = 0; num_bands_log2 < 4; num_bands_log2++) { + const bool isSingleBand = (num_bands_log2 == 0); + shift_bits_ = cur_bit_depth - num_bands_log2; + params_.ref_func(src_y_, dst_ref_, 0, 0, pic_width_, pic_height_, + src_cls_, offset_buf_, src_y_stride_, dst_stride_, + y_uv_hscale_, y_uv_vscale_, thr_, neg_thr_, src_loc_, + max_val_, CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), + isSingleBand, shift_bits_, edge_clf_, 0); + + ASM_REGISTER_STATE_CHECK( + params_.tst_func(src_y_, dst_tst_, 0, 0, pic_width_, pic_height_, + src_cls_, offset_buf_, src_y_stride_, dst_stride_, + y_uv_hscale_, y_uv_vscale_, thr_, neg_thr_, src_loc_, + max_val_, CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), + isSingleBand, shift_bits_, edge_clf_, 0)); + + for (int r = 0; r < blk_size_; ++r) { + for (int c = 0; c < blk_size_; ++c) { + ASSERT_EQ(dst_ref_[r * dst_stride_ + c], + dst_tst_[r * dst_stride_ + c]) + << "num_bands=" << (1 << num_bands_log2) << " r=" << r + << " c=" << c; + } } } } @@ -153,20 +162,23 @@ class CCSOWOBUFTest : public CCSOFilterTest { thr_ = quant_sz[this->rng_(4)]; neg_thr_ = -1 * thr_; - const int max_band_log2 = 3; - const uint8_t bit_depth[2] = { 8, 10 }; - const uint8_t cur_bit_depth = bit_depth[this->rng_(2)]; + // src_y_ is filled with rng_(1 << 10), so the bit depth has to be 10 + // for the samples to stay within max_val_. + const uint8_t cur_bit_depth = 10; max_val_ = (1 << cur_bit_depth) - 1; int num_planes = 2; for (int plane = 0; plane < num_planes; plane++) { - for (int isSingleBand = 0; isSingleBand < 2; isSingleBand++) { - shift_bits_ = - isSingleBand ? cur_bit_depth : cur_bit_depth - max_band_log2; - y_uv_hscale_ = y_uv_vscale_ = plane; - pic_width_ = (MAX_SB_SIZE) >> y_uv_hscale_; - pic_height_ = (MAX_SB_SIZE) >> y_uv_vscale_; - blk_size_ = pic_width_; + y_uv_hscale_ = y_uv_vscale_ = plane; + pic_width_ = (MAX_SB_SIZE) >> y_uv_hscale_; + pic_height_ = (MAX_SB_SIZE) >> y_uv_vscale_; + blk_size_ = pic_width_; + + // 1, 2, 4 and 8 bands. + for (int num_bands_log2 = 0; num_bands_log2 < 4; num_bands_log2++) { + const int num_bands = 1 << num_bands_log2; + const bool isSingleBand = (num_bands == 1); + shift_bits_ = cur_bit_depth - num_bands_log2; avm_usec_timer timer; avm_usec_timer_start(&timer); @@ -197,10 +209,10 @@ class CCSOWOBUFTest : public CCSOFilterTest { (kSpeedIterations * blk_size_ * blk_size_); float scaling = c_time_per_pixel / opt_time_per_pixel; printf( - "%3dx%-3d: plane=%d, isSingleBand=%d " + "%3dx%-3d: plane=%d, num_bands=%d " "c_time_per_pixel=%10.5f, " "opt_time_per_pixel=%10.5f, scaling=%f \n", - blk_size_, blk_size_, plane, isSingleBand, c_time_per_pixel, + blk_size_, blk_size_, plane, num_bands, c_time_per_pixel, opt_time_per_pixel, scaling); } } @@ -383,23 +395,93 @@ class CCSOWITHBUFTest : public CCSOFilterTest { protected: void Execute() { ccso_stride_ = src_y_stride_ - (CCSO_PADDING_SIZE << 1); - params_.ref_func(src_y_, dst_ref_, src_cls0_, src_cls1_, src_y_stride_, - dst_stride_, ccso_stride_, 0, 0, pic_width_, pic_height_, - offset_buf_, CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), - y_uv_hscale_, y_uv_vscale_, max_val_, shift_bits_, 0); + const int cur_bit_depth = 10; + max_val_ = (1 << cur_bit_depth) - 1; + // Test for number of bands = 1, 2, 4, 8. + for (int band_log2 = 0; band_log2 < 4; band_log2++) { + shift_bits_ = cur_bit_depth - band_log2; + params_.ref_func(src_y_, dst_ref_, src_cls0_, src_cls1_, src_y_stride_, + dst_stride_, ccso_stride_, 0, 0, pic_width_, pic_height_, + offset_buf_, CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), + y_uv_hscale_, y_uv_vscale_, max_val_, shift_bits_, 0); + + ASM_REGISTER_STATE_CHECK(params_.tst_func( + src_y_, dst_tst_, src_cls0_, src_cls1_, src_y_stride_, dst_stride_, + ccso_stride_, 0, 0, pic_width_, pic_height_, offset_buf_, + CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), y_uv_hscale_, + y_uv_vscale_, max_val_, shift_bits_, 0)); + + for (int r = 0; r < blk_size_; ++r) { + for (int c = 0; c < blk_size_; ++c) { + ASSERT_EQ(dst_ref_[r * dst_stride_ + c], + dst_tst_[r * dst_stride_ + c]) + << "num_bands=" << (1 << band_log2) << " r=" << r << " c=" << c; + } + } + } + } - ASM_REGISTER_STATE_CHECK(params_.tst_func( - src_y_, dst_tst_, src_cls0_, src_cls1_, src_y_stride_, dst_stride_, - ccso_stride_, 0, 0, pic_width_, pic_height_, offset_buf_, - CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), y_uv_hscale_, y_uv_vscale_, - max_val_, shift_bits_, 0)); + void RunSpeedTest() { + src_y_stride_ = MAX_SB_SIZE + (CCSO_PADDING_SIZE << 1); + dst_stride_ = MAX_SB_SIZE; + ccso_stride_ = src_y_stride_ - (CCSO_PADDING_SIZE << 1); - for (int r = 0; r < blk_size_; ++r) { - for (int c = 0; c < blk_size_; ++c) { - ASSERT_EQ(dst_ref_[r * dst_stride_ + c], dst_tst_[r * dst_stride_ + c]); + // src_y_ is filled with rng_(1 << 10), so the bit depth has to be 10 + // for the samples to stay within max_val_. + const uint8_t cur_bit_depth = 10; + max_val_ = (1 << cur_bit_depth) - 1; + const int num_planes = 2; + const uint8_t ccso_bo_only = 0; + + for (int plane = 0; plane < num_planes; plane++) { + y_uv_hscale_ = y_uv_vscale_ = plane; + pic_width_ = (MAX_SB_SIZE) >> y_uv_hscale_; + pic_height_ = (MAX_SB_SIZE) >> y_uv_vscale_; + blk_size_ = pic_width_; + + // 1, 2, 4 and 8 bands. + for (int band_log2 = 0; band_log2 < 4; band_log2++) { + const int num_bands = 1 << band_log2; + shift_bits_ = cur_bit_depth - band_log2; + + avm_usec_timer timer; + avm_usec_timer_start(&timer); + for (int i = 0; i < kSpeedIterations; ++i) { + params_.ref_func( + src_y_, dst_ref_, src_cls0_, src_cls1_, src_y_stride_, + dst_stride_, ccso_stride_, 0, 0, pic_width_, pic_height_, + offset_buf_, CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), + y_uv_hscale_, y_uv_vscale_, max_val_, shift_bits_, ccso_bo_only); + } + avm_usec_timer_mark(&timer); + auto elapsed_time_c = avm_usec_timer_elapsed(&timer); + + avm_usec_timer_start(&timer); + for (int i = 0; i < kSpeedIterations; ++i) { + params_.tst_func( + src_y_, dst_tst_, src_cls0_, src_cls1_, src_y_stride_, + dst_stride_, ccso_stride_, 0, 0, pic_width_, pic_height_, + offset_buf_, CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), + y_uv_hscale_, y_uv_vscale_, max_val_, shift_bits_, ccso_bo_only); + } + avm_usec_timer_mark(&timer); + auto elapsed_time_opt = avm_usec_timer_elapsed(&timer); + + float c_time_per_pixel = (float)1000.0 * elapsed_time_c / + (kSpeedIterations * blk_size_ * blk_size_); + float opt_time_per_pixel = (float)1000.0 * elapsed_time_opt / + (kSpeedIterations * blk_size_ * blk_size_); + float scaling = c_time_per_pixel / opt_time_per_pixel; + printf( + "%3dx%-3d: plane=%d, num_bands=%d " + "c_time_per_pixel=%10.5f, " + "opt_time_per_pixel=%10.5f, scaling=%f \n", + blk_size_, blk_size_, plane, num_bands, c_time_per_pixel, + opt_time_per_pixel, scaling); } } } + uint8_t src_cls0_[kBufSize]; uint8_t src_cls1_[kBufSize]; int ccso_stride_; @@ -424,6 +506,23 @@ TEST_P(CCSOWITHBUFTest, RandomValues) { Common(); } } + +TEST_P(CCSOWITHBUFTest, DISABLED_Speed) { + const int hi = 1 << 10; + for (int i = 0; i < kBufSize; ++i) { + dst_ref_[i] = 0; + dst_tst_[i] = 0; + src_cls0_[i] = rng_(3); + src_cls1_[i] = rng_(3); + src_y_[i] = rng_(hi); + } + const int ccso_offset[8] = { -10, -7, -3, -1, 0, 1, 3, 7 }; + + for (int i = 0; i < CCSO_BAND_NUM * 16; i++) { + offset_buf_[i] = ccso_offset[rng_(8)]; + } + RunSpeedTest(); +} ////////////////////////////////////////////////////////////////////////////// // ccso_derive_src_block_avx2 ////////////////////////////////////////////////////////////////////////////// @@ -464,6 +563,65 @@ class CCSODeriveSrcTest : public CCSOFilterTest { } } } + + void RunSpeedTest() { + src_y_stride_ = + this->rng_(kMaxWidth + 1 - 32) + 32 + (CCSO_PADDING_SIZE << 1); + ccso_stride_ = src_y_stride_ - (CCSO_PADDING_SIZE << 1); + + const uint8_t quant_sz[4] = { 16, 8, 32, 64 }; + thr_ = quant_sz[this->rng_(4)]; + neg_thr_ = -1 * thr_; + + const int num_planes = 2; + + for (int plane = 0; plane < num_planes; plane++) { + for (int edge_clf = 0; edge_clf < 2; edge_clf++) { + filter_sup_ = this->rng_(7); + derive_ccso_sample_pos(src_loc_, src_y_stride_, filter_sup_); + y_uv_hscale_ = y_uv_vscale_ = plane; + pic_width_ = (MAX_SB_SIZE) >> y_uv_hscale_; + pic_height_ = (MAX_SB_SIZE) >> y_uv_vscale_; + blk_size_ = pic_width_; + + avm_usec_timer timer; + avm_usec_timer_start(&timer); + for (int i = 0; i < kSpeedIterations; ++i) { + params_.ref_func(src_y_, src_cls0_ref, src_cls1_ref, src_y_stride_, + ccso_stride_, 0, 0, pic_width_, pic_height_, + y_uv_hscale_, y_uv_vscale_, thr_, neg_thr_, src_loc_, + CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), + edge_clf); + } + avm_usec_timer_mark(&timer); + auto elapsed_time_c = avm_usec_timer_elapsed(&timer); + + avm_usec_timer_start(&timer); + for (int i = 0; i < kSpeedIterations; ++i) { + params_.tst_func(src_y_, src_cls0_tst, src_cls1_tst, src_y_stride_, + ccso_stride_, 0, 0, pic_width_, pic_height_, + y_uv_hscale_, y_uv_vscale_, thr_, neg_thr_, src_loc_, + CCSO_BLK_SIZE_PARAMS(blk_size_, blk_size_), + edge_clf); + } + avm_usec_timer_mark(&timer); + auto elapsed_time_opt = avm_usec_timer_elapsed(&timer); + + float c_time_per_pixel = (float)1000.0 * elapsed_time_c / + (kSpeedIterations * blk_size_ * blk_size_); + float opt_time_per_pixel = (float)1000.0 * elapsed_time_opt / + (kSpeedIterations * blk_size_ * blk_size_); + float scaling = c_time_per_pixel / opt_time_per_pixel; + printf( + "%3dx%-3d: plane=%d, edge_clf=%d " + "c_time_per_pixel=%10.5f, " + "opt_time_per_pixel=%10.5f, scaling=%f\n", + blk_size_, blk_size_, plane, edge_clf, c_time_per_pixel, + opt_time_per_pixel, scaling); + } + } + } + uint8_t src_cls0_ref[kBufSize]; uint8_t src_cls1_ref[kBufSize]; uint8_t src_cls0_tst[kBufSize]; @@ -488,6 +646,18 @@ TEST_P(CCSODeriveSrcTest, RandomValues) { } } +TEST_P(CCSODeriveSrcTest, DISABLED_Speed) { + const int hi = 1 << 10; + for (int i = 0; i < kBufSize; ++i) { + src_cls0_ref[i] = 0; + src_cls1_ref[i] = 0; + src_cls0_tst[i] = 0; + src_cls1_tst[i] = 0; + src_y_[i] = rng_(hi); + } + RunSpeedTest(); +} + ////////////////////////////////////////////////////////////////////////////// // compute_distortion_block_avx2 ////////////////////////////////////////////////////////////////////////////// @@ -496,7 +666,8 @@ typedef uint64_t (*CCSO_Dist_Block)(const uint16_t *org, const int org_stride, const int x, const int y, const int log2_filter_unit_size_y, const int log2_filter_unit_size_x, - const int height, const int width); + const int height, const int width, + const int bd); typedef libavm_test::FuncParam TestFuncsCCSO_Dist_Block; @@ -511,16 +682,65 @@ class CCSODistBlockTest : public CCSOFilterTest { log2_filter_unit_size_x_ = 1 - y_uv_hscale_ + 7; height_ = pic_height_; width_ = pic_width_; - uint64_t ref = params_.ref_func(org_, org_stride_, rec16_, rec_stride_, 0, - 0, log2_filter_unit_size_y_, - log2_filter_unit_size_x_, height_, width_); + const int bd = 10; + uint64_t ref = params_.ref_func( + org_, org_stride_, rec16_, rec_stride_, 0, 0, log2_filter_unit_size_y_, + log2_filter_unit_size_x_, height_, width_, bd); uint64_t tst; ASM_REGISTER_STATE_CHECK( tst = params_.tst_func(org_, org_stride_, rec16_, rec_stride_, 0, 0, log2_filter_unit_size_y_, - log2_filter_unit_size_x_, height_, width_)); + log2_filter_unit_size_x_, height_, width_, bd)); ASSERT_EQ(ref, tst); } + + void RunSpeedTest() { + const int bd = 10; + org_ = src_y_; + rec16_ = dst_ref_; + org_stride_ = kMaxWidth; + rec_stride_ = kMaxWidth; + pic_width_ = (MAX_SB_SIZE * 2); + pic_height_ = (MAX_SB_SIZE * 2); + + // Test for 32x32, 64x64, 128x128, 256x256. + for (int log2_proc_unit_size = MIN_SB_SIZE_LOG2 - 1; + log2_proc_unit_size <= MAX_SB_SIZE_LOG2; ++log2_proc_unit_size) { + const int proc_unit_size = 1 << log2_proc_unit_size; + const int pixels = proc_unit_size * proc_unit_size; + + avm_usec_timer timer; + avm_usec_timer_start(&timer); + for (int i = 0; i < kSpeedIterations; ++i) { + params_.ref_func(org_, org_stride_, rec16_, rec_stride_, 0, 0, + log2_proc_unit_size, log2_proc_unit_size, pic_height_, + pic_width_, bd); + } + avm_usec_timer_mark(&timer); + const auto elapsed_time_c = avm_usec_timer_elapsed(&timer); + + avm_usec_timer_start(&timer); + for (int i = 0; i < kSpeedIterations; ++i) { + params_.tst_func(org_, org_stride_, rec16_, rec_stride_, 0, 0, + log2_proc_unit_size, log2_proc_unit_size, pic_height_, + pic_width_, bd); + } + avm_usec_timer_mark(&timer); + const auto elapsed_time_opt = avm_usec_timer_elapsed(&timer); + + const float c_time_per_pixel = + (float)1000.0 * elapsed_time_c / (kSpeedIterations * pixels); + const float opt_time_per_pixel = + (float)1000.0 * elapsed_time_opt / (kSpeedIterations * pixels); + const float scaling = c_time_per_pixel / opt_time_per_pixel; + printf( + "%3dx%-3d: c_time_per_pixel=%10.5f, " + "opt_time_per_pixel=%10.5f, scaling=%f \n", + proc_unit_size, proc_unit_size, c_time_per_pixel, opt_time_per_pixel, + scaling); + } + } + uint16_t *org_; int org_stride_; uint16_t *rec16_; @@ -543,6 +763,15 @@ TEST_P(CCSODistBlockTest, RandomValues) { } } +TEST_P(CCSODistBlockTest, DISABLED_Speed) { + const int hi = 1 << 10; + for (int i = 0; i < kBufSize; ++i) { + dst_ref_[i] = rng_(hi); + src_y_[i] = rng_(hi); + } + RunSpeedTest(); +} + #if HAVE_AVX2 INSTANTIATE_TEST_SUITE_P( AVX2, CCSODistBlockTest,