pub(crate) const SCALE: u32 = 12;
const COLUMN_SHIFT: u32 = SCALE - 2;
const ROW_SHIFT: u32 = SCALE + 5;
pub(crate) const fn rounding(bits: u32) -> i64 {
1 << (bits - 1)
}
const fn fixed(value: f64) -> i64 {
(value * (1_i64 << SCALE) as f64 + 0.5) as i64
}
pub(crate) const C_0_298: i64 = fixed(0.298_631_336);
pub(crate) const C_0_390: i64 = fixed(0.390_180_644);
pub(crate) const C_0_541: i64 = fixed(0.541_196_100);
pub(crate) const C_0_765: i64 = fixed(0.765_366_865);
pub(crate) const C_0_899: i64 = fixed(0.899_976_223);
pub(crate) const C_1_175: i64 = fixed(1.175_875_602);
pub(crate) const C_1_501: i64 = fixed(1.501_321_110);
pub(crate) const C_1_847: i64 = fixed(1.847_759_065);
pub(crate) const C_1_961: i64 = fixed(1.961_570_560);
pub(crate) const C_2_053: i64 = fixed(2.053_119_869);
pub(crate) const C_2_562: i64 = fixed(2.562_915_447);
pub(crate) const C_3_072: i64 = fixed(3.072_711_026);
fn transform([s0, s1, s2, s3, s4, s5, s6, s7]: [i64; 8]) -> [i64; 8] {
let shared = (s2 + s6) * C_0_541;
let even2 = shared - (s6 * C_1_847);
let even3 = shared + (s2 * C_0_765);
let sum = (s0 + s4) << SCALE;
let difference = (s0 - s4) << SCALE;
let x0 = sum + even3;
let x3 = sum - even3;
let x1 = difference + even2;
let x2 = difference - even2;
let a = s7 + s3;
let b = s5 + s1;
let c = s7 + s1;
let d = s5 + s3;
let common = (c + d) * C_1_175;
let p1 = common - (c * C_0_899);
let p2 = common - (d * C_2_562);
let p3 = -(a * C_1_961);
let p4 = -(b * C_0_390);
let y0 = (s7 * C_0_298) + p1 + p3;
let y1 = (s5 * C_2_053) + p2 + p4;
let y2 = (s3 * C_3_072) + p2 + p3;
let y3 = (s1 * C_1_501) + p1 + p4;
[
x0 + y3,
x1 + y2,
x2 + y1,
x3 + y0,
x3 - y0,
x2 - y1,
x1 - y2,
x0 - y3,
]
}
pub fn block(coefficients: &[i32; 64], out: &mut [u8], offset: usize, stride: usize) {
let mut columns = [0_i64; 64];
for column in 0..8 {
let input: [i64; 8] =
std::array::from_fn(|row| i64::from(*coefficients.get(column + row * 8).unwrap_or(&0)));
let [dc, rest @ ..] = input;
if rest.iter().all(|&value| value == 0) {
let flat = ((dc << SCALE) + rounding(COLUMN_SHIFT)) >> COLUMN_SHIFT;
for row in 0..8 {
if let Some(slot) = columns.get_mut(column + row * 8) {
*slot = flat;
}
}
continue;
}
let output = transform(input);
for (row, value) in output.into_iter().enumerate() {
if let Some(slot) = columns.get_mut(column + row * 8) {
*slot = (value + rounding(COLUMN_SHIFT)) >> COLUMN_SHIFT;
}
}
}
for row in 0..8 {
let input: [i64; 8] =
std::array::from_fn(|column| *columns.get(column + row * 8).unwrap_or(&0));
let output = transform(input);
let Some(target) = out.get_mut(offset + row * stride..) else {
continue;
};
for (slot, value) in target.iter_mut().take(8).zip(output) {
let shifted = (value + rounding(ROW_SHIFT)) >> ROW_SHIFT;
*slot = (shifted + 128).clamp(0, 255) as u8;
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
#[non_exhaustive]
pub enum Scale {
Eighth,
Quarter,
Half,
#[default]
Full,
}
impl Scale {
pub const ALL: [Self; 4] = [Self::Eighth, Self::Quarter, Self::Half, Self::Full];
#[must_use]
pub const fn block_size(self) -> u32 {
match self {
Self::Eighth => 1,
Self::Quarter => 2,
Self::Half => 4,
Self::Full => 8,
}
}
#[must_use]
pub const fn apply(self, source: u32) -> u32 {
let scaled = source.div_ceil(8 / self.block_size());
if scaled == 0 { 1 } else { scaled }
}
#[must_use]
pub fn fitting(source: (u32, u32), target: (u32, u32)) -> Self {
Self::ALL
.into_iter()
.find(|scale| scale.apply(source.0) >= target.0 && scale.apply(source.1) >= target.1)
.unwrap_or(Self::Full)
}
}
const BOX_1: [[i64; 1]; 8] = [[2896], [0], [0], [0], [0], [0], [0], [0]];
const BOX_2: [[i64; 2]; 8] = [
[2896, 2896],
[2624, -2624],
[0, 0],
[-922, 922],
[0, 0],
[616, -616],
[0, 0],
[-522, 522],
];
const BOX_4: [[i64; 4]; 8] = [
[2896, 2896, 2896, 2896],
[3711, 1537, -1537, -3711],
[2676, -2676, -2676, 2676],
[1303, -3146, 3146, -1303],
[0, 0, 0, 0],
[-871, 2102, -2102, 871],
[-1108, 1108, 1108, -1108],
[-738, -306, 306, 738],
];
fn reduced<const M: usize>(
coefficients: &[i32; 64],
basis: &[[i64; M]; 8],
out: &mut [u8],
offset: usize,
stride: usize,
) {
let mut rows = [[0_i64; M]; 8];
for (v, row) in rows.iter_mut().enumerate() {
for (u, weights) in basis.iter().enumerate() {
let coefficient = i64::from(coefficients.get(v * 8 + u).copied().unwrap_or(0));
if coefficient == 0 {
continue;
}
for (slot, &weight) in row.iter_mut().zip(weights.iter()) {
*slot += coefficient * weight;
}
}
}
for y in 0..M {
let Some(target) = out.get_mut(offset + y * stride..) else {
continue;
};
for (x, slot) in target.iter_mut().take(M).enumerate() {
let mut sum = 0_i64;
for (v, weights) in basis.iter().enumerate() {
let value = rows.get(v).and_then(|row| row.get(x)).copied().unwrap_or(0);
if value != 0 {
sum += value * weights.get(y).copied().unwrap_or(0);
}
}
let shifted = (sum + rounding(2 * SCALE + 2)) >> (2 * SCALE + 2);
*slot = (shifted + 128).clamp(0, 255) as u8;
}
}
}
pub fn scaled_block(
coefficients: &[i32; 64],
scale: Scale,
out: &mut [u8],
offset: usize,
stride: usize,
) {
match scale {
Scale::Eighth => reduced(coefficients, &BOX_1, out, offset, stride),
Scale::Quarter => reduced(coefficients, &BOX_2, out, offset, stride),
Scale::Half => reduced(coefficients, &BOX_4, out, offset, stride),
Scale::Full => block(coefficients, out, offset, stride),
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::indexing_slicing,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
fn reference(coefficients: &[i32; 64]) -> [f64; 64] {
let mut out = [0.0; 64];
for y in 0..8 {
for x in 0..8 {
let mut sum = 0.0;
for v in 0..8 {
for u in 0..8 {
let cu = if u == 0 { 1.0 / 2.0_f64.sqrt() } else { 1.0 };
let cv = if v == 0 { 1.0 / 2.0_f64.sqrt() } else { 1.0 };
sum += cu
* cv
* f64::from(coefficients[v * 8 + u])
* (((2 * x + 1) as f64 * u as f64 * std::f64::consts::PI) / 16.0).cos()
* (((2 * y + 1) as f64 * v as f64 * std::f64::consts::PI) / 16.0).cos();
}
}
out[y * 8 + x] = sum / 4.0 + 128.0;
}
}
out
}
fn decode(coefficients: &[i32; 64]) -> Vec<u8> {
let mut out = vec![0_u8; 64];
block(coefficients, &mut out, 0, 8);
out
}
#[test]
fn a_dc_only_block_is_a_flat_grey() {
let mut coefficients = [0_i32; 64];
coefficients[0] = 8 * 16;
assert!(decode(&coefficients).iter().all(|&v| v == 144));
assert!(decode(&[0; 64]).iter().all(|&v| v == 128));
}
#[test]
fn output_matches_the_reference_idct_within_one_step() {
let mut cases: Vec<[i32; 64]> = Vec::new();
for index in [0, 1, 8, 9, 27, 63] {
let mut block = [0_i32; 64];
block[index] = 200;
cases.push(block);
let mut negative = [0_i32; 64];
negative[index] = -300;
cases.push(negative);
}
let mut ramp = [0_i32; 64];
for (index, slot) in ramp.iter_mut().enumerate() {
*slot = (index as i32 % 7) * 20 - 60;
}
ramp[0] = 400;
cases.push(ramp);
let mut noisy = [0_i32; 64];
let mut state = 12_345_u32;
for slot in &mut noisy {
state = state.wrapping_mul(1_103_515_245).wrapping_add(12_345);
*slot = ((state >> 16) as i32 % 512) - 256;
}
cases.push(noisy);
for coefficients in cases {
let ours = decode(&coefficients);
let theirs = reference(&coefficients);
for (index, (&got, &want)) in ours.iter().zip(theirs.iter()).enumerate() {
let want = want.clamp(0.0, 255.0);
assert!(
(f64::from(got) - want).abs() <= 1.0,
"sample {index}: got {got}, reference {want:.3}"
);
}
}
}
#[test]
fn the_flat_column_shortcut_agrees_with_the_full_transform() {
let mut coefficients = [0_i32; 64];
coefficients[0] = 300;
coefficients[3] = -120;
let ours = decode(&coefficients);
let theirs = reference(&coefficients);
for (&got, &want) in ours.iter().zip(theirs.iter()) {
assert!((f64::from(got) - want.clamp(0.0, 255.0)).abs() <= 1.0);
}
}
#[test]
fn extreme_coefficients_clamp_instead_of_wrapping() {
let mut coefficients = [0_i32; 64];
coefficients[0] = -32_768;
assert!(decode(&coefficients).iter().all(|&v| v == 0));
coefficients[0] = 32_767;
assert!(decode(&coefficients).iter().all(|&v| v == 255));
}
#[test]
fn the_basis_tables_match_the_averages_they_stand_for() {
let expect = |m: usize, u: usize, x: usize| -> i64 {
let group = 8 / m;
let c = if u == 0 { 1.0 / 2.0_f64.sqrt() } else { 1.0 };
let sum: f64 = (0..group)
.map(|k| {
let t = x * group + k;
(((2 * t + 1) as f64 * u as f64 * std::f64::consts::PI) / 16.0).cos()
})
.sum();
(c * (m as f64 / 8.0) * sum * f64::from(1_i32 << SCALE)).round() as i64
};
for (u, row) in BOX_1.iter().enumerate() {
for (x, &value) in row.iter().enumerate() {
assert_eq!(value, expect(1, u, x), "BOX_1[{u}][{x}]");
}
}
for (u, row) in BOX_2.iter().enumerate() {
for (x, &value) in row.iter().enumerate() {
assert_eq!(value, expect(2, u, x), "BOX_2[{u}][{x}]");
}
}
for (u, row) in BOX_4.iter().enumerate() {
for (x, &value) in row.iter().enumerate() {
assert_eq!(value, expect(4, u, x), "BOX_4[{u}][{x}]");
}
}
}
#[test]
fn a_scaled_block_averages_to_what_the_full_block_averages_to() {
let mut state = 5_150_u32;
for _ in 0..40 {
let mut coefficients = [0_i32; 64];
for slot in &mut coefficients {
state = state.wrapping_mul(1_103_515_245).wrapping_add(12_345);
*slot = ((state >> 20) as i32 % 50) - 25;
}
coefficients[0] = 200;
let mut full = vec![0_u8; 64];
block(&coefficients, &mut full, 0, 8);
assert!(
full.iter().all(|&v| v > 0 && v < 255),
"the test block clamped; the mean comparison would be vacuous"
);
let mean = full.iter().map(|&v| u32::from(v)).sum::<u32>() as f64 / 64.0;
for scale in [Scale::Eighth, Scale::Quarter, Scale::Half] {
let m = scale.block_size() as usize;
let mut small = vec![0_u8; m * m];
scaled_block(&coefficients, scale, &mut small, 0, m);
let got = small.iter().map(|&v| u32::from(v)).sum::<u32>() as f64 / (m * m) as f64;
assert!(
(got - mean).abs() <= 1.5,
"{scale:?}: mean {got:.2} against the full block's {mean:.2}"
);
}
}
}
#[test]
fn a_flat_block_stays_flat_at_every_scale() {
let mut coefficients = [0_i32; 64];
coefficients[0] = 8 * 16;
for scale in Scale::ALL {
let m = scale.block_size() as usize;
let mut out = vec![0_u8; m * m];
scaled_block(&coefficients, scale, &mut out, 0, m);
assert!(
out.iter().all(|&v| v == 144),
"{scale:?}: {out:?} is not a flat 144"
);
}
}
#[test]
fn a_scaled_block_is_a_box_downsample_of_the_full_one() {
let mut coefficients = [0_i32; 64];
coefficients[0] = 400;
coefficients[1] = -180;
coefficients[8] = 120;
coefficients[9] = 60;
coefficients[2] = 40;
coefficients[16] = -30;
let mut full = vec![0_u8; 64];
block(&coefficients, &mut full, 0, 8);
for scale in [Scale::Quarter, Scale::Half] {
let m = scale.block_size() as usize;
let factor = 8 / m;
let mut small = vec![0_u8; m * m];
scaled_block(&coefficients, scale, &mut small, 0, m);
for y in 0..m {
for x in 0..m {
let mut sum = 0_u32;
for dy in 0..factor {
for dx in 0..factor {
sum += u32::from(full[(y * factor + dy) * 8 + x * factor + dx]);
}
}
let boxed = sum as f64 / (factor * factor) as f64;
let got = f64::from(small[y * m + x]);
assert!(
(got - boxed).abs() <= 1.5,
"{scale:?} at ({x},{y}): {got} against box average {boxed:.1}"
);
}
}
}
}
#[test]
fn scales_map_sizes_the_way_the_format_counts_blocks() {
assert_eq!(Scale::Eighth.apply(64), 8);
assert_eq!(Scale::Quarter.apply(64), 16);
assert_eq!(Scale::Half.apply(64), 32);
assert_eq!(Scale::Full.apply(64), 64);
assert_eq!(Scale::Eighth.apply(61), 8);
assert_eq!(Scale::Quarter.apply(61), 16);
assert_eq!(Scale::Eighth.apply(1), 1);
assert_eq!(Scale::Eighth.apply(3), 1);
}
#[test]
fn fitting_never_decodes_below_the_target() {
let source = (4000_u32, 3000_u32);
assert_eq!(Scale::fitting(source, (200, 150)), Scale::Eighth);
assert_eq!(Scale::fitting(source, (600, 450)), Scale::Quarter);
assert_eq!(Scale::fitting(source, (1200, 900)), Scale::Half);
assert_eq!(Scale::fitting(source, (3000, 2250)), Scale::Full);
assert_eq!(Scale::fitting(source, (9000, 9000)), Scale::Full);
assert_eq!(Scale::fitting(source, (200, 400)), Scale::Quarter);
}
#[test]
fn blocks_write_at_a_stride_and_clip_at_the_end_of_the_buffer() {
let mut coefficients = [0_i32; 64];
coefficients[0] = 8 * 16;
let mut plane = vec![0_u8; 16 * 8];
block(&coefficients, &mut plane, 8, 16);
for row in 0..8 {
assert_eq!(&plane[row * 16..row * 16 + 8], &[0_u8; 8]);
assert_eq!(&plane[row * 16 + 8..row * 16 + 16], &[144_u8; 8]);
}
let mut short = vec![0_u8; 8 * 3];
block(&coefficients, &mut short, 0, 8);
assert!(short.iter().all(|&v| v == 144));
}
}