#[inline]
#[allow(clippy::too_many_arguments)]
fn sad_scalar(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
let mut src_p = src_off;
let mut ref_p = ref_off;
let mut sad = 0u32;
for _ in 0..h {
for c in 0..w {
let s = src[src_p + c] as i32;
let r = reference[ref_p + c] as i32;
sad += (s - r).unsigned_abs();
}
src_p += src_stride;
ref_p += ref_stride;
}
sad
}
#[inline]
#[allow(clippy::too_many_arguments)]
pub(crate) fn sad(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
debug_assert!(w == 0 || h == 0 || src_off + (h - 1) * src_stride + w <= src.len());
debug_assert!(w == 0 || h == 0 || ref_off + (h - 1) * ref_stride + w <= reference.len());
dispatch(
src, src_off, src_stride, reference, ref_off, ref_stride, w, h,
)
}
#[cfg(target_arch = "aarch64")]
#[inline]
#[allow(clippy::too_many_arguments)]
fn dispatch(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
unsafe {
aarch64::sad_neon(
src, src_off, src_stride, reference, ref_off, ref_stride, w, h,
)
}
}
#[cfg(target_arch = "x86_64")]
#[inline]
#[allow(clippy::too_many_arguments)]
fn dispatch(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
if x86::avx2_available() {
unsafe {
x86::sad_avx2(
src, src_off, src_stride, reference, ref_off, ref_stride, w, h,
)
}
} else {
unsafe {
x86::sad_sse2(
src, src_off, src_stride, reference, ref_off, ref_stride, w, h,
)
}
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
#[inline]
#[allow(clippy::too_many_arguments)]
fn dispatch(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
sad_scalar(
src, src_off, src_stride, reference, ref_off, ref_stride, w, h,
)
}
#[cfg(target_arch = "aarch64")]
mod aarch64 {
use std::arch::aarch64::*;
#[target_feature(enable = "neon")]
#[allow(clippy::too_many_arguments)]
pub(super) unsafe fn sad_neon(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
let src_ptr = src.as_ptr();
let ref_ptr = reference.as_ptr();
let mut acc16 = vdupq_n_u16(0);
let mut tail: u32 = 0;
let mut src_p = src_off;
let mut ref_p = ref_off;
for _ in 0..h {
let mut c = 0;
while c + 16 <= w {
let s = vld1q_u8(src_ptr.add(src_p + c));
let r = vld1q_u8(ref_ptr.add(ref_p + c));
acc16 = vpadalq_u8(acc16, vabdq_u8(s, r));
c += 16;
}
if c + 8 <= w {
let s = vld1_u8(src_ptr.add(src_p + c));
let r = vld1_u8(ref_ptr.add(ref_p + c));
let d = vabd_u8(s, r);
acc16 = vaddq_u16(acc16, vmovl_u8(d));
c += 8;
}
while c < w {
let s = *src_ptr.add(src_p + c) as i32;
let r = *ref_ptr.add(ref_p + c) as i32;
tail += (s - r).unsigned_abs();
c += 1;
}
src_p += src_stride;
ref_p += ref_stride;
}
vaddlvq_u16(acc16) + tail
}
}
#[cfg(target_arch = "x86_64")]
mod x86 {
use std::arch::x86_64::*;
use std::sync::atomic::{AtomicU8, Ordering};
static AVX2: AtomicU8 = AtomicU8::new(0);
#[inline]
pub(super) fn avx2_available() -> bool {
match AVX2.load(Ordering::Relaxed) {
1 => true,
2 => false,
_ => {
let has = is_x86_feature_detected!("avx2");
AVX2.store(if has { 1 } else { 2 }, Ordering::Relaxed);
has
}
}
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn hsum_epi64x4(v: __m256i) -> u32 {
let lo = _mm256_castsi256_si128(v);
let hi = _mm256_extracti128_si256(v, 1);
let s = _mm_add_epi64(lo, hi);
let s = _mm_add_epi64(s, _mm_unpackhi_epi64(s, s));
_mm_cvtsi128_si32(s) as u32
}
#[inline]
#[target_feature(enable = "sse2")]
unsafe fn hsum_epi64x2(v: __m128i) -> u32 {
let s = _mm_add_epi64(v, _mm_unpackhi_epi64(v, v));
_mm_cvtsi128_si32(s) as u32
}
#[target_feature(enable = "sse2")]
#[allow(clippy::too_many_arguments)]
pub(super) unsafe fn sad_sse2(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
let src_ptr = src.as_ptr();
let ref_ptr = reference.as_ptr();
let mut acc = _mm_setzero_si128();
let mut tail: u32 = 0;
let mut src_p = src_off;
let mut ref_p = ref_off;
for _ in 0..h {
let mut c = 0;
while c + 16 <= w {
let s = _mm_loadu_si128(src_ptr.add(src_p + c) as *const __m128i);
let r = _mm_loadu_si128(ref_ptr.add(ref_p + c) as *const __m128i);
acc = _mm_add_epi64(acc, _mm_sad_epu8(s, r));
c += 16;
}
if c + 8 <= w {
let s = _mm_loadl_epi64(src_ptr.add(src_p + c) as *const __m128i);
let r = _mm_loadl_epi64(ref_ptr.add(ref_p + c) as *const __m128i);
acc = _mm_add_epi64(acc, _mm_sad_epu8(s, r));
c += 8;
}
while c < w {
let s = *src_ptr.add(src_p + c) as i32;
let r = *ref_ptr.add(ref_p + c) as i32;
tail += (s - r).unsigned_abs();
c += 1;
}
src_p += src_stride;
ref_p += ref_stride;
}
hsum_epi64x2(acc) + tail
}
#[target_feature(enable = "avx2")]
#[allow(clippy::too_many_arguments)]
pub(super) unsafe fn sad_avx2(
src: &[u8],
src_off: usize,
src_stride: usize,
reference: &[u8],
ref_off: usize,
ref_stride: usize,
w: usize,
h: usize,
) -> u32 {
if w != 16 || h < 2 {
return sad_sse2(
src, src_off, src_stride, reference, ref_off, ref_stride, w, h,
);
}
let src_ptr = src.as_ptr();
let ref_ptr = reference.as_ptr();
let mut acc = _mm256_setzero_si256();
let pairs = h / 2;
let mut src_p = src_off;
let mut ref_p = ref_off;
for _ in 0..pairs {
let s0 = _mm_loadu_si128(src_ptr.add(src_p) as *const __m128i);
let s1 = _mm_loadu_si128(src_ptr.add(src_p + src_stride) as *const __m128i);
let r0 = _mm_loadu_si128(ref_ptr.add(ref_p) as *const __m128i);
let r1 = _mm_loadu_si128(ref_ptr.add(ref_p + ref_stride) as *const __m128i);
let s = _mm256_set_m128i(s1, s0);
let r = _mm256_set_m128i(r1, r0);
acc = _mm256_add_epi64(acc, _mm256_sad_epu8(s, r));
src_p += 2 * src_stride;
ref_p += 2 * ref_stride;
}
let mut total = hsum_epi64x4(acc);
if h & 1 == 1 {
total += sad_sse2(src, src_p, src_stride, reference, ref_p, ref_stride, w, 1);
}
total
}
}
#[cfg(test)]
mod tests {
use super::*;
struct XorShift32(u32);
impl XorShift32 {
fn next(&mut self) -> u32 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
self.0 = x;
x
}
fn byte(&mut self) -> u8 {
(self.next() & 0xff) as u8
}
}
#[test]
fn simd_matches_scalar_all_sizes() {
let mut rng = XorShift32(0x1234_5678);
let sizes = [
(16usize, 16usize),
(8, 8),
(16, 8),
(8, 16),
(13, 7),
(16, 1),
];
for &(w, h) in &sizes {
let src_stride = w + 5;
let ref_stride = w + 11;
let src_off = 3;
let ref_off = 7;
let src_len = src_off + h * src_stride + 16;
let ref_len = ref_off + h * ref_stride + 16;
for trial in 0..64 {
let mut src = vec![0u8; src_len];
let mut reference = vec![0u8; ref_len];
match trial {
0 => { }
1 => reference.iter_mut().for_each(|b| *b = 255),
2 => src.iter_mut().for_each(|b| *b = 255),
3 => {
src.iter_mut().for_each(|b| *b = 255);
reference.iter_mut().for_each(|b| *b = 255);
}
_ => {
src.iter_mut().for_each(|b| *b = rng.byte());
reference.iter_mut().for_each(|b| *b = rng.byte());
}
}
let expected = sad_scalar(
&src, src_off, src_stride, &reference, ref_off, ref_stride, w, h,
);
let got = sad(
&src, src_off, src_stride, &reference, ref_off, ref_stride, w, h,
);
assert_eq!(
got, expected,
"mismatch for {w}x{h} trial {trial}: simd={got} scalar={expected}"
);
}
}
}
#[test]
fn sad_known_value() {
let w = 8;
let h = 8;
let src = vec![10u8; w * h];
let reference = vec![3u8; w * h];
assert_eq!(sad(&src, 0, w, &reference, 0, w, w, h), 448);
}
}