use crate::tables::{DCT32, DST4_PAD, LEVEL_SCALE};
use rusty_h265_accel as accel;
pub const COEFF_MIN: i32 = -32768;
pub const COEFF_MAX: i32 = 32767;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TransformKind {
Dct,
Dst,
Skip,
Bypass,
}
pub fn dequant(coeffs: &mut [i32], n: usize, nz_w: usize, nz_h: usize, qp: i32, bit_depth: u8, m: Option<&[u8]>) {
let log2n = n.trailing_zeros() as i32;
let bd_shift = bit_depth as i32 + log2n - 5;
let scale = LEVEL_SCALE[(qp % 6) as usize] << (qp / 6);
let add = 1i64 << (bd_shift - 1);
for y in 0..nz_h {
for x in 0..nz_w {
let i = y * n + x;
let c = coeffs[i];
if c == 0 {
continue;
}
let f = m.map_or(16, |t| t[i] as i32);
let v = ((c as i64 * f as i64 * scale as i64) + add) >> bd_shift;
coeffs[i] = v.clamp(COEFF_MIN as i64, COEFF_MAX as i64) as i32;
}
}
}
fn no_dc_fast() -> bool {
static F: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*F.get_or_init(|| std::env::var_os("RH265_NO_DC_FAST").is_some())
}
fn idct_sums(src: &[i32], s_in: usize, n: usize, nz: usize, out: &mut [i32]) {
let levels = n.trailing_zeros() as usize - 2; let mut nzs = [0usize; 4];
let mut z = nz.min(n);
let mut m = n;
for slot in nzs.iter_mut().take(levels + 1) {
*slot = z;
if m > 4 {
m /= 2;
z = z.div_ceil(2).min(m);
}
}
let tab = DCT32.as_flattened();
accel::itx::accum(&mut out[..4], src, s_in << levels, tab, 8, 0, 1, nzs[levels], 4);
for d in (0..levels).rev() {
let m = n >> d;
accel::itx::accum_butterfly(&mut out[..m], src, s_in << d, tab, 32 >> m.trailing_zeros(), nzs[d], m);
}
}
#[cfg(test)]
fn idct_sums_naive(src: &[i32], s_in: usize, n: usize, nz: usize, out: &mut [i32]) {
let step = 32 >> n.trailing_zeros();
for (j, o) in out[..n].iter_mut().enumerate() {
let mut sum = 0i32;
for k in 0..nz.min(n) {
let c = src[k * s_in];
if c != 0 {
sum += c * DCT32[k * step][j] as i32;
}
}
*o = sum;
}
}
#[inline]
fn idct_1d<const CLIP: bool>(src: &[i32], s_in: usize, dst: &mut [i32], s_out: usize, n: usize, nz: usize, shift: u32) {
let mut sums = [0i32; 32];
idct_sums(src, s_in, n, nz, &mut sums[..n]);
if s_out == 1 {
accel::itx::shift_clip::<CLIP>(&mut dst[..n], &sums[..n], n, shift, COEFF_MIN, COEFF_MAX);
} else {
let add = 1i32 << (shift - 1);
for i in 0..n {
let v = (sums[i] + add) >> shift;
dst[i * s_out] = if CLIP { v.clamp(COEFF_MIN, COEFF_MAX) } else { v };
}
}
}
#[inline]
fn idst_1d<const CLIP: bool>(src: &[i32], s_in: usize, dst: &mut [i32], s_out: usize, nz: usize, shift: u32) {
let mut sum = [0i32; 4];
accel::itx::accum(&mut sum, src, s_in, DST4_PAD.as_flattened(), 1, 0, 1, nz.min(4), 4);
if s_out == 1 {
accel::itx::shift_clip::<CLIP>(&mut dst[..4], &sum, 4, shift, COEFF_MIN, COEFF_MAX);
} else {
let add = 1i32 << (shift - 1);
for i in 0..4 {
let v = (sum[i] + add) >> shift;
dst[i * s_out] = if CLIP { v.clamp(COEFF_MIN, COEFF_MAX) } else { v };
}
}
}
pub fn inverse_transform(d: &mut [i32], tmp: &mut [i32], n: usize, nz_w: usize, nz_h: usize, bit_depth: u8, kind: TransformKind) {
let bd_shift = 20 - bit_depth as u32;
match kind {
TransformKind::Bypass => {}
TransformKind::Skip => rusty_h265_accel::pixel::transform_skip(d, n, bd_shift),
TransformKind::Dct if nz_w.max(1) == 1 && nz_h.max(1) == 1 && !no_dc_fast() => {
if accel::census::ALWAYS {
accel::census::arm(&accel::census::RT_TX_DC_ONLY);
}
let v1 = (((d[0] as i64 * 64 + 64) >> 7) as i32).clamp(COEFF_MIN, COEFF_MAX);
let add = 1i64 << (bd_shift - 1);
let out = ((v1 as i64 * 64 + add) >> bd_shift) as i32;
d[..n * n].fill(out);
}
TransformKind::Dct | TransformKind::Dst => {
if accel::census::ALWAYS {
accel::census::arm(&accel::census::RT_TX_GENERAL);
}
let nz_w = nz_w.clamp(1, n);
let nz_h = nz_h.clamp(1, n);
for x in 0..nz_w {
if kind == TransformKind::Dst {
idst_1d::<true>(&d[x..], n, &mut tmp[x..], n, nz_h, 7);
} else {
idct_1d::<true>(&d[x..], n, &mut tmp[x..], n, n, nz_h, 7);
}
}
for y in 0..n {
if kind == TransformKind::Dst {
idst_1d::<false>(&tmp[y * n..], 1, &mut d[y * n..], 1, nz_w, bd_shift);
} else {
idct_1d::<false>(&tmp[y * n..], 1, &mut d[y * n..], 1, n, nz_w, bd_shift);
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dc_only_block_is_flat() {
let mut tmp = vec![0i32; 32 * 32];
for &n in &[4usize, 8, 16, 32] {
let mut d = vec![0i32; n * n];
d[0] = 64 * 8;
inverse_transform(&mut d, &mut tmp, n, 1, 1, 8, TransformKind::Dct);
let v = d[0];
assert!(d.iter().all(|&x| x == v), "n={n}");
assert_eq!(v, 4, "n={n}");
}
}
#[test]
fn sparse_bounds_match_dense() {
let mut tmp = vec![0i32; 32 * 32];
let mut st = 0x1234_5678u32;
let rnd = |s: &mut u32| {
*s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((*s >> 16) as i32 & 0x1ff) - 256
};
for &n in &[4usize, 8, 16, 32] {
for &(nz_w, nz_h) in &[(1usize, 1usize), (2, 3), (4, 4), (n / 2, n / 4), (n, n)] {
let (nz_w, nz_h) = (nz_w.min(n).max(1), nz_h.min(n).max(1));
let mut dense = vec![0i32; n * n];
for y in 0..nz_h {
for x in 0..nz_w {
dense[y * n + x] = rnd(&mut st);
}
}
let mut sparse = dense.clone();
for &kind in &[TransformKind::Dct, TransformKind::Dst] {
if kind == TransformKind::Dst && n != 4 {
continue;
}
let mut a = dense.clone();
let mut b = sparse.clone();
inverse_transform(&mut a, &mut tmp, n, n, n, 8, kind);
inverse_transform(&mut b, &mut tmp, n, nz_w, nz_h, 8, kind);
assert_eq!(a, b, "n={n} nz=({nz_w},{nz_h}) {kind:?}");
}
sparse[0] = dense[0];
}
}
}
#[test]
fn dequant_flat_qp() {
let mut c = vec![0i32; 16];
c[0] = 10;
dequant(&mut c, 4, 4, 4, 4, 8, None);
assert_eq!(c[0], 320);
}
#[test]
fn dequant_bounds_match_dense() {
let mut a = vec![0i32; 64];
let mut b = vec![0i32; 64];
for (i, (x, y)) in a.iter_mut().zip(b.iter_mut()).enumerate() {
if i % 8 < 3 && i / 8 < 2 {
*x = (i as i32) - 30;
*y = *x;
}
}
dequant(&mut a, 8, 8, 8, 17, 8, None);
dequant(&mut b, 8, 3, 2, 17, 8, None);
assert_eq!(a, b);
}
#[test]
fn transform_skip_scales() {
let mut tmp = vec![0i32; 32 * 32];
let mut d = vec![0i32; 16];
d[5] = 3;
inverse_transform(&mut d, &mut tmp, 4, 4, 4, 8, TransformKind::Skip);
assert_eq!(d[5], 0);
let mut d = vec![0i32; 16];
d[5] = 40;
inverse_transform(&mut d, &mut tmp, 4, 4, 4, 8, TransformKind::Skip);
assert_eq!(d[5], (40 * 128 + 2048) >> 12);
}
#[test]
fn transform_matrix_has_the_butterfly_symmetry() {
for &n in &[4usize, 8, 16, 32] {
let step = 32 / n;
for k in 0..n {
for j in 0..n {
let a = DCT32[k * step][j];
let b = DCT32[k * step][n - 1 - j];
if k % 2 == 0 {
assert_eq!(a, b, "even row {k} of the {n}-point matrix is not symmetric at {j}");
} else {
assert_eq!(a, -b, "odd row {k} of the {n}-point matrix is not antisymmetric at {j}");
}
}
}
}
}
#[test]
fn butterfly_matches_naive() {
let mut st = 0x51ee_7a11u32;
let rnd = |s: &mut u32| {
*s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((*s >> 8) as i32 % 65536) - 32768
};
for &n in &[4usize, 8, 16, 32] {
for nz in 1..=n {
for stride in [1usize, n, n + 3] {
let src: Vec<i32> = (0..n * stride + 8).map(|_| rnd(&mut st)).collect();
let mut a = vec![0i32; n];
let mut b = vec![0i32; n];
idct_sums_naive(&src, stride, n, nz, &mut a);
idct_sums(&src, stride, n, nz, &mut b);
assert_eq!(a, b, "n={n} nz={nz} stride={stride}");
}
}
}
}
#[test]
fn dc_only_matches_the_general_transform() {
let mut st = 0x0dc0_0001u32;
let mut tmp = vec![0i32; 32 * 32];
for &bd in &[8u8, 10] {
for &n in &[4usize, 8, 16, 32] {
for _ in 0..32 {
st = st.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
let dc = ((st >> 8) as i32 % 65536) - 32768;
let mut want = vec![0i32; n * n];
let mut got = vec![0i32; n * n];
want[0] = dc;
got[0] = dc;
inverse_transform(&mut want, &mut tmp, n, 2.min(n), 2.min(n), bd, TransformKind::Dct);
inverse_transform(&mut got, &mut tmp, n, 1, 1, bd, TransformKind::Dct);
assert_eq!(want, got, "n={n} bd={bd} dc={dc}");
}
}
}
}
#[test]
fn transform_accumulator_fits_i32() {
let worst_col = (0..32).map(|j| (0..32).map(|k| (DCT32[k][j] as i64).abs()).sum::<i64>()).max().unwrap();
let worst = COEFF_MAX.max(-COEFF_MIN) as i64 * worst_col;
assert!(worst <= i32::MAX as i64, "accumulator needs i64: worst case {worst} against {}", i32::MAX);
let headroom = i32::MAX as i64 / worst;
assert!(headroom >= 4, "only {headroom}x headroom left; re-derive before narrowing further");
eprintln!("transform accumulator: worst {worst}, headroom {headroom}x");
}
}