use core::arch::x86_64::*;
use crate::reconstruct::{can_reconstruct_full_block, narrow_i32_to_i16_scalar, sample_max};
#[inline(always)]
fn supported_n(n: usize) -> bool {
matches!(n, 2 | 4 | 8 | 16 | 32)
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_u16x2(src: &[u16]) -> __m128i {
debug_assert!(src.len() >= 2);
unsafe { _mm_castps_si128(_mm_load_ss(src.as_ptr().cast())) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_u16x4(src: &[u16]) -> __m128i {
debug_assert!(src.len() >= 4);
unsafe { _mm_loadl_epi64(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_u16x8(src: &[u16]) -> __m128i {
debug_assert!(src.len() >= 8);
unsafe { _mm_loadu_si128(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_u16x16(src: &[u16]) -> __m256i {
debug_assert!(src.len() >= 16);
unsafe { _mm256_loadu_si256(src.as_ptr().cast::<__m256i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_i16x2(src: &[i16]) -> __m128i {
debug_assert!(src.len() >= 2);
unsafe { _mm_castps_si128(_mm_load_ss(src.as_ptr().cast())) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_i16x4(src: &[i16]) -> __m128i {
debug_assert!(src.len() >= 4);
unsafe { _mm_loadl_epi64(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_i16x8(src: &[i16]) -> __m128i {
debug_assert!(src.len() >= 8);
unsafe { _mm_loadu_si128(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_i16x16(src: &[i16]) -> __m256i {
debug_assert!(src.len() >= 16);
unsafe { _mm256_loadu_si256(src.as_ptr().cast::<__m256i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_i32x2(src: &[i32]) -> __m128i {
debug_assert!(src.len() >= 2);
unsafe { _mm_loadl_epi64(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_i32x4(src: &[i32]) -> __m128i {
debug_assert!(src.len() >= 4);
unsafe { _mm_loadu_si128(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn load_i32x8(src: &[i32]) -> __m256i {
debug_assert!(src.len() >= 8);
unsafe { _mm256_loadu_si256(src.as_ptr().cast::<__m256i>()) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn zext128(v: __m128i) -> __m256i {
_mm256_inserti128_si256::<0>(_mm256_setzero_si256(), v)
}
#[inline]
#[target_feature(enable = "avx2")]
fn store_u16x2(dst: &mut [u16], v: __m128i) {
debug_assert!(dst.len() >= 2);
unsafe { _mm_store_ss(dst.as_mut_ptr().cast(), _mm_castsi128_ps(v)) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn store_u16x4(dst: &mut [u16], v: __m128i) {
debug_assert!(dst.len() >= 4);
unsafe { _mm_storel_epi64(dst.as_mut_ptr().cast::<__m128i>(), v) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn store_u16x8(dst: &mut [u16], v: __m128i) {
debug_assert!(dst.len() >= 8);
unsafe { _mm_storeu_si128(dst.as_mut_ptr().cast::<__m128i>(), v) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn store_u16x16(dst: &mut [u16], v: __m256i) {
debug_assert!(dst.len() >= 16);
unsafe { _mm256_storeu_si256(dst.as_mut_ptr().cast::<__m256i>(), v) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn pack_u16x8(sum: __m256i) -> __m128i {
let packed = _mm256_packus_epi32(sum, _mm256_setzero_si256());
let packed = _mm256_permute4x64_epi64::<0xD8>(packed);
_mm256_castsi256_si128(packed)
}
#[inline]
#[target_feature(enable = "avx2")]
fn pack_u16x16(lo: __m256i, hi: __m256i) -> __m256i {
_mm256_permute4x64_epi64::<0xD8>(_mm256_packus_epi32(lo, hi))
}
#[inline]
#[target_feature(enable = "avx2")]
fn add_clip8_i32(pred: __m128i, res: __m256i, zero: __m256i, max: __m256i) -> __m256i {
let pred = _mm256_cvtepu16_epi32(pred);
let sum = _mm256_add_epi32(pred, res);
_mm256_min_epi32(_mm256_max_epi32(sum, zero), max)
}
#[inline]
#[target_feature(enable = "avx2")]
fn add_clip16_i32(
pred: __m256i,
res_lo: __m256i,
res_hi: __m256i,
zero: __m256i,
max: __m256i,
) -> (__m256i, __m256i) {
let pred_lo = _mm256_cvtepu16_epi32(_mm256_castsi256_si128(pred));
let pred_hi = _mm256_cvtepu16_epi32(_mm256_extracti128_si256::<1>(pred));
let sum_lo = _mm256_add_epi32(pred_lo, res_lo);
let sum_hi = _mm256_add_epi32(pred_hi, res_hi);
(
_mm256_min_epi32(_mm256_max_epi32(sum_lo, zero), max),
_mm256_min_epi32(_mm256_max_epi32(sum_hi, zero), max),
)
}
#[inline]
#[target_feature(enable = "avx2")]
fn add_clip_row_avx2(dst: &mut [u16], pred: &[u16], res: &[i32], n: usize, max: __m256i) {
let zero = _mm256_setzero_si256();
if n == 2 {
let sum = add_clip8_i32(load_u16x2(pred), zext128(load_i32x2(res)), zero, max);
store_u16x2(dst, pack_u16x8(sum));
return;
}
if n == 4 {
let sum = add_clip8_i32(load_u16x4(pred), zext128(load_i32x4(res)), zero, max);
store_u16x4(dst, pack_u16x8(sum));
return;
}
if n == 8 {
let sum = add_clip8_i32(load_u16x8(pred), load_i32x8(res), zero, max);
store_u16x8(dst, pack_u16x8(sum));
return;
}
let (pred16, _) = pred[..n].as_chunks::<16>();
let (res16, _) = res[..n].as_chunks::<16>();
let (dst16, _) = dst[..n].as_chunks_mut::<16>();
for ((pred, res), dst) in pred16.iter().zip(res16.iter()).zip(dst16.iter_mut()) {
let pred = load_u16x16(pred);
let (res_lo, res_hi) = res.split_at(8);
let (sum_lo, sum_hi) =
add_clip16_i32(pred, load_i32x8(res_lo), load_i32x8(res_hi), zero, max);
store_u16x16(dst, pack_u16x16(sum_lo, sum_hi));
}
}
#[inline]
#[target_feature(enable = "avx2")]
fn add_clip_row_avx2_16(dst: &mut [u16], pred: &[u16], res: &[i16], n: usize, max: __m256i) {
let zero = _mm256_setzero_si256();
if n == 2 {
let pred = zext128(load_u16x2(pred));
let res = zext128(load_i16x2(res));
let sum = _mm256_min_epi16(_mm256_max_epi16(_mm256_adds_epi16(pred, res), zero), max);
store_u16x2(dst, _mm256_castsi256_si128(sum));
return;
}
if n == 4 {
let pred = zext128(load_u16x4(pred));
let res = zext128(load_i16x4(res));
let sum = _mm256_min_epi16(_mm256_max_epi16(_mm256_adds_epi16(pred, res), zero), max);
store_u16x4(dst, _mm256_castsi256_si128(sum));
return;
}
if n == 8 {
let pred = zext128(load_u16x8(pred));
let res = zext128(load_i16x8(res));
let sum = _mm256_min_epi16(_mm256_max_epi16(_mm256_adds_epi16(pred, res), zero), max);
store_u16x8(dst, _mm256_castsi256_si128(sum));
return;
}
let (pred16, _) = pred[..n].as_chunks::<16>();
let (res16, _) = res[..n].as_chunks::<16>();
let (dst16, _) = dst[..n].as_chunks_mut::<16>();
for ((pred, res), dst) in pred16.iter().zip(res16.iter()).zip(dst16.iter_mut()) {
let pred = load_u16x16(pred);
let res = load_i16x16(res);
let sum = _mm256_min_epi16(_mm256_max_epi16(_mm256_adds_epi16(pred, res), zero), max);
store_u16x16(dst, sum);
}
}
#[inline]
#[target_feature(enable = "avx2")]
fn add_residual_into_avx2_impl(
dst: &mut [u16],
stride: usize,
pred: &[u16],
res: &[i32],
n: usize,
bit_depth: u8,
) {
debug_assert!(supported_n(n));
let Some(n2) = n.checked_mul(n) else {
return;
};
let Some(pred) = pred.get(..n2) else {
return;
};
let Some(res) = res.get(..n2) else {
return;
};
let max = _mm256_set1_epi32(sample_max(bit_depth));
let dst_rows = dst.chunks_mut(stride).take(n);
let pred_rows = pred.chunks_exact(n);
let res_rows = res.chunks_exact(n);
for ((dst_row, pred_row), res_row) in dst_rows.zip(pred_rows).zip(res_rows) {
add_clip_row_avx2(&mut dst_row[..n], pred_row, res_row, n, max);
}
}
#[inline]
#[target_feature(enable = "avx2")]
fn add_residual_into_avx2_impl_16(
dst: &mut [u16],
stride: usize,
pred: &[u16],
res: &[i16],
n: usize,
bit_depth: u8,
) {
debug_assert!(supported_n(n));
let Some(n2) = n.checked_mul(n) else {
return;
};
let Some(pred) = pred.get(..n2) else {
return;
};
let Some(res) = res.get(..n2) else {
return;
};
let max = _mm256_set1_epi16(sample_max(bit_depth) as i16);
let dst_rows = dst.chunks_mut(stride).take(n);
let pred_rows = pred.chunks_exact(n);
let res_rows = res.chunks_exact(n);
for ((dst_row, pred_row), res_row) in dst_rows.zip(pred_rows).zip(res_rows) {
add_clip_row_avx2_16(&mut dst_row[..n], pred_row, res_row, n, max);
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn add_residual_into_avx2(
dst: &mut [u16],
stride: usize,
pred: &[u16],
res: &[i32],
n: usize,
valid_w: usize,
valid_h: usize,
bit_depth: u8,
) {
if !supported_n(n)
|| !can_reconstruct_full_block(dst, stride, pred, res, n, valid_w, valid_h, bit_depth)
{
return;
}
unsafe { add_residual_into_avx2_impl(dst, stride, pred, res, n, bit_depth) }
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn add_residual_into_avx2_16(
dst: &mut [u16],
stride: usize,
pred: &[u16],
res: &[i16],
n: usize,
valid_w: usize,
valid_h: usize,
bit_depth: u8,
) {
if !supported_n(n)
|| sample_max(bit_depth) > 32767
|| !can_reconstruct_full_block(dst, stride, pred, res, n, valid_w, valid_h, bit_depth)
{
return;
}
unsafe { add_residual_into_avx2_impl_16(dst, stride, pred, res, n, bit_depth) }
}
#[inline]
#[target_feature(enable = "avx2")]
fn narrow_i32x16_avx2(src: &[i32; 16]) -> __m256i {
let (src8, _) = src.as_chunks::<8>();
_mm256_permute4x64_epi64::<0xD8>(_mm256_packs_epi32(
load_i32x8(&src8[0]),
load_i32x8(&src8[1]),
))
}
#[inline]
#[target_feature(enable = "avx2")]
fn narrow_i32_to_i16_avx2_impl(src: &[i32], dst: &mut [i16]) {
debug_assert_eq!(src.len(), dst.len());
debug_assert_eq!(src.len() & 15, 0);
let (src64, src_rem) = src.as_chunks::<64>();
let (dst64, dst_rem) = dst.as_chunks_mut::<64>();
for (src, dst) in src64.iter().zip(dst64.iter_mut()) {
let (src16, _) = src.as_chunks::<16>();
let (dst16, _) = dst.as_chunks_mut::<16>();
let p0 = narrow_i32x16_avx2(&src16[0]);
let p1 = narrow_i32x16_avx2(&src16[1]);
let p2 = narrow_i32x16_avx2(&src16[2]);
let p3 = narrow_i32x16_avx2(&src16[3]);
unsafe {
_mm256_storeu_si256(dst16[0].as_mut_ptr().cast::<__m256i>(), p0);
_mm256_storeu_si256(dst16[1].as_mut_ptr().cast::<__m256i>(), p1);
_mm256_storeu_si256(dst16[2].as_mut_ptr().cast::<__m256i>(), p2);
_mm256_storeu_si256(dst16[3].as_mut_ptr().cast::<__m256i>(), p3);
}
}
let (src16, src_tail) = src_rem.as_chunks::<16>();
let (dst16, dst_tail) = dst_rem.as_chunks_mut::<16>();
debug_assert!(src_tail.is_empty());
debug_assert!(dst_tail.is_empty());
for (src, dst) in src16.iter().zip(dst16.iter_mut()) {
let packed = narrow_i32x16_avx2(src);
unsafe { _mm256_storeu_si256(dst.as_mut_ptr().cast::<__m256i>(), packed) };
}
}
pub(crate) fn narrow_i32_to_i16_avx2(src: &[i32], dst: &mut [i16]) {
let len = src.len().min(dst.len());
let simd_len = len & !15;
let (src_simd, src_tail) = src[..len].split_at(simd_len);
let (dst_simd, dst_tail) = dst[..len].split_at_mut(simd_len);
if simd_len != 0 {
unsafe { narrow_i32_to_i16_avx2_impl(src_simd, dst_simd) };
}
narrow_i32_to_i16_scalar(src_tail, dst_tail);
}