#[allow(unsafe_code)]
pub fn inverse_dct_4x4(coeffs: &mut [i32; 16]) {
#[cfg(target_arch = "aarch64")]
{
unsafe {
inverse_dct_4x4_neon(coeffs);
}
return;
}
#[cfg(target_arch = "x86_64")]
{
if yscv_cpu::host_cpu().features.sse2 {
unsafe {
inverse_dct_4x4_sse2(coeffs);
}
return;
}
}
#[allow(unreachable_code)]
inverse_dct_4x4_scalar(coeffs);
}
fn inverse_dct_4x4_scalar(coeffs: &mut [i32; 16]) {
for i in 0..4 {
let base = i * 4;
let s0 = coeffs[base];
let s1 = coeffs[base + 1];
let s2 = coeffs[base + 2];
let s3 = coeffs[base + 3];
let e0 = s0 + s2;
let e1 = s0 - s2;
let e2 = (s1 >> 1) - s3;
let e3 = s1 + (s3 >> 1);
coeffs[base] = e0 + e3;
coeffs[base + 1] = e1 + e2;
coeffs[base + 2] = e1 - e2;
coeffs[base + 3] = e0 - e3;
}
for j in 0..4 {
let s0 = coeffs[j];
let s1 = coeffs[4 + j];
let s2 = coeffs[8 + j];
let s3 = coeffs[12 + j];
let e0 = s0 + s2;
let e1 = s0 - s2;
let e2 = (s1 >> 1) - s3;
let e3 = s1 + (s3 >> 1);
coeffs[j] = (e0 + e3 + 32) >> 6;
coeffs[4 + j] = (e1 + e2 + 32) >> 6;
coeffs[8 + j] = (e1 - e2 + 32) >> 6;
coeffs[12 + j] = (e0 - e3 + 32) >> 6;
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[allow(unsafe_code, unsafe_op_in_unsafe_fn)]
unsafe fn inverse_dct_4x4_neon(coeffs: &mut [i32; 16]) {
use std::arch::aarch64::*;
let ptr = coeffs.as_mut_ptr();
for i in 0..4 {
let row = vld1q_s32(ptr.add(i * 4));
let s0 = vgetq_lane_s32(row, 0);
let s1 = vgetq_lane_s32(row, 1);
let s2 = vgetq_lane_s32(row, 2);
let s3 = vgetq_lane_s32(row, 3);
let e0 = s0 + s2;
let e1 = s0 - s2;
let e2 = (s1 >> 1) - s3;
let e3 = s1 + (s3 >> 1);
let out = [e0 + e3, e1 + e2, e1 - e2, e0 - e3];
vst1q_s32(ptr.add(i * 4), vld1q_s32(out.as_ptr()));
}
let r0 = vld1q_s32(ptr);
let r1 = vld1q_s32(ptr.add(4));
let r2 = vld1q_s32(ptr.add(8));
let r3 = vld1q_s32(ptr.add(12));
let t01_lo = vzipq_s32(r0, r2); let t01_hi = vzipq_s32(r1, r3); let col0 = vzipq_s32(t01_lo.0, t01_hi.0).0;
let col1 = vzipq_s32(t01_lo.0, t01_hi.0).1;
let col2 = vzipq_s32(t01_lo.1, t01_hi.1).0;
let col3 = vzipq_s32(t01_lo.1, t01_hi.1).1;
let _round = vdupq_n_s32(32);
for (col_vec, j) in [(col0, 0), (col1, 1), (col2, 2), (col3, 3)] {
let s0 = vgetq_lane_s32(col_vec, 0);
let s1 = vgetq_lane_s32(col_vec, 1);
let s2 = vgetq_lane_s32(col_vec, 2);
let s3 = vgetq_lane_s32(col_vec, 3);
let e0 = s0 + s2;
let e1 = s0 - s2;
let e2 = (s1 >> 1) - s3;
let e3 = s1 + (s3 >> 1);
*ptr.add(j) = (e0 + e3 + 32) >> 6;
*ptr.add(4 + j) = (e1 + e2 + 32) >> 6;
*ptr.add(8 + j) = (e1 - e2 + 32) >> 6;
*ptr.add(12 + j) = (e0 - e3 + 32) >> 6;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse2")]
#[allow(unsafe_code, unsafe_op_in_unsafe_fn)]
unsafe fn inverse_dct_4x4_sse2(coeffs: &mut [i32; 16]) {
use std::arch::x86_64::*;
let ptr = coeffs.as_mut_ptr();
for i in 0..4 {
let row = _mm_loadu_si128(ptr.add(i * 4) as *const __m128i);
let s0 = _mm_extract_epi32::<0>(row);
let s1 = _mm_extract_epi32::<1>(row);
let s2 = _mm_extract_epi32::<2>(row);
let s3 = _mm_extract_epi32::<3>(row);
let e0 = s0 + s2;
let e1 = s0 - s2;
let e2 = (s1 >> 1) - s3;
let e3 = s1 + (s3 >> 1);
let out = _mm_set_epi32(e0 - e3, e1 - e2, e1 + e2, e0 + e3);
_mm_storeu_si128(ptr.add(i * 4) as *mut __m128i, out);
}
for j in 0..4 {
let s0 = *ptr.add(j);
let s1 = *ptr.add(4 + j);
let s2 = *ptr.add(8 + j);
let s3 = *ptr.add(12 + j);
let e0 = s0 + s2;
let e1 = s0 - s2;
let e2 = (s1 >> 1) - s3;
let e3 = s1 + (s3 >> 1);
*ptr.add(j) = (e0 + e3 + 32) >> 6;
*ptr.add(4 + j) = (e1 + e2 + 32) >> 6;
*ptr.add(8 + j) = (e1 - e2 + 32) >> 6;
*ptr.add(12 + j) = (e0 - e3 + 32) >> 6;
}
}
const DEQUANT_SCALE: [[i32; 16]; 6] = [
[
10, 13, 10, 13, 13, 16, 13, 16, 10, 13, 10, 13, 13, 16, 13, 16,
],
[
11, 14, 11, 14, 14, 18, 14, 18, 11, 14, 11, 14, 14, 18, 14, 18,
],
[
13, 16, 13, 16, 16, 20, 16, 20, 13, 16, 13, 16, 16, 20, 16, 20,
],
[
14, 18, 14, 18, 18, 23, 18, 23, 14, 18, 14, 18, 18, 23, 18, 23,
],
[
16, 20, 16, 20, 20, 25, 20, 25, 16, 20, 16, 20, 20, 25, 20, 25,
],
[
18, 23, 18, 23, 23, 29, 23, 29, 18, 23, 18, 23, 23, 29, 23, 29,
],
];
pub(crate) const MAX_LEVEL: i32 = 1 << 17;
pub(crate) const MAX_COEFF: i32 = 1 << 15;
pub(crate) fn clamp_coeffs(coeffs: &mut [i32], bound: i32) {
for c in coeffs {
*c = (*c).clamp(-bound, bound);
}
}
#[allow(unsafe_code)]
pub fn dequant_4x4(coeffs: &mut [i32; 16], qp: i32) {
let qp = qp.clamp(0, 51);
let shift = (qp / 6) as u32;
let scale = &DEQUANT_SCALE[(qp % 6) as usize];
#[cfg(any(target_arch = "aarch64", all(target_arch = "arm", feature = "neon-v7")))]
if yscv_cpu::host_cpu().features.neon {
unsafe { dequant_4x4_neon(coeffs, scale, shift) };
return;
}
#[cfg(target_arch = "x86_64")]
{
let features = yscv_cpu::host_cpu().features;
if features.avx512f {
unsafe { dequant_4x4_avx512(coeffs, scale, shift) };
return;
}
if features.avx2 {
unsafe { dequant_4x4_avx2(coeffs, scale, shift) };
return;
}
}
dequant_4x4_scalar(coeffs, scale, shift);
}
fn dequant_4x4_scalar(coeffs: &mut [i32; 16], scale: &[i32; 16], shift: u32) {
for (c, &s) in coeffs.iter_mut().zip(scale) {
*c = (c.wrapping_mul(s) << shift).clamp(-MAX_COEFF, MAX_COEFF);
}
}
#[cfg(any(target_arch = "aarch64", all(target_arch = "arm", feature = "neon-v7")))]
#[target_feature(enable = "neon")]
#[allow(unsafe_code, unsafe_op_in_unsafe_fn)]
unsafe fn dequant_4x4_neon(coeffs: &mut [i32; 16], scale: &[i32; 16], shift: u32) {
#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
#[cfg(target_arch = "arm")]
use std::arch::arm::*;
let shift = vdupq_n_s32(shift as i32);
let (hi, lo) = (vdupq_n_s32(MAX_COEFF), vdupq_n_s32(-MAX_COEFF));
for i in (0..16).step_by(4) {
let c = vld1q_s32(coeffs.as_ptr().add(i));
let d = vshlq_s32(vmulq_s32(c, vld1q_s32(scale.as_ptr().add(i))), shift);
vst1q_s32(coeffs.as_mut_ptr().add(i), vmaxq_s32(vminq_s32(d, hi), lo));
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[allow(unsafe_code, unsafe_op_in_unsafe_fn)]
unsafe fn dequant_4x4_avx2(coeffs: &mut [i32; 16], scale: &[i32; 16], shift: u32) {
use std::arch::x86_64::*;
let count = _mm_cvtsi32_si128(shift as i32);
let (hi, lo) = (_mm256_set1_epi32(MAX_COEFF), _mm256_set1_epi32(-MAX_COEFF));
for i in [0, 8] {
let c = _mm256_loadu_si256(coeffs.as_ptr().add(i).cast());
let s = _mm256_loadu_si256(scale.as_ptr().add(i).cast());
let d = _mm256_sll_epi32(_mm256_mullo_epi32(c, s), count);
let d = _mm256_max_epi32(_mm256_min_epi32(d, hi), lo);
_mm256_storeu_si256(coeffs.as_mut_ptr().add(i).cast(), d);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
#[allow(unsafe_code, unsafe_op_in_unsafe_fn)]
unsafe fn dequant_4x4_avx512(coeffs: &mut [i32; 16], scale: &[i32; 16], shift: u32) {
use std::arch::x86_64::*;
let c = _mm512_loadu_si512(coeffs.as_ptr().cast());
let s = _mm512_loadu_si512(scale.as_ptr().cast());
let d = _mm512_sll_epi32(_mm512_mullo_epi32(c, s), _mm_cvtsi32_si128(shift as i32));
let d = _mm512_max_epi32(
_mm512_min_epi32(d, _mm512_set1_epi32(MAX_COEFF)),
_mm512_set1_epi32(-MAX_COEFF),
);
_mm512_storeu_si512(coeffs.as_mut_ptr().cast(), d);
}
pub(crate) const ZIGZAG_4X4: [(usize, usize); 16] = [
(0, 0),
(0, 1),
(1, 0),
(2, 0),
(1, 1),
(0, 2),
(0, 3),
(1, 2),
(2, 1),
(3, 0),
(3, 1),
(2, 2),
(1, 3),
(2, 3),
(3, 2),
(3, 3),
];
pub fn inverse_dct_8x8(coeffs: &mut [i32; 64]) {
#[inline]
const fn idct8_1d(e: [i32; 8]) -> [i32; 8] {
let a0 = e[0] + e[4];
let a4 = e[0] - e[4];
let a2 = (e[2] >> 1) - e[6];
let a6 = e[2] + (e[6] >> 1);
let a1 = -e[3] + e[5] - e[7] - (e[7] >> 1);
let a3 = e[1] + e[7] - e[3] - (e[3] >> 1);
let a5 = -e[1] + e[7] + e[5] + (e[5] >> 1);
let a7 = e[3] + e[5] + e[1] + (e[1] >> 1);
let b0 = a0 + a6;
let b2 = a4 + a2;
let b4 = a4 - a2;
let b6 = a0 - a6;
let b1 = a1 + (a7 >> 2);
let b3 = a3 + (a5 >> 2);
let b5 = (a3 >> 2) - a5;
let b7 = a7 - (a1 >> 2);
[
b0 + b7,
b2 + b5,
b4 + b3,
b6 + b1,
b6 - b1,
b4 - b3,
b2 - b5,
b0 - b7,
]
}
for i in 0..8 {
let base = i * 8;
let row = [
coeffs[base],
coeffs[base + 1],
coeffs[base + 2],
coeffs[base + 3],
coeffs[base + 4],
coeffs[base + 5],
coeffs[base + 6],
coeffs[base + 7],
];
let g = idct8_1d(row);
coeffs[base..base + 8].copy_from_slice(&g);
}
for j in 0..8 {
let col = [
coeffs[j],
coeffs[8 + j],
coeffs[16 + j],
coeffs[24 + j],
coeffs[32 + j],
coeffs[40 + j],
coeffs[48 + j],
coeffs[56 + j],
];
let g = idct8_1d(col);
for (k, &v) in g.iter().enumerate() {
coeffs[k * 8 + j] = (v + 32) >> 6;
}
}
}
const DEQUANT_SCALE_8X8: [[i32; 64]; 6] = [
[
20, 19, 25, 19, 20, 19, 25, 19, 19, 18, 24, 18, 19, 18, 24, 18, 25, 24, 32, 24, 25, 24, 32,
24, 19, 18, 24, 18, 19, 18, 24, 18, 20, 19, 25, 19, 20, 19, 25, 19, 19, 18, 24, 18, 19, 18,
24, 18, 25, 24, 32, 24, 25, 24, 32, 24, 19, 18, 24, 18, 19, 18, 24, 18,
],
[
22, 21, 28, 21, 22, 21, 28, 21, 21, 19, 26, 19, 21, 19, 26, 19, 28, 26, 35, 26, 28, 26, 35,
26, 21, 19, 26, 19, 21, 19, 26, 19, 22, 21, 28, 21, 22, 21, 28, 21, 21, 19, 26, 19, 21, 19,
26, 19, 28, 26, 35, 26, 28, 26, 35, 26, 21, 19, 26, 19, 21, 19, 26, 19,
],
[
26, 24, 33, 24, 26, 24, 33, 24, 24, 23, 31, 23, 24, 23, 31, 23, 33, 31, 42, 31, 33, 31, 42,
31, 24, 23, 31, 23, 24, 23, 31, 23, 26, 24, 33, 24, 26, 24, 33, 24, 24, 23, 31, 23, 24, 23,
31, 23, 33, 31, 42, 31, 33, 31, 42, 31, 24, 23, 31, 23, 24, 23, 31, 23,
],
[
28, 26, 35, 26, 28, 26, 35, 26, 26, 25, 33, 25, 26, 25, 33, 25, 35, 33, 45, 33, 35, 33, 45,
33, 26, 25, 33, 25, 26, 25, 33, 25, 28, 26, 35, 26, 28, 26, 35, 26, 26, 25, 33, 25, 26, 25,
33, 25, 35, 33, 45, 33, 35, 33, 45, 33, 26, 25, 33, 25, 26, 25, 33, 25,
],
[
32, 30, 40, 30, 32, 30, 40, 30, 30, 28, 38, 28, 30, 28, 38, 28, 40, 38, 51, 38, 40, 38, 51,
38, 30, 28, 38, 28, 30, 28, 38, 28, 32, 30, 40, 30, 32, 30, 40, 30, 30, 28, 38, 28, 30, 28,
38, 28, 40, 38, 51, 38, 40, 38, 51, 38, 30, 28, 38, 28, 30, 28, 38, 28,
],
[
36, 34, 46, 34, 36, 34, 46, 34, 34, 32, 43, 32, 34, 32, 43, 32, 46, 43, 58, 43, 46, 43, 58,
43, 34, 32, 43, 32, 34, 32, 43, 32, 36, 34, 46, 34, 36, 34, 46, 34, 34, 32, 43, 32, 34, 32,
43, 32, 46, 43, 58, 43, 46, 43, 58, 43, 34, 32, 43, 32, 34, 32, 43, 32,
],
];
pub fn dequant_8x8(coeffs: &mut [i32; 64], qp: i32) {
let qp = qp.clamp(0, 51);
let shift = (qp / 6) as u32;
let scale = &DEQUANT_SCALE_8X8[(qp % 6) as usize];
for i in 0..64 {
let qmul = (16i64 * scale[i] as i64) << shift;
let d = (coeffs[i] as i64 * qmul + 32) >> 6;
coeffs[i] = d.clamp(-MAX_COEFF as i64, MAX_COEFF as i64) as i32;
}
}
pub(crate) const ZIGZAG_8X8: [usize; 64] = [
0, 1, 8, 16, 9, 2, 3, 10, 17, 24, 32, 25, 18, 11, 4, 5, 12, 19, 26, 33, 40, 48, 41, 34, 27, 20,
13, 6, 7, 14, 21, 28, 35, 42, 49, 56, 57, 50, 43, 36, 29, 22, 15, 23, 30, 37, 44, 51, 58, 59,
52, 45, 38, 31, 39, 46, 53, 60, 61, 54, 47, 55, 62, 63,
];
pub(crate) fn unscan_8x8(scan_coeffs: &[i32; 64], out: &mut [i32; 64]) {
*out = [0i32; 64];
for (scan_idx, &val) in scan_coeffs.iter().enumerate() {
out[ZIGZAG_8X8[scan_idx]] = val;
}
}
pub(crate) fn unscan_4x4(scan_coeffs: &[i32], out: &mut [i32; 16]) {
*out = [0i32; 16];
for (scan_idx, &val) in scan_coeffs.iter().enumerate().take(16) {
let (r, c) = ZIGZAG_4X4[scan_idx];
out[r * 4 + c] = val;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[allow(unsafe_code)]
fn dequant_4x4_paths_match_the_scalar_reference() {
let levels = [
0,
1,
-1,
3276,
-2048,
1 << 17,
-(1 << 20),
i32::MAX,
i32::MIN,
0x5555_5555,
];
for qp in 0..=51 {
let (shift, scale) = ((qp / 6) as u32, &DEQUANT_SCALE[(qp % 6) as usize]);
for k in 0..levels.len() {
let input: [i32; 16] = std::array::from_fn(|i| levels[(k + i) % levels.len()]);
let mut expected = input;
dequant_4x4_scalar(&mut expected, scale, shift);
assert!(expected.iter().all(|v| v.abs() <= MAX_COEFF));
let mut got = input;
dequant_4x4(&mut got, qp);
assert_eq!(got, expected, "dispatch, qp {qp}");
#[cfg(any(target_arch = "aarch64", all(target_arch = "arm", feature = "neon-v7")))]
if yscv_cpu::host_cpu().features.neon {
let mut got = input;
unsafe { dequant_4x4_neon(&mut got, scale, shift) };
assert_eq!(got, expected, "neon, qp {qp}");
}
#[cfg(target_arch = "x86_64")]
{
if yscv_cpu::host_cpu().features.avx2 {
let mut got = input;
unsafe { dequant_4x4_avx2(&mut got, scale, shift) };
assert_eq!(got, expected, "avx2, qp {qp}");
}
if yscv_cpu::host_cpu().features.avx512f {
let mut got = input;
unsafe { dequant_4x4_avx512(&mut got, scale, shift) };
assert_eq!(got, expected, "avx512, qp {qp}");
}
}
}
}
}
}