#[cfg(target_arch = "x86")]
use core::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
use crate::reconstruct::{add_residual_into_scalar, can_reconstruct_full_block, sample_max};
#[inline]
fn supported_n(n: usize) -> bool {
matches!(n, 2 | 4 | 8 | 16 | 32)
}
#[inline]
#[target_feature(enable = "sse4.1")]
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 = "sse4.1")]
fn load_u16x4(src: &[u16]) -> __m128i {
debug_assert!(src.len() >= 4);
unsafe { _mm_loadl_epi64(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "sse4.1")]
fn load_u16x8(src: &[u16]) -> __m128i {
debug_assert!(src.len() >= 8);
unsafe { _mm_loadu_si128(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "sse4.1")]
fn load_i32x2(src: &[i32]) -> __m128i {
debug_assert!(src.len() >= 2);
unsafe { _mm_loadl_epi64(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "sse4.1")]
fn load_i32x4(src: &[i32]) -> __m128i {
debug_assert!(src.len() >= 4);
unsafe { _mm_loadu_si128(src.as_ptr().cast::<__m128i>()) }
}
#[inline]
#[target_feature(enable = "sse4.1")]
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 = "sse4.1")]
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 = "sse4.1")]
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 = "sse4.1")]
fn add_clip4_sse41(pred: __m128i, res: __m128i, zero: __m128i, max: __m128i) -> __m128i {
let pred = _mm_cvtepu16_epi32(pred);
let sum = _mm_add_epi32(pred, res);
_mm_min_epi32(_mm_max_epi32(sum, zero), max)
}
#[inline]
#[target_feature(enable = "sse4.1")]
fn add_clip8_sse41(pred: __m128i, res: &[i32], zero: __m128i, max: __m128i) -> __m128i {
let lo = add_clip4_sse41(pred, load_i32x4(res), zero, max);
let hi = add_clip4_sse41(_mm_srli_si128::<8>(pred), load_i32x4(&res[4..]), zero, max);
_mm_packus_epi32(lo, hi)
}
#[inline]
#[target_feature(enable = "sse4.1")]
fn add_clip_row_sse41(dst: &mut [u16], pred: &[u16], res: &[i32], n: usize, max: __m128i) {
let zero = _mm_setzero_si128();
if n == 2 {
let pred = load_u16x2(pred);
let sum = add_clip4_sse41(pred, load_i32x2(res), zero, max);
store_u16x2(dst, _mm_packus_epi32(sum, zero));
return;
}
if n == 4 {
let pred = load_u16x4(pred);
let sum = add_clip4_sse41(pred, load_i32x4(res), zero, max);
store_u16x4(dst, _mm_packus_epi32(sum, zero));
return;
}
let mut x = 0usize;
while x < n {
let pred = load_u16x8(&pred[x..]);
let out = add_clip8_sse41(pred, &res[x..], zero, max);
store_u16x8(&mut dst[x..], out);
x += 8;
}
}
#[inline]
#[target_feature(enable = "sse4.1")]
fn add_residual_into_sse41_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 = _mm_set1_epi32(sample_max(bit_depth));
for y in 0..n {
let row_off = y * n;
let dst_off = y * stride;
let Some(dst_row) = dst.get_mut(dst_off..dst_off.saturating_add(n)) else {
break;
};
let Some(pred_row) = pred.get(row_off..row_off + n) else {
break;
};
let Some(res_row) = res.get(row_off..row_off + n) else {
break;
};
add_clip_row_sse41(dst_row, pred_row, res_row, n, max);
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn add_residual_into_sse41(
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)
{
add_residual_into_scalar(dst, stride, pred, res, n, valid_w, valid_h, bit_depth);
return;
}
unsafe { add_residual_into_sse41_impl(dst, stride, pred, res, n, bit_depth) }
}