#![allow(clippy::needless_range_loop)]
use crate::dct::{
SQRT2, WC4_0, WC4_1, WC8_0, WC8_1, WC8_2, WC8_3, WC16_0, WC16_1, WC16_2, WC16_3, WC16_4,
WC16_5, WC16_6, WC16_7, WC32,
};
use std::arch::aarch64::*;
use std::mem::MaybeUninit;
#[derive(Clone, Copy)]
struct I32x8 {
lo: int32x4_t,
hi: int32x4_t,
}
impl I32x8 {
#[inline]
#[target_feature(enable = "neon")]
fn zero() -> Self {
Self {
lo: vdupq_n_s32(0),
hi: vdupq_n_s32(0),
}
}
#[inline]
#[target_feature(enable = "neon")]
fn add(self, rhs: I32x8) -> I32x8 {
I32x8 {
lo: vaddq_s32(self.lo, rhs.lo),
hi: vaddq_s32(self.hi, rhs.hi),
}
}
#[inline]
#[target_feature(enable = "neon")]
fn sub(self, rhs: I32x8) -> I32x8 {
I32x8 {
lo: vsubq_s32(self.lo, rhs.lo),
hi: vsubq_s32(self.hi, rhs.hi),
}
}
#[inline]
#[target_feature(enable = "neon")]
fn muls_q16(self, coeff: i32) -> I32x8 {
let c = vdup_n_s32(coeff);
I32x8 {
lo: vcombine_s32(
vshrn_n_s64(vmull_s32(vget_low_s32(self.lo), c), 16),
vshrn_n_s64(vmull_s32(vget_high_s32(self.lo), c), 16),
),
hi: vcombine_s32(
vshrn_n_s64(vmull_s32(vget_low_s32(self.hi), c), 16),
vshrn_n_s64(vmull_s32(vget_high_s32(self.hi), c), 16),
),
}
}
#[inline]
#[target_feature(enable = "neon")]
fn fma_sqrt2(self, b: I32x8) -> I32x8 {
self.muls_q16(SQRT2).add(b)
}
#[inline]
#[target_feature(enable = "neon")]
fn shr<const N: i32>(self) -> I32x8 {
I32x8 {
lo: vshrq_n_s32(self.lo, N),
hi: vshrq_n_s32(self.hi, N),
}
}
#[inline]
#[target_feature(enable = "neon")]
fn shl<const N: i32>(self) -> I32x8 {
I32x8 {
lo: vshlq_n_s32(self.lo, N),
hi: vshlq_n_s32(self.hi, N),
}
}
#[inline]
#[target_feature(enable = "neon")]
fn shr_round<const N: i32>(self) -> I32x8 {
I32x8 {
lo: vrshrq_n_s32(self.lo, N),
hi: vrshrq_n_s32(self.hi, N),
}
}
}
#[inline]
#[target_feature(enable = "neon")]
fn transpose_4x4_i32(
r0: int32x4_t,
r1: int32x4_t,
r2: int32x4_t,
r3: int32x4_t,
) -> (int32x4_t, int32x4_t, int32x4_t, int32x4_t) {
let v0 = vtrn1q_s32(r0, r1);
let v1 = vtrn2q_s32(r0, r1);
let v2 = vtrn1q_s32(r2, r3);
let v3 = vtrn2q_s32(r2, r3);
let c0 = vreinterpretq_s32_s64(vtrn1q_s64(
vreinterpretq_s64_s32(v0),
vreinterpretq_s64_s32(v2),
));
let c1 = vreinterpretq_s32_s64(vtrn1q_s64(
vreinterpretq_s64_s32(v1),
vreinterpretq_s64_s32(v3),
));
let c2 = vreinterpretq_s32_s64(vtrn2q_s64(
vreinterpretq_s64_s32(v0),
vreinterpretq_s64_s32(v2),
));
let c3 = vreinterpretq_s32_s64(vtrn2q_s64(
vreinterpretq_s64_s32(v1),
vreinterpretq_s64_s32(v3),
));
(c0, c1, c2, c3)
}
#[inline]
#[target_feature(enable = "neon")]
fn transpose_8x8_i32(c: &mut [I32x8; 8]) {
let (a0, a1, a2, a3) = transpose_4x4_i32(c[0].lo, c[1].lo, c[2].lo, c[3].lo);
let (b0, b1, b2, b3) = transpose_4x4_i32(c[0].hi, c[1].hi, c[2].hi, c[3].hi);
let (cc0, cc1, cc2, cc3) = transpose_4x4_i32(c[4].lo, c[5].lo, c[6].lo, c[7].lo);
let (d0, d1, d2, d3) = transpose_4x4_i32(c[4].hi, c[5].hi, c[6].hi, c[7].hi);
c[0] = I32x8 { lo: a0, hi: cc0 };
c[1] = I32x8 { lo: a1, hi: cc1 };
c[2] = I32x8 { lo: a2, hi: cc2 };
c[3] = I32x8 { lo: a3, hi: cc3 };
c[4] = I32x8 { lo: b0, hi: d0 };
c[5] = I32x8 { lo: b1, hi: d1 };
c[6] = I32x8 { lo: b2, hi: d2 };
c[7] = I32x8 { lo: b3, hi: d3 };
}
#[inline]
#[target_feature(enable = "neon")]
fn dct1d_4_v_i32(c: &mut [I32x8; 4]) {
let t0 = c[0].add(c[3]);
let t1 = c[1].add(c[2]);
let sum = t0.add(t1);
let diff = t0.sub(t1);
let t2 = c[0].sub(c[3]).muls_q16(WC4_0);
let t3 = c[1].sub(c[2]).muls_q16(WC4_1);
let t2p = t2.add(t3);
let t3p = t2.sub(t3);
let t2pp = t2p.fma_sqrt2(t3p);
c[0] = sum;
c[1] = t2pp;
c[2] = diff;
c[3] = t3p;
}
#[inline]
#[target_feature(enable = "neon")]
fn dct1d_8_v_i32(c: &mut [I32x8; 8]) {
let mut evens = [
c[0].add(c[7]),
c[1].add(c[6]),
c[2].add(c[5]),
c[3].add(c[4]),
];
dct1d_4_v_i32(&mut evens);
let mut odds = [
c[0].sub(c[7]).muls_q16(WC8_0),
c[1].sub(c[6]).muls_q16(WC8_1),
c[2].sub(c[5]).muls_q16(WC8_2),
c[3].sub(c[4]).muls_q16(WC8_3),
];
dct1d_4_v_i32(&mut odds);
odds[0] = odds[0].fma_sqrt2(odds[1]); odds[1] = odds[1].add(odds[2]);
odds[2] = odds[2].add(odds[3]);
c[0] = evens[0];
c[1] = odds[0];
c[2] = evens[1];
c[3] = odds[1];
c[4] = evens[2];
c[5] = odds[2];
c[6] = evens[3];
c[7] = odds[3];
}
#[inline]
#[target_feature(enable = "neon")]
fn dct1d_16_v_i32(c: &mut [I32x8; 16]) {
let mut evens = [
c[0].add(c[15]),
c[1].add(c[14]),
c[2].add(c[13]),
c[3].add(c[12]),
c[4].add(c[11]),
c[5].add(c[10]),
c[6].add(c[9]),
c[7].add(c[8]),
];
let mut odds = [
c[0].sub(c[15]).muls_q16(WC16_0),
c[1].sub(c[14]).muls_q16(WC16_1),
c[2].sub(c[13]).muls_q16(WC16_2),
c[3].sub(c[12]).muls_q16(WC16_3),
c[4].sub(c[11]).muls_q16(WC16_4),
c[5].sub(c[10]).muls_q16(WC16_5),
c[6].sub(c[9]).muls_q16(WC16_6),
c[7].sub(c[8]).muls_q16(WC16_7),
];
dct1d_8_v_i32(&mut evens);
dct1d_8_v_i32(&mut odds);
odds[0] = odds[0].fma_sqrt2(odds[1]);
odds[1] = odds[1].add(odds[2]);
odds[2] = odds[2].add(odds[3]);
odds[3] = odds[3].add(odds[4]);
odds[4] = odds[4].add(odds[5]);
odds[5] = odds[5].add(odds[6]);
odds[6] = odds[6].add(odds[7]);
c[0] = evens[0];
c[1] = odds[0];
c[2] = evens[1];
c[3] = odds[1];
c[4] = evens[2];
c[5] = odds[2];
c[6] = evens[3];
c[7] = odds[3];
c[8] = evens[4];
c[9] = odds[4];
c[10] = evens[5];
c[11] = odds[5];
c[12] = evens[6];
c[13] = odds[6];
c[14] = evens[7];
c[15] = odds[7];
}
#[inline]
#[target_feature(enable = "neon")]
fn load8_i32(ptr: &[i32], stride: usize) -> [I32x8; 8] {
unsafe {
let row = |y: usize| {
let p = ptr.get_unchecked(y * stride..);
I32x8 {
lo: vld1q_s32(p.as_ptr()),
hi: vld1q_s32(p.get_unchecked(4..).as_ptr()),
}
};
std::array::from_fn(row)
}
}
#[inline]
#[target_feature(enable = "neon")]
fn load16_i32(ptr: &[i32], stride: usize) -> [I32x8; 16] {
unsafe {
let row = |y: usize| {
let p = ptr.get_unchecked(y * stride..);
I32x8 {
lo: vld1q_s32(p.as_ptr()),
hi: vld1q_s32(p.get_unchecked(4..).as_ptr()),
}
};
std::array::from_fn(row)
}
}
#[inline]
#[target_feature(enable = "neon")]
fn mul_q16_vec(data: int32x4_t, coeff: int32x4_t) -> int32x4_t {
vcombine_s32(
vshrn_n_s64::<16>(vmull_s32(vget_low_s32(data), vget_low_s32(coeff))),
vshrn_n_s64::<16>(vmull_s32(vget_high_s32(data), vget_high_s32(coeff))),
)
}
#[inline]
fn quant_flat<const N: usize>(coeffs: &[i32; N], dc_q: i32, ac_q: i32, out: &mut [i32; N]) {
let mq = |a: i32, b: i32| {
let prod = (a as i64) * (b as i64);
let mag = prod.unsigned_abs();
if mag < 65536 {
return 0;
}
let lvl = ((mag + 32768) >> 16) as i32;
if prod >= 0 { lvl } else { -lvl }
};
out[0] = mq(coeffs[0], dc_q);
for i in 1..N {
out[i] = mq(coeffs[i], ac_q);
}
}
#[target_feature(enable = "neon")]
pub(crate) fn dct8x8_neon_coeffs(input: &[i32; 64]) -> [i32; 64] {
let mut cols = load8_i32(input, 8);
dct1d_8_v_i32(&mut cols);
transpose_8x8_i32(&mut cols);
dct1d_8_v_i32(&mut cols);
let mut out = MaybeUninit::<[i32; 64]>::uninit();
for (k, col) in cols.iter().enumerate() {
unsafe {
let dst_ptr = out.as_mut_ptr() as *mut i32;
vst1q_s32(dst_ptr.add(k * 8), col.lo);
vst1q_s32(dst_ptr.add(k * 8 + 4), col.hi);
}
}
unsafe { out.assume_init() }
}
#[target_feature(enable = "neon")]
pub(crate) fn dct16x16_neon_coeffs(input: &[i32; 256]) -> [i32; 256] {
unsafe {
let mut c_l = load16_i32(input, 16);
let mut c_r = load16_i32(&input[8..], 16);
dct1d_16_v_i32(&mut c_l);
dct1d_16_v_i32(&mut c_r);
let mut top_l: [I32x8; 8] = c_l[..8].try_into().unwrap();
let mut bot_l: [I32x8; 8] = c_l[8..16].try_into().unwrap();
let mut top_r: [I32x8; 8] = c_r[0..8].try_into().unwrap();
let mut bot_r: [I32x8; 8] = c_r[8..16].try_into().unwrap();
transpose_8x8_i32(&mut top_l);
transpose_8x8_i32(&mut bot_l);
transpose_8x8_i32(&mut top_r);
transpose_8x8_i32(&mut bot_r);
let mut d_a = [I32x8::zero(); 16];
let mut d_b = [I32x8::zero(); 16];
d_a[0..8].copy_from_slice(&top_l);
d_a[8..16].copy_from_slice(&top_r);
d_b[0..8].copy_from_slice(&bot_l);
d_b[8..16].copy_from_slice(&bot_r);
dct1d_16_v_i32(&mut d_a);
dct1d_16_v_i32(&mut d_b);
let mut out = MaybeUninit::<[i32; 256]>::uninit();
for u in 0..16usize {
let na = d_a[u].shr::<1>();
let nb = d_b[u].shr::<1>();
let dst_ptr = out.as_mut_ptr() as *mut i32;
let base = dst_ptr.add(u * 16);
vst1q_s32(base, na.lo);
vst1q_s32(base.add(4), na.hi);
vst1q_s32(base.add(8), nb.lo);
vst1q_s32(base.add(12), nb.hi);
}
out.assume_init()
}
}
#[target_feature(enable = "neon")]
pub(crate) fn dct16x16_neon_i32(input: &mut [i32; 256], dc_q: i32, ac_q: i32) {
let coeffs = dct16x16_neon_coeffs(input);
quant_flat(&coeffs, dc_q, ac_q, input);
}
#[inline]
#[target_feature(enable = "neon")]
fn dct1d_32_v_i32(c: &mut [I32x8; 32]) {
let mut evens = std::array::from_fn::<I32x8, 16, _>(|i| c[i].add(c[31 - i]));
let mut odds = std::array::from_fn::<I32x8, 16, _>(|i| c[i].sub(c[31 - i]));
dct1d_16_v_i32(&mut evens);
for i in 0..16 {
odds[i] = odds[i].muls_q16(WC32[i]);
}
dct1d_16_v_i32(&mut odds);
odds[0] = odds[0].fma_sqrt2(odds[1]);
odds[1] = odds[1].add(odds[2]);
odds[2] = odds[2].add(odds[3]);
odds[3] = odds[3].add(odds[4]);
odds[4] = odds[4].add(odds[5]);
odds[5] = odds[5].add(odds[6]);
odds[6] = odds[6].add(odds[7]);
odds[7] = odds[7].add(odds[8]);
odds[8] = odds[8].add(odds[9]);
odds[9] = odds[9].add(odds[10]);
odds[10] = odds[10].add(odds[11]);
odds[11] = odds[11].add(odds[12]);
odds[12] = odds[12].add(odds[13]);
odds[13] = odds[13].add(odds[14]);
odds[14] = odds[14].add(odds[15]);
for i in 0..16 {
c[2 * i] = evens[i];
c[2 * i + 1] = odds[i];
}
}
#[inline]
#[target_feature(enable = "neon")]
fn load32_i32(ptr: &[i32], stride: usize) -> [I32x8; 32] {
unsafe {
std::array::from_fn(|y| {
let p = &ptr[y * stride..];
I32x8 {
lo: vld1q_s32(p.as_ptr()),
hi: vld1q_s32(p[4..].as_ptr()),
}
})
}
}
#[target_feature(enable = "neon")]
pub(crate) fn dct32x32_neon_coeffs(input: &[i32; 1024]) -> [i32; 1024] {
let mut tmp_u = MaybeUninit::<[i32; 1024]>::uninit();
for group in 0..4usize {
let col_start = group * 8;
let mut cols = load32_i32(&input[col_start..], 32);
for c in cols.iter_mut() {
*c = c.shl::<6>();
}
dct1d_32_v_i32(&mut cols);
for v in 0..32usize {
let dst_ptr = tmp_u.as_mut_ptr() as *mut i32;
let base = unsafe { dst_ptr.add(v * 32 + col_start) };
unsafe {
vst1q_s32(base, cols[v].lo);
vst1q_s32(base.add(4), cols[v].hi);
}
}
}
let tmp = unsafe { tmp_u.assume_init() };
let mut out = MaybeUninit::<[i32; 1024]>::uninit();
for group in 0..4usize {
let row_start = group * 8;
let base_off = row_start * 32;
let mut seg_a = load8_i32(&tmp[base_off..], 32);
let mut seg_b = load8_i32(&tmp[base_off + 8..], 32);
let mut seg_c = load8_i32(&tmp[base_off + 16..], 32);
let mut seg_d = load8_i32(&tmp[base_off + 24..], 32);
transpose_8x8_i32(&mut seg_a);
transpose_8x8_i32(&mut seg_b);
transpose_8x8_i32(&mut seg_c);
transpose_8x8_i32(&mut seg_d);
let mut rows = [I32x8::zero(); 32];
rows[..8].copy_from_slice(&seg_a);
rows[8..16].copy_from_slice(&seg_b);
rows[16..24].copy_from_slice(&seg_c);
rows[24..32].copy_from_slice(&seg_d);
dct1d_32_v_i32(&mut rows);
for u in 0..32usize {
let n = rows[u].shr_round::<8>();
unsafe {
let dst_ptr = out.as_mut_ptr() as *mut i32;
let base = dst_ptr.add(u * 32 + row_start);
vst1q_s32(base, n.lo);
vst1q_s32(base.add(4), n.hi);
}
}
}
unsafe { out.assume_init() }
}
#[target_feature(enable = "neon")]
pub(crate) fn dct32x32_neon_i32(input: &mut [i32; 1024], dc_q: i32, ac_q: i32) {
let coeffs = dct32x32_neon_coeffs(input);
quant_flat(&coeffs, dc_q, ac_q, input);
}
#[target_feature(enable = "neon")]
pub(crate) fn dct8x16_neon_coeffs(input: &[i32; 128]) -> [i32; 128] {
let mut rows = load16_i32(input, 8); dct1d_16_v_i32(&mut rows);
let mut a: [I32x8; 8] = rows[0..8].try_into().unwrap();
let mut b: [I32x8; 8] = rows[8..16].try_into().unwrap();
transpose_8x8_i32(&mut a);
transpose_8x8_i32(&mut b);
dct1d_8_v_i32(&mut a); dct1d_8_v_i32(&mut b);
let nrm = vdupq_n_s32(46341);
let mut out = MaybeUninit::<[i32; 128]>::uninit();
for fx in 0..8usize {
let a_lo = mul_q16_vec(a[fx].lo, nrm);
let a_hi = mul_q16_vec(a[fx].hi, nrm);
let b_lo = mul_q16_vec(b[fx].lo, nrm);
let b_hi = mul_q16_vec(b[fx].hi, nrm);
unsafe {
let dst_ptr = out.as_mut_ptr() as *mut i32;
let base = dst_ptr.add(fx * 16);
vst1q_s32(base, a_lo);
vst1q_s32(base.add(4), a_hi);
vst1q_s32(base.add(8), b_lo);
vst1q_s32(base.add(12), b_hi);
}
}
unsafe { out.assume_init() }
}
#[allow(unused)]
#[target_feature(enable = "neon")]
pub(crate) fn dct8x16_neon_i32(input: &mut [i32; 128], dc_q: i32, ac_q: i32) {
let coeffs = dct8x16_neon_coeffs(input);
quant_flat(&coeffs, dc_q, ac_q, input);
}
#[cfg(test)]
mod neon_vs_scalar {
use crate::dct::{dct8x16_i32_scalar, dct16x16_scalar, dct32x32_scalar};
use crate::neon::{dct8x16_neon_i32, dct16x16_neon_i32, dct32x32_neon_i32};
fn lcg(state: &mut u32) -> i32 {
*state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((*state >> 16) as i32 & 0x3FF) - 512
}
fn fill_lcg(buf: &mut [i32], seed: u32) {
let mut s = seed;
for v in buf.iter_mut() {
*v = lcg(&mut s);
}
}
fn fill_ramp(buf: &mut [i32]) {
for (i, v) in buf.iter_mut().enumerate() {
*v = (i % 256) as i32 - 128;
}
}
fn fill_alt(buf: &mut [i32]) {
for (i, v) in buf.iter_mut().enumerate() {
*v = if i % 2 == 0 { 64 } else { -64 };
}
}
const QUANT_PAIRS: &[(i32, i32)] = &[
(65536, 65536), (65536, 46341), (32768, 32768), ];
fn run_16x16(input: [i32; 256], dc_q: i32, ac_q: i32) -> ([i32; 256], [i32; 256]) {
let mut scalar = input;
dct16x16_scalar(&mut scalar, dc_q, ac_q);
let mut neon = input;
unsafe { dct16x16_neon_i32(&mut neon, dc_q, ac_q) };
(scalar, neon)
}
#[test]
fn test_16x16_zeros() {
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_16x16([0i32; 256], dc_q, ac_q);
assert_eq!(s, n, "16x16 zeros dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_16x16_constant() {
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_16x16([32i32; 256], dc_q, ac_q);
assert_eq!(s, n, "16x16 constant dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_16x16_ramp() {
let mut input = [0i32; 256];
fill_ramp(&mut input);
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_16x16(input, dc_q, ac_q);
assert_eq!(s, n, "16x16 ramp dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_16x16_alternating() {
let mut input = [0i32; 256];
fill_alt(&mut input);
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_16x16(input, dc_q, ac_q);
assert_eq!(s, n, "16x16 alternating dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_16x16_random_seed1() {
let mut input = [0i32; 256];
fill_lcg(&mut input, 0x1234_5678);
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_16x16(input, dc_q, ac_q);
assert_eq!(s, n, "16x16 rand(12345678) dc_q={dc_q} ac_q={ac_q}");
}
}
fn run_32x32(input: [i32; 1024], dc_q: i32, ac_q: i32) -> ([i32; 1024], [i32; 1024]) {
let mut scalar = input;
dct32x32_scalar(&mut scalar, dc_q, ac_q);
let mut neon = input;
unsafe { dct32x32_neon_i32(&mut neon, dc_q, ac_q) };
(scalar, neon)
}
#[test]
fn test_32x32_zeros() {
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_32x32([0i32; 1024], dc_q, ac_q);
assert_eq!(s, n, "32x32 zeros dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_32x32_constant() {
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_32x32([8i32; 1024], dc_q, ac_q);
assert_eq!(s, n, "32x32 constant dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_32x32_ramp() {
let mut input = [0i32; 1024];
for (i, v) in input.iter_mut().enumerate() {
*v = (i % 128) as i32 - 64;
}
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_32x32(input, dc_q, ac_q);
let first = s.iter().zip(n.iter()).position(|(a, b)| a != b);
assert_eq!(
s, n,
"32x32 ramp dc_q={dc_q} ac_q={ac_q}: mismatch at {first:?}"
);
}
}
#[test]
fn test_32x32_random_seed0() {
let mut input = [0i32; 1024];
let mut s = 0xDEAD_BEEFu32;
for v in input.iter_mut() {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
*v = (s >> 16) as i32 % 128 - 64;
}
for &(dc_q, ac_q) in QUANT_PAIRS {
let (sc, n) = run_32x32(input, dc_q, ac_q);
let first = sc.iter().zip(n.iter()).position(|(a, b)| a != b);
assert_eq!(
sc, n,
"32x32 rand(DEADBEEF) dc_q={dc_q} ac_q={ac_q}: mismatch at {first:?}"
);
}
}
#[test]
fn test_32x32_random_seed1() {
let mut input = [0i32; 1024];
let mut s = 0x1234_5678u32;
for v in input.iter_mut() {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
*v = (s >> 16) as i32 % 128 - 64;
}
for &(dc_q, ac_q) in QUANT_PAIRS {
let (sc, n) = run_32x32(input, dc_q, ac_q);
assert_eq!(sc, n, "32x32 rand(12345678) dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_32x32_high_amplitude_parity() {
let amp = 4095i32;
let mk = |f: &dyn Fn(usize, usize) -> i32| {
let mut input = [0i32; 1024];
for y in 0..32 {
for x in 0..32 {
input[y * 32 + x] = f(x, y);
}
}
input
};
let checker = mk(&|x, y| if (x + y) & 1 == 0 { amp } else { -amp });
let vstripe = mk(&|x, _| if x & 1 == 0 { amp } else { -amp });
let mut rnd = [0i32; 1024];
let mut s = 0xABCD_1234u32;
for v in rnd.iter_mut() {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
*v = (s >> 16) as i32 % (2 * amp + 1) - amp;
}
for input in [checker, vstripe, rnd] {
for &(dc_q, ac_q) in QUANT_PAIRS {
let (sc, n) = run_32x32(input, dc_q, ac_q);
assert_eq!(sc, n, "32x32 high-amplitude dc_q={dc_q} ac_q={ac_q}");
}
}
}
fn run_8x16(input: [i32; 128], dc_q: i32, ac_q: i32) -> ([i32; 128], [i32; 128]) {
let mut scalar = input;
dct8x16_i32_scalar(&mut scalar, dc_q, ac_q);
let mut neon = input;
unsafe { dct8x16_neon_i32(&mut neon, dc_q, ac_q) };
(scalar, neon)
}
#[test]
fn test_8x16_zeros() {
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_8x16([0i32; 128], dc_q, ac_q);
assert_eq!(s, n, "8x16 zeros dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_8x16_constant() {
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_8x16([64i32; 128], dc_q, ac_q);
assert_eq!(s, n, "8x16 constant dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_8x16_ramp() {
let mut input = [0i32; 128];
fill_ramp(&mut input);
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_8x16(input, dc_q, ac_q);
assert_eq!(s, n, "8x16 ramp dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_8x16_alternating() {
let mut input = [0i32; 128];
fill_alt(&mut input);
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_8x16(input, dc_q, ac_q);
assert_eq!(s, n, "8x16 alternating dc_q={dc_q} ac_q={ac_q}");
}
}
#[test]
fn test_8x16_random_seed0() {
let mut input = [0i32; 128];
fill_lcg(&mut input, 0xDEAD_BEEF);
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_8x16(input, dc_q, ac_q);
let first = s.iter().zip(n.iter()).position(|(a, b)| a != b);
assert_eq!(
s, n,
"8x16 rand(DEADBEEF) dc_q={dc_q} ac_q={ac_q}: mismatch at {first:?}"
);
}
}
#[test]
fn test_8x16_random_seed1() {
let mut input = [0i32; 128];
fill_lcg(&mut input, 0x1234_5678);
for &(dc_q, ac_q) in QUANT_PAIRS {
let (s, n) = run_8x16(input, dc_q, ac_q);
assert_eq!(s, n, "8x16 rand(12345678) dc_q={dc_q} ac_q={ac_q}");
}
}
}