av1: fix ASAN overflow in AVX512 convolve_2d_sr Fix num_z for h=4 blocks and clamp row index in LoadAndCombine to prevent out-of-bounds frame buffer reads in Highway AVX512. Change-Id: I5b5515a55581183dad5ad6569a8354382007fbae
diff --git a/av1/common/convolve_2d_sr_hwy.h b/av1/common/convolve_2d_sr_hwy.h index af680c1..6547f26 100644 --- a/av1/common/convolve_2d_sr_hwy.h +++ b/av1/common/convolve_2d_sr_hwy.h
@@ -246,6 +246,26 @@ } } +template <class D, class V128> +HWY_ATTR HWY_INLINE hn::VFromD<D> Combine4(D d, V128 r0, V128 r1, V128 r2, + V128 r3) { + using DHalf = hn::Half<D>; + const DHalf d_half; + using DQuarter = hn::Half<DHalf>; + + auto r01 = hn::Combine(d_half, hn::BitCast(DQuarter(), r1), + hn::BitCast(DQuarter(), r0)); + auto r23 = hn::Combine(d_half, hn::BitCast(DQuarter(), r3), + hn::BitCast(DQuarter(), r2)); + return hn::Combine(d, r23, r01); +} + +template <class D, class V128> +HWY_ATTR HWY_INLINE hn::VFromD<D> Combine2(D d, V128 r0, V128 r1) { + using DHalf = hn::Half<D>; + return hn::Combine(d, hn::BitCast(DHalf(), r1), hn::BitCast(DHalf(), r0)); +} + template <int num_rows, class D, class D128> HWY_ATTR HWY_INLINE hn::VFromD<D> LoadAndCombine(D d, D128 d128, const uint8_t *src, @@ -255,27 +275,45 @@ auto r1 = hn::LoadU(d128, src + 1 * src_stride); auto r2 = hn::LoadU(d128, src + 2 * src_stride); auto r3 = hn::LoadU(d128, src + 3 * src_stride); - - using DHalf = hn::Half<D>; - const DHalf d_half; - using DQuarter = hn::Half<DHalf>; - - auto r01 = hn::Combine(d_half, hn::BitCast(DQuarter(), r1), - hn::BitCast(DQuarter(), r0)); - auto r23 = hn::Combine(d_half, hn::BitCast(DQuarter(), r3), - hn::BitCast(DQuarter(), r2)); - return hn::Combine(d, r23, r01); + return Combine4(d, r0, r1, r2, r3); } HWY_IF_CONSTEXPR(num_rows == 2) { auto r0 = hn::LoadU(d128, src + 0 * src_stride); auto r1 = hn::LoadU(d128, src + 1 * src_stride); - - using DHalf = hn::Half<D>; - return hn::Combine(d, hn::BitCast(DHalf(), r1), hn::BitCast(DHalf(), r0)); + return Combine2(d, r0, r1); } return hn::ResizeBitCast(d, hn::LoadU(d128, src)); } +template <int num_rows, class D, class D128> +HWY_ATTR HWY_INLINE hn::VFromD<D> LoadAndCombineClamped(D d, D128 d128, + const uint8_t *src, + int src_stride, + int row_idx, + int max_rows) { + HWY_IF_CONSTEXPR(num_rows == 4) { + int r0 = row_idx; + int r1 = std::min(row_idx + 1, max_rows - 1); + int r2 = std::min(row_idx + 2, max_rows - 1); + int r3 = std::min(row_idx + 3, max_rows - 1); + + auto v0 = hn::LoadU(d128, src + r0 * src_stride); + auto v1 = hn::LoadU(d128, src + r1 * src_stride); + auto v2 = hn::LoadU(d128, src + r2 * src_stride); + auto v3 = hn::LoadU(d128, src + r3 * src_stride); + return Combine4(d, v0, v1, v2, v3); + } + HWY_IF_CONSTEXPR(num_rows == 2) { + int r0 = row_idx; + int r1 = std::min(row_idx + 1, max_rows - 1); + + auto v0 = hn::LoadU(d128, src + r0 * src_stride); + auto v1 = hn::LoadU(d128, src + r1 * src_stride); + return Combine2(d, v0, v1); + } + return hn::ResizeBitCast(d, hn::LoadU(d128, src + row_idx * src_stride)); +} + /// Helper function for the vertical convolve pass. // Template parameters: // - `taps_y_val`: Number of vertical filter taps (2, 4, 6, 8, or 12). @@ -296,7 +334,7 @@ const IVec16 *coeffs_v, IVec32 round_const_y, int round_1, int bits, IdxTbl idx_tbl, ComputeHRowBlock compute_h_row_block = nullptr) { constexpr int num_coeffs = taps_y_val / 2; - constexpr int num_z = 8 / 2 + num_coeffs - 1; + constexpr int num_z = (h_ge_8 ? 8 : 4) / 2 + num_coeffs - 1; constexpr int num_h = num_z / 2 + 1; hn::Half<decltype(int16xN_tag)> d16_16; @@ -445,7 +483,8 @@ (is_taps_y_12 || is_taps_x_12) ? false : (h == 4 || h == 8); if (!skip_strip_im_buf) { - for (int i = 0; i < im_h; i += kNumRows) { + int i = 0; + for (; i < (im_h & ~(kNumRows - 1)); i += kNumRows) { IVec16 res; if (is_taps_x_12) { auto data1 = LoadAndCombine<kNumRows>( @@ -476,6 +515,37 @@ hn::ShiftRightSame(hn::Add(res, round_const_h), round_0 - 1); hn::StoreU(shifted_res, int16xN_tag, strip_im_buf + i * 8); } + for (; i < im_h; i += kNumRows) { + IVec16 res; + if (is_taps_x_12) { + auto data1 = LoadAndCombineClamped<kNumRows>( + uint8xN_tag, uint8x16_capped_tag, src_ptr + j + 0, src_stride, i, + im_h); + auto data2 = LoadAndCombineClamped<kNumRows>( + uint8xN_tag, uint8x16_capped_tag, src_ptr + j + 6, src_stride, i, + im_h); + HWY_IF_CONSTEXPR(taps_x_const == 12) { + res = ConvolveLowbdX12TapShuffleFree( + int16xN_tag, data1, data2, coeffs_h_12[0], coeffs_h_12[1], + coeffs_h_12[2], coeffs_h_12[3], coeffs_h_12[4], coeffs_h_12[5]); + } + else { + res = ConvolveLowbdX12Tap(int16xN_tag, data1, data2, mask1, mask2, + mask3, coeffs_h_12[0], coeffs_h_12[1], + coeffs_h_12[2], coeffs_h_12[3], + coeffs_h_12[4], coeffs_h_12[5]); + } + } else { + auto data = + LoadAndCombineClamped<kNumRows>(uint8xN_tag, uint8x16_capped_tag, + src_ptr + j, src_stride, i, im_h); + res = convolve_horizontal(data); + } + + auto shifted_res = + hn::ShiftRightSame(hn::Add(res, round_const_h), round_0 - 1); + hn::StoreU(shifted_res, int16xN_tag, strip_im_buf + i * 8); + } } const int bits = FILTER_BITS * 2 - round_0 - round_1; @@ -500,8 +570,8 @@ auto compute_h_row_block = [&](int row_idx) { auto data = - LoadAndCombine<4>(uint8xN_tag, uint8x16_capped_tag, - src_ptr + row_idx * src_stride + j, src_stride); + LoadAndCombineClamped<4>(uint8xN_tag, uint8x16_capped_tag, + src_ptr + j, src_stride, row_idx, im_h); auto res = convolve_horizontal(data); return hn::ShiftRightSame(hn::Add(res, round_const_h), round_0 - 1); };